diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index af165897..8604ce65 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -228,7 +228,7 @@ async def save_partial_message(conv_mgr, thread_id, full_msg=None, error_message extra_metadata = { "error_type": error_type, "is_error": True, - "error_message": error_message or f"发生错误: {error_type}" + "error_message": error_message or f"发生错误: {error_type}", } if full_msg: # 保存部分生成的AI消息 @@ -522,10 +522,7 @@ async def chat_agent( # Input guard if conf.enable_content_guard and await content_guard.check(query): yield make_chunk( - status="error", - error_type="content_guard_blocked", - error_message="输入内容包含敏感词", - meta=meta + status="error", error_type="content_guard_blocked", error_message="输入内容包含敏感词", meta=meta ) return @@ -537,7 +534,7 @@ async def chat_agent( status="error", error_type="agent_error", error_message=f"智能体 {agent_id} 获取失败: {str(e)}", - meta=meta + meta=meta, ) return @@ -679,12 +676,7 @@ async def chat_agent( error_type=error_type, ) - yield make_chunk( - status="error", - error_type=error_type, - error_message=error_msg, - meta=meta - ) + yield make_chunk(status="error", error_type=error_type, error_message=error_msg, meta=meta) return StreamingResponse(stream_messages(), media_type="application/json") @@ -1043,7 +1035,9 @@ async def create_thread( @chat.get("/threads", response_model=list[ThreadResponse]) -async def list_threads(agent_id: str, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_required_user)): +async def list_threads( + agent_id: str, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_required_user) +): """获取用户的所有对话线程 (使用新存储系统)""" assert agent_id, "agent_id 不能为空" @@ -1071,7 +1065,9 @@ async def list_threads(agent_id: str, db: AsyncSession = Depends(get_db), curren @chat.delete("/thread/{thread_id}") -async def delete_thread(thread_id: str, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_required_user)): +async def delete_thread( + thread_id: str, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_required_user) +): """删除对话线程 (使用新存储系统)""" # Use new storage system conv_manager = ConversationManager(db) @@ -1288,7 +1284,9 @@ async def get_message_feedback( """Get feedback status for a specific message (for current user)""" try: # Get user's feedback for this message - feedback_result = await db.execute(select(MessageFeedback).filter_by(message_id=message_id, user_id=str(current_user.id))) + feedback_result = await db.execute( + select(MessageFeedback).filter_by(message_id=message_id, user_id=str(current_user.id)) + ) feedback = feedback_result.scalar_one_or_none() if not feedback: diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index e9655fd3..782be605 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -173,7 +173,7 @@ async def update_database_info( name: str = Body(...), description: str = Body(...), llm_info: dict = Body(None), - additional_params: dict = Body({}), # Now accepts a dict + additional_params: dict = Body({}), # Now accepts a dict current_user: User = Depends(get_admin_user), ): """更新知识库信息""" @@ -187,7 +187,7 @@ async def update_database_info( name, description, llm_info, - additional_params=additional_params, # Pass the dict to the manager + additional_params=additional_params, # Pass the dict to the manager ) return {"message": "更新成功", "database": database} except Exception as e: diff --git a/server/utils/auth_middleware.py b/server/utils/auth_middleware.py index abf1ed97..2be55704 100644 --- a/server/utils/auth_middleware.py +++ b/server/utils/auth_middleware.py @@ -60,6 +60,7 @@ async def get_current_user(token: str | None = Depends(oauth2_scheme), db: Async # 查找用户(异步版本) from sqlalchemy import select + result = await db.execute(select(User).filter(User.id == user_id)) user = result.scalar_one_or_none() if user is None: @@ -78,6 +79,7 @@ async def get_required_user(user: User | None = Depends(get_current_user)): ) return user + # 获取管理员用户 async def get_admin_user(current_user: User = Depends(get_required_user)): if current_user.role not in ["admin", "superadmin"]: @@ -87,6 +89,7 @@ async def get_admin_user(current_user: User = Depends(get_required_user)): ) return current_user + # 获取超级管理员用户 async def get_superadmin_user(current_user: User = Depends(get_required_user)): if current_user.role != "superadmin": diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index 38f62dc5..a5f38945 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -1,13 +1,12 @@ """MCP Client setup and management for LangGraph ReAct Agent.""" -import os + +import traceback from collections.abc import Callable from typing import Any, cast -import traceback from langchain_mcp_adapters.client import MultiServerMCPClient - from src.utils import logger # Global MCP tools cache @@ -37,6 +36,7 @@ MCP_SERVERS = { # 更多用法参考:https://xerrors.github.io/Yuxi-Know/latest/advanced/agents-config.html#内置工具与-mcp-集成 } + async def get_mcp_client( server_configs: dict[str, Any] | None = None, ) -> MultiServerMCPClient | None: diff --git a/src/agents/deep_agent/context.py b/src/agents/deep_agent/context.py index 92438d8f..8fa18959 100644 --- a/src/agents/deep_agent/context.py +++ b/src/agents/deep_agent/context.py @@ -4,7 +4,6 @@ from dataclasses import dataclass, field from src.agents.common.context import BaseContext - DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。 你应该做的第一件事是把原始的用户问题写入 `question.txt`,以便你有一个记录。 @@ -90,6 +89,7 @@ DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的 你可以使用一些工具。 """ + @dataclass class DeepContext(BaseContext): """ diff --git a/src/agents/deep_agent/graph.py b/src/agents/deep_agent/graph.py index 4d58de88..57f3615d 100644 --- a/src/agents/deep_agent/graph.py +++ b/src/agents/deep_agent/graph.py @@ -1,7 +1,7 @@ """Deep Agent - 基于create_deep_agent的深度分析智能体""" from deepagents import create_deep_agent -from langchain.agents.middleware import dynamic_prompt, ModelRequest +from langchain.agents.middleware import ModelRequest, dynamic_prompt from src.agents.common import BaseAgent, load_chat_model from src.agents.common.middlewares import context_based_model, inject_attachment_context @@ -9,7 +9,6 @@ from src.agents.common.tools import search from .prompts import DEEP_PROMPT - search_tools = [search] @@ -54,6 +53,7 @@ def context_aware_prompt(request: ModelRequest) -> str: """从 runtime context 动态生成系统提示词""" return DEEP_PROMPT + "\n\n\n" + request.runtime.context.system_prompt + class DeepAgent(BaseAgent): name = "深度分析智能体" description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务" diff --git a/src/agents/deep_agent/prompts.py b/src/agents/deep_agent/prompts.py index b28550f8..f98d8185 100644 --- a/src/agents/deep_agent/prompts.py +++ b/src/agents/deep_agent/prompts.py @@ -1,4 +1,3 @@ - DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。 你应该做的第一件事是把原始的用户问题写入 `question.txt`,以便你有一个记录。 @@ -82,4 +81,4 @@ DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的 你可以使用一些工具。 -""" \ No newline at end of file +""" diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 7753a840..11f64cec 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -6,13 +6,7 @@ from src.agents.common.middlewares import context_aware_prompt, context_based_mo from src.agents.common.toolkits.mysql import get_mysql_tools from src.utils import logger -_mcp_servers = { - "mcp-server-chart": { - "command": "npx", - "args": ["-y", "@antv/mcp-server-chart"], - "transport": "stdio" - } -} +_mcp_servers = {"mcp-server-chart": {"command": "npx", "args": ["-y", "@antv/mcp-server-chart"], "transport": "stdio"}} class SqlReporterAgent(BaseAgent): diff --git a/src/config/static/models.py b/src/config/static/models.py index 53b5fcbf..e5aef3e8 100644 --- a/src/config/static/models.py +++ b/src/config/static/models.py @@ -31,6 +31,7 @@ class EmbedModelInfo(BaseModel): api_key: str = Field(..., description="API Key 或环境变量名") model_id: str | None = Field(None, description="可选的模型 ID") + class RerankerInfo(BaseModel): """重排序模型配置""" diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index 0b22d674..d2d70dae 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -164,6 +164,7 @@ class MilvusKB(KnowledgeBase): # 检查是否有 model_id 字段,优先使用 select_embedding_model if embed_info and "model_id" in embed_info: from src.models.embed import select_embedding_model + return select_embedding_model(embed_info["model_id"]) # 使用原有的逻辑(兼容模式)) diff --git a/src/knowledge/indexing.py b/src/knowledge/indexing.py index ed1eb54c..848ec4d3 100644 --- a/src/knowledge/indexing.py +++ b/src/knowledge/indexing.py @@ -1,8 +1,8 @@ import asyncio import os import re -import zipfile import xml.etree.ElementTree as ET +import zipfile from pathlib import Path import aiofiles diff --git a/src/plugins/mineru_parser.py b/src/plugins/mineru_parser.py index 63dd728b..30e7b6f1 100644 --- a/src/plugins/mineru_parser.py +++ b/src/plugins/mineru_parser.py @@ -159,7 +159,9 @@ class MinerUParser(BaseDocumentProcessor): ) # 检查响应状态 - logger.debug(f"MinerU 响应状态: {response.status_code}, Content-Type: {response.headers.get('content-type', 'unknown')}") + logger.debug( + f"MinerU 响应状态: {response.status_code}, Content-Type: {response.headers.get('content-type')}" + ) if response.status_code != 200: error_detail = "未知错误"