style: 统一代码格式并修复格式问题

This commit is contained in:
Wenjie Zhang 2025-12-04 10:12:01 +08:00
parent 2126f45493
commit 100f9dac7f
12 changed files with 32 additions and 34 deletions

View File

@ -228,7 +228,7 @@ async def save_partial_message(conv_mgr, thread_id, full_msg=None, error_message
extra_metadata = { extra_metadata = {
"error_type": error_type, "error_type": error_type,
"is_error": True, "is_error": True,
"error_message": error_message or f"发生错误: {error_type}" "error_message": error_message or f"发生错误: {error_type}",
} }
if full_msg: if full_msg:
# 保存部分生成的AI消息 # 保存部分生成的AI消息
@ -522,10 +522,7 @@ async def chat_agent(
# Input guard # Input guard
if conf.enable_content_guard and await content_guard.check(query): if conf.enable_content_guard and await content_guard.check(query):
yield make_chunk( yield make_chunk(
status="error", status="error", error_type="content_guard_blocked", error_message="输入内容包含敏感词", meta=meta
error_type="content_guard_blocked",
error_message="输入内容包含敏感词",
meta=meta
) )
return return
@ -537,7 +534,7 @@ async def chat_agent(
status="error", status="error",
error_type="agent_error", error_type="agent_error",
error_message=f"智能体 {agent_id} 获取失败: {str(e)}", error_message=f"智能体 {agent_id} 获取失败: {str(e)}",
meta=meta meta=meta,
) )
return return
@ -679,12 +676,7 @@ async def chat_agent(
error_type=error_type, error_type=error_type,
) )
yield make_chunk( yield make_chunk(status="error", error_type=error_type, error_message=error_msg, meta=meta)
status="error",
error_type=error_type,
error_message=error_msg,
meta=meta
)
return StreamingResponse(stream_messages(), media_type="application/json") return StreamingResponse(stream_messages(), media_type="application/json")
@ -1043,7 +1035,9 @@ async def create_thread(
@chat.get("/threads", response_model=list[ThreadResponse]) @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 不能为空" 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}") @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 # Use new storage system
conv_manager = ConversationManager(db) conv_manager = ConversationManager(db)
@ -1288,7 +1284,9 @@ async def get_message_feedback(
"""Get feedback status for a specific message (for current user)""" """Get feedback status for a specific message (for current user)"""
try: try:
# Get user's feedback for this message # 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() feedback = feedback_result.scalar_one_or_none()
if not feedback: if not feedback:

View File

@ -60,6 +60,7 @@ async def get_current_user(token: str | None = Depends(oauth2_scheme), db: Async
# 查找用户(异步版本) # 查找用户(异步版本)
from sqlalchemy import select from sqlalchemy import select
result = await db.execute(select(User).filter(User.id == user_id)) result = await db.execute(select(User).filter(User.id == user_id))
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
if user is None: if user is None:
@ -78,6 +79,7 @@ async def get_required_user(user: User | None = Depends(get_current_user)):
) )
return user return user
# 获取管理员用户 # 获取管理员用户
async def get_admin_user(current_user: User = Depends(get_required_user)): async def get_admin_user(current_user: User = Depends(get_required_user)):
if current_user.role not in ["admin", "superadmin"]: 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 return current_user
# 获取超级管理员用户 # 获取超级管理员用户
async def get_superadmin_user(current_user: User = Depends(get_required_user)): async def get_superadmin_user(current_user: User = Depends(get_required_user)):
if current_user.role != "superadmin": if current_user.role != "superadmin":

View File

@ -1,13 +1,12 @@
"""MCP Client setup and management for LangGraph ReAct Agent.""" """MCP Client setup and management for LangGraph ReAct Agent."""
import os
import traceback
from collections.abc import Callable from collections.abc import Callable
from typing import Any, cast from typing import Any, cast
import traceback
from langchain_mcp_adapters.client import MultiServerMCPClient from langchain_mcp_adapters.client import MultiServerMCPClient
from src.utils import logger from src.utils import logger
# Global MCP tools cache # Global MCP tools cache
@ -37,6 +36,7 @@ MCP_SERVERS = {
# 更多用法参考https://xerrors.github.io/Yuxi-Know/latest/advanced/agents-config.html#内置工具与-mcp-集成 # 更多用法参考https://xerrors.github.io/Yuxi-Know/latest/advanced/agents-config.html#内置工具与-mcp-集成
} }
async def get_mcp_client( async def get_mcp_client(
server_configs: dict[str, Any] | None = None, server_configs: dict[str, Any] | None = None,
) -> MultiServerMCPClient | None: ) -> MultiServerMCPClient | None:

View File

@ -4,7 +4,6 @@ from dataclasses import dataclass, field
from src.agents.common.context import BaseContext from src.agents.common.context import BaseContext
DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。 DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。
你应该做的第一件事是把原始的用户问题写入 `question.txt`以便你有一个记录 你应该做的第一件事是把原始的用户问题写入 `question.txt`以便你有一个记录
@ -90,6 +89,7 @@ DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的
你可以使用一些工具 你可以使用一些工具
""" """
@dataclass @dataclass
class DeepContext(BaseContext): class DeepContext(BaseContext):
""" """

View File

@ -1,7 +1,7 @@
"""Deep Agent - 基于create_deep_agent的深度分析智能体""" """Deep Agent - 基于create_deep_agent的深度分析智能体"""
from deepagents import 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 import BaseAgent, load_chat_model
from src.agents.common.middlewares import context_based_model, inject_attachment_context 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 from .prompts import DEEP_PROMPT
search_tools = [search] search_tools = [search]
@ -54,6 +53,7 @@ def context_aware_prompt(request: ModelRequest) -> str:
"""从 runtime context 动态生成系统提示词""" """从 runtime context 动态生成系统提示词"""
return DEEP_PROMPT + "\n\n\n" + request.runtime.context.system_prompt return DEEP_PROMPT + "\n\n\n" + request.runtime.context.system_prompt
class DeepAgent(BaseAgent): class DeepAgent(BaseAgent):
name = "深度分析智能体" name = "深度分析智能体"
description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务" description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务"

View File

@ -1,4 +1,3 @@
DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。 DEEP_PROMPT = """你是一位专家级研究员。你的工作是进行彻底的研究,然后撰写一份精美的报告。
你应该做的第一件事是把原始的用户问题写入 `question.txt`以便你有一个记录 你应该做的第一件事是把原始的用户问题写入 `question.txt`以便你有一个记录

View File

@ -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.agents.common.toolkits.mysql import get_mysql_tools
from src.utils import logger from src.utils import logger
_mcp_servers = { _mcp_servers = {"mcp-server-chart": {"command": "npx", "args": ["-y", "@antv/mcp-server-chart"], "transport": "stdio"}}
"mcp-server-chart": {
"command": "npx",
"args": ["-y", "@antv/mcp-server-chart"],
"transport": "stdio"
}
}
class SqlReporterAgent(BaseAgent): class SqlReporterAgent(BaseAgent):

View File

@ -31,6 +31,7 @@ class EmbedModelInfo(BaseModel):
api_key: str = Field(..., description="API Key 或环境变量名") api_key: str = Field(..., description="API Key 或环境变量名")
model_id: str | None = Field(None, description="可选的模型 ID") model_id: str | None = Field(None, description="可选的模型 ID")
class RerankerInfo(BaseModel): class RerankerInfo(BaseModel):
"""重排序模型配置""" """重排序模型配置"""

View File

@ -164,6 +164,7 @@ class MilvusKB(KnowledgeBase):
# 检查是否有 model_id 字段,优先使用 select_embedding_model # 检查是否有 model_id 字段,优先使用 select_embedding_model
if embed_info and "model_id" in embed_info: if embed_info and "model_id" in embed_info:
from src.models.embed import select_embedding_model from src.models.embed import select_embedding_model
return select_embedding_model(embed_info["model_id"]) return select_embedding_model(embed_info["model_id"])
# 使用原有的逻辑(兼容模式)) # 使用原有的逻辑(兼容模式))

View File

@ -1,8 +1,8 @@
import asyncio import asyncio
import os import os
import re import re
import zipfile
import xml.etree.ElementTree as ET import xml.etree.ElementTree as ET
import zipfile
from pathlib import Path from pathlib import Path
import aiofiles import aiofiles

View File

@ -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: if response.status_code != 200:
error_detail = "未知错误" error_detail = "未知错误"