# tests/test_ai_chat.py — Unit tests for provider-agnostic tool calling (Phase B) """Tests for backend.ai_chat: payload building, parsing, and tools fallback.""" import httpx import pytest from backend import ai_chat from backend.ai_chat import ToolCall, _gemini_tools, _parse_arguments, chat_completion OPENAI_CFG = { "name": "deepseek", "api_key": "sk-test-key", "base_url": "https://api.deepseek.com/v1", "model": "deepseek-chat", "auth_header": "Bearer {api_key}", } GEMINI_CFG = { "name": "gemini", "api_key": "gem-test", "base_url": "https://generativelanguage.googleapis.com/v1beta", "model": "gemini-2.0-flash", "auth_header": None, } TOOLS = [{ "type": "function", "function": { "name": "read_file", "description": "Read a file", "parameters": {"type": "object", "properties": {"vault": {"type": "string"}}, "required": ["vault"]}, }, }] class _Capture: """Fake ``_post_json`` capturing calls and returning scripted responses.""" def __init__(self, responses): self.responses = list(responses) self.calls = [] async def __call__(self, url, headers, payload): self.calls.append({"url": url, "headers": headers, "payload": payload}) result = self.responses.pop(0) if isinstance(result, Exception): raise result return result def _http_error(status: int) -> httpx.HTTPStatusError: request = httpx.Request("POST", "https://example.test") response = httpx.Response(status, request=request) return httpx.HTTPStatusError("error", request=request, response=response) # ═══════════════════════════════════════════════════════════════════ # Argument parsing # ═══════════════════════════════════════════════════════════════════ class TestParseArguments: def test_dict_passthrough(self): assert _parse_arguments({"a": 1}) == {"a": 1} def test_json_string(self): assert _parse_arguments('{"a": 1}') == {"a": 1} def test_invalid_string(self): assert _parse_arguments("not json") == {"_raw": "not json"} def test_empty_string(self): assert _parse_arguments("") == {} def test_none(self): assert _parse_arguments(None) == {} # ═══════════════════════════════════════════════════════════════════ # OpenAI-compatible path # ═══════════════════════════════════════════════════════════════════ class TestOpenAIChat: @pytest.mark.asyncio async def test_content_only(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) cap = _Capture([{"choices": [{"message": {"content": "bonjour"}}]}]) monkeypatch.setattr(ai_chat, "_post_json", cap) result = await chat_completion([{"role": "user", "content": "hi"}]) assert result.content == "bonjour" assert result.tool_calls == [] assert result.provider == "deepseek" @pytest.mark.asyncio async def test_tool_calls_parsed(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) cap = _Capture([{ "choices": [{ "message": { "content": None, "tool_calls": [{ "id": "call_1", "type": "function", "function": {"name": "read_file", "arguments": '{"vault": "V", "path": "a.md"}'}, }], }, }], }]) monkeypatch.setattr(ai_chat, "_post_json", cap) result = await chat_completion([{"role": "user", "content": "hi"}]) assert result.content is None assert result.has_tool_calls assert result.tool_calls[0] == ToolCall(id="call_1", name="read_file", arguments={"vault": "V", "path": "a.md"}) @pytest.mark.asyncio async def test_tools_included_in_payload(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) cap = _Capture([{"choices": [{"message": {"content": "ok"}}]}]) monkeypatch.setattr(ai_chat, "_post_json", cap) await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS) payload = cap.calls[0]["payload"] assert payload["tools"] == TOOLS assert payload["tool_choice"] == "auto" @pytest.mark.asyncio async def test_fallback_when_tools_rejected(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) cap = _Capture([ _http_error(400), {"choices": [{"message": {"content": "sans outils"}}]}, ]) monkeypatch.setattr(ai_chat, "_post_json", cap) result = await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS) assert result.content == "sans outils" assert "tools" in cap.calls[0]["payload"] assert "tools" not in cap.calls[1]["payload"] @pytest.mark.asyncio async def test_http_error_without_tools_propagates(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) cap = _Capture([_http_error(500)]) monkeypatch.setattr(ai_chat, "_post_json", cap) with pytest.raises(httpx.HTTPStatusError): await chat_completion([{"role": "user", "content": "hi"}]) # ═══════════════════════════════════════════════════════════════════ # Gemini path # ═══════════════════════════════════════════════════════════════════ class TestGeminiChat: @pytest.mark.asyncio async def test_content_and_system_instruction(self, monkeypatch): monkeypatch.setitem(ai_chat.PROVIDERS, "gemini", GEMINI_CFG) monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: GEMINI_CFG) cap = _Capture([{"candidates": [{"content": {"parts": [{"text": "salut"}]}}]}]) monkeypatch.setattr(ai_chat, "_post_json", cap) result = await chat_completion([ {"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}, ]) assert result.content == "salut" assert result.provider == "gemini" payload = cap.calls[0]["payload"] assert payload["system_instruction"]["parts"][0]["text"] == "sys" @pytest.mark.asyncio async def test_function_call_parsed(self, monkeypatch): monkeypatch.setitem(ai_chat.PROVIDERS, "gemini", GEMINI_CFG) monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: GEMINI_CFG) cap = _Capture([{ "candidates": [{"content": {"parts": [ {"functionCall": {"name": "read_file", "args": {"vault": "V", "path": "a.md"}}}, ]}}], }]) monkeypatch.setattr(ai_chat, "_post_json", cap) result = await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS) assert result.tool_calls[0].name == "read_file" assert result.tool_calls[0].arguments == {"vault": "V", "path": "a.md"} # Tools converted to Gemini declarations decls = cap.calls[0]["payload"]["tools"][0]["functionDeclarations"] assert decls[0]["name"] == "read_file" class TestGeminiToolsConversion: def test_none(self): assert _gemini_tools(None) is None def test_conversion(self): converted = _gemini_tools(TOOLS) assert converted[0]["functionDeclarations"][0]["name"] == "read_file" assert converted[0]["functionDeclarations"][0]["description"] == "Read a file" # ═══════════════════════════════════════════════════════════════════ # Streaming (B4) # ═══════════════════════════════════════════════════════════════════ class _StreamResponse: """Fake httpx streaming response yielding pre-scripted SSE lines.""" def __init__(self, lines, status=200): self.lines = lines self.status_code = status def raise_for_status(self): if self.status_code >= 400: raise _http_error(self.status_code) async def aiter_lines(self): for line in self.lines: yield line class _FakeAsyncClient: """Minimal stand-in for ``httpx.AsyncClient`` used in stream mode.""" captured: dict = {} response: _StreamResponse | None = None def __init__(self, *args, **kwargs): pass async def __aenter__(self): return self async def __aexit__(self, *exc): return False def stream(self, method, url, headers=None, json=None): _FakeAsyncClient.captured = {"method": method, "url": url, "headers": headers, "json": json} resp = _FakeAsyncClient.response class _Ctx: async def __aenter__(self_inner): return resp async def __aexit__(self_inner, *exc): return False return _Ctx() class TestStreamCompletion: @pytest.mark.asyncio async def test_openai_stream_yields_deltas(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) monkeypatch.setattr(ai_chat.httpx, "AsyncClient", _FakeAsyncClient) _FakeAsyncClient.response = _StreamResponse([ 'data: {"choices":[{"delta":{"content":"Bon"}}]}', "", 'data: {"choices":[{"delta":{"content":"jour"}}]}', "data: not-json", "data: [DONE]", ]) chunks = [c async for c in ai_chat.stream_completion([{"role": "user", "content": "hi"}])] assert "".join(chunks) == "Bonjour" assert _FakeAsyncClient.captured["json"]["stream"] is True @pytest.mark.asyncio async def test_gemini_stream_yields_text(self, monkeypatch): monkeypatch.setitem(ai_chat.PROVIDERS, "gemini", GEMINI_CFG) monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: GEMINI_CFG) monkeypatch.setattr(ai_chat.httpx, "AsyncClient", _FakeAsyncClient) _FakeAsyncClient.response = _StreamResponse([ 'data: {"candidates":[{"content":{"parts":[{"text":"sa"}]}}]}', 'data: {"candidates":[{"content":{"parts":[{"text":"lut"}]}}]}', ]) chunks = [c async for c in ai_chat.stream_completion([{"role": "user", "content": "hi"}])] assert "".join(chunks) == "salut" assert "streamGenerateContent" in _FakeAsyncClient.captured["url"] @pytest.mark.asyncio async def test_stream_raises_on_http_error(self, monkeypatch): monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG) monkeypatch.setattr(ai_chat.httpx, "AsyncClient", _FakeAsyncClient) _FakeAsyncClient.response = _StreamResponse([], status=500) with pytest.raises(httpx.HTTPStatusError): async for _ in ai_chat.stream_completion([{"role": "user", "content": "hi"}]): pass