219 lines
6.3 KiB
Python
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
|