CI / lint (push) Successful in 57s
CI / security (push) Successful in 39s
CI / test (push) Successful in 1m13s
CI / build (push) Successful in 36s
CI / e2e (push) Successful in 10m13s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
- backend/ai_chat.py: chat_completion provider-agnostique (OpenAI-compat tools/tool_calls + Gemini functionDeclarations/functionCall), retry sans tools si rejete - backend/agent/loop.py: run_agent multi-etapes (limite 10, truncation, confirmation two-step), LLM injectable - endpoint opt-in POST /api/ai/bookslm/agent (events SSE tool/message/confirmation), extraction _resolve_system_prompt - tests: test_ai_chat.py, test_agent_loop.py + 3 tests endpoint (728 passed au total) - ROADMAP B1/B2/B3/B7 livres ; B4/B5/B6 restants
202 lines
8.2 KiB
Python
202 lines
8.2 KiB
Python
# 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"
|