dp-zp-agent/test/test_rag_runner_routing.py

219 lines
6.3 KiB
Python

from __future__ import annotations
from typing import Any
import evaluation.rag_runner as runner
def tool() -> dict[str, Any]:
return {
"type": "function",
"function": {
"name": "retrieve_zpwiki_context",
"description": "RAG",
"parameters": {
"type": "object",
"properties": {
"query": {"type": "string"},
},
"required": ["query"],
},
},
}
def test_initial_messages_force_retrieval_intent() -> None:
messages = runner.build_initial_messages(
"S akým systémom sa má chatbot integrovať?"
)
assert messages[0]["role"] == "system"
assert "najprv použi" in messages[0]["content"]
assert "nežiadaj" in messages[0]["content"]
assert messages[1]["role"] == "user"
def test_forced_tool_choice_names_exact_operation() -> None:
choice = runner.forced_tool_choice(
"retrieve_zpwiki_context"
)
assert choice == {
"type": "function",
"function": {
"name": "retrieve_zpwiki_context",
},
}
def test_run_question_uses_tool_and_preserves_original_query(
monkeypatch,
) -> None:
calls: list[dict[str, Any]] = []
def fake_request_json(
url: str,
**kwargs: Any,
) -> dict[str, Any]:
calls.append(
{
"url": url,
**kwargs,
}
)
if len(calls) == 1:
return {
"model": "model120-fast",
"usage": {},
"choices": [
{
"message": {
"content": "",
"tool_calls": [
{
"id": "call-1",
"type": "function",
"function": {
"name": "retrieve_zpwiki_context",
"arguments": (
'{"query": "Ján Pták databáza backend"}'
),
},
}
],
}
}
],
}
if url == runner.LOCAL_RAG_URL:
return {
"query": "Ján Pták databáza backend",
"sources": [
{
"source_url": "https://example.test/jan_ptak",
"text": "SQLite a FastAPI",
}
],
}
return {
"model": "model120-fast",
"usage": {},
"choices": [
{
"message": {
"content": (
"Používal SQLite a FastAPI.\n\n"
"Zdroj: https://example.test/jan_ptak"
)
}
}
],
}
monkeypatch.setattr(
runner,
"request_json",
fake_request_json,
)
result = runner.run_question(
"Akú databázu a backend už Ján Pták podľa stavu používal?",
model="model120-fast",
operation_id="retrieve_zpwiki_context",
rag_tool=tool(),
openwebui_api_key="secret-openwebui",
search_api_key="secret-search",
timeout=30,
max_attempts=4,
backoff_base=0,
backoff_max=0,
)
first_payload = calls[0]["payload"]
assert first_payload["tool_choice"] == "auto"
assert first_payload["messages"][0]["role"] == "system"
assert calls[1]["payload"]["query"] == (
"Akú databázu a backend už Ján Pták podľa stavu používal?"
)
assert result["tool_called"] is True
assert result["rag_source_urls"] == ["https://example.test/jan_ptak"]
def test_run_question_falls_back_to_rag_when_model_skips_tool(monkeypatch) -> None:
calls = []
def fake_request_json(url, **kwargs):
calls.append((url, kwargs))
if len(calls) == 1:
return {
"model": "model120-fast",
"choices": [
{
"message": {
"role": "assistant",
"content": "Potrebujem viac kontextu.",
}
}
],
"usage": {},
}
if len(calls) == 2:
assert url == runner.LOCAL_RAG_URL
assert kwargs["payload"]["query"] == (
"Koľko otázok má vzniknúť pre každý odsek pri anotácii otázok?"
)
return {
"sources": [
{
"source_url": "https://zp.kemt.fei.tuke.sk/topics/question",
"text": "Pre každý odsek má vzniknúť 5 otázok.",
}
]
}
return {
"model": "model120-fast",
"choices": [
{
"message": {
"role": "assistant",
"content": (
"Pre každý odsek má vzniknúť 5 otázok.\n\n"
"Zdroj: https://zp.kemt.fei.tuke.sk/topics/question"
),
}
}
],
"usage": {},
}
monkeypatch.setattr(runner, "request_json", fake_request_json)
result = runner.run_question(
"Koľko otázok má vzniknúť pre každý odsek pri anotácii otázok?",
model="model120-fast",
operation_id="retrieve_zpwiki_context",
rag_tool={
"type": "function",
"function": {
"name": "retrieve_zpwiki_context",
"parameters": {"type": "object"},
},
},
openwebui_api_key="secret",
search_api_key="search",
timeout=10,
)
assert result["tool_called"] is True
assert result["tool_call_count"] == 1
assert result["tool_calls"][0]["arguments"]["query"] == (
"Koľko otázok má vzniknúť pre každý odsek pri anotácii otázok?"
)
assert len(calls) == 3