优化项目结构
This commit is contained in:
@@ -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`。
|
||||
|
||||
@@ -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
@@ -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
@@ -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 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(
|
||||
@@ -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)
|
||||
|
||||
@@ -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 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):
|
||||
@@ -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,
|
||||
@@ -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 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):
|
||||
@@ -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):
|
||||
@@ -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
|
||||
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()]
|
||||
@@ -92,13 +78,3 @@ def _tool_error(code: str, message: str) -> ToolExecutionResult:
|
||||
),
|
||||
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
|
||||
|
||||
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
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user