From 8073dd63a764a43bcf65ab1ae21ef123d18222ce Mon Sep 17 00:00:00 2001 From: Mplan Date: Fri, 3 Jul 2026 20:49:36 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E7=9B=AE=E5=BD=95=E7=BB=93?= =?UTF-8?q?=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 10 +++++- app.py | 79 ++++++------------------------------------ models.py => domain.py | 63 ++------------------------------- routes.py | 66 +++++++++++++++++++++++++++++++++++ schemas.py | 65 ++++++++++++++++++++++++++++++++++ service.py | 27 ++++++++++++--- storage.py | 2 +- tests/test_agent.py | 44 ++++++++++++++++------- tests/test_tools.py | 33 +++++++----------- tools.py | 69 +----------------------------------- 10 files changed, 220 insertions(+), 238 deletions(-) rename models.py => domain.py (51%) create mode 100644 routes.py create mode 100644 schemas.py diff --git a/README.md b/README.md index 3267a31..595fc5c 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,15 @@ 一个使用 FastAPI 和 DeepSeek API 的多轮 Agent 服务。每个会话的系统提示词、消息历史和工具调用过程保存在本地 JSON 文件中。 +## 代码结构 + +- `routes.py`:集中定义所有 HTTP 路由。 +- `schemas.py`:集中定义 API 请求和响应结构。 +- `domain.py`:定义会话、消息和工具调用的持久化模型。 +- `tools.py`:集中定义工具注册表、参数结构和工具函数。 +- `service.py`:处理对话、Agent 循环和业务规则。 +- `app.py`:创建 FastAPI 应用并管理生命周期。 + ## 启动 ```bash @@ -38,7 +47,6 @@ API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完 模型可以按需调用服务端白名单工具: - `get_current_time`:查询指定 IANA 时区的当前时间。 -- `calculate`:计算受限的基础算术表达式,不执行任意 Python 代码。 工具调用无需增加请求参数。服务端会执行工具并把结果返回模型,直到模型生成最终回答。发送消息接口通过 `tools_use` 返回本轮使用的工具名称,完整调用过程可通过历史接口查询;未调用工具时 `tools_use` 为空数组。 diff --git a/app.py b/app.py index 141d9d4..2b04c27 100644 --- a/app.py +++ b/app.py @@ -1,19 +1,13 @@ -from contextlib import asynccontextmanager from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager -from fastapi import FastAPI, HTTPException, Request, status +from fastapi import FastAPI from openai import AsyncOpenAI from config import Settings -from models import ( - CreateSessionRequest, - CreateSessionResponse, - SendMessageRequest, - SendMessageResponse, - SessionHistoryResponse, -) -from service import ChatProviderError, ChatService -from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError +from routes import router +from service import ChatService +from storage import JsonSessionStorage from tools import ToolRegistry @@ -22,7 +16,6 @@ def create_app( client: AsyncOpenAI | None = None, ) -> FastAPI: resolved_settings = settings or Settings.from_env() - storage = JsonSessionStorage(resolved_settings.data_dir) @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: @@ -34,75 +27,23 @@ def create_app( api_key=resolved_settings.api_key, base_url=resolved_settings.base_url, ) + app.state.chat_service = ChatService( - storage=storage, + storage=JsonSessionStorage(resolved_settings.data_dir), client=resolved_client, model=resolved_settings.model, + default_system_prompt=resolved_settings.default_system_prompt, tool_registry=ToolRegistry(resolved_settings.tool_timeout_seconds), max_tool_rounds=resolved_settings.max_tool_rounds, max_tool_calls_per_turn=resolved_settings.max_tool_calls_per_turn, ) yield + if client is None: await resolved_client.close() app = FastAPI(title="Simple Chat API", version="0.1.0", lifespan=lifespan) - - @app.post( - "/sessions", - response_model=CreateSessionResponse, - status_code=status.HTTP_201_CREATED, - ) - async def create_session( - body: CreateSessionRequest | None = None, - ) -> CreateSessionResponse: - system_prompt = ( - body.system_prompt - if body is not None and body.system_prompt is not None - else resolved_settings.default_system_prompt - ) - try: - session = await storage.create(system_prompt) - except SessionStorageError as exc: - raise HTTPException(status_code=500, detail=str(exc)) from exc - return CreateSessionResponse( - session_id=session.session_id, - system_prompt=session.system_prompt, - created_at=session.created_at, - ) - - @app.post( - "/sessions/{session_id}/messages", - response_model=SendMessageResponse, - ) - async def send_message( - session_id: str, body: SendMessageRequest, request: Request - ) -> SendMessageResponse: - service: ChatService = request.app.state.chat_service - try: - return await service.generate_response(session_id, body.content) - except SessionNotFoundError as exc: - raise HTTPException(status_code=404, detail="session not found") from exc - except ChatProviderError as exc: - raise HTTPException(status_code=502, detail=str(exc)) from exc - except SessionStorageError as exc: - raise HTTPException(status_code=500, detail=str(exc)) from exc - - @app.get( - "/sessions/{session_id}/messages", - response_model=SessionHistoryResponse, - ) - async def get_session_history( - session_id: str, request: Request - ) -> SessionHistoryResponse: - service: ChatService = request.app.state.chat_service - try: - return await service.get_session_history(session_id) - except SessionNotFoundError as exc: - raise HTTPException(status_code=404, detail="session not found") from exc - except SessionStorageError as exc: - raise HTTPException(status_code=500, detail=str(exc)) from exc - + app.include_router(router) return app diff --git a/models.py b/domain.py similarity index 51% rename from models.py rename to domain.py index e013ecf..e746700 100644 --- a/models.py +++ b/domain.py @@ -1,7 +1,7 @@ from datetime import datetime from typing import Annotated, Literal -from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic import BaseModel, ConfigDict, Field, model_validator class ToolFunctionCall(BaseModel): @@ -41,6 +41,8 @@ Message = Annotated[ UserMessage | AssistantMessage | ToolMessage, Field(discriminator="role"), ] + + class Session(BaseModel): model_config = ConfigDict(extra="forbid") @@ -63,62 +65,3 @@ class Session(BaseModel): if isinstance(message, dict): message.setdefault("created_at", fallback) return data - - -class CreateSessionRequest(BaseModel): - model_config = ConfigDict(extra="forbid") - - system_prompt: str | None = None - - @field_validator("system_prompt") - @classmethod - def validate_system_prompt(cls, value: str | None) -> str | None: - if value is not None and not value.strip(): - raise ValueError("system_prompt must not be blank") - return value - - -class CreateSessionResponse(BaseModel): - session_id: str - system_prompt: str - created_at: datetime - - -class SendMessageRequest(BaseModel): - model_config = ConfigDict(extra="forbid") - - content: str - - @field_validator("content") - @classmethod - def validate_content(cls, value: str) -> str: - if not value.strip(): - raise ValueError("content must not be blank") - return value - - -class TokenUsage(BaseModel): - prompt_tokens: int - completion_tokens: int - total_tokens: int - - -class FinalAssistantMessage(BaseModel): - role: Literal["assistant"] = "assistant" - content: str - created_at: datetime - - -class SendMessageResponse(BaseModel): - session_id: str - message: FinalAssistantMessage - tools_use: list[str] = Field(default_factory=list) - usage: TokenUsage | None - - -class SessionHistoryResponse(BaseModel): - session_id: str - system_prompt: str - created_at: datetime - updated_at: datetime - messages: list[Message] diff --git a/routes.py b/routes.py new file mode 100644 index 0000000..b69e185 --- /dev/null +++ b/routes.py @@ -0,0 +1,66 @@ +from fastapi import APIRouter, HTTPException, Request, status + +from schemas import ( + CreateSessionRequest, + CreateSessionResponse, + SendMessageRequest, + SendMessageResponse, + SessionHistoryResponse, +) +from service import ChatProviderError, ChatService +from storage import SessionNotFoundError, SessionStorageError + + +router = APIRouter(prefix="/sessions", tags=["sessions"]) + + +def get_chat_service(request: Request) -> ChatService: + return request.app.state.chat_service + + +@router.post( + "", + response_model=CreateSessionResponse, + status_code=status.HTTP_201_CREATED, +) +async def create_session( + request: Request, + body: CreateSessionRequest | None = None, +) -> CreateSessionResponse: + try: + return await get_chat_service(request).create_session( + body.system_prompt if body is not None else None + ) + except SessionStorageError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + + +@router.post("/{session_id}/messages", response_model=SendMessageResponse) +async def send_message( + session_id: str, + body: SendMessageRequest, + request: Request, +) -> SendMessageResponse: + try: + return await get_chat_service(request).generate_response( + session_id, body.content + ) + except SessionNotFoundError as exc: + raise HTTPException(status_code=404, detail="session not found") from exc + except ChatProviderError as exc: + raise HTTPException(status_code=502, detail=str(exc)) from exc + except SessionStorageError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + + +@router.get("/{session_id}/messages", response_model=SessionHistoryResponse) +async def get_session_history( + session_id: str, + request: Request, +) -> SessionHistoryResponse: + try: + return await get_chat_service(request).get_session_history(session_id) + except SessionNotFoundError as exc: + raise HTTPException(status_code=404, detail="session not found") from exc + except SessionStorageError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc diff --git a/schemas.py b/schemas.py new file mode 100644 index 0000000..e7ee4f2 --- /dev/null +++ b/schemas.py @@ -0,0 +1,65 @@ +from datetime import datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator + +from domain import Message + + +class CreateSessionRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + system_prompt: str | None = None + + @field_validator("system_prompt") + @classmethod + def validate_system_prompt(cls, value: str | None) -> str | None: + if value is not None and not value.strip(): + raise ValueError("system_prompt must not be blank") + return value + + +class CreateSessionResponse(BaseModel): + session_id: str + system_prompt: str + created_at: datetime + + +class SendMessageRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + content: str + + @field_validator("content") + @classmethod + def validate_content(cls, value: str) -> str: + if not value.strip(): + raise ValueError("content must not be blank") + return value + + +class TokenUsage(BaseModel): + prompt_tokens: int + completion_tokens: int + total_tokens: int + + +class FinalAssistantMessage(BaseModel): + role: Literal["assistant"] = "assistant" + content: str + created_at: datetime + + +class SendMessageResponse(BaseModel): + session_id: str + message: FinalAssistantMessage + tools_use: list[str] = Field(default_factory=list) + usage: TokenUsage | None + + +class SessionHistoryResponse(BaseModel): + session_id: str + system_prompt: str + created_at: datetime + updated_at: datetime + messages: list[Message] diff --git a/service.py b/service.py index 374d8f7..c398712 100644 --- a/service.py +++ b/service.py @@ -4,19 +4,22 @@ from typing import Any from openai import APIError, AsyncOpenAI -from models import ( +from domain import ( AssistantMessage, - FinalAssistantMessage, Message, - SendMessageResponse, Session, - SessionHistoryResponse, - TokenUsage, ToolCall, ToolFunctionCall, ToolMessage, UserMessage, ) +from schemas import ( + CreateSessionResponse, + FinalAssistantMessage, + SendMessageResponse, + SessionHistoryResponse, + TokenUsage, +) from storage import JsonSessionStorage from tools import ToolRegistry @@ -31,6 +34,7 @@ class ChatService: storage: JsonSessionStorage, client: AsyncOpenAI, model: str, + default_system_prompt: str = "You are a helpful assistant.", tool_registry: ToolRegistry | None = None, max_tool_rounds: int = 5, max_tool_calls_per_turn: int = 10, @@ -40,11 +44,24 @@ class ChatService: self.storage = storage self.client = client self.model = model + self.default_system_prompt = default_system_prompt self.tool_registry = tool_registry or ToolRegistry() self.max_tool_rounds = max_tool_rounds self.max_tool_calls_per_turn = max_tool_calls_per_turn self._locks: dict[str, asyncio.Lock] = {} + async def create_session( + self, system_prompt: str | None = None + ) -> CreateSessionResponse: + session = await self.storage.create( + system_prompt or self.default_system_prompt + ) + return CreateSessionResponse( + session_id=session.session_id, + system_prompt=session.system_prompt, + created_at=session.created_at, + ) + async def generate_response( self, session_id: str, content: str ) -> SendMessageResponse: diff --git a/storage.py b/storage.py index 6f43ada..d6e35f4 100644 --- a/storage.py +++ b/storage.py @@ -8,7 +8,7 @@ from uuid import UUID, uuid4 from pydantic import ValidationError -from models import Session +from domain import Session class SessionNotFoundError(Exception): diff --git a/tests/test_agent.py b/tests/test_agent.py index fd9d17f..6435565 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -58,20 +58,26 @@ def test_agent_executes_tool_and_persists_full_history(tmp_path: Path) -> None: [ completion( content=None, - tool_calls=[tool_call("call-1", "calculate", '{"expression":"2+2"}')], + tool_calls=[ + tool_call( + "call-1", + "get_current_time", + '{"timezone":"UTC"}', + ) + ], ), - completion(content="结果是 4。", prompt_tokens=20, completion_tokens=3), + completion(content="当前时间已查询。", prompt_tokens=20, completion_tokens=3), ] ) storage = JsonSessionStorage(tmp_path) session = await storage.create("Use tools when needed.") service = ChatService(storage, provider, "test-model") # type: ignore[arg-type] - response = await service.generate_response(session.session_id, "2+2 等于多少?") + response = await service.generate_response(session.session_id, "现在几点?") - assert response.message.content == "结果是 4。" + assert response.message.content == "当前时间已查询。" assert "tool_calls" not in response.message.model_dump() - assert response.tools_use == ["calculate"] + assert response.tools_use == ["get_current_time"] assert response.usage is not None assert response.usage.model_dump() == { "prompt_tokens": 30, @@ -86,7 +92,7 @@ def test_agent_executes_tool_and_persists_full_history(tmp_path: Path) -> None: "assistant", "tool", ] - assert json.loads(second_messages[-1]["content"])["result"] == 4 + assert json.loads(second_messages[-1]["content"])["timezone"] == "UTC" saved = await storage.read(session.session_id) assert [message.role for message in saved.messages] == [ @@ -108,11 +114,15 @@ def test_agent_executes_multiple_tools_in_order(tmp_path: Path) -> None: completion( content=None, tool_calls=[ - tool_call("call-1", "calculate", '{"expression":"6*7"}'), + tool_call( + "call-1", + "get_current_time", + '{"timezone":"UTC"}', + ), tool_call( "call-2", "get_current_time", - '{"timezone":"UTC"}', + '{"timezone":"Asia/Shanghai"}', ), ], ), @@ -130,7 +140,7 @@ def test_agent_executes_multiple_tools_in_order(tmp_path: Path) -> None: "call-1", "call-2", ] - assert response.tools_use == ["calculate", "get_current_time"] + assert response.tools_use == ["get_current_time"] asyncio.run(scenario()) @@ -141,7 +151,7 @@ def test_tool_error_is_returned_to_model(tmp_path: Path) -> None: [ completion( content=None, - tool_calls=[tool_call("bad-1", "calculate", "not-json")], + tool_calls=[tool_call("bad-1", "get_current_time", "not-json")], ), completion(content="参数无效。"), ] @@ -164,13 +174,21 @@ def test_tool_error_is_returned_to_model(tmp_path: Path) -> None: @pytest.mark.parametrize( ("max_rounds", "max_calls", "calls"), [ - (1, 10, [tool_call("call-1", "calculate", '{"expression":"1+1"}')]), + ( + 1, + 10, + [tool_call("call-1", "get_current_time", '{"timezone":"UTC"}')], + ), ( 5, 1, [ - tool_call("call-1", "calculate", '{"expression":"1+1"}'), - tool_call("call-2", "calculate", '{"expression":"2+2"}'), + tool_call("call-1", "get_current_time", '{"timezone":"UTC"}'), + tool_call( + "call-2", + "get_current_time", + '{"timezone":"Asia/Shanghai"}', + ), ], ), ], diff --git a/tests/test_tools.py b/tests/test_tools.py index 4c2bb37..920216a 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -2,7 +2,7 @@ import asyncio import json import time -from tools import CalculatorArguments, ToolRegistry, ToolSpec +from tools import CurrentTimeArguments, ToolRegistry, ToolSpec def test_builtin_tool_definitions() -> None: @@ -11,37 +11,28 @@ def test_builtin_tool_definitions() -> None: definitions = registry.definitions() assert [item["function"]["name"] for item in definitions] == [ # type: ignore[index] - "get_current_time", - "calculate", + "get_current_time" ] -def test_calculator_and_time_tools() -> None: +def test_current_time_tool() -> None: async def scenario() -> None: registry = ToolRegistry() - calculation = await registry.execute( - "calculate", '{"expression":"(2 + 3) * 4"}' - ) current_time = await registry.execute( "get_current_time", '{"timezone":"Asia/Shanghai"}' ) - assert calculation.is_error is False - assert json.loads(calculation.content)["result"] == 20 assert current_time.is_error is False assert json.loads(current_time.content)["timezone"] == "Asia/Shanghai" asyncio.run(scenario()) -def test_tool_argument_and_expression_errors_are_safe() -> None: +def test_tool_argument_and_execution_errors_are_safe() -> None: async def scenario() -> None: registry = ToolRegistry() - invalid_json = await registry.execute("calculate", "not-json") + invalid_json = await registry.execute("get_current_time", "not-json") unknown = await registry.execute("missing", "{}") - unsafe = await registry.execute( - "calculate", '{"expression":"__import__(\\"os\\").system(\\"id\\")"}' - ) invalid_timezone = await registry.execute( "get_current_time", '{"timezone":"Not/A_Real_Zone"}' ) @@ -50,8 +41,6 @@ def test_tool_argument_and_expression_errors_are_safe() -> None: assert "invalid_arguments" in invalid_json.content assert unknown.is_error is True assert "unknown_tool" in unknown.content - assert unsafe.is_error is True - assert "execution_failed" in unsafe.content assert invalid_timezone.is_error is True assert "execution_failed" in invalid_timezone.content @@ -61,18 +50,20 @@ def test_tool_argument_and_expression_errors_are_safe() -> None: def test_tool_timeout_returns_error() -> None: def slow_handler(arguments): time.sleep(0.05) - return {"result": 1} + return {"timezone": "UTC"} async def scenario() -> None: registry = ToolRegistry(timeout_seconds=0.01) - registry._specs["calculate"] = ToolSpec( - name="calculate", + registry._specs["get_current_time"] = ToolSpec( + name="get_current_time", description="slow", - arguments_model=CalculatorArguments, + arguments_model=CurrentTimeArguments, handler=slow_handler, ) - result = await registry.execute("calculate", '{"expression":"1+1"}') + result = await registry.execute( + "get_current_time", '{"timezone":"UTC"}' + ) assert result.is_error is True assert "timeout" in result.content diff --git a/tools.py b/tools.py index 39587ee..b55be67 100644 --- a/tools.py +++ b/tools.py @@ -1,12 +1,8 @@ -import ast import asyncio from collections.abc import Callable from dataclasses import dataclass from datetime import datetime import json -import math -import operator -from typing import Any from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from pydantic import BaseModel, ConfigDict, Field, ValidationError @@ -18,12 +14,6 @@ class CurrentTimeArguments(BaseModel): timezone: str = Field(default="UTC", description="IANA timezone, e.g. Asia/Shanghai") -class CalculatorArguments(BaseModel): - model_config = ConfigDict(extra="forbid") - - expression: str = Field(description="Arithmetic expression to evaluate") - - @dataclass(frozen=True, slots=True) class ToolSpec: name: str @@ -59,13 +49,7 @@ class ToolRegistry: description="Get the current date and time in an IANA timezone.", arguments_model=CurrentTimeArguments, handler=_get_current_time, - ), - ToolSpec( - name="calculate", - description="Safely evaluate a basic arithmetic expression.", - arguments_model=CalculatorArguments, - handler=_calculate, - ), + ) ] self._specs = {spec.name: spec for spec in specs} @@ -118,54 +102,3 @@ def _get_current_time(arguments: BaseModel) -> dict[str, object]: raise ValueError("unknown timezone") from exc now = datetime.now(timezone) return {"timezone": arguments.timezone, "datetime": now.isoformat()} - - -_BINARY_OPERATORS: dict[type[ast.operator], Callable[[Any, Any], Any]] = { - ast.Add: operator.add, - ast.Sub: operator.sub, - ast.Mult: operator.mul, - ast.Div: operator.truediv, - ast.FloorDiv: operator.floordiv, - ast.Mod: operator.mod, - ast.Pow: operator.pow, -} -_UNARY_OPERATORS: dict[type[ast.unaryop], Callable[[Any], Any]] = { - ast.UAdd: operator.pos, - ast.USub: operator.neg, -} -_MAX_EXPRESSION_LENGTH = 200 -_MAX_ABSOLUTE_RESULT = 1e100 -_MAX_EXPONENT = 100 - - -def _calculate(arguments: BaseModel) -> dict[str, object]: - assert isinstance(arguments, CalculatorArguments) - expression = arguments.expression.strip() - if not expression or len(expression) > _MAX_EXPRESSION_LENGTH: - raise ValueError("expression is empty or too long") - tree = ast.parse(expression, mode="eval") - result = _evaluate_node(tree.body) - if isinstance(result, float) and not math.isfinite(result): - raise ValueError("result is not finite") - if abs(result) > _MAX_ABSOLUTE_RESULT: - raise ValueError("result is too large") - return {"expression": expression, "result": result} - - -def _evaluate_node(node: ast.AST) -> int | float: - if isinstance(node, ast.Constant): - if isinstance(node.value, bool) or not isinstance(node.value, (int, float)): - raise ValueError("only numeric constants are allowed") - return node.value - if isinstance(node, ast.UnaryOp) and type(node.op) in _UNARY_OPERATORS: - return _UNARY_OPERATORS[type(node.op)](_evaluate_node(node.operand)) - if isinstance(node, ast.BinOp) and type(node.op) in _BINARY_OPERATORS: - left = _evaluate_node(node.left) - right = _evaluate_node(node.right) - if isinstance(node.op, ast.Pow) and abs(right) > _MAX_EXPONENT: - raise ValueError("exponent is too large") - result = _BINARY_OPERATORS[type(node.op)](left, right) - if abs(result) > _MAX_ABSOLUTE_RESULT: - raise ValueError("intermediate result is too large") - return result - raise ValueError("unsupported expression")