diff --git a/scripts/rag_query_evidence.py b/scripts/rag_query_evidence.py new file mode 100644 index 0000000..fc2a3a6 --- /dev/null +++ b/scripts/rag_query_evidence.py @@ -0,0 +1,1198 @@ +from __future__ import annotations + +import json +import re +import sqlite3 +from pathlib import Path +from typing import Any + +from scripts.search_utils import normalize_for_compare + + +WORD_RE = re.compile(r"[^\W_]+", re.UNICODE) +QUOTE_RE = re.compile(r'[„“”\"\']([^„“”\"\']{4,})[„“”\"\']') + +QUERY_STOPWORDS = { + "a", "aj", "aky", "aka", "ake", "akym", "ako", "alebo", "bol", "bola", "bolo", + "by", "co", "dalsi", "dalsieho", "dalsim", "do", "je", "k", "kazdy", + "kazdeho", "kolko", "kto", "ktory", "ktora", "ktore", "mal", + "mala", "ma", "na", "nad", "najdi", "nazov", "nazvom", "o", "od", + "osoba", "osobe", "osobou", "po", "pod", "podla", "pre", "pri", "praca", + "prace", "pracu", "projekt", "projektu", "roku", "sa", "s", "so", + "student", "studenta", "studentovi", "tema", "temou", "typ", "typom", "typu", + "viacero", "dokumentov", "dokumentoch", "riesi", "riesia", "ktore", "ktory", + "uvedena", "uvedeny", "uvedene", "v", "vo", "vznikne", "vzniknut", "z", "za", "zo", +} + +MIN_EVIDENCE_SCORE = 24.0 +DEFAULT_TOP_EVIDENCE_CHUNKS = 4 +FOCUS_LINE_RADIUS = 7 +MAX_FOCUS_CHARACTERS = 2200 + +CANONICAL_TOKEN_ALIASES = { + "anotacia": "annotation", + "anotacii": "annotation", + "anotacny": "annotation", + "annotation": "annotation", + + "ciel": "goal", + "ciele": "goal", + "cielov": "goal", + "goal": "goal", + "goals": "goal", + + "databaza": "database", + "databazu": "database", + "database": "database", + + "backend": "backend", + + "odsek": "paragraph", + "odseku": "paragraph", + "odseky": "paragraph", + "paragraph": "paragraph", + "paragraphs": "paragraph", + + "otazka": "question", + "otazky": "question", + "otazok": "question", + "question": "question", + "questions": "question", + + "stav": "status", + "stavu": "status", + "status": "status", + + "nazov": "title", + "nazvu": "title", + "title": "title", + + "tema": "topic", + "temy": "topic", + "theme": "topic", + + "dp": "diploma", + "diplomova": "diploma", + "diplomovej": "diploma", + "diplomovu": "diploma", + "diploma": "diploma", + + "navrh": "proposal", + "navrhu": "proposal", +} + + +def sqlite_table_exists( + conn: sqlite3.Connection, + table_name: str, +) -> bool: + row = conn.execute( + "SELECT 1 FROM sqlite_master " + "WHERE type IN ('table', 'view') AND name = ? LIMIT 1", + (table_name,), + ).fetchone() + + return row is not None + + +def sqlite_table_columns( + conn: sqlite3.Connection, + table_name: str, +) -> set[str]: + if not sqlite_table_exists( + conn, + table_name, + ): + return set() + + rows = conn.execute( + f"PRAGMA table_info({table_name})" + ).fetchall() + + return { + str(row[1]) + for row in rows + if len(row) > 1 + } + + +def parse_heading_paths_json( + value: Any, +) -> list[Any]: + if isinstance(value, list): + return value + + if not value: + return [] + + try: + parsed = json.loads( + str(value) + ) + except ( + json.JSONDecodeError, + TypeError, + ValueError, + ): + return [] + + if not isinstance(parsed, list): + return [] + + return parsed + + +def flatten_heading_paths( + value: Any, +) -> str: + paths = parse_heading_paths_json( + value + ) + + parts: list[str] = [] + + for item in paths: + if isinstance(item, str): + if item.strip(): + parts.append( + item.strip() + ) + + elif isinstance( + item, + (list, tuple), + ): + parts.extend( + str(part).strip() + for part in item + if str(part).strip() + ) + + else: + text = str(item).strip() + + if text: + parts.append(text) + + return " > ".join(parts) + + +def canonical_token( + token: str, +) -> str: + return CANONICAL_TOKEN_ALIASES.get( + token, + token, + ) + + +def normalized_tokens( + value: str, +) -> list[str]: + return [ + canonical_token(token) + for token in WORD_RE.findall( + normalize_for_compare( + value + ) + ) + ] + + +def query_content_tokens( + query: str, +) -> list[str]: + tokens = [ + token + for token in normalized_tokens( + query + ) + if ( + token not in QUERY_STOPWORDS + and ( + token.isdigit() + or len(token) >= 2 + ) + ) + ] + + return list( + dict.fromkeys(tokens) + ) + + +def quoted_query_phrases( + query: str, +) -> list[str]: + phrases = [ + normalize_for_compare( + match.group(1) + ) + for match in QUOTE_RE.finditer( + query + ) + if len( + normalize_for_compare( + match.group(1) + ).split() + ) + >= 2 + ] + + return list( + dict.fromkeys(phrases) + ) + + +def token_matches( + left: str, + right: str, +) -> bool: + if left == right: + return True + + if ( + left.isdigit() + or right.isdigit() + ): + return False + + shorter = min( + len(left), + len(right), + ) + + if shorter <= 2: + return False + + required = ( + 3 + if shorter <= 4 + else 4 + if shorter <= 6 + else 5 + ) + + return ( + left[:required] + == right[:required] + ) + + +def count_token_matches( + query_tokens: list[str], + value: str, +) -> int: + value_tokens = normalized_tokens( + value + ) + + return sum( + 1 + for query_token in query_tokens + if any( + token_matches( + query_token, + value_token, + ) + for value_token in value_tokens + ) + ) + + +def coverage_ratio( + query_tokens: list[str], + value: str, +) -> float: + if not query_tokens: + return 0.0 + + return ( + count_token_matches( + query_tokens, + value, + ) + / len(query_tokens) + ) + + +def structural_evidence_bonus( + query: str, + text: str, + heading_paths: Any = None, +) -> float: + query_tokens = set( + normalized_tokens(query) + ) + + text_tokens = set( + normalized_tokens(text) + ) + + heading_tokens = set( + normalized_tokens( + flatten_heading_paths( + heading_paths + ) + ) + ) + + all_tokens = ( + text_tokens + | heading_tokens + ) + + normalized_text = ( + normalize_for_compare( + text + ) + ) + + bonus = 0.0 + + # Otázky typu: + # "Aký je názov diplomovej práce..." + # "Aká téma diplomovej práce..." + # + # Explicitný "Návrh na názov DP" musí poraziť + # chunk, ktorý obsahuje iba heading "Diplomová práca 2021". + if ( + ( + "title" in query_tokens + or "topic" in query_tokens + ) + and "diploma" in query_tokens + ): + if { + "proposal", + "title", + "diploma", + } <= all_tokens: + bonus += 110.0 + + elif "title" in all_tokens: + bonus += 60.0 + + # Slovenská otázka "aké ciele" musí spoľahlivo + # nájsť aj anglický blok "Goals". + if ( + "goal" in query_tokens + and "goal" in all_tokens + ): + bonus += 100.0 + + # Slovenské "koľko otázok na odsek" + # musí vedieť nájsť anglické: + # "Output: 5 questions for each paragraph". + if ( + "question" in query_tokens + and "paragraph" in query_tokens + ): + if ( + "question" in all_tokens + and "paragraph" in all_tokens + ): + bonus += 100.0 + + if re.search( + r"(? float: + normalized_text = ( + normalize_for_compare( + text + ) + ) + + normalized_heading = ( + normalize_for_compare( + flatten_heading_paths( + heading_paths + ) + ) + ) + + query_tokens = ( + query_content_tokens( + query + ) + ) + + score = 0.0 + + for phrase in quoted_query_phrases( + query + ): + if phrase in normalized_text: + score += 120.0 + + elif phrase in normalized_heading: + score += 100.0 + + else: + phrase_tokens = ( + normalized_tokens( + phrase + ) + ) + + score += ( + 55.0 + * coverage_ratio( + phrase_tokens, + text, + ) + ) + + if query_tokens: + score += ( + 7.0 + * count_token_matches( + query_tokens, + text, + ) + ) + + score += ( + 5.0 + * count_token_matches( + query_tokens, + normalized_heading, + ) + ) + + score += ( + 20.0 + * coverage_ratio( + query_tokens, + text, + ) + ) + + normalized_query = ( + normalize_for_compare( + query + ) + ) + + if ( + normalized_query + and normalized_query + in normalized_text + ): + score += 30.0 + + score += ( + structural_evidence_bonus( + query, + text, + heading_paths, + ) + ) + + return round( + score, + 6, + ) + + +def _line_score( + query: str, + line: str, +) -> float: + normalized_line = ( + normalize_for_compare( + line + ) + ) + + score = ( + 5.0 + * count_token_matches( + query_content_tokens( + query + ), + line, + ) + ) + + for phrase in quoted_query_phrases( + query + ): + if phrase in normalized_line: + score += 100.0 + + return score + + +def build_focus_excerpt( + query: str, + text: str, + *, + radius: int = FOCUS_LINE_RADIUS, + max_characters: int = MAX_FOCUS_CHARACTERS, +) -> str: + clean_text = text.strip() + + if not clean_text: + return "" + + lines = clean_text.splitlines() + + scored = [ + ( + _line_score( + query, + line, + ), + index, + ) + for index, line + in enumerate(lines) + ] + + best_score, best_index = max( + scored, + key=lambda item: ( + item[0], + -item[1], + ), + ) + + if best_score <= 0: + return "" + + start = max( + 0, + best_index - radius, + ) + + end = min( + len(lines), + best_index + radius + 1, + ) + + excerpt = "\n".join( + lines[start:end] + ).strip() + + if ( + len(excerpt) + <= max_characters + ): + return excerpt + + return excerpt[ + :max_characters + ].rstrip() + + +def query_evidence_metadata( + *, + applied: bool = False, + document_path: str | None = None, + primary_chunk_id: str | None = None, + evidence_chunk_id: str | None = None, + evidence_chunk_index: int | None = None, + score: float | None = None, + same_as_primary: bool = False, + evidence_chunks: list[ + dict[str, Any] + ] + | None = None, +) -> dict[str, Any]: + return { + "strategy": ( + "within_document_query_evidence" + ), + "applied": applied, + "document_path": document_path, + "primary_chunk_id": ( + primary_chunk_id + ), + "evidence_chunk_id": ( + evidence_chunk_id + ), + "evidence_chunk_index": ( + evidence_chunk_index + ), + "score": score, + "same_as_primary": ( + same_as_primary + ), + "evidence_chunks": ( + evidence_chunks + or [] + ), + } + + +def _load_document_chunks( + conn: sqlite3.Connection, + document_path: str, + *, + published_only: bool, +) -> list[dict[str, Any]]: + columns = sqlite_table_columns( + conn, + "chunks", + ) + + required_columns = { + "chunk_id", + "document_path", + "chunk_index", + "heading_paths_json", + "text", + } + + if not ( + required_columns + <= columns + ): + return [] + + where_parts = [ + "document_path = ?", + ] + + parameters: list[Any] = [ + document_path, + ] + + if ( + published_only + and "published" in columns + ): + where_parts.append( + "published = 1" + ) + + order_parts = [ + "chunk_index ASC", + ] + + if "id" in columns: + order_parts.append( + "id ASC" + ) + + sql = ( + "SELECT " + "chunk_id, " + "document_path, " + "chunk_index, " + "heading_paths_json, " + "text " + "FROM chunks " + "WHERE " + + " AND ".join( + where_parts + ) + + " ORDER BY " + + ", ".join( + order_parts + ) + ) + + rows = conn.execute( + sql, + parameters, + ).fetchall() + + result: list[ + dict[str, Any] + ] = [] + + for row in rows: + item = dict(row) + + item[ + "heading_paths" + ] = ( + parse_heading_paths_json( + item.get( + "heading_paths_json" + ) + ) + ) + + result.append(item) + + return result + + +def find_top_query_evidence_chunks( + conn: sqlite3.Connection, + document_path: str, + query: str, + *, + published_only: bool = False, + top_k: int = DEFAULT_TOP_EVIDENCE_CHUNKS, +) -> list[dict[str, Any]]: + if top_k <= 0: + return [] + + chunks = _load_document_chunks( + conn, + document_path, + published_only=published_only, + ) + + if not chunks: + return [] + + scored: list[ + dict[str, Any] + ] = [] + + for item in chunks: + score = evidence_score( + query, + text=str( + item.get("text") + or "" + ), + heading_paths=( + item.get( + "heading_paths" + ) + ), + ) + + focus_text = ( + build_focus_excerpt( + query, + str( + item.get("text") + or "" + ), + ) + ) + + candidate = dict(item) + + candidate[ + "evidence_score" + ] = score + + candidate[ + "focus_text" + ] = focus_text + + scored.append(candidate) + + scored.sort( + key=lambda item: ( + -float( + item.get( + "evidence_score" + ) + or 0.0 + ), + int( + item.get( + "chunk_index" + ) + or 0 + ), + ) + ) + + selected: list[ + dict[str, Any] + ] = [] + + seen_focus: set[str] = set() + + for item in scored: + score = float( + item.get( + "evidence_score" + ) + or 0.0 + ) + + focus_text = str( + item.get( + "focus_text" + ) + or "" + ).strip() + + if ( + score + < MIN_EVIDENCE_SCORE + or not focus_text + ): + continue + + normalized_focus = ( + normalize_for_compare( + focus_text + ) + ) + + if ( + normalized_focus + in seen_focus + ): + continue + + seen_focus.add( + normalized_focus + ) + + selected.append(item) + + if ( + len(selected) + >= top_k + ): + break + + return selected + + +def find_best_query_evidence_chunk( + conn: sqlite3.Connection, + document_path: str, + query: str, + *, + published_only: bool = False, +) -> dict[str, Any] | None: + chunks = _load_document_chunks( + conn, + document_path, + published_only=published_only, + ) + + if not chunks: + return None + + scored: list[ + dict[str, Any] + ] = [] + + for item in chunks: + candidate = dict(item) + + candidate[ + "evidence_score" + ] = evidence_score( + query, + text=str( + candidate.get("text") + or "" + ), + heading_paths=( + candidate.get( + "heading_paths" + ) + ), + ) + + candidate[ + "focus_text" + ] = build_focus_excerpt( + query, + str( + candidate.get("text") + or "" + ), + ) + + scored.append(candidate) + + scored.sort( + key=lambda item: ( + -float( + item.get( + "evidence_score" + ) + or 0.0 + ), + int( + item.get( + "chunk_index" + ) + or 0 + ), + ) + ) + + return scored[0] + + +def expand_results_with_query_evidence( + db_path: Path, + query: str, + results: list[ + dict[str, Any] + ], + *, + published_only: bool = False, + top_k: int = DEFAULT_TOP_EVIDENCE_CHUNKS, +) -> list[dict[str, Any]]: + base_results = [ + dict(item) + for item in results + ] + + if ( + not base_results + or not query.strip() + or not db_path.exists() + ): + return base_results + + with sqlite3.connect( + db_path, + timeout=5.0, + ) as conn: + conn.row_factory = ( + sqlite3.Row + ) + + conn.execute( + "PRAGMA query_only = ON" + ) + + if not sqlite_table_exists( + conn, + "chunks", + ): + return base_results + + expanded: list[ + dict[str, Any] + ] = [] + + for result in base_results: + item = dict(result) + + document_path = str( + item.get( + "document_path" + ) + or "" + ).strip() + + primary_chunk_id = ( + str( + item.get( + "chunk_id" + ) + or "" + ).strip() + or None + ) + + if not document_path: + item[ + "query_evidence" + ] = ( + query_evidence_metadata( + primary_chunk_id=( + primary_chunk_id + ) + ) + ) + + expanded.append(item) + continue + + top_chunks = ( + find_top_query_evidence_chunks( + conn, + document_path, + query, + published_only=( + published_only + ), + top_k=top_k, + ) + ) + + if not top_chunks: + item[ + "query_evidence" + ] = ( + query_evidence_metadata( + document_path=( + document_path + ), + primary_chunk_id=( + primary_chunk_id + ), + ) + ) + + expanded.append(item) + continue + + blocks: list[ + dict[str, Any] + ] = [] + + for candidate in top_chunks: + blocks.append( + { + "chunk_id": ( + candidate.get( + "chunk_id" + ) + ), + "chunk_index": int( + candidate.get( + "chunk_index" + ) + or 0 + ), + "heading_paths": ( + candidate.get( + "heading_paths" + ) + or [] + ), + "score": round( + float( + candidate.get( + "evidence_score" + ) + or 0.0 + ), + 6, + ), + "focus_text": str( + candidate.get( + "focus_text" + ) + or "" + ).strip(), + "text": str( + candidate.get( + "text" + ) + or "" + ).strip(), + } + ) + + best = blocks[0] + + best_chunk_id = ( + str( + best.get( + "chunk_id" + ) + or "" + ).strip() + or None + ) + + same_as_primary = ( + best_chunk_id + == primary_chunk_id + ) + + item[ + "query_evidence" + ] = ( + query_evidence_metadata( + applied=True, + document_path=( + document_path + ), + primary_chunk_id=( + primary_chunk_id + ), + evidence_chunk_id=( + best_chunk_id + ), + evidence_chunk_index=int( + best.get( + "chunk_index" + ) + or 0 + ), + score=float( + best.get( + "score" + ) + or 0.0 + ), + same_as_primary=( + same_as_primary + ), + evidence_chunks=( + blocks + ), + ) + ) + + item[ + "query_evidence_blocks" + ] = blocks + + # Spätná kompatibilita s E4 a build_source_text(). + item[ + "query_evidence_text" + ] = best[ + "text" + ] + + item[ + "query_focus_text" + ] = best[ + "focus_text" + ] + + item[ + "query_evidence_heading_paths" + ] = best[ + "heading_paths" + ] + + expanded.append(item) + + return expanded