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 = {
"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:

View File

@ -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:

View File

@ -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":

View File

@ -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:

View File

@ -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):
"""

View File

@ -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 = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务"

View File

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

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.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):

View File

@ -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):
"""重排序模型配置"""

View File

@ -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"])
# 使用原有的逻辑(兼容模式))

View File

@ -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

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:
error_detail = "未知错误"