优化项目结构

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 路由 采用 `src/` 布局,所有业务代码收纳在 `src/chat_api/` 包内,按职责分模块/子包,根目录只保留 `main.py` 入口
- `auth.py`:根据 `user_id` 识别当前用户,校验请求 `X-API-Key`
- `schemas.py`:集中定义 API 请求和响应结构。 ```
- `users.py`:基于 SQLite 的用户注册、登录与 API Key 持久化。 src/chat_api/
- `domain.py`:定义会话、消息和工具调用的持久化模型。 ├── app.py 应用工厂 + 生命周期(组装各层)
- `tools.py`:集中定义工具注册表、参数结构和工具函数。 ├── config.py Settings:从 .env 读取配置
- `service.py`:处理对话、Agent 循环和业务规则。 ├── auth.py API Key 鉴权依赖项(CurrentUser
- `app.py`:创建 FastAPI 应用并管理生命周期。 ├── 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 uv sync
cp .env.example .env cp .env.example .env
# 编辑 .env 并填写 DEEPSEEK_API_KEY # 编辑 .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` 服务默认运行在 `http://127.0.0.1:8000`,交互式 API 文档位于 `/docs`
+9 -3
View File
@@ -1,8 +1,14 @@
from app import app from chat_api import app
from config import Settings from chat_api.config import Settings
if __name__ == "__main__": if __name__ == "__main__":
import uvicorn import uvicorn
settings = Settings.from_env() 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", "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] [dependency-groups]
dev = [ dev = [
"pytest>=8.3.0", "pytest>=8.3.0",
@@ -21,4 +31,4 @@ dev = [
[tool.pytest.ini_options] [tool.pytest.ini_options]
testpaths = ["tests"] 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 fastapi import FastAPI
from openai import AsyncOpenAI from openai import AsyncOpenAI
from auth import ApiKeyAuthenticator from .auth import ApiKeyAuthenticator
from config import Settings from .config import Settings
from routes import router from .routes import router
from service import ChatService from .service import ChatService
from storage import JsonSessionStorage from .storage import JsonSessionStorage, UserStore
from tools import ToolRegistry from .tools import ToolRegistry
from users import UserStore
def create_app( 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 import Depends, HTTPException, Request, Security, status
from fastapi.security import APIKeyHeader 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) 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 datetime import datetime
from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator from pydantic import BaseModel, ConfigDict, Field, model_validator
from .messages import Message
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"),
]
class Session(BaseModel): class Session(BaseModel):
@@ -69,4 +31,4 @@ class Session(BaseModel):
for message in messages: for message in messages:
if isinstance(message, dict): if isinstance(message, dict):
message.setdefault("created_at", fallback) 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 fastapi import APIRouter, HTTPException, Request, status
from auth import CurrentUser from .auth import CurrentUser
from schemas import ( from .schemas import (
CreateSessionRequest, CreateSessionRequest,
CreateSessionResponse, CreateSessionResponse,
LoginRequest, LoginRequest,
@@ -13,10 +13,11 @@ from schemas import (
SessionHistoryResponse, SessionHistoryResponse,
UserUsageResponse, UserUsageResponse,
) )
from service import ChatProviderError, ChatService from .service import ChatProviderError, ChatService
from storage import SessionNotFoundError, SessionStorageError from .storage import (
from users import (
InvalidCredentialsError, InvalidCredentialsError,
SessionNotFoundError,
SessionStorageError,
UserStore, UserStore,
UserStoreError, UserStoreError,
UsernameExistsError, 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 openai import APIError, AsyncOpenAI
from domain import ( from ..domain import (
AssistantMessage, AssistantMessage,
Message, Message,
Session, Session,
@@ -13,7 +13,7 @@ from domain import (
ToolMessage, ToolMessage,
UserMessage, UserMessage,
) )
from schemas import ( from ..schemas import (
CreateSessionResponse, CreateSessionResponse,
FinalAssistantMessage, FinalAssistantMessage,
SendMessageResponse, SendMessageResponse,
@@ -21,8 +21,8 @@ from schemas import (
TokenUsage, TokenUsage,
UserUsageResponse, UserUsageResponse,
) )
from storage import JsonSessionStorage, SessionNotFoundError from ..storage import JsonSessionStorage, SessionNotFoundError
from tools import ToolRegistry from ..tools import ToolRegistry
class ChatProviderError(Exception): class ChatProviderError(Exception):
@@ -251,4 +251,4 @@ class ChatService:
def _add_usage(total: TokenUsage, usage: object) -> None: def _add_usage(total: TokenUsage, usage: object) -> None:
total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined] total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined]
total.completion_tokens += usage.completion_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 pydantic import ValidationError
from domain import Session from ..domain import Session
class SessionNotFoundError(Exception): class SessionNotFoundError(Exception):
@@ -128,4 +128,4 @@ class JsonSessionStorage:
try: try:
Path(temporary_path).unlink(missing_ok=True) Path(temporary_path).unlink(missing_ok=True)
except OSError: 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 import asyncio
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime
import json import json
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from pydantic import BaseModel, ConfigDict, Field, ValidationError from pydantic import BaseModel, ValidationError
class CurrentTimeArguments(BaseModel):
model_config = ConfigDict(extra="forbid")
timezone: str = Field(default="UTC", description="IANA timezone, e.g. Asia/Shanghai")
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -43,15 +35,9 @@ class ToolRegistry:
if timeout_seconds <= 0: if timeout_seconds <= 0:
raise ValueError("tool timeout must be greater than zero") raise ValueError("tool timeout must be greater than zero")
self.timeout_seconds = timeout_seconds self.timeout_seconds = timeout_seconds
specs = [ from .time_tools import build_builtin_specs
ToolSpec(
name="get_current_time", self._specs = {spec.name: spec for spec in build_builtin_specs()}
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}
def definitions(self) -> list[dict[str, object]]: def definitions(self) -> list[dict[str, object]]:
return [spec.api_definition() for spec in self._specs.values()] 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, ensure_ascii=False,
), ),
is_error=True, 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 import pytest
from service import ChatProviderError, ChatService from chat_api.service import ChatProviderError, ChatService
from storage import JsonSessionStorage from chat_api.storage import JsonSessionStorage
def completion( def completion(
+5 -5
View File
@@ -10,11 +10,11 @@ from fastapi.testclient import TestClient
from openai import APIConnectionError from openai import APIConnectionError
import pytest import pytest
from app import create_app from chat_api import create_app
from config import Settings from chat_api.config import Settings
from service import ChatService from chat_api.service import ChatService
from storage import JsonSessionStorage from chat_api.storage import JsonSessionStorage
from users import UserStore from chat_api.storage import UserStore
@dataclass @dataclass
+1 -1
View File
@@ -2,7 +2,7 @@ import asyncio
import json import json
import time import time
from tools import CurrentTimeArguments, ToolRegistry, ToolSpec from chat_api.tools import CurrentTimeArguments, ToolRegistry, ToolSpec
def test_builtin_tool_definitions() -> None: def test_builtin_tool_definitions() -> None:
Generated
+1 -1
View File
@@ -421,7 +421,7 @@ wheels = [
[[package]] [[package]]
name = "simple-chat-api" name = "simple-chat-api"
version = "0.1.0" version = "0.1.0"
source = { virtual = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "aiosqlite" }, { name = "aiosqlite" },
{ name = "bcrypt" }, { name = "bcrypt" },