增加工具使用功能

This commit is contained in:
2026-07-03 20:33:22 +08:00
parent 8c05311628
commit 04fefcf488
10 changed files with 662 additions and 71 deletions
+3
View File
@@ -4,3 +4,6 @@ DEEPSEEK_BASE_URL=https://api.deepseek.com
DEEPSEEK_MODEL=deepseek-v4-flash DEEPSEEK_MODEL=deepseek-v4-flash
DEFAULT_SYSTEM_PROMPT=You are a helpful assistant. DEFAULT_SYSTEM_PROMPT=You are a helpful assistant.
CHAT_DATA_DIR=data CHAT_DATA_DIR=data
MAX_TOOL_ROUNDS=5
MAX_TOOL_CALLS_PER_TURN=10
TOOL_TIMEOUT_SECONDS=5
+9 -2
View File
@@ -1,6 +1,6 @@
# Simple Chat API # 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` 服务默认运行在 `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`。只有模型完整响应成功后,本轮消息才会写入会话历史。 API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完整响应成功后,本轮消息才会写入会话历史。
模型可以按需调用服务端白名单工具:
- `get_current_time`:查询指定 IANA 时区的当前时间。
- `calculate`:计算受限的基础算术表达式,不执行任意 Python 代码。
工具调用无需增加请求参数。服务端会执行工具并把结果返回模型,直到模型生成最终回答。发送消息接口通过 `tools_use` 返回本轮使用的工具名称,完整调用过程可通过历史接口查询;未调用工具时 `tools_use` 为空数组。
获取指定会话的完整历史: 获取指定会话的完整历史:
```bash ```bash
+4
View File
@@ -14,6 +14,7 @@ from models import (
) )
from service import ChatProviderError, ChatService from service import ChatProviderError, ChatService
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
from tools import ToolRegistry
def create_app( def create_app(
@@ -37,6 +38,9 @@ def create_app(
storage=storage, storage=storage,
client=resolved_client, client=resolved_client,
model=resolved_settings.model, 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 yield
if client is None: if client is None:
+8
View File
@@ -13,6 +13,9 @@ class Settings:
model: str model: str
default_system_prompt: str default_system_prompt: str
data_dir: Path data_dir: Path
max_tool_rounds: int = 5
max_tool_calls_per_turn: int = 10
tool_timeout_seconds: float = 5.0
@classmethod @classmethod
def from_env(cls) -> "Settings": def from_env(cls) -> "Settings":
@@ -26,4 +29,9 @@ class Settings:
"DEFAULT_SYSTEM_PROMPT", "You are a helpful assistant." "DEFAULT_SYSTEM_PROMPT", "You are a helpful assistant."
), ),
data_dir=Path(os.getenv("CHAT_DATA_DIR", "data")), 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")),
) )
+42 -4
View File
@@ -1,15 +1,46 @@
from datetime import datetime from datetime import datetime
from typing import Literal from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
class Message(BaseModel): class ToolFunctionCall(BaseModel):
role: Literal["user", "assistant"] name: str
arguments: str
class ToolCall(BaseModel):
id: str
type: Literal["function"] = "function"
function: ToolFunctionCall
class UserMessage(BaseModel):
role: Literal["user"] = "user"
content: str content: str
created_at: datetime 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):
model_config = ConfigDict(extra="forbid") model_config = ConfigDict(extra="forbid")
@@ -72,9 +103,16 @@ class TokenUsage(BaseModel):
total_tokens: int total_tokens: int
class FinalAssistantMessage(BaseModel):
role: Literal["assistant"] = "assistant"
content: str
created_at: datetime
class SendMessageResponse(BaseModel): class SendMessageResponse(BaseModel):
session_id: str session_id: str
message: Message message: FinalAssistantMessage
tools_use: list[str] = Field(default_factory=list)
usage: TokenUsage | None usage: TokenUsage | None
+137 -61
View File
@@ -1,15 +1,24 @@
import asyncio import asyncio
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import Any
from openai import APIError, AsyncOpenAI from openai import APIError, AsyncOpenAI
from models import ( from models import (
AssistantMessage,
FinalAssistantMessage,
Message, Message,
SendMessageResponse, SendMessageResponse,
Session, Session,
SessionHistoryResponse, SessionHistoryResponse,
TokenUsage, TokenUsage,
ToolCall,
ToolFunctionCall,
ToolMessage,
UserMessage,
) )
from storage import JsonSessionStorage from storage import JsonSessionStorage
from tools import ToolRegistry
class ChatProviderError(Exception): class ChatProviderError(Exception):
@@ -22,10 +31,18 @@ class ChatService:
storage: JsonSessionStorage, storage: JsonSessionStorage,
client: AsyncOpenAI, client: AsyncOpenAI,
model: str, model: str,
tool_registry: ToolRegistry | None = None,
max_tool_rounds: int = 5,
max_tool_calls_per_turn: int = 10,
) -> None: ) -> 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.storage = storage
self.client = client self.client = client
self.model = model 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] = {} self._locks: dict[str, asyncio.Lock] = {}
async def generate_response( async def generate_response(
@@ -33,34 +50,89 @@ class ChatService:
) -> SendMessageResponse: ) -> SendMessageResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock()) lock = self._locks.setdefault(session_id, asyncio.Lock())
async with 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) 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: for _ in range(self.max_tool_rounds):
completion = await self.client.chat.completions.create( completion = await self._request_completion(api_messages)
model=self.model, if completion.usage is not None:
messages=api_messages, # type: ignore[arg-type] self._add_usage(usage, completion.usage)
extra_body={"thinking": {"type": "disabled"}}, 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"
) )
assistant_content = completion.choices[0].message.content final_message = AssistantMessage(
except (APIError, IndexError, AttributeError) as exc: content=assistant_content,
raise ChatProviderError("Upstream chat request failed") from exc created_at=datetime.now(UTC),
if assistant_content is None:
raise ChatProviderError("Upstream chat provider returned an empty response")
assistant_message = await self._save_exchange(
session, content, user_created_at, assistant_content
) )
usage = self._parse_usage(completion.usage) pending_messages.append(final_message)
session.messages.extend(pending_messages)
session.updated_at = final_message.created_at
await self.storage.write(session)
return SendMessageResponse( return SendMessageResponse(
session_id=session_id, session_id=session_id,
message=assistant_message, message=FinalAssistantMessage(
usage=usage, 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),
)
pending_messages.append(assistant_tool_message)
api_messages.append(self._message_to_api(assistant_tool_message))
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))
raise ChatProviderError("Tool round limit exceeded")
async def get_session_history(self, session_id: str) -> SessionHistoryResponse: async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock()) lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock: async with lock:
@@ -73,50 +145,54 @@ class ChatService:
messages=session.messages, messages=session.messages,
) )
@staticmethod async def _request_completion(
def _build_api_messages( self, messages: list[dict[str, Any]]
session: Session, content: str ) -> Any:
) -> list[dict[str, str]]: try:
messages = [{"role": "system", "content": session.system_prompt}] completion = await self.client.chat.completions.create(
messages.extend( model=self.model,
{"role": message.role, "content": message.content} messages=messages, # type: ignore[arg-type]
for message in session.messages tools=self.tool_registry.definitions(), # type: ignore[arg-type]
extra_body={"thinking": {"type": "disabled"}},
) )
messages.append({"role": "user", "content": content}) 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 return messages
async def _save_exchange( @staticmethod
self, def _message_to_api(message: Message) -> dict[str, Any]:
session: Session, if isinstance(message, UserMessage):
user_content: str, return {"role": "user", "content": message.content}
user_created_at: datetime, if isinstance(message, AssistantMessage):
assistant_content: str, result: dict[str, Any] = {
) -> Message: "role": "assistant",
assistant_message = Message( "content": message.content,
role="assistant", }
content=assistant_content, if message.tool_calls:
created_at=datetime.now(UTC), result["tool_calls"] = [
) tool_call.model_dump() for tool_call in message.tool_calls
session.messages.extend(
[
Message(
role="user",
content=user_content,
created_at=user_created_at,
),
assistant_message,
] ]
) return result
session.updated_at = assistant_message.created_at return {
await self.storage.write(session) "role": "tool",
return assistant_message "content": message.content,
"tool_call_id": message.tool_call_id,
}
@staticmethod @staticmethod
def _parse_usage(usage: object | None) -> TokenUsage | None: def _add_usage(total: TokenUsage, usage: object) -> None:
if usage is None: total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined]
return None total.completion_tokens += usage.completion_tokens # type: ignore[attr-defined]
return TokenUsage( total.total_tokens += usage.total_tokens # type: ignore[attr-defined]
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]
)
+202
View File
@@ -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())
+2
View File
@@ -136,7 +136,9 @@ def test_multi_turn_request_contains_full_history(
"completion_tokens": 2, "completion_tokens": 2,
"total_tokens": 12, "total_tokens": 12,
} }
assert second.json()["tools_use"] == []
datetime.fromisoformat(second.json()["message"]["created_at"]) 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()) data = json.loads((tmp_path / f"{session_id}.json").read_text())
assert [message["content"] for message in data["messages"]] == [ assert [message["content"] for message in data["messages"]] == [
"first", "first",
+80
View File
@@ -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())
+171
View File
@@ -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")