diff --git a/test/test_rag_runner_routing.py b/test/test_rag_runner_routing.py new file mode 100644 index 0000000..cfa2b98 --- /dev/null +++ b/test/test_rag_runner_routing.py @@ -0,0 +1,218 @@ +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