471 lines
9.2 KiB
Python
471 lines
9.2 KiB
Python
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"
|
|
]
|
|
)
|