pridanie testov tool promptu

This commit is contained in:
Ján Pták 2026-09-27 14:57:05 +02:00
parent cd18ef0fa6
commit 67d66a9f3e

View File

@ -0,0 +1,470 @@
from __future__ import annotations
from typing import Any
import pytest
import evaluation.rag_runner as rag_runner
def test_tool_routing_prompt_preserves_query_intent() -> None:
prompt = (
rag_runner
.TOOL_ROUTING_SYSTEM_PROMPT
)
assert (
"zachovaj všetky významovo"
in prompt
)
assert (
"meno osoby"
in prompt
)
assert (
"typ práce"
in prompt
)
assert (
"požadovaný atribút"
in prompt
)
assert (
"rok"
in prompt
)
assert (
"kľúčové slová"
in prompt
)
assert (
"Nezredukuj otázku iba na meno osoby"
in prompt
)
def test_build_initial_messages_contains_system_and_user() -> None:
question = (
"Aká téma diplomovej práce je uvedená "
"pri osobe Test Student?"
)
messages = (
rag_runner
.build_initial_messages(
question
)
)
assert (
len(
messages
)
== 2
)
assert (
messages[
0
][
"role"
]
== "system"
)
assert (
messages[
0
][
"content"
]
== (
rag_runner
.TOOL_ROUTING_SYSTEM_PROMPT
)
)
assert (
messages[
1
]
== {
"role": "user",
"content": question,
}
)
def test_build_initial_messages_rejects_empty_question() -> None:
with pytest.raises(
RuntimeError,
match="Otázka je prázdna",
):
(
rag_runner
.build_initial_messages(
" "
)
)
def test_run_question_sends_routing_prompt_and_falls_back_to_rag(
monkeypatch: pytest.MonkeyPatch,
) -> None:
captured_calls: list[
dict[str, Any]
] = []
def fake_request_json(
url: str,
*,
method: str = "GET",
headers: dict[
str,
str,
]
| None = None,
payload: dict[
str,
Any,
]
| None = None,
timeout: int = (
rag_runner
.DEFAULT_TIMEOUT
),
max_attempts: int = 1,
backoff_base: float = (
rag_runner
.DEFAULT_BACKOFF_BASE
),
backoff_max: float = (
rag_runner
.DEFAULT_BACKOFF_MAX
),
) -> dict[str, Any]:
assert (
payload
is not None
)
captured_calls.append(
{
"url": url,
"method": method,
"headers": headers,
"payload": payload,
"timeout": timeout,
"max_attempts": (
max_attempts
),
"backoff_base": (
backoff_base
),
"backoff_max": (
backoff_max
),
}
)
if (
len(
captured_calls
)
== 1
):
return {
"model": (
"model120-fast"
),
"usage": {
"prompt_tokens": 10,
"completion_tokens": 2,
"total_tokens": 12,
},
"choices": [
{
"message": {
"role": (
"assistant"
),
"content": (
"Bez tool callu."
),
},
},
],
}
if (
url
== (
rag_runner
.LOCAL_RAG_URL
)
):
assert (
payload[
"query"
]
== (
"Aká téma diplomovej práce "
"je uvedená pri osobe "
"Test Student?"
)
)
return {
"query": (
payload[
"query"
]
),
"sources": [
{
"source_url": (
"https://example.test/"
"test_student"
),
"text": (
"Názov diplomovej práce: "
"Testovacia téma."
),
},
],
}
return {
"model": (
"model120-fast"
),
"usage": {
"prompt_tokens": 20,
"completion_tokens": 5,
"total_tokens": 25,
},
"choices": [
{
"message": {
"role": (
"assistant"
),
"content": (
"Téma diplomovej práce "
"je Testovacia téma.\n\n"
"Zdroj: "
"https://example.test/"
"test_student"
),
},
},
],
}
monkeypatch.setattr(
rag_runner,
"request_json",
fake_request_json,
)
question = (
"Aká téma diplomovej práce "
"je uvedená pri osobe "
"Test Student?"
)
result = (
rag_runner
.run_question(
question,
model=(
"model120-fast"
),
operation_id=(
"retrieve_zpwiki_context"
),
rag_tool={
"type": "function",
"function": {
"name": (
"retrieve_zpwiki_context"
),
"description": (
"Test tool"
),
"parameters": {
"type": (
"object"
),
"properties": {
"query": {
"type": (
"string"
),
},
},
"required": [
"query",
],
},
},
},
openwebui_api_key=(
"openwebui-test-key"
),
search_api_key=(
"search-test-key"
),
timeout=30,
max_attempts=1,
backoff_base=0,
backoff_max=0,
)
)
# E5 fallback:
# 1. prvý model
# 2. lokálny /rag
# 3. finálny model
assert (
len(
captured_calls
)
== 3
)
first_payload = (
captured_calls[
0
][
"payload"
]
)
assert (
first_payload[
"model"
]
== "model120-fast"
)
assert (
first_payload[
"tool_choice"
]
== "auto"
)
assert (
first_payload[
"messages"
][
0
]
== {
"role": "system",
"content": (
rag_runner
.TOOL_ROUTING_SYSTEM_PROMPT
),
}
)
assert (
first_payload[
"messages"
][
1
]
== {
"role": "user",
"content": question,
}
)
assert (
captured_calls[
1
][
"url"
]
== (
rag_runner
.LOCAL_RAG_URL
)
)
assert (
captured_calls[
1
][
"payload"
][
"query"
]
== question
)
assert (
captured_calls[
2
][
"url"
]
== (
rag_runner
.OPENWEBUI_URL
)
)
assert (
result[
"tool_called"
]
is True
)
assert (
result[
"tool_call_count"
]
== 1
)
assert (
result[
"tool_calls"
][
0
][
"arguments"
][
"query"
]
== question
)
assert (
result[
"rag_source_urls"
]
== [
(
"https://example.test/"
"test_student"
),
]
)
assert (
"Testovacia téma"
in result[
"answer"
]
)