优化项目结构

This commit is contained in:
2026-07-03 22:37:29 +08:00
parent 98a92b83fa
commit 3a73c345f1
29 changed files with 388 additions and 237 deletions
+28 -9
View File
@@ -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 UserStoreSQLite 用户 + bcrypt + AuthenticatedUser
├── service/
│ └── chat.py ChatServiceAgent 循环与业务规则
└── 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`
+9 -3
View File
@@ -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"],
)
+11 -1
View File
@@ -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"]
-126
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
from .app import app, create_app
__all__ = ["app", "create_app"]
+6 -7
View File
@@ -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(
+1 -1
View File
@@ -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)
+19
View File
@@ -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",
]
+43
View File
@@ -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"),
]
+2 -40
View File
@@ -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
+6 -5
View File
@@ -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,
+33
View File
@@ -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",
]
+55
View File
@@ -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
+36
View File
@@ -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
+32
View File
@@ -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]
+12
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
from .chat import ChatProviderError, ChatService
__all__ = ["ChatProviderError", "ChatService"]
+5 -5
View File
@@ -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]
+25
View File
@@ -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",
]
@@ -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
+9
View File
@@ -0,0 +1,9 @@
from .registry import ToolExecutionResult, ToolRegistry, ToolSpec
from .time_tools import CurrentTimeArguments
__all__ = [
"CurrentTimeArguments",
"ToolExecutionResult",
"ToolRegistry",
"ToolSpec",
]
+5 -29
View File
@@ -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()}
)
+34
View File
@@ -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,
)
]
+2 -2
View File
@@ -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(
+5 -5
View File
@@ -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
+1 -1
View File
@@ -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:
Generated
+1 -1
View File
@@ -421,7 +421,7 @@ wheels = [
[[package]]
name = "simple-chat-api"
version = "0.1.0"
source = { virtual = "." }
source = { editable = "." }
dependencies = [
{ name = "aiosqlite" },
{ name = "bcrypt" },