diff --git a/.env.example b/.env.example index 6449e65..ec90496 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/README.md b/README.md index 875130e..3267a31 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/app.py b/app.py index f881d09..141d9d4 100644 --- a/app.py +++ b/app.py @@ -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: diff --git a/config.py b/config.py index 83018d9..8e48490 100644 --- a/config.py +++ b/config.py @@ -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")), ) diff --git a/models.py b/models.py index 5164e22..e013ecf 100644 --- a/models.py +++ b/models.py @@ -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 diff --git a/service.py b/service.py index e90d0ea..374d8f7 100644 --- a/service.py +++ b/service.py @@ -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] diff --git a/tests/test_agent.py b/tests/test_agent.py new file mode 100644 index 0000000..fd9d17f --- /dev/null +++ b/tests/test_agent.py @@ -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()) diff --git a/tests/test_api.py b/tests/test_api.py index 81c3f17..efc11b5 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -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", diff --git a/tests/test_tools.py b/tests/test_tools.py new file mode 100644 index 0000000..4c2bb37 --- /dev/null +++ b/tests/test_tools.py @@ -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()) diff --git a/tools.py b/tools.py new file mode 100644 index 0000000..39587ee --- /dev/null +++ b/tools.py @@ -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")