增加工具使用功能
This commit is contained in:
@@ -4,3 +4,6 @@ DEEPSEEK_BASE_URL=https://api.deepseek.com
|
||||
DEEPSEEK_MODEL=deepseek-v4-flash
|
||||
DEFAULT_SYSTEM_PROMPT=You are a helpful assistant.
|
||||
CHAT_DATA_DIR=data
|
||||
MAX_TOOL_ROUNDS=5
|
||||
MAX_TOOL_CALLS_PER_TURN=10
|
||||
TOOL_TIMEOUT_SECONDS=5
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Simple Chat API
|
||||
|
||||
一个使用 FastAPI 和 DeepSeek API 的最小多轮对话服务。每个会话的系统提示词与消息历史保存在本地 JSON 文件中。
|
||||
一个使用 FastAPI 和 DeepSeek API 的多轮 Agent 服务。每个会话的系统提示词、消息历史和工具调用过程保存在本地 JSON 文件中。
|
||||
|
||||
## 启动
|
||||
|
||||
@@ -13,7 +13,7 @@ uv run uvicorn main:app --reload
|
||||
|
||||
服务默认运行在 `http://127.0.0.1:8000`,交互式 API 文档位于 `/docs`。
|
||||
|
||||
服务会自动加载项目根目录的 `.env`,已有系统环境变量优先级更高。可配置 `DEEPSEEK_BASE_URL`、`DEEPSEEK_MODEL`、`DEFAULT_SYSTEM_PROMPT` 和 `CHAT_DATA_DIR`,完整示例见 `.env.example`。第一版应只使用一个 Uvicorn worker。
|
||||
服务会自动加载项目根目录的 `.env`,已有系统环境变量优先级更高。可配置 `DEEPSEEK_BASE_URL`、`DEEPSEEK_MODEL`、`DEFAULT_SYSTEM_PROMPT`、`CHAT_DATA_DIR` 和工具调用限制,完整示例见 `.env.example`。第一版应只使用一个 Uvicorn worker。
|
||||
|
||||
## 使用
|
||||
|
||||
@@ -35,6 +35,13 @@ curl -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \
|
||||
|
||||
API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完整响应成功后,本轮消息才会写入会话历史。
|
||||
|
||||
模型可以按需调用服务端白名单工具:
|
||||
|
||||
- `get_current_time`:查询指定 IANA 时区的当前时间。
|
||||
- `calculate`:计算受限的基础算术表达式,不执行任意 Python 代码。
|
||||
|
||||
工具调用无需增加请求参数。服务端会执行工具并把结果返回模型,直到模型生成最终回答。发送消息接口通过 `tools_use` 返回本轮使用的工具名称,完整调用过程可通过历史接口查询;未调用工具时 `tools_use` 为空数组。
|
||||
|
||||
获取指定会话的完整历史:
|
||||
|
||||
```bash
|
||||
|
||||
@@ -14,6 +14,7 @@ from models import (
|
||||
)
|
||||
from service import ChatProviderError, ChatService
|
||||
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
|
||||
from tools import ToolRegistry
|
||||
|
||||
|
||||
def create_app(
|
||||
@@ -37,6 +38,9 @@ def create_app(
|
||||
storage=storage,
|
||||
client=resolved_client,
|
||||
model=resolved_settings.model,
|
||||
tool_registry=ToolRegistry(resolved_settings.tool_timeout_seconds),
|
||||
max_tool_rounds=resolved_settings.max_tool_rounds,
|
||||
max_tool_calls_per_turn=resolved_settings.max_tool_calls_per_turn,
|
||||
)
|
||||
yield
|
||||
if client is None:
|
||||
|
||||
@@ -13,6 +13,9 @@ class Settings:
|
||||
model: str
|
||||
default_system_prompt: str
|
||||
data_dir: Path
|
||||
max_tool_rounds: int = 5
|
||||
max_tool_calls_per_turn: int = 10
|
||||
tool_timeout_seconds: float = 5.0
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "Settings":
|
||||
@@ -26,4 +29,9 @@ class Settings:
|
||||
"DEFAULT_SYSTEM_PROMPT", "You are a helpful assistant."
|
||||
),
|
||||
data_dir=Path(os.getenv("CHAT_DATA_DIR", "data")),
|
||||
max_tool_rounds=int(os.getenv("MAX_TOOL_ROUNDS", "5")),
|
||||
max_tool_calls_per_turn=int(
|
||||
os.getenv("MAX_TOOL_CALLS_PER_TURN", "10")
|
||||
),
|
||||
tool_timeout_seconds=float(os.getenv("TOOL_TIMEOUT_SECONDS", "5")),
|
||||
)
|
||||
|
||||
@@ -1,15 +1,46 @@
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
|
||||
class Message(BaseModel):
|
||||
role: Literal["user", "assistant"]
|
||||
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):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
@@ -72,9 +103,16 @@ class TokenUsage(BaseModel):
|
||||
total_tokens: int
|
||||
|
||||
|
||||
class FinalAssistantMessage(BaseModel):
|
||||
role: Literal["assistant"] = "assistant"
|
||||
content: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class SendMessageResponse(BaseModel):
|
||||
session_id: str
|
||||
message: Message
|
||||
message: FinalAssistantMessage
|
||||
tools_use: list[str] = Field(default_factory=list)
|
||||
usage: TokenUsage | None
|
||||
|
||||
|
||||
|
||||
+141
-65
@@ -1,15 +1,24 @@
|
||||
import asyncio
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from openai import APIError, AsyncOpenAI
|
||||
|
||||
from models import (
|
||||
AssistantMessage,
|
||||
FinalAssistantMessage,
|
||||
Message,
|
||||
SendMessageResponse,
|
||||
Session,
|
||||
SessionHistoryResponse,
|
||||
TokenUsage,
|
||||
ToolCall,
|
||||
ToolFunctionCall,
|
||||
ToolMessage,
|
||||
UserMessage,
|
||||
)
|
||||
from storage import JsonSessionStorage
|
||||
from tools import ToolRegistry
|
||||
|
||||
|
||||
class ChatProviderError(Exception):
|
||||
@@ -22,10 +31,18 @@ class ChatService:
|
||||
storage: JsonSessionStorage,
|
||||
client: AsyncOpenAI,
|
||||
model: str,
|
||||
tool_registry: ToolRegistry | None = None,
|
||||
max_tool_rounds: int = 5,
|
||||
max_tool_calls_per_turn: int = 10,
|
||||
) -> None:
|
||||
if max_tool_rounds <= 0 or max_tool_calls_per_turn <= 0:
|
||||
raise ValueError("tool limits must be greater than zero")
|
||||
self.storage = storage
|
||||
self.client = client
|
||||
self.model = model
|
||||
self.tool_registry = tool_registry or ToolRegistry()
|
||||
self.max_tool_rounds = max_tool_rounds
|
||||
self.max_tool_calls_per_turn = max_tool_calls_per_turn
|
||||
self._locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
async def generate_response(
|
||||
@@ -33,33 +50,88 @@ class ChatService:
|
||||
) -> SendMessageResponse:
|
||||
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
||||
async with lock:
|
||||
user_created_at = datetime.now(UTC)
|
||||
user_message = UserMessage(content=content, created_at=datetime.now(UTC))
|
||||
session = await self.storage.read(session_id)
|
||||
api_messages = self._build_api_messages(session, content)
|
||||
api_messages = self._build_api_messages(session)
|
||||
api_messages.append(self._message_to_api(user_message))
|
||||
pending_messages: list[Message] = [user_message]
|
||||
tools_use: list[str] = []
|
||||
usage = TokenUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
|
||||
has_usage = False
|
||||
tool_call_count = 0
|
||||
|
||||
try:
|
||||
completion = await self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=api_messages, # type: ignore[arg-type]
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
for _ in range(self.max_tool_rounds):
|
||||
completion = await self._request_completion(api_messages)
|
||||
if completion.usage is not None:
|
||||
self._add_usage(usage, completion.usage)
|
||||
has_usage = True
|
||||
|
||||
provider_message = completion.choices[0].message
|
||||
provider_tool_calls = getattr(provider_message, "tool_calls", None) or []
|
||||
if not provider_tool_calls:
|
||||
assistant_content = provider_message.content
|
||||
if not assistant_content:
|
||||
raise ChatProviderError(
|
||||
"Upstream chat provider returned an empty response"
|
||||
)
|
||||
final_message = AssistantMessage(
|
||||
content=assistant_content,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
pending_messages.append(final_message)
|
||||
session.messages.extend(pending_messages)
|
||||
session.updated_at = final_message.created_at
|
||||
await self.storage.write(session)
|
||||
return SendMessageResponse(
|
||||
session_id=session_id,
|
||||
message=FinalAssistantMessage(
|
||||
content=assistant_content,
|
||||
created_at=final_message.created_at,
|
||||
),
|
||||
tools_use=tools_use,
|
||||
usage=usage if has_usage else None,
|
||||
)
|
||||
|
||||
tool_call_count += len(provider_tool_calls)
|
||||
if tool_call_count > self.max_tool_calls_per_turn:
|
||||
raise ChatProviderError("Tool call limit exceeded")
|
||||
|
||||
tool_calls = [
|
||||
ToolCall(
|
||||
id=tool_call.id,
|
||||
function=ToolFunctionCall(
|
||||
name=tool_call.function.name,
|
||||
arguments=tool_call.function.arguments,
|
||||
),
|
||||
)
|
||||
for tool_call in provider_tool_calls
|
||||
]
|
||||
assistant_tool_message = AssistantMessage(
|
||||
content=provider_message.content,
|
||||
tool_calls=tool_calls,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
assistant_content = completion.choices[0].message.content
|
||||
except (APIError, IndexError, AttributeError) as exc:
|
||||
raise ChatProviderError("Upstream chat request failed") from exc
|
||||
pending_messages.append(assistant_tool_message)
|
||||
api_messages.append(self._message_to_api(assistant_tool_message))
|
||||
|
||||
if assistant_content is None:
|
||||
raise ChatProviderError("Upstream chat provider returned an empty response")
|
||||
for tool_call in tool_calls:
|
||||
if tool_call.function.name not in tools_use:
|
||||
tools_use.append(tool_call.function.name)
|
||||
result = await self.tool_registry.execute(
|
||||
tool_call.function.name,
|
||||
tool_call.function.arguments,
|
||||
)
|
||||
tool_message = ToolMessage(
|
||||
content=result.content,
|
||||
tool_call_id=tool_call.id,
|
||||
name=tool_call.function.name,
|
||||
is_error=result.is_error,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
pending_messages.append(tool_message)
|
||||
api_messages.append(self._message_to_api(tool_message))
|
||||
|
||||
assistant_message = await self._save_exchange(
|
||||
session, content, user_created_at, assistant_content
|
||||
)
|
||||
usage = self._parse_usage(completion.usage)
|
||||
|
||||
return SendMessageResponse(
|
||||
session_id=session_id,
|
||||
message=assistant_message,
|
||||
usage=usage,
|
||||
)
|
||||
raise ChatProviderError("Tool round limit exceeded")
|
||||
|
||||
async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
|
||||
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
||||
@@ -73,50 +145,54 @@ class ChatService:
|
||||
messages=session.messages,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_api_messages(
|
||||
session: Session, content: str
|
||||
) -> list[dict[str, str]]:
|
||||
messages = [{"role": "system", "content": session.system_prompt}]
|
||||
messages.extend(
|
||||
{"role": message.role, "content": message.content}
|
||||
for message in session.messages
|
||||
)
|
||||
messages.append({"role": "user", "content": content})
|
||||
async def _request_completion(
|
||||
self, messages: list[dict[str, Any]]
|
||||
) -> Any:
|
||||
try:
|
||||
completion = await self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=messages, # type: ignore[arg-type]
|
||||
tools=self.tool_registry.definitions(), # type: ignore[arg-type]
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
)
|
||||
if not completion.choices:
|
||||
raise ChatProviderError("Upstream chat provider returned no choices")
|
||||
return completion
|
||||
except ChatProviderError:
|
||||
raise
|
||||
except (APIError, IndexError, AttributeError) as exc:
|
||||
raise ChatProviderError("Upstream chat request failed") from exc
|
||||
|
||||
@classmethod
|
||||
def _build_api_messages(cls, session: Session) -> list[dict[str, Any]]:
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": session.system_prompt}
|
||||
]
|
||||
messages.extend(cls._message_to_api(message) for message in session.messages)
|
||||
return messages
|
||||
|
||||
async def _save_exchange(
|
||||
self,
|
||||
session: Session,
|
||||
user_content: str,
|
||||
user_created_at: datetime,
|
||||
assistant_content: str,
|
||||
) -> Message:
|
||||
assistant_message = Message(
|
||||
role="assistant",
|
||||
content=assistant_content,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
session.messages.extend(
|
||||
[
|
||||
Message(
|
||||
role="user",
|
||||
content=user_content,
|
||||
created_at=user_created_at,
|
||||
),
|
||||
assistant_message,
|
||||
]
|
||||
)
|
||||
session.updated_at = assistant_message.created_at
|
||||
await self.storage.write(session)
|
||||
return assistant_message
|
||||
@staticmethod
|
||||
def _message_to_api(message: Message) -> dict[str, Any]:
|
||||
if isinstance(message, UserMessage):
|
||||
return {"role": "user", "content": message.content}
|
||||
if isinstance(message, AssistantMessage):
|
||||
result: dict[str, Any] = {
|
||||
"role": "assistant",
|
||||
"content": message.content,
|
||||
}
|
||||
if message.tool_calls:
|
||||
result["tool_calls"] = [
|
||||
tool_call.model_dump() for tool_call in message.tool_calls
|
||||
]
|
||||
return result
|
||||
return {
|
||||
"role": "tool",
|
||||
"content": message.content,
|
||||
"tool_call_id": message.tool_call_id,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _parse_usage(usage: object | None) -> TokenUsage | None:
|
||||
if usage is None:
|
||||
return None
|
||||
return TokenUsage(
|
||||
prompt_tokens=usage.prompt_tokens, # type: ignore[attr-defined]
|
||||
completion_tokens=usage.completion_tokens, # type: ignore[attr-defined]
|
||||
total_tokens=usage.total_tokens, # type: ignore[attr-defined]
|
||||
)
|
||||
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]
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
import asyncio
|
||||
import json
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from service import ChatProviderError, ChatService
|
||||
from storage import JsonSessionStorage
|
||||
|
||||
|
||||
def completion(
|
||||
*,
|
||||
content: str | None,
|
||||
tool_calls: list[SimpleNamespace] | None = None,
|
||||
prompt_tokens: int = 10,
|
||||
completion_tokens: int = 2,
|
||||
) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
choices=[
|
||||
SimpleNamespace(
|
||||
message=SimpleNamespace(content=content, tool_calls=tool_calls)
|
||||
)
|
||||
],
|
||||
usage=SimpleNamespace(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def tool_call(call_id: str, name: str, arguments: str) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id=call_id,
|
||||
type="function",
|
||||
function=SimpleNamespace(name=name, arguments=arguments),
|
||||
)
|
||||
|
||||
|
||||
class ScriptedClient:
|
||||
def __init__(self, responses: list[SimpleNamespace]) -> None:
|
||||
self.responses = responses
|
||||
self.calls: list[dict[str, Any]] = []
|
||||
self.chat = SimpleNamespace(
|
||||
completions=SimpleNamespace(create=self.create)
|
||||
)
|
||||
|
||||
async def create(self, **kwargs: Any) -> SimpleNamespace:
|
||||
self.calls.append(kwargs)
|
||||
return self.responses.pop(0)
|
||||
|
||||
|
||||
def test_agent_executes_tool_and_persists_full_history(tmp_path: Path) -> None:
|
||||
async def scenario() -> None:
|
||||
provider = ScriptedClient(
|
||||
[
|
||||
completion(
|
||||
content=None,
|
||||
tool_calls=[tool_call("call-1", "calculate", '{"expression":"2+2"}')],
|
||||
),
|
||||
completion(content="结果是 4。", prompt_tokens=20, completion_tokens=3),
|
||||
]
|
||||
)
|
||||
storage = JsonSessionStorage(tmp_path)
|
||||
session = await storage.create("Use tools when needed.")
|
||||
service = ChatService(storage, provider, "test-model") # type: ignore[arg-type]
|
||||
|
||||
response = await service.generate_response(session.session_id, "2+2 等于多少?")
|
||||
|
||||
assert response.message.content == "结果是 4。"
|
||||
assert "tool_calls" not in response.message.model_dump()
|
||||
assert response.tools_use == ["calculate"]
|
||||
assert response.usage is not None
|
||||
assert response.usage.model_dump() == {
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 35,
|
||||
}
|
||||
assert "tools" in provider.calls[0]
|
||||
second_messages = provider.calls[1]["messages"]
|
||||
assert [message["role"] for message in second_messages] == [
|
||||
"system",
|
||||
"user",
|
||||
"assistant",
|
||||
"tool",
|
||||
]
|
||||
assert json.loads(second_messages[-1]["content"])["result"] == 4
|
||||
|
||||
saved = await storage.read(session.session_id)
|
||||
assert [message.role for message in saved.messages] == [
|
||||
"user",
|
||||
"assistant",
|
||||
"tool",
|
||||
"assistant",
|
||||
]
|
||||
assert saved.messages[1].tool_calls[0].id == "call-1" # type: ignore[union-attr]
|
||||
assert saved.messages[2].tool_call_id == "call-1" # type: ignore[union-attr]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_agent_executes_multiple_tools_in_order(tmp_path: Path) -> None:
|
||||
async def scenario() -> None:
|
||||
provider = ScriptedClient(
|
||||
[
|
||||
completion(
|
||||
content=None,
|
||||
tool_calls=[
|
||||
tool_call("call-1", "calculate", '{"expression":"6*7"}'),
|
||||
tool_call(
|
||||
"call-2",
|
||||
"get_current_time",
|
||||
'{"timezone":"UTC"}',
|
||||
),
|
||||
],
|
||||
),
|
||||
completion(content="完成"),
|
||||
]
|
||||
)
|
||||
storage = JsonSessionStorage(tmp_path)
|
||||
session = await storage.create("system")
|
||||
service = ChatService(storage, provider, "test-model") # type: ignore[arg-type]
|
||||
|
||||
response = await service.generate_response(session.session_id, "计算并查询时间")
|
||||
|
||||
messages = provider.calls[1]["messages"]
|
||||
assert [message.get("tool_call_id") for message in messages[-2:]] == [
|
||||
"call-1",
|
||||
"call-2",
|
||||
]
|
||||
assert response.tools_use == ["calculate", "get_current_time"]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_tool_error_is_returned_to_model(tmp_path: Path) -> None:
|
||||
async def scenario() -> None:
|
||||
provider = ScriptedClient(
|
||||
[
|
||||
completion(
|
||||
content=None,
|
||||
tool_calls=[tool_call("bad-1", "calculate", "not-json")],
|
||||
),
|
||||
completion(content="参数无效。"),
|
||||
]
|
||||
)
|
||||
storage = JsonSessionStorage(tmp_path)
|
||||
session = await storage.create("system")
|
||||
service = ChatService(storage, provider, "test-model") # type: ignore[arg-type]
|
||||
|
||||
await service.generate_response(session.session_id, "test")
|
||||
|
||||
saved = await storage.read(session.session_id)
|
||||
tool_message = saved.messages[2]
|
||||
assert tool_message.role == "tool"
|
||||
assert tool_message.is_error is True # type: ignore[union-attr]
|
||||
assert "invalid_arguments" in tool_message.content
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("max_rounds", "max_calls", "calls"),
|
||||
[
|
||||
(1, 10, [tool_call("call-1", "calculate", '{"expression":"1+1"}')]),
|
||||
(
|
||||
5,
|
||||
1,
|
||||
[
|
||||
tool_call("call-1", "calculate", '{"expression":"1+1"}'),
|
||||
tool_call("call-2", "calculate", '{"expression":"2+2"}'),
|
||||
],
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_agent_limits_do_not_persist_partial_turn(
|
||||
tmp_path: Path,
|
||||
max_rounds: int,
|
||||
max_calls: int,
|
||||
calls: list[SimpleNamespace],
|
||||
) -> None:
|
||||
async def scenario() -> None:
|
||||
provider = ScriptedClient([completion(content=None, tool_calls=calls)])
|
||||
storage = JsonSessionStorage(tmp_path)
|
||||
session = await storage.create("system")
|
||||
service = ChatService(
|
||||
storage,
|
||||
provider, # type: ignore[arg-type]
|
||||
"test-model",
|
||||
max_tool_rounds=max_rounds,
|
||||
max_tool_calls_per_turn=max_calls,
|
||||
)
|
||||
|
||||
with pytest.raises(ChatProviderError):
|
||||
await service.generate_response(session.session_id, "test")
|
||||
|
||||
saved = await storage.read(session.session_id)
|
||||
assert saved.messages == []
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -136,7 +136,9 @@ def test_multi_turn_request_contains_full_history(
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 12,
|
||||
}
|
||||
assert second.json()["tools_use"] == []
|
||||
datetime.fromisoformat(second.json()["message"]["created_at"])
|
||||
assert "tool_calls" not in second.json()["message"]
|
||||
data = json.loads((tmp_path / f"{session_id}.json").read_text())
|
||||
assert [message["content"] for message in data["messages"]] == [
|
||||
"first",
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
|
||||
from tools import CalculatorArguments, ToolRegistry, ToolSpec
|
||||
|
||||
|
||||
def test_builtin_tool_definitions() -> None:
|
||||
registry = ToolRegistry()
|
||||
|
||||
definitions = registry.definitions()
|
||||
|
||||
assert [item["function"]["name"] for item in definitions] == [ # type: ignore[index]
|
||||
"get_current_time",
|
||||
"calculate",
|
||||
]
|
||||
|
||||
|
||||
def test_calculator_and_time_tools() -> None:
|
||||
async def scenario() -> None:
|
||||
registry = ToolRegistry()
|
||||
calculation = await registry.execute(
|
||||
"calculate", '{"expression":"(2 + 3) * 4"}'
|
||||
)
|
||||
current_time = await registry.execute(
|
||||
"get_current_time", '{"timezone":"Asia/Shanghai"}'
|
||||
)
|
||||
|
||||
assert calculation.is_error is False
|
||||
assert json.loads(calculation.content)["result"] == 20
|
||||
assert current_time.is_error is False
|
||||
assert json.loads(current_time.content)["timezone"] == "Asia/Shanghai"
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_tool_argument_and_expression_errors_are_safe() -> None:
|
||||
async def scenario() -> None:
|
||||
registry = ToolRegistry()
|
||||
invalid_json = await registry.execute("calculate", "not-json")
|
||||
unknown = await registry.execute("missing", "{}")
|
||||
unsafe = await registry.execute(
|
||||
"calculate", '{"expression":"__import__(\\"os\\").system(\\"id\\")"}'
|
||||
)
|
||||
invalid_timezone = await registry.execute(
|
||||
"get_current_time", '{"timezone":"Not/A_Real_Zone"}'
|
||||
)
|
||||
|
||||
assert invalid_json.is_error is True
|
||||
assert "invalid_arguments" in invalid_json.content
|
||||
assert unknown.is_error is True
|
||||
assert "unknown_tool" in unknown.content
|
||||
assert unsafe.is_error is True
|
||||
assert "execution_failed" in unsafe.content
|
||||
assert invalid_timezone.is_error is True
|
||||
assert "execution_failed" in invalid_timezone.content
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_tool_timeout_returns_error() -> None:
|
||||
def slow_handler(arguments):
|
||||
time.sleep(0.05)
|
||||
return {"result": 1}
|
||||
|
||||
async def scenario() -> None:
|
||||
registry = ToolRegistry(timeout_seconds=0.01)
|
||||
registry._specs["calculate"] = ToolSpec(
|
||||
name="calculate",
|
||||
description="slow",
|
||||
arguments_model=CalculatorArguments,
|
||||
handler=slow_handler,
|
||||
)
|
||||
|
||||
result = await registry.execute("calculate", '{"expression":"1+1"}')
|
||||
|
||||
assert result.is_error is True
|
||||
assert "timeout" in result.content
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -0,0 +1,171 @@
|
||||
import ast
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
import json
|
||||
import math
|
||||
import operator
|
||||
from typing import Any
|
||||
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")
|
||||
|
||||
|
||||
class CalculatorArguments(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
expression: str = Field(description="Arithmetic expression to evaluate")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolSpec:
|
||||
name: str
|
||||
description: str
|
||||
arguments_model: type[BaseModel]
|
||||
handler: Callable[[BaseModel], dict[str, object]]
|
||||
|
||||
def api_definition(self) -> dict[str, object]:
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.arguments_model.model_json_schema(),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolExecutionResult:
|
||||
content: str
|
||||
is_error: bool
|
||||
|
||||
|
||||
class ToolRegistry:
|
||||
def __init__(self, timeout_seconds: float = 5.0) -> None:
|
||||
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,
|
||||
),
|
||||
ToolSpec(
|
||||
name="calculate",
|
||||
description="Safely evaluate a basic arithmetic expression.",
|
||||
arguments_model=CalculatorArguments,
|
||||
handler=_calculate,
|
||||
),
|
||||
]
|
||||
self._specs = {spec.name: spec for spec in specs}
|
||||
|
||||
def definitions(self) -> list[dict[str, object]]:
|
||||
return [spec.api_definition() for spec in self._specs.values()]
|
||||
|
||||
async def execute(self, name: str, arguments: str) -> ToolExecutionResult:
|
||||
spec = self._specs.get(name)
|
||||
if spec is None:
|
||||
return _tool_error("unknown_tool", f"Unknown tool: {name}")
|
||||
|
||||
try:
|
||||
raw_arguments = json.loads(arguments)
|
||||
if not isinstance(raw_arguments, dict):
|
||||
raise ValueError("arguments must be a JSON object")
|
||||
parsed_arguments = spec.arguments_model.model_validate(raw_arguments)
|
||||
except (json.JSONDecodeError, ValidationError, ValueError) as exc:
|
||||
return _tool_error("invalid_arguments", str(exc))
|
||||
|
||||
try:
|
||||
result = await asyncio.wait_for(
|
||||
asyncio.to_thread(spec.handler, parsed_arguments),
|
||||
timeout=self.timeout_seconds,
|
||||
)
|
||||
return ToolExecutionResult(
|
||||
content=json.dumps(result, ensure_ascii=False),
|
||||
is_error=False,
|
||||
)
|
||||
except TimeoutError:
|
||||
return _tool_error("timeout", f"Tool {name} timed out")
|
||||
except Exception:
|
||||
return _tool_error("execution_failed", f"Tool {name} failed")
|
||||
|
||||
|
||||
def _tool_error(code: str, message: str) -> ToolExecutionResult:
|
||||
return ToolExecutionResult(
|
||||
content=json.dumps(
|
||||
{"error": {"code": code, "message": message}},
|
||||
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()}
|
||||
|
||||
|
||||
_BINARY_OPERATORS: dict[type[ast.operator], Callable[[Any, Any], Any]] = {
|
||||
ast.Add: operator.add,
|
||||
ast.Sub: operator.sub,
|
||||
ast.Mult: operator.mul,
|
||||
ast.Div: operator.truediv,
|
||||
ast.FloorDiv: operator.floordiv,
|
||||
ast.Mod: operator.mod,
|
||||
ast.Pow: operator.pow,
|
||||
}
|
||||
_UNARY_OPERATORS: dict[type[ast.unaryop], Callable[[Any], Any]] = {
|
||||
ast.UAdd: operator.pos,
|
||||
ast.USub: operator.neg,
|
||||
}
|
||||
_MAX_EXPRESSION_LENGTH = 200
|
||||
_MAX_ABSOLUTE_RESULT = 1e100
|
||||
_MAX_EXPONENT = 100
|
||||
|
||||
|
||||
def _calculate(arguments: BaseModel) -> dict[str, object]:
|
||||
assert isinstance(arguments, CalculatorArguments)
|
||||
expression = arguments.expression.strip()
|
||||
if not expression or len(expression) > _MAX_EXPRESSION_LENGTH:
|
||||
raise ValueError("expression is empty or too long")
|
||||
tree = ast.parse(expression, mode="eval")
|
||||
result = _evaluate_node(tree.body)
|
||||
if isinstance(result, float) and not math.isfinite(result):
|
||||
raise ValueError("result is not finite")
|
||||
if abs(result) > _MAX_ABSOLUTE_RESULT:
|
||||
raise ValueError("result is too large")
|
||||
return {"expression": expression, "result": result}
|
||||
|
||||
|
||||
def _evaluate_node(node: ast.AST) -> int | float:
|
||||
if isinstance(node, ast.Constant):
|
||||
if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
|
||||
raise ValueError("only numeric constants are allowed")
|
||||
return node.value
|
||||
if isinstance(node, ast.UnaryOp) and type(node.op) in _UNARY_OPERATORS:
|
||||
return _UNARY_OPERATORS[type(node.op)](_evaluate_node(node.operand))
|
||||
if isinstance(node, ast.BinOp) and type(node.op) in _BINARY_OPERATORS:
|
||||
left = _evaluate_node(node.left)
|
||||
right = _evaluate_node(node.right)
|
||||
if isinstance(node.op, ast.Pow) and abs(right) > _MAX_EXPONENT:
|
||||
raise ValueError("exponent is too large")
|
||||
result = _BINARY_OPERATORS[type(node.op)](left, right)
|
||||
if abs(result) > _MAX_ABSOLUTE_RESULT:
|
||||
raise ValueError("intermediate result is too large")
|
||||
return result
|
||||
raise ValueError("unsupported expression")
|
||||
Reference in New Issue
Block a user