import asyncio import json from pathlib import Path from types import SimpleNamespace from typing import Any import pytest from chat_api.service import ChatProviderError, ChatService from chat_api.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", "get_current_time", '{"timezone":"UTC"}', ) ], ), completion(content="当前时间已查询。", 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, "现在几点?") assert response.message.content == "当前时间已查询。" assert "tool_calls" not in response.message.model_dump() assert response.tools_use == ["get_current_time"] 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"])["timezone"] == "UTC" 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", "get_current_time", '{"timezone":"UTC"}', ), tool_call( "call-2", "get_current_time", '{"timezone":"Asia/Shanghai"}', ), ], ), 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 == ["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", "get_current_time", "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", "get_current_time", '{"timezone":"UTC"}')], ), ( 5, 1, [ tool_call("call-1", "get_current_time", '{"timezone":"UTC"}'), tool_call( "call-2", "get_current_time", '{"timezone":"Asia/Shanghai"}', ), ], ), ], ) 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())