From 3a73c345f1bd78b497e4175520b22982c4401afb Mon Sep 17 00:00:00 2001 From: Mplan Date: Fri, 3 Jul 2026 22:37:29 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E9=A1=B9=E7=9B=AE=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 | 37 +++-- main.py | 12 +- pyproject.toml | 12 +- schemas.py | 126 ------------------ src/chat_api/__init__.py | 3 + app.py => src/chat_api/app.py | 13 +- auth.py => src/chat_api/auth.py | 2 +- config.py => src/chat_api/config.py | 0 src/chat_api/domain/__init__.py | 19 +++ src/chat_api/domain/messages.py | 43 ++++++ domain.py => src/chat_api/domain/session.py | 42 +----- routes.py => src/chat_api/routes.py | 11 +- src/chat_api/schemas/__init__.py | 33 +++++ src/chat_api/schemas/auth.py | 55 ++++++++ src/chat_api/schemas/message.py | 36 +++++ src/chat_api/schemas/session.py | 32 +++++ src/chat_api/schemas/usage.py | 12 ++ src/chat_api/service/__init__.py | 3 + service.py => src/chat_api/service/chat.py | 10 +- src/chat_api/storage/__init__.py | 25 ++++ .../chat_api/storage/sessions.py | 4 +- users.py => src/chat_api/storage/users.py | 0 src/chat_api/tools/__init__.py | 9 ++ tools.py => src/chat_api/tools/registry.py | 34 +---- src/chat_api/tools/time_tools.py | 34 +++++ tests/test_agent.py | 4 +- tests/test_api.py | 10 +- tests/test_tools.py | 2 +- uv.lock | 2 +- 29 files changed, 388 insertions(+), 237 deletions(-) delete mode 100644 schemas.py create mode 100644 src/chat_api/__init__.py rename app.py => src/chat_api/app.py (88%) rename auth.py => src/chat_api/auth.py (95%) rename config.py => src/chat_api/config.py (100%) create mode 100644 src/chat_api/domain/__init__.py create mode 100644 src/chat_api/domain/messages.py rename domain.py => src/chat_api/domain/session.py (53%) rename routes.py => src/chat_api/routes.py (95%) create mode 100644 src/chat_api/schemas/__init__.py create mode 100644 src/chat_api/schemas/auth.py create mode 100644 src/chat_api/schemas/message.py create mode 100644 src/chat_api/schemas/session.py create mode 100644 src/chat_api/schemas/usage.py create mode 100644 src/chat_api/service/__init__.py rename service.py => src/chat_api/service/chat.py (98%) create mode 100644 src/chat_api/storage/__init__.py rename storage.py => src/chat_api/storage/sessions.py (98%) rename users.py => src/chat_api/storage/users.py (100%) create mode 100644 src/chat_api/tools/__init__.py rename tools.py => src/chat_api/tools/registry.py (69%) create mode 100644 src/chat_api/tools/time_tools.py diff --git a/README.md b/README.md index 9f2a14d..84a6852 100644 --- a/README.md +++ b/README.md @@ -4,14 +4,33 @@ ## 代码结构 -- `routes.py`:集中定义所有 HTTP 路由。 -- `auth.py`:根据 `user_id` 识别当前用户,校验请求 `X-API-Key`。 -- `schemas.py`:集中定义 API 请求和响应结构。 -- `users.py`:基于 SQLite 的用户注册、登录与 API Key 持久化。 -- `domain.py`:定义会话、消息和工具调用的持久化模型。 -- `tools.py`:集中定义工具注册表、参数结构和工具函数。 -- `service.py`:处理对话、Agent 循环和业务规则。 -- `app.py`:创建 FastAPI 应用并管理生命周期。 +采用 `src/` 布局,所有业务代码收纳在 `src/chat_api/` 包内,按职责分模块/子包,根目录只保留 `main.py` 入口。 + +``` +src/chat_api/ +├── app.py 应用工厂 + 生命周期(组装各层) +├── config.py Settings:从 .env 读取配置 +├── auth.py API Key 鉴权依赖项(CurrentUser) +├── routes.py 所有 HTTP 路由 +├── domain/ 持久化模型(不依赖任何其它层) +│ ├── messages.py UserMessage / AssistantMessage / ToolMessage / ToolCall +│ └── session.py Session(含旧数据兼容校验器) +├── schemas/ API 请求/响应模型(按接口分组) +│ ├── session.py 创建会话 / 历史查询 +│ ├── message.py 发送消息 / TokenUsage +│ ├── auth.py 注册 / 登录 +│ └── usage.py 用量查询 +├── storage/ 持久化实现 +│ ├── sessions.py JsonSessionStorage(会话 JSON) +│ └── users.py UserStore(SQLite 用户 + bcrypt + AuthenticatedUser) +├── service/ +│ └── chat.py ChatService:Agent 循环与业务规则 +└── tools/ + ├── registry.py ToolRegistry / ToolSpec 注册表基础设施 + └── time_tools.py get_current_time 内置工具 +``` + +分层依赖方向:`domain` ← `schemas`/`storage` ← `service`/`auth` ← `routes` ← `app`。各子包 `__init__.py` 做重导出,跨层引用形如 `from chat_api.storage import JsonSessionStorage, UserStore`。 ## 启动 @@ -19,7 +38,7 @@ uv sync cp .env.example .env # 编辑 .env 并填写 DEEPSEEK_API_KEY -uv run uvicorn main:app --reload +uv run python main.py # 或:uv run uvicorn main:app --reload --reload-dir src ``` 服务默认运行在 `http://127.0.0.1:8000`,交互式 API 文档位于 `/docs`。 diff --git a/main.py b/main.py index b76f231..8de5133 100644 --- a/main.py +++ b/main.py @@ -1,8 +1,14 @@ -from app import app -from config import Settings +from chat_api import app +from chat_api.config import Settings if __name__ == "__main__": import uvicorn settings = Settings.from_env() - uvicorn.run("main:app", host="0.0.0.0", port=settings.port, reload=True) + uvicorn.run( + "chat_api:app", + host="0.0.0.0", + port=settings.port, + reload=True, + reload_dirs=["src"], + ) \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 7ed0b56..362f7fc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,6 +14,16 @@ dependencies = [ "uvicorn[standard]>=0.32.0", ] +[build-system] +requires = ["setuptools>=68.0"] +build-backend = "setuptools.build_meta" + +[tool.setuptools] +package-dir = {"" = "src"} + +[tool.setuptools.packages.find] +where = ["src"] + [dependency-groups] dev = [ "pytest>=8.3.0", @@ -21,4 +31,4 @@ dev = [ [tool.pytest.ini_options] testpaths = ["tests"] -pythonpath = ["."] +pythonpath = ["src"] diff --git a/schemas.py b/schemas.py deleted file mode 100644 index dee2522..0000000 --- a/schemas.py +++ /dev/null @@ -1,126 +0,0 @@ -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] - - -class UserUsageResponse(BaseModel): - user_id: str - api_calls: int - prompt_tokens: int - completion_tokens: int - total_tokens: int - updated_at: datetime - - -class RegisterRequest(BaseModel): - model_config = ConfigDict(extra="forbid") - - username: str - password: str - - @field_validator("username") - @classmethod - def validate_username(cls, value: str) -> str: - if not value.strip(): - raise ValueError("username must not be blank") - return value - - @field_validator("password") - @classmethod - def validate_password(cls, value: str) -> str: - if not value.strip(): - raise ValueError("password must not be blank") - return value - - -class RegisterResponse(BaseModel): - user_id: str - created_at: datetime - - -class LoginRequest(BaseModel): - model_config = ConfigDict(extra="forbid") - - username: str - password: str - - @field_validator("username") - @classmethod - def validate_username(cls, value: str) -> str: - if not value.strip(): - raise ValueError("username must not be blank") - return value - - @field_validator("password") - @classmethod - def validate_password(cls, value: str) -> str: - if not value.strip(): - raise ValueError("password must not be blank") - return value - - -class LoginResponse(BaseModel): - user_id: str - api_key: str \ No newline at end of file diff --git a/src/chat_api/__init__.py b/src/chat_api/__init__.py new file mode 100644 index 0000000..9bbeb77 --- /dev/null +++ b/src/chat_api/__init__.py @@ -0,0 +1,3 @@ +from .app import app, create_app + +__all__ = ["app", "create_app"] \ No newline at end of file diff --git a/app.py b/src/chat_api/app.py similarity index 88% rename from app.py rename to src/chat_api/app.py index 5eba108..c52ba40 100644 --- a/app.py +++ b/src/chat_api/app.py @@ -4,13 +4,12 @@ from contextlib import asynccontextmanager from fastapi import FastAPI from openai import AsyncOpenAI -from auth import ApiKeyAuthenticator -from config import Settings -from routes import router -from service import ChatService -from storage import JsonSessionStorage -from tools import ToolRegistry -from users import UserStore +from .auth import ApiKeyAuthenticator +from .config import Settings +from .routes import router +from .service import ChatService +from .storage import JsonSessionStorage, UserStore +from .tools import ToolRegistry def create_app( diff --git a/auth.py b/src/chat_api/auth.py similarity index 95% rename from auth.py rename to src/chat_api/auth.py index f308f5b..ec75552 100644 --- a/auth.py +++ b/src/chat_api/auth.py @@ -3,7 +3,7 @@ from typing import Annotated from fastapi import Depends, HTTPException, Request, Security, status from fastapi.security import APIKeyHeader -from users import AuthenticatedUser, UserStore +from .storage import AuthenticatedUser, UserStore api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False) diff --git a/config.py b/src/chat_api/config.py similarity index 100% rename from config.py rename to src/chat_api/config.py diff --git a/src/chat_api/domain/__init__.py b/src/chat_api/domain/__init__.py new file mode 100644 index 0000000..194f9e3 --- /dev/null +++ b/src/chat_api/domain/__init__.py @@ -0,0 +1,19 @@ +from .messages import ( + AssistantMessage, + Message, + ToolCall, + ToolFunctionCall, + ToolMessage, + UserMessage, +) +from .session import Session + +__all__ = [ + "AssistantMessage", + "Message", + "Session", + "ToolCall", + "ToolFunctionCall", + "ToolMessage", + "UserMessage", +] \ No newline at end of file diff --git a/src/chat_api/domain/messages.py b/src/chat_api/domain/messages.py new file mode 100644 index 0000000..bdaed64 --- /dev/null +++ b/src/chat_api/domain/messages.py @@ -0,0 +1,43 @@ +from datetime import datetime +from typing import Annotated, Literal + +from pydantic import BaseModel, Field + + +class ToolFunctionCall(BaseModel): + name: str + arguments: str + + +class ToolCall(BaseModel): + id: str + type: Literal["function"] = "function" + function: ToolFunctionCall + + +class UserMessage(BaseModel): + role: Literal["user"] = "user" + content: str + created_at: datetime + + +class AssistantMessage(BaseModel): + role: Literal["assistant"] = "assistant" + content: str | None = None + tool_calls: list[ToolCall] = Field(default_factory=list) + created_at: datetime + + +class ToolMessage(BaseModel): + role: Literal["tool"] = "tool" + content: str + tool_call_id: str + name: str + is_error: bool = False + created_at: datetime + + +Message = Annotated[ + UserMessage | AssistantMessage | ToolMessage, + Field(discriminator="role"), +] \ No newline at end of file diff --git a/domain.py b/src/chat_api/domain/session.py similarity index 53% rename from domain.py rename to src/chat_api/domain/session.py index d092dad..a09b19a 100644 --- a/domain.py +++ b/src/chat_api/domain/session.py @@ -1,46 +1,8 @@ from datetime import datetime -from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, model_validator - -class ToolFunctionCall(BaseModel): - name: str - arguments: str - - -class ToolCall(BaseModel): - id: str - type: Literal["function"] = "function" - function: ToolFunctionCall - - -class UserMessage(BaseModel): - role: Literal["user"] = "user" - content: str - created_at: datetime - - -class AssistantMessage(BaseModel): - role: Literal["assistant"] = "assistant" - content: str | None = None - tool_calls: list[ToolCall] = Field(default_factory=list) - created_at: datetime - - -class ToolMessage(BaseModel): - role: Literal["tool"] = "tool" - content: str - tool_call_id: str - name: str - is_error: bool = False - created_at: datetime - - -Message = Annotated[ - UserMessage | AssistantMessage | ToolMessage, - Field(discriminator="role"), -] +from .messages import Message class Session(BaseModel): @@ -69,4 +31,4 @@ class Session(BaseModel): for message in messages: if isinstance(message, dict): message.setdefault("created_at", fallback) - return data + return data \ No newline at end of file diff --git a/routes.py b/src/chat_api/routes.py similarity index 95% rename from routes.py rename to src/chat_api/routes.py index 3afeeb6..6687efb 100644 --- a/routes.py +++ b/src/chat_api/routes.py @@ -1,7 +1,7 @@ from fastapi import APIRouter, HTTPException, Request, status -from auth import CurrentUser -from schemas import ( +from .auth import CurrentUser +from .schemas import ( CreateSessionRequest, CreateSessionResponse, LoginRequest, @@ -13,10 +13,11 @@ from schemas import ( SessionHistoryResponse, UserUsageResponse, ) -from service import ChatProviderError, ChatService -from storage import SessionNotFoundError, SessionStorageError -from users import ( +from .service import ChatProviderError, ChatService +from .storage import ( InvalidCredentialsError, + SessionNotFoundError, + SessionStorageError, UserStore, UserStoreError, UsernameExistsError, diff --git a/src/chat_api/schemas/__init__.py b/src/chat_api/schemas/__init__.py new file mode 100644 index 0000000..78c976e --- /dev/null +++ b/src/chat_api/schemas/__init__.py @@ -0,0 +1,33 @@ +from .auth import ( + LoginRequest, + LoginResponse, + RegisterRequest, + RegisterResponse, +) +from .message import ( + FinalAssistantMessage, + SendMessageRequest, + SendMessageResponse, + TokenUsage, +) +from .session import ( + CreateSessionRequest, + CreateSessionResponse, + SessionHistoryResponse, +) +from .usage import UserUsageResponse + +__all__ = [ + "CreateSessionRequest", + "CreateSessionResponse", + "FinalAssistantMessage", + "LoginRequest", + "LoginResponse", + "RegisterRequest", + "RegisterResponse", + "SendMessageRequest", + "SendMessageResponse", + "SessionHistoryResponse", + "TokenUsage", + "UserUsageResponse", +] \ No newline at end of file diff --git a/src/chat_api/schemas/auth.py b/src/chat_api/schemas/auth.py new file mode 100644 index 0000000..40f7814 --- /dev/null +++ b/src/chat_api/schemas/auth.py @@ -0,0 +1,55 @@ +from datetime import datetime + +from pydantic import BaseModel, ConfigDict, field_validator + + +class RegisterRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + username: str + password: str + + @field_validator("username") + @classmethod + def validate_username(cls, value: str) -> str: + if not value.strip(): + raise ValueError("username must not be blank") + return value + + @field_validator("password") + @classmethod + def validate_password(cls, value: str) -> str: + if not value.strip(): + raise ValueError("password must not be blank") + return value + + +class RegisterResponse(BaseModel): + user_id: str + created_at: datetime + + +class LoginRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + username: str + password: str + + @field_validator("username") + @classmethod + def validate_username(cls, value: str) -> str: + if not value.strip(): + raise ValueError("username must not be blank") + return value + + @field_validator("password") + @classmethod + def validate_password(cls, value: str) -> str: + if not value.strip(): + raise ValueError("password must not be blank") + return value + + +class LoginResponse(BaseModel): + user_id: str + api_key: str \ No newline at end of file diff --git a/src/chat_api/schemas/message.py b/src/chat_api/schemas/message.py new file mode 100644 index 0000000..1e2762b --- /dev/null +++ b/src/chat_api/schemas/message.py @@ -0,0 +1,36 @@ +from datetime import datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator + + +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 \ No newline at end of file diff --git a/src/chat_api/schemas/session.py b/src/chat_api/schemas/session.py new file mode 100644 index 0000000..bab4b3f --- /dev/null +++ b/src/chat_api/schemas/session.py @@ -0,0 +1,32 @@ +from datetime import datetime + +from pydantic import BaseModel, ConfigDict, 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 SessionHistoryResponse(BaseModel): + session_id: str + system_prompt: str + created_at: datetime + updated_at: datetime + messages: list[Message] \ No newline at end of file diff --git a/src/chat_api/schemas/usage.py b/src/chat_api/schemas/usage.py new file mode 100644 index 0000000..787bef0 --- /dev/null +++ b/src/chat_api/schemas/usage.py @@ -0,0 +1,12 @@ +from datetime import datetime + +from pydantic import BaseModel + + +class UserUsageResponse(BaseModel): + user_id: str + api_calls: int + prompt_tokens: int + completion_tokens: int + total_tokens: int + updated_at: datetime \ No newline at end of file diff --git a/src/chat_api/service/__init__.py b/src/chat_api/service/__init__.py new file mode 100644 index 0000000..6b86c69 --- /dev/null +++ b/src/chat_api/service/__init__.py @@ -0,0 +1,3 @@ +from .chat import ChatProviderError, ChatService + +__all__ = ["ChatProviderError", "ChatService"] \ No newline at end of file diff --git a/service.py b/src/chat_api/service/chat.py similarity index 98% rename from service.py rename to src/chat_api/service/chat.py index e53818d..fd11050 100644 --- a/service.py +++ b/src/chat_api/service/chat.py @@ -4,7 +4,7 @@ from typing import Any from openai import APIError, AsyncOpenAI -from domain import ( +from ..domain import ( AssistantMessage, Message, Session, @@ -13,7 +13,7 @@ from domain import ( ToolMessage, UserMessage, ) -from schemas import ( +from ..schemas import ( CreateSessionResponse, FinalAssistantMessage, SendMessageResponse, @@ -21,8 +21,8 @@ from schemas import ( TokenUsage, UserUsageResponse, ) -from storage import JsonSessionStorage, SessionNotFoundError -from tools import ToolRegistry +from ..storage import JsonSessionStorage, SessionNotFoundError +from ..tools import ToolRegistry class ChatProviderError(Exception): @@ -251,4 +251,4 @@ class ChatService: def _add_usage(total: TokenUsage, usage: object) -> None: total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined] total.completion_tokens += usage.completion_tokens # type: ignore[attr-defined] - total.total_tokens += usage.total_tokens # type: ignore[attr-defined] + total.total_tokens += usage.total_tokens # type: ignore[attr-defined] \ No newline at end of file diff --git a/src/chat_api/storage/__init__.py b/src/chat_api/storage/__init__.py new file mode 100644 index 0000000..b7dcf13 --- /dev/null +++ b/src/chat_api/storage/__init__.py @@ -0,0 +1,25 @@ +from .sessions import ( + JsonSessionStorage, + SessionNotFoundError, + SessionStorageError, +) +from .users import ( + AuthenticatedUser, + InvalidCredentialsError, + RegisteredUser, + UserStore, + UserStoreError, + UsernameExistsError, +) + +__all__ = [ + "AuthenticatedUser", + "InvalidCredentialsError", + "JsonSessionStorage", + "RegisteredUser", + "SessionNotFoundError", + "SessionStorageError", + "UserStore", + "UserStoreError", + "UsernameExistsError", +] \ No newline at end of file diff --git a/storage.py b/src/chat_api/storage/sessions.py similarity index 98% rename from storage.py rename to src/chat_api/storage/sessions.py index e074dc4..f03e9a0 100644 --- a/storage.py +++ b/src/chat_api/storage/sessions.py @@ -8,7 +8,7 @@ from uuid import UUID, uuid4 from pydantic import ValidationError -from domain import Session +from ..domain import Session class SessionNotFoundError(Exception): @@ -128,4 +128,4 @@ class JsonSessionStorage: try: Path(temporary_path).unlink(missing_ok=True) except OSError: - pass + pass \ No newline at end of file diff --git a/users.py b/src/chat_api/storage/users.py similarity index 100% rename from users.py rename to src/chat_api/storage/users.py diff --git a/src/chat_api/tools/__init__.py b/src/chat_api/tools/__init__.py new file mode 100644 index 0000000..099375e --- /dev/null +++ b/src/chat_api/tools/__init__.py @@ -0,0 +1,9 @@ +from .registry import ToolExecutionResult, ToolRegistry, ToolSpec +from .time_tools import CurrentTimeArguments + +__all__ = [ + "CurrentTimeArguments", + "ToolExecutionResult", + "ToolRegistry", + "ToolSpec", +] \ No newline at end of file diff --git a/tools.py b/src/chat_api/tools/registry.py similarity index 69% rename from tools.py rename to src/chat_api/tools/registry.py index b55be67..b8c5a73 100644 --- a/tools.py +++ b/src/chat_api/tools/registry.py @@ -1,17 +1,9 @@ import asyncio from collections.abc import Callable from dataclasses import dataclass -from datetime import datetime import json -from zoneinfo import ZoneInfo, ZoneInfoNotFoundError -from pydantic import BaseModel, ConfigDict, Field, ValidationError - - -class CurrentTimeArguments(BaseModel): - model_config = ConfigDict(extra="forbid") - - timezone: str = Field(default="UTC", description="IANA timezone, e.g. Asia/Shanghai") +from pydantic import BaseModel, ValidationError @dataclass(frozen=True, slots=True) @@ -43,15 +35,9 @@ class ToolRegistry: if timeout_seconds <= 0: raise ValueError("tool timeout must be greater than zero") self.timeout_seconds = timeout_seconds - specs = [ - ToolSpec( - name="get_current_time", - description="Get the current date and time in an IANA timezone.", - arguments_model=CurrentTimeArguments, - handler=_get_current_time, - ) - ] - self._specs = {spec.name: spec for spec in specs} + from .time_tools import build_builtin_specs + + self._specs = {spec.name: spec for spec in build_builtin_specs()} def definitions(self) -> list[dict[str, object]]: return [spec.api_definition() for spec in self._specs.values()] @@ -91,14 +77,4 @@ def _tool_error(code: str, message: str) -> ToolExecutionResult: ensure_ascii=False, ), is_error=True, - ) - - -def _get_current_time(arguments: BaseModel) -> dict[str, object]: - assert isinstance(arguments, CurrentTimeArguments) - try: - timezone = ZoneInfo(arguments.timezone) - except ZoneInfoNotFoundError as exc: - raise ValueError("unknown timezone") from exc - now = datetime.now(timezone) - return {"timezone": arguments.timezone, "datetime": now.isoformat()} + ) \ No newline at end of file diff --git a/src/chat_api/tools/time_tools.py b/src/chat_api/tools/time_tools.py new file mode 100644 index 0000000..9ce70e4 --- /dev/null +++ b/src/chat_api/tools/time_tools.py @@ -0,0 +1,34 @@ +from collections.abc import Iterable +from datetime import datetime +from zoneinfo import ZoneInfo, ZoneInfoNotFoundError + +from pydantic import BaseModel, ConfigDict, Field + +from .registry import ToolSpec + + +class CurrentTimeArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + + timezone: str = Field(default="UTC", description="IANA timezone, e.g. Asia/Shanghai") + + +def _get_current_time(arguments: BaseModel) -> dict[str, object]: + assert isinstance(arguments, CurrentTimeArguments) + try: + timezone = ZoneInfo(arguments.timezone) + except ZoneInfoNotFoundError as exc: + raise ValueError("unknown timezone") from exc + now = datetime.now(timezone) + return {"timezone": arguments.timezone, "datetime": now.isoformat()} + + +def build_builtin_specs() -> Iterable[ToolSpec]: + return [ + ToolSpec( + name="get_current_time", + description="Get the current date and time in an IANA timezone.", + arguments_model=CurrentTimeArguments, + handler=_get_current_time, + ) + ] \ No newline at end of file diff --git a/tests/test_agent.py b/tests/test_agent.py index 6435565..356014c 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -6,8 +6,8 @@ from typing import Any import pytest -from service import ChatProviderError, ChatService -from storage import JsonSessionStorage +from chat_api.service import ChatProviderError, ChatService +from chat_api.storage import JsonSessionStorage def completion( diff --git a/tests/test_api.py b/tests/test_api.py index c4979d0..a8bcea8 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -10,11 +10,11 @@ from fastapi.testclient import TestClient from openai import APIConnectionError import pytest -from app import create_app -from config import Settings -from service import ChatService -from storage import JsonSessionStorage -from users import UserStore +from chat_api import create_app +from chat_api.config import Settings +from chat_api.service import ChatService +from chat_api.storage import JsonSessionStorage +from chat_api.storage import UserStore @dataclass diff --git a/tests/test_tools.py b/tests/test_tools.py index 920216a..ce9d13b 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -2,7 +2,7 @@ import asyncio import json import time -from tools import CurrentTimeArguments, ToolRegistry, ToolSpec +from chat_api.tools import CurrentTimeArguments, ToolRegistry, ToolSpec def test_builtin_tool_definitions() -> None: diff --git a/uv.lock b/uv.lock index a611008..32f6166 100644 --- a/uv.lock +++ b/uv.lock @@ -421,7 +421,7 @@ wheels = [ [[package]] name = "simple-chat-api" version = "0.1.0" -source = { virtual = "." } +source = { editable = "." } dependencies = [ { name = "aiosqlite" }, { name = "bcrypt" },