优化项目结构
This commit is contained in:
@@ -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 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
|
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`。
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .app import app, create_app
|
||||||
|
|
||||||
|
__all__ = ["app", "create_app"]
|
||||||
@@ -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(
|
||||||
@@ -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)
|
||||||
|
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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"),
|
||||||
|
]
|
||||||
@@ -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):
|
||||||
@@ -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,
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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]
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .chat import ChatProviderError, ChatService
|
||||||
|
|
||||||
|
__all__ = ["ChatProviderError", "ChatService"]
|
||||||
@@ -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):
|
||||||
@@ -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):
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
from .registry import ToolExecutionResult, ToolRegistry, ToolSpec
|
||||||
|
from .time_tools import CurrentTimeArguments
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CurrentTimeArguments",
|
||||||
|
"ToolExecutionResult",
|
||||||
|
"ToolRegistry",
|
||||||
|
"ToolSpec",
|
||||||
|
]
|
||||||
@@ -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()]
|
||||||
@@ -92,13 +78,3 @@ def _tool_error(code: str, message: str) -> ToolExecutionResult:
|
|||||||
),
|
),
|
||||||
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()}
|
|
||||||
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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" },
|
||||||
|
|||||||
Reference in New Issue
Block a user