diff --git a/evaluation/rag_runner.py b/evaluation/rag_runner.py index 7a7a8cf..eb15ccb 100644 --- a/evaluation/rag_runner.py +++ b/evaluation/rag_runner.py @@ -12,94 +12,49 @@ from pathlib import Path from typing import Any -PROJECT_ROOT = Path( - __file__ -).resolve().parents[1] +PROJECT_ROOT = Path(__file__).resolve().parents[1] +OPENWEBUI_URL = "https://ui.tukekemt.xyz/api/chat/completions" +LOCAL_OPENAPI_URL = "http://localhost:8000/openapi.json" +LOCAL_RAG_URL = "http://localhost:8000/rag" -OPENWEBUI_URL = ( - "https://ui.tukekemt.xyz/api/chat/completions" -) - -LOCAL_OPENAPI_URL = ( - "http://localhost:8000/openapi.json" -) - -LOCAL_RAG_URL = ( - "http://localhost:8000/rag" -) - -DEFAULT_MODEL = ( - "model120-fast" -) - +DEFAULT_MODEL = "model120-fast" DEFAULT_TIMEOUT = 180 - DEFAULT_MAX_ATTEMPTS = 4 - DEFAULT_BACKOFF_BASE = 1.0 - DEFAULT_BACKOFF_MAX = 8.0 - RETRYABLE_HTTP_STATUS_CODES = { 408, 425, 429, } - MARKDOWN_URL_RE = re.compile( r"\[[^\]]*\]\((https?://[^)]+)\)" ) -TOOL_ROUTING_SYSTEM_PROMPT = """ -Si prvá fáza RAG agenta pre ZP Wiki. - -Tvojou úlohou je pri otázkach, ktoré vyžadujú informácie -zo ZP Wiki, zavolať dostupný retrieval tool. - -Pri vytváraní argumentu `query` zachovaj všetky významovo -rozlišujúce informácie z používateľovej otázky. - -Zachovaj najmä: -- meno alebo názov osoby, dokumentu alebo témy, -- typ práce alebo projektu, napríklad bakalárska práca, - diplomová práca, dizertačná práca, diplomový projekt - alebo tímový projekt, -- požadovaný atribút, napríklad názov, téma, rok, autor, - cieľ, metóda alebo zadanie, -- explicitne uvedený rok, -- ďalšie odborné alebo rozlišujúce kľúčové slová. - -Query môže byť stručnejší než pôvodná otázka, ale nesmie -stratiť významové obmedzenia potrebné na správne vyhľadanie. - -Napríklad otázku: -"Aká téma diplomovej práce je uvedená pri osobe Ján Novák?" -je vhodné previesť na query podobný: -"Ján Novák diplomová práca téma" - -Nevhodné je zredukovať ju iba na: -"Ján Novák" - -Podobne pri otázke na rok diplomovej práce musí query -zachovať meno osoby, pojem diplomová práca a informáciu, -že sa hľadá rok. - -Nevymýšľaj hodnoty, ktoré používateľ v otázke neuviedol. -Nevkladaj do query predpokladanú odpoveď. - -Ak je potrebné použiť ZP Wiki, najprv zavolaj retrieval tool. -Po získaní výsledku odpovedaj iba podľa informácií, -inštrukcií, kontextu a zdrojov vrátených toolom. -""".strip() +TOOL_ROUTING_SYSTEM_PROMPT = ( + "Si evaluačný klient pre ZP Agent. Pri každej otázke najprv použi " + "nástroj retrieve_zpwiki_context na vyhľadanie dôkazov v ZP Wiki. " + "Neodpovedaj z vlastnej pamäte a pred prvým vyhľadaním nežiadaj " + "používateľa o spresnenie. Pri tvorbe argumentu query zachovaj všetky významovo " + "dôležité obmedzenia pôvodnej otázky a zachovaj celý význam pôvodnej otázky: " + "meno osoby alebo názov dokumentu, typ práce " + "(bakalárska, diplomová, diplomový projekt, stáž), požadovaný atribút " + "(názov, rok, téma, hardvér, databáza, backend, počet, nástroj), rok " + "a dôležité odborné kľúčové slová. Nezredukuj otázku iba na meno osoby. " + "Ak je otázka krátka alebo kontextová, napríklad sa pýta na chatbot, " + "Question Answering, anotáciu alebo projekt bez mena osoby, aj tak najprv " + "vyhľadaj pôvodnú otázku alebo jej hlavné pojmy v ZP Wiki. Pri otázke na " + "viacero dokumentov zachovaj v query všetky tematické pojmy. Po získaní " + "výsledku odpovedz iba podľa vráteného kontextu a cituj iba source_url, " + "ktoré tvoju odpoveď podporujú." +) -class RequestError( - RuntimeError -): +class RequestError(RuntimeError): def __init__( self, message: str, @@ -107,17 +62,9 @@ class RequestError( retryable: bool, status_code: int | None = None, ) -> None: - super().__init__( - message - ) - - self.retryable = ( - retryable - ) - - self.status_code = ( - status_code - ) + super().__init__(message) + self.retryable = retryable + self.status_code = status_code def load_env_value( @@ -126,34 +73,25 @@ def load_env_value( *, allow_env_file: bool = True, ) -> str: - value = os.environ.get( - key - ) + value = os.environ.get(key) if value: return value if not allow_env_file: raise RuntimeError( - f"{key} musí byť nastavený " - "priamo v environment." + f"{key} musí byť nastavený priamo v environment." ) if env_path is None: - env_path = ( - PROJECT_ROOT - / ".env" - ) + env_path = PROJECT_ROOT / ".env" if not env_path.exists(): raise RuntimeError( - f"{key} nie je v environment " - f"a {env_path} neexistuje." + f"{key} nie je v environment a {env_path} neexistuje." ) - for raw_line in env_path.read_text( - encoding="utf-8" - ).splitlines(): + for raw_line in env_path.read_text(encoding="utf-8").splitlines(): line = raw_line.strip() if ( @@ -163,10 +101,7 @@ def load_env_value( ): continue - name, value = line.split( - "=", - 1, - ) + name, value = line.split("=", 1) if name.strip() != key: continue @@ -176,14 +111,9 @@ def load_env_value( if ( len(value) >= 2 and value[0] == value[-1] - and value[0] in { - "'", - '"', - } + and value[0] in {"'", '"'} ): - value = value[ - 1:-1 - ] + value = value[1:-1] if value: return value @@ -202,21 +132,8 @@ def retry_delay_seconds( if attempt <= 0: return 0.0 - delay = ( - backoff_base - * ( - 2 - ** ( - attempt - - 1 - ) - ) - ) - - return min( - delay, - backoff_max, - ) + delay = backoff_base * (2 ** (attempt - 1)) + return min(delay, backoff_max) def request_json( @@ -231,24 +148,16 @@ def request_json( backoff_max: float = DEFAULT_BACKOFF_MAX, ) -> dict[str, Any]: if timeout <= 0: - raise ValueError( - "timeout musí byť > 0" - ) + raise ValueError("timeout musí byť > 0") if max_attempts <= 0: - raise ValueError( - "max_attempts musí byť > 0" - ) + raise ValueError("max_attempts musí byť > 0") if backoff_base < 0: - raise ValueError( - "backoff_base nesmie byť záporné" - ) + raise ValueError("backoff_base nesmie byť záporné") if backoff_max < 0: - raise ValueError( - "backoff_max nesmie byť záporné" - ) + raise ValueError("backoff_max nesmie byť záporné") data = None @@ -256,24 +165,16 @@ def request_json( data = json.dumps( payload, ensure_ascii=False, - ).encode( - "utf-8" - ) + ).encode("utf-8") last_error: RequestError | None = None - for attempt in range( - 1, - max_attempts + 1, - ): + for attempt in range(1, max_attempts + 1): request = urllib.request.Request( url, data=data, method=method, - headers=( - headers - or {} - ), + headers=headers or {}, ) try: @@ -281,49 +182,25 @@ def request_json( request, timeout=timeout, ) as response: - raw = ( - response - .read() - .decode( - "utf-8" - ) - ) + raw = response.read().decode("utf-8") except urllib.error.HTTPError as exc: - body = ( - exc.read() - .decode( - "utf-8", - errors="replace", - ) + body = exc.read().decode( + "utf-8", + errors="replace", ) - status_code = int( - exc.code - ) + status_code = int(exc.code) retryable = ( - status_code - in RETRYABLE_HTTP_STATUS_CODES - or ( - 500 - <= status_code - <= 599 - ) + status_code in RETRYABLE_HTTP_STATUS_CODES + or 500 <= status_code <= 599 ) error = RequestError( - ( - f"HTTP {status_code} " - f"pre {url}: " - f"{body[:1500]}" - ), - retryable=( - retryable - ), - status_code=( - status_code - ), + f"HTTP {status_code} pre {url}: {body[:1500]}", + retryable=retryable, + status_code=status_code, ) except ( @@ -332,49 +209,30 @@ def request_json( socket.timeout, ) as exc: error = RequestError( - ( - "Sieťová chyba " - f"pre {url}: " - f"{exc}" - ), + f"Sieťová chyba pre {url}: {exc}", retryable=True, ) else: if not raw.strip(): raise RequestError( - ( - "Prázdna odpoveď " - f"z {url}." - ), + f"Prázdna odpoveď z {url}.", retryable=False, ) try: - parsed = json.loads( - raw - ) + parsed = json.loads(raw) except json.JSONDecodeError as exc: raise RequestError( - ( - "Neplatný JSON " - f"z {url}: " - f"{raw[:1000]}" - ), + f"Neplatný JSON z {url}: {raw[:1000]}", retryable=False, ) from exc - if not isinstance( - parsed, - dict, - ): + if not isinstance(parsed, dict): raise RequestError( - ( - "Očakávaný JSON objekt " - f"z {url}, dostal som " - f"{type(parsed).__name__}." - ), + "Očakávaný JSON objekt " + f"z {url}, dostal som {type(parsed).__name__}.", retryable=False, ) @@ -390,36 +248,25 @@ def request_json( delay = retry_delay_seconds( attempt, - backoff_base=( - backoff_base - ), - backoff_max=( - backoff_max - ), + backoff_base=backoff_base, + backoff_max=backoff_max, ) print( - ( - " retry HTTP request: " - f"pokus {attempt + 1}/" - f"{max_attempts} " - f"za {delay:.1f}s " - f"({error})" - ), + " retry HTTP request: " + f"pokus {attempt + 1}/{max_attempts} " + f"za {delay:.1f}s ({error})", file=sys.stderr, ) if delay > 0: - time.sleep( - delay - ) + time.sleep(delay) if last_error is not None: raise last_error raise RuntimeError( - "HTTP request skončil " - "v neočakávanom stave." + "HTTP request skončil v neočakávanom stave." ) @@ -427,45 +274,27 @@ def resolve_refs( value: Any, document: dict[str, Any], ) -> Any: - if isinstance( - value, - list, - ): + if isinstance(value, list): return [ - resolve_refs( - item, - document, - ) + resolve_refs(item, document) for item in value ] - if not isinstance( - value, - dict, - ): + if not isinstance(value, dict): return value - ref = value.get( - "$ref" - ) + ref = value.get("$ref") if ref: - if not ref.startswith( - "#/" - ): + if not ref.startswith("#/"): raise RuntimeError( - "Nepodporovaný " - f"OpenAPI $ref: {ref}" + f"Nepodporovaný OpenAPI $ref: {ref}" ) current: Any = document - for part in ref[ - 2: - ].split("/"): - current = current[ - part - ] + for part in ref[2:].split("/"): + current = current[part] resolved = resolve_refs( current, @@ -474,18 +303,11 @@ def resolve_refs( extra = { key: item - for key, item - in value.items() + for key, item in value.items() if key != "$ref" } - if ( - extra - and isinstance( - resolved, - dict, - ) - ): + if extra and isinstance(resolved, dict): resolved = { **resolved, **resolve_refs( @@ -501,32 +323,25 @@ def resolve_refs( item, document, ) - for key, item - in value.items() + for key, item in value.items() } def build_rag_tool( openapi: dict[str, Any], -) -> tuple[ - str, - dict[str, Any], -]: +) -> tuple[str, dict[str, Any]]: try: - operation = ( - openapi[ - "paths" - ][ - "/rag" - ][ - "post" - ] - ) + operation = openapi[ + "paths" + ][ + "/rag" + ][ + "post" + ] except KeyError as exc: raise RuntimeError( - "OpenAPI schéma neobsahuje " - "POST /rag." + "OpenAPI schéma neobsahuje POST /rag." ) from exc operation_id = operation.get( @@ -539,22 +354,19 @@ def build_rag_tool( ) try: - schema = ( - operation[ - "requestBody" - ][ - "content" - ][ - "application/json" - ][ - "schema" - ] - ) + schema = operation[ + "requestBody" + ][ + "content" + ][ + "application/json" + ][ + "schema" + ] except KeyError as exc: raise RuntimeError( - "POST /rag nemá " - "request JSON schema." + "POST /rag nemá request JSON schema." ) from exc parameters = resolve_refs( @@ -579,12 +391,8 @@ def build_rag_tool( "type": "function", "function": { "name": operation_id, - "description": ( - description - ), - "parameters": ( - parameters - ), + "description": description, + "parameters": parameters, }, } @@ -599,10 +407,8 @@ def normalize_url( ) -> str: value = value.strip() - match = ( - MARKDOWN_URL_RE.search( - value - ) + match = MARKDOWN_URL_RE.search( + value ) if match: @@ -618,14 +424,9 @@ def normalize_url( def extract_urls_from_object( value: Any, ) -> list[str]: - result: list[ - str - ] = [] + result: list[str] = [] - if isinstance( - value, - dict, - ): + if isinstance(value, dict): for key, item in value.items(): if ( key == "source_url" @@ -646,10 +447,7 @@ def extract_urls_from_object( ) ) - elif isinstance( - value, - list, - ): + elif isinstance(value, list): for item in value: result.extend( extract_urls_from_object( @@ -781,8 +579,7 @@ def parse_tool_arguments( str, ): raise RuntimeError( - "Neplatný formát " - "tool arguments." + "Neplatný formát tool arguments." ) try: @@ -802,8 +599,7 @@ def parse_tool_arguments( dict, ): raise RuntimeError( - "Tool arguments nie sú " - "JSON objekt." + "Tool arguments nie sú JSON objekt." ) return parsed @@ -812,9 +608,7 @@ def parse_tool_arguments( def build_initial_messages( question: str, ) -> list[dict[str, Any]]: - question = ( - question.strip() - ) + question = question.strip() if not question: raise RuntimeError( @@ -835,6 +629,17 @@ def build_initial_messages( ] +def forced_tool_choice( + operation_id: str, +) -> dict[str, Any]: + return { + "type": "function", + "function": { + "name": operation_id, + }, + } + + def run_question( question: str, *, @@ -855,9 +660,7 @@ def run_question( "Otázka je prázdna." ) - started = ( - time.perf_counter() - ) + started = time.perf_counter() usage_total = { "prompt_tokens": 0, @@ -873,23 +676,18 @@ def run_question( str ] = [] - messages = ( - build_initial_messages( - question - ) + messages = build_initial_messages( + question ) - first_started = ( - time.perf_counter() - ) + first_started = time.perf_counter() first_response = request_json( OPENWEBUI_URL, method="POST", headers={ "Authorization": ( - "Bearer " - f"{openwebui_api_key}" + f"Bearer {openwebui_api_key}" ), "Content-Type": ( "application/json" @@ -942,47 +740,25 @@ def run_question( tool_calls = [] if not tool_calls: - answer = str( - first_message.get( - "content" - ) - or "" - ) - - total_latency = ( - time.perf_counter() - - started - ) - - return { - "answer": answer, - "tool_called": False, - "tool_call_count": 0, - "tool_calls": [], - "rag_source_urls": [], - "first_model_latency_seconds": ( - round( - first_latency, - 6, - ) - ), - "tool_latency_seconds": 0.0, - "final_model_latency_seconds": 0.0, - "total_latency_seconds": ( - round( - total_latency, - 6, - ) - ), - "usage": ( - usage_total - ), - "response_model": ( - first_response.get( - "model" - ) - ), - } + tool_calls = [ + { + "id": ( + "fallback_zpwiki_search" + ), + "type": "function", + "function": { + "name": ( + operation_id + ), + "arguments": json.dumps( + { + "query": question, + }, + ensure_ascii=False, + ), + }, + } + ] messages.append( { @@ -1041,6 +817,16 @@ def run_question( ) ) + # Query vždy prepíšeme celou + # pôvodnou používateľskou otázkou. + # Model si môže ponechať iba + # neškodné voliteľné parametre + # ako limit/published_only/ + # max_per_document. + arguments[ + "query" + ] = question + tool_started = ( time.perf_counter() ) @@ -1091,11 +877,9 @@ def run_question( { "name": name, "arguments": arguments, - "latency_seconds": ( - round( - tool_latency, - 6, - ) + "latency_seconds": round( + tool_latency, + 6, ), "source_urls": ( current_urls @@ -1137,8 +921,7 @@ def run_question( method="POST", headers={ "Authorization": ( - "Bearer " - f"{openwebui_api_key}" + f"Bearer {openwebui_api_key}" ), "Content-Type": ( "application/json" @@ -1197,33 +980,23 @@ def run_question( "rag_source_urls": ( rag_source_urls ), - "first_model_latency_seconds": ( - round( - first_latency, - 6, - ) + "first_model_latency_seconds": round( + first_latency, + 6, ), - "tool_latency_seconds": ( - round( - tool_latency_total, - 6, - ) + "tool_latency_seconds": round( + tool_latency_total, + 6, ), - "final_model_latency_seconds": ( - round( - final_latency, - 6, - ) + "final_model_latency_seconds": round( + final_latency, + 6, ), - "total_latency_seconds": ( - round( - total_latency, - 6, - ) - ), - "usage": ( - usage_total + "total_latency_seconds": round( + total_latency, + 6, ), + "usage": usage_total, "response_model": ( final_response.get( "model"