diff --git a/server/main.py b/server/main.py index 947a83b8..2534dcb7 100644 --- a/server/main.py +++ b/server/main.py @@ -12,6 +12,7 @@ from server.routers import router from server.utils.lifespan import lifespan from server.utils.auth_middleware import is_public_path from server.utils.common_utils import setup_logging +from server.utils.access_log_middleware import AccessLogMiddleware # 设置日志配置 setup_logging() @@ -115,6 +116,9 @@ class AuthMiddleware(BaseHTTPMiddleware): return await call_next(request) +# 添加访问日志中间件(记录请求处理时间) +app.add_middleware(AccessLogMiddleware) + # 添加鉴权中间件 app.add_middleware(LoginRateLimitMiddleware) app.add_middleware(AuthMiddleware) diff --git a/server/utils/access_log_middleware.py b/server/utils/access_log_middleware.py new file mode 100644 index 00000000..af72768b --- /dev/null +++ b/server/utils/access_log_middleware.py @@ -0,0 +1,67 @@ +"""访问日志中间件 - 记录请求处理时间""" + +import time +import logging +from collections.abc import Callable + +from fastapi import Request, Response +from starlette.middleware.base import BaseHTTPMiddleware + +# 创建专用的访问日志记录器 +access_logger = logging.getLogger("access_logger") + +# 设置访问日志记录器 +if not access_logger.handlers: + handler = logging.StreamHandler() + formatter = logging.Formatter(fmt="%(asctime)s %(levelname)s: %(message)s", datefmt="%m-%d %H:%M:%S") + handler.setFormatter(formatter) + access_logger.addHandler(handler) + access_logger.setLevel(logging.INFO) + # 避免传播到根日志记录器,防止重复日志 + access_logger.propagate = False + + +def _extract_client_ip(request: Request) -> str: + """提取客户端IP地址""" + forwarded_for = request.headers.get("x-forwarded-for") + if forwarded_for: + return forwarded_for.split(",")[0].strip() + if request.client: + return request.client.host + return "unknown" + + +class AccessLogMiddleware(BaseHTTPMiddleware): + """访问日志中间件 - 记录请求处理时间""" + + def __init__(self, app, logger: logging.Logger = None): + super().__init__(app) + self.logger = logger or access_logger + + async def dispatch(self, request: Request, call_next: Callable) -> Response: + """处理请求并记录访问日志""" + # 记录请求开始时间 + start_time = time.perf_counter() + + # 获取客户端IP + client_ip = _extract_client_ip(request) + + # 处理请求 + response = await call_next(request) + + # 计算处理时间 + process_time = time.perf_counter() - start_time + process_time_ms = int(process_time * 1000) # 转换为毫秒 + + # 格式化日志消息,添加处理时间 + log_message = ( + f"{client_ip}:{request.client.port if request.client else 'unknown'} - " + f'"{request.method} {request.url.path}{"?" + request.url.query if request.url.query else ""} ' + f'HTTP/{request.scope["http_version"]}" ' + f"{response.status_code} - {process_time_ms}ms" + ) + + # 记录日志 + self.logger.info(log_message) + + return response diff --git a/server/utils/common_utils.py b/server/utils/common_utils.py index ba6b8bb3..02d68576 100644 --- a/server/utils/common_utils.py +++ b/server/utils/common_utils.py @@ -19,14 +19,15 @@ def setup_logging(): uvicorn_logger = logging.getLogger("uvicorn") uvicorn_access_logger = logging.getLogger("uvicorn.access") + # 禁用默认的uvicorn访问日志(因为我们使用自定义中间件) + uvicorn_access_logger.handlers.clear() + # 创建格式化器 formatter = logging.Formatter(fmt="%(asctime)s %(levelname)s: %(message)s", datefmt="%m-%d %H:%M:%S") - # 为所有处理器设置格式化器 + # 为uvicorn主日志设置格式化器 for handler in uvicorn_logger.handlers: handler.setFormatter(formatter) - for handler in uvicorn_access_logger.handlers: - handler.setFormatter(formatter) async def log_operation(db: Session, user_id: int, operation: str, details: str = None, request: Request = None): diff --git a/src/agents/common/models.py b/src/agents/common/models.py index b40b9946..e1b715c0 100644 --- a/src/agents/common/models.py +++ b/src/agents/common/models.py @@ -32,7 +32,6 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: logger.debug(f"[offical] Loading model {model_spec} with kwargs {kwargs}") return init_chat_model(model_spec, **kwargs) - elif provider in ["dashscope"]: from langchain_deepseek import ChatDeepSeek diff --git a/src/agents/deep_agent/graph.py b/src/agents/deep_agent/graph.py index 97760209..a95fb572 100644 --- a/src/agents/deep_agent/graph.py +++ b/src/agents/deep_agent/graph.py @@ -1,14 +1,11 @@ """Deep Agent - 基于create_deep_agent的深度分析智能体""" -from langchain.agents.middleware import ModelRequest, dynamic_prompt, SummarizationMiddleware - -from langchain.agents import create_agent -from langchain.agents.middleware import TodoListMiddleware -from langchain_anthropic.middleware import AnthropicPromptCachingMiddleware - from deepagents.middleware.filesystem import FilesystemMiddleware from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware from deepagents.middleware.subagents import SubAgentMiddleware +from langchain.agents import create_agent +from langchain.agents.middleware import ModelRequest, SummarizationMiddleware, TodoListMiddleware, dynamic_prompt +from langchain_anthropic.middleware import AnthropicPromptCachingMiddleware from src.agents.common import BaseAgent, load_chat_model from src.agents.common.middlewares import context_based_model, inject_attachment_context @@ -21,9 +18,7 @@ search_tools = [search] research_sub_agent = { "name": "research-agent", - "description": ( - "利用搜索工具,用于研究更深入的问题。" - ), + "description": ("利用搜索工具,用于研究更深入的问题。"), "system_prompt": ( "你是一位专注的研究员。你的工作是根据用户的问题进行研究。" "进行彻底的研究,然后用详细的答案回复用户的问题,只有你的最终答案会被传递给用户。" diff --git a/src/knowledge/implementations/lightrag.py b/src/knowledge/implementations/lightrag.py index cdddd5f2..d62d1f44 100644 --- a/src/knowledge/implementations/lightrag.py +++ b/src/knowledge/implementations/lightrag.py @@ -235,7 +235,7 @@ class LightRagKB(KnowledgeBase): model=model_name, api_key=config_dict["api_key"], base_url=config_dict["base_url"].replace("/embeddings", ""), - ) + ), ) async def add_content(self, db_id: str, items: list[str], params: dict | None = None) -> list[dict]: diff --git a/src/knowledge/utils/kb_utils.py b/src/knowledge/utils/kb_utils.py index d955fffd..75dbec58 100644 --- a/src/knowledge/utils/kb_utils.py +++ b/src/knowledge/utils/kb_utils.py @@ -319,7 +319,7 @@ def get_embedding_config(embed_info: dict) -> dict: "model": embed_info["name"], "api_key": os.getenv(embed_info["api_key"]) or embed_info["api_key"], "base_url": embed_info["base_url"], - "dimension": embed_info.get("dimension", 1024) + "dimension": embed_info.get("dimension", 1024), } logger.debug(f"Embedding config from dict: {config_dict}") return config_dict