diff --git a/scripts/rag_supplemental_expansion.py b/scripts/rag_supplemental_expansion.py new file mode 100644 index 0000000..306ed84 --- /dev/null +++ b/scripts/rag_supplemental_expansion.py @@ -0,0 +1,821 @@ +from __future__ import annotations + +import json +import re +import sqlite3 +from pathlib import Path +from typing import Any + +from scripts.rag_query_evidence import ( + build_focus_excerpt, + evidence_score, + parse_heading_paths_json, + query_content_tokens, +) +from scripts.search_utils import make_source_url, normalize_for_compare + + +WORD_RE = re.compile(r"[^\W_]+", re.UNICODE) + +SUPPLEMENTAL_STOPWORDS = { + "autor", + "autora", + "dokument", + "dokumenty", + "informacna", + "informacnej", + "projektova", + "projektovej", + "stranka", + "stranku", + "student", + "studenta", + "tema", + "teme", + "temou", +} + +CANONICAL_TOKEN_ALIASES = { + "multilingualna": "multilingual", + "multilingualny": "multilingual", + "multilingual": "multilingual", + "viacjazycna": "multilingual", + "viacjazycneho": "multilingual", + "viacjazycny": "multilingual", + "extrakcia": "extract", + "extrakcie": "extract", + "extrahovanie": "extract", + "extrahovania": "extract", + "extraction": "extract", + "extract": "extract", + "trojic": "triplet", + "trojice": "triplet", + "trojicami": "triplet", + "triplet": "triplet", + "triplets": "triplet", + "medicinskych": "medical", + "medicinsky": "medical", + "medicinske": "medical", + "lekarskych": "medical", + "lekarsky": "medical", + "medical": "medical", + "dat": "data", + "data": "data", + "strojovy": "machine", + "strojoveho": "machine", + "machine": "machine", + "preklad": "translation", + "prekladu": "translation", + "translation": "translation", + "automaticky": "automatic", + "automatickeho": "automatic", + "automatic": "automatic", + "pomenovane": "named", + "pomenovanych": "named", + "named": "named", + "entity": "entity", + "entities": "entity", + "entit": "entity", + "anotacia": "annotation", + "anotacii": "annotation", + "annotation": "annotation", + "znalostny": "knowledge", + "znalostne": "knowledge", + "knowledge": "knowledge", + "graf": "graph", + "grafy": "graph", + "graph": "graph", + "graphy": "graph", + "otazok": "question", + "otazky": "question", + "question": "question", + "questions": "question", + "odpovedi": "answer", + "odpovede": "answer", + "answer": "answer", + "answers": "answer", + "system": "system", + "systemom": "system", + "systemu": "system", + "integrovat": "integrate", + "integracia": "integrate", + "integrate": "integrate", +} + +MIN_SINGLE_GLOBAL_COVERAGE = 0.50 +MIN_MULTI_GLOBAL_COVERAGE = 0.45 +MIN_GLOBAL_MATCHES = 3 +MAX_FAMILY_EVIDENCE = 3 + + +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 normalized_words(value: str) -> list[str]: + return WORD_RE.findall(normalize_for_compare(value)) + + +def canonical_token(token: str) -> str: + return CANONICAL_TOKEN_ALIASES.get(token, token) + + +def canonical_tokens(value: str) -> list[str]: + result: list[str] = [] + + for token in normalized_words(value): + if token in SUPPLEMENTAL_STOPWORDS: + continue + result.append(canonical_token(token)) + + return list(dict.fromkeys(result)) + + +def canonical_query_tokens(query: str) -> list[str]: + base = [ + token + for token in query_content_tokens(query) + if token not in SUPPLEMENTAL_STOPWORDS + ] + return list( + dict.fromkeys( + canonical_token(token) + for token in base + ) + ) + + +def canonical_coverage(query_tokens: list[str], value: str) -> tuple[int, float]: + if not query_tokens: + return 0, 0.0 + + value_tokens = set(canonical_tokens(value)) + matches = sum( + 1 + for token in query_tokens + if token in value_tokens + ) + return matches, matches / len(query_tokens) + + +def canonical_student_main_path(document_path: str) -> str: + parts = Path(document_path).parts + + if ( + len(parts) >= 5 + and parts[0] == "pages" + and parts[1] == "students" + ): + return "/".join( + [ + parts[0], + parts[1], + parts[2], + parts[3], + "README.md", + ] + ) + + return document_path + + +def same_student_family(left: str, right: str) -> bool: + return ( + canonical_student_main_path(left) + == canonical_student_main_path(right) + ) + + +def is_multi_document_query(query: str) -> bool: + normalized = normalize_for_compare(query) + return any( + marker in normalized + for marker in ( + "viacero dokument", + "dokumenty", + "ktore dokumenty", + "najdi dokumenty", + ) + ) + + +def _load_documents(conn: sqlite3.Connection) -> list[dict[str, Any]]: + columns = sqlite_table_columns(conn, "documents") + if not columns: + return [] + + path_column = ( + "path" + if "path" in columns + else "document_path" + if "document_path" in columns + else None + ) + if path_column is None: + return [] + + selected = [ + f"{path_column} AS document_path", + ] + + for optional in ("title", "author", "published"): + if optional in columns: + selected.append(optional) + + query = ( + "SELECT " + + ", ".join(selected) + + " FROM documents" + ) + return [ + dict(row) + for row in conn.execute(query).fetchall() + ] + + +def _load_chunks( + conn: sqlite3.Connection, + *, + published_only: bool, +) -> list[dict[str, Any]]: + columns = sqlite_table_columns(conn, "chunks") + required = { + "chunk_id", + "document_path", + "chunk_index", + "heading_paths_json", + "text", + } + if not required <= columns: + return [] + + selected = [ + "chunk_id", + "document_path", + "chunk_index", + "heading_paths_json", + "text", + ] + + for optional in ("title", "author", "published"): + if optional in columns: + selected.append(optional) + + where = "" + if published_only and "published" in columns: + where = " WHERE published = 1" + + order = " ORDER BY document_path ASC, chunk_index ASC" + if "id" in columns: + order += ", id ASC" + + query = ( + "SELECT " + + ", ".join(selected) + + " FROM chunks" + + where + + order + ) + + rows = [ + dict(row) + for row in conn.execute(query).fetchall() + ] + + for row in rows: + row["heading_paths"] = parse_heading_paths_json( + row.get("heading_paths_json") + ) + + return rows + + +def _document_metadata_by_path( + documents: list[dict[str, Any]], +) -> dict[str, dict[str, Any]]: + return { + str(item.get("document_path") or ""): item + for item in documents + if str(item.get("document_path") or "").strip() + } + + +def _best_main_metadata( + canonical_path: str, + metadata_by_path: dict[str, dict[str, Any]], + fallback_chunk: dict[str, Any], +) -> dict[str, Any]: + direct = metadata_by_path.get(canonical_path) + if direct is not None: + return direct + + evidence_path = str( + fallback_chunk.get("document_path") or "" + ) + evidence_meta = metadata_by_path.get(evidence_path) + if evidence_meta is not None: + return evidence_meta + + return { + "document_path": canonical_path, + "title": fallback_chunk.get("title"), + "author": fallback_chunk.get("author"), + "published": fallback_chunk.get("published"), + } + + +def _source_from_chunk( + canonical_path: str, + chunk: dict[str, Any], + metadata_by_path: dict[str, dict[str, Any]], + *, + reason: str, + score: float, +) -> dict[str, Any]: + meta = _best_main_metadata( + canonical_path, + metadata_by_path, + chunk, + ) + + return { + "chunk_id": chunk.get("chunk_id"), + "document_path": canonical_path, + "title": meta.get("title") or chunk.get("title"), + "author": meta.get("author") or chunk.get("author"), + "published": ( + meta.get("published") + if meta.get("published") is not None + else chunk.get("published") + ), + "chunk_index": int(chunk.get("chunk_index") or 0), + "heading_paths": chunk.get("heading_paths") or [], + "text": str(chunk.get("text") or ""), + "source_url": make_source_url(canonical_path), + "match_strategy": "rag_supplemental", + "fts_rank": None, + "vector_rank": None, + "vector_score": None, + "hybrid_score": None, + "supplemental_expansion": { + "strategy": "global_or_entity_anchor", + "applied": True, + "reason": reason, + "score": round(score, 6), + "evidence_document_path": chunk.get("document_path"), + "evidence_chunk_id": chunk.get("chunk_id"), + }, + } + + +def _title_anchor_candidates( + query: str, + documents: list[dict[str, Any]], + chunks: list[dict[str, Any]], +) -> list[tuple[float, str, dict[str, Any], str]]: + normalized_query = normalize_for_compare(query) + chunks_by_path: dict[str, list[dict[str, Any]]] = {} + + for chunk in chunks: + chunks_by_path.setdefault( + str(chunk.get("document_path") or ""), + [], + ).append(chunk) + + candidates: list[tuple[float, str, dict[str, Any], str]] = [] + + for document in documents: + title = str(document.get("title") or "").strip() + path = str(document.get("document_path") or "").strip() + + if not title or not path: + continue + + normalized_title = normalize_for_compare(title) + title_tokens = normalized_words(title) + + if len(title_tokens) < 2: + continue + + if normalized_title not in normalized_query: + continue + + family_path = canonical_student_main_path(path) + family_chunks = [ + chunk + for chunk in chunks + if canonical_student_main_path( + str(chunk.get("document_path") or "") + ) == family_path + ] + path_chunks = chunks_by_path.get(path, []) + pool = family_chunks or path_chunks + + if not pool: + continue + + best = max( + pool, + key=lambda chunk: ( + evidence_score( + query, + text=str(chunk.get("text") or ""), + heading_paths=chunk.get("heading_paths"), + ), + -int(chunk.get("chunk_index") or 0), + ), + ) + score = 500.0 + 10.0 * len(title_tokens) + candidates.append( + (score, family_path, best, "document_title_in_query") + ) + + return candidates + + +def _global_chunk_candidates( + query: str, + chunks: list[dict[str, Any]], + *, + multi_document: bool, +) -> list[tuple[float, str, dict[str, Any], str]]: + query_tokens = canonical_query_tokens(query) + if not query_tokens: + return [] + + threshold = ( + MIN_MULTI_GLOBAL_COVERAGE + if multi_document + else MIN_SINGLE_GLOBAL_COVERAGE + ) + + normalized_query = normalize_for_compare(query) + explicitly_entity_or_student = any( + marker in normalized_query + for marker in ( + "student", + "osob", + "kto", + "autor", + ) + ) + + best_by_source: dict[ + str, + tuple[float, str, dict[str, Any], str], + ] = {} + + for chunk in chunks: + text = str(chunk.get("text") or "") + headings = json.dumps( + chunk.get("heading_paths") or [], + ensure_ascii=False, + ) + combined = f"{headings}\n{text}" + matches, coverage = canonical_coverage( + query_tokens, + combined, + ) + + raw_path = str(chunk.get("document_path") or "") + topic_page = raw_path.startswith("pages/topics/") + required_coverage = threshold + required_matches = ( + 2 + if multi_document + else MIN_GLOBAL_MATCHES + ) + + # Generic factual questions without a named person often belong to a + # canonical project/topic page. Allow such pages to enter with two + # strong concept matches, then rank them with an explicit topic-page + # bonus. This handles short follow-up-like questions without changing + # the locked global retrieval scorer. + if topic_page and not explicitly_entity_or_student: + required_coverage = min(required_coverage, 0.33) + required_matches = min(required_matches, 2) + + if ( + coverage < required_coverage + or matches < required_matches + ): + continue + + lexical = evidence_score( + query, + text=text, + heading_paths=chunk.get("heading_paths"), + ) + score = ( + coverage * 160.0 + + matches * 12.0 + + lexical + ) + + if topic_page and not explicitly_entity_or_student: + score += 25.0 + + topic_parts = Path(raw_path).parts + title_matches, _ = canonical_coverage( + query_tokens, + str(chunk.get("title") or ""), + ) + if ( + title_matches >= 1 + and len(topic_parts) == 4 + and topic_parts[0] == "pages" + and topic_parts[1] == "topics" + and topic_parts[3] == "README.md" + ): + score += 45.0 + + if ( + len(topic_parts) == 4 + and topic_parts[0] == "pages" + and topic_parts[1] == "topics" + and topic_parts[3] == "README.md" + ): + score += 20.0 + canonical_path = canonical_student_main_path( + str(chunk.get("document_path") or "") + ) + candidate = ( + score, + canonical_path, + chunk, + "global_query_match", + ) + + previous = best_by_source.get(canonical_path) + if previous is None or candidate[0] > previous[0]: + best_by_source[canonical_path] = candidate + + return sorted( + best_by_source.values(), + key=lambda item: (-item[0], item[1]), + ) + + +def _family_evidence_blocks( + query: str, + canonical_path: str, + chunks: list[dict[str, Any]], + *, + max_blocks: int = MAX_FAMILY_EVIDENCE, +) -> list[dict[str, Any]]: + family_chunks = [ + chunk + for chunk in chunks + if same_student_family( + str(chunk.get("document_path") or ""), + canonical_path, + ) + ] + + scored: list[tuple[float, dict[str, Any]]] = [] + + for chunk in family_chunks: + score = evidence_score( + query, + text=str(chunk.get("text") or ""), + heading_paths=chunk.get("heading_paths"), + ) + focus = build_focus_excerpt( + query, + str(chunk.get("text") or ""), + ) + + if score < 20.0 or not focus: + continue + + item = dict(chunk) + item["score"] = score + item["focus_text"] = focus + scored.append((score, item)) + + scored.sort( + key=lambda pair: ( + -pair[0], + str(pair[1].get("document_path") or ""), + int(pair[1].get("chunk_index") or 0), + ) + ) + + blocks: list[dict[str, Any]] = [] + seen: set[str] = set() + + for _, chunk in scored: + focus = str(chunk.get("focus_text") or "").strip() + key = normalize_for_compare(focus) + if key in seen: + continue + + seen.add(key) + blocks.append( + { + "document_path": chunk.get("document_path"), + "chunk_id": chunk.get("chunk_id"), + "chunk_index": int(chunk.get("chunk_index") or 0), + "heading_paths": chunk.get("heading_paths") or [], + "score": round(float(chunk.get("score") or 0.0), 6), + "focus_text": focus, + "text": str(chunk.get("text") or "").strip(), + } + ) + + if len(blocks) >= max_blocks: + break + + return blocks + + +def expand_results_with_supplemental_sources( + db_path: Path, + query: str, + results: list[dict[str, Any]], + *, + published_only: bool = False, + limit: int = 5, +) -> list[dict[str, Any]]: + base_results = [dict(item) for item in results] + + if ( + not query.strip() + or not db_path.exists() + or limit <= 0 + ): + 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 + + documents = _load_documents(conn) + chunks = _load_chunks( + conn, + published_only=published_only, + ) + + if not chunks: + return base_results + + metadata_by_path = _document_metadata_by_path(documents) + + # Attach strongest evidence from child documents belonging to the same student. + for item in base_results: + path = str(item.get("document_path") or "").strip() + if not path: + continue + + canonical_path = canonical_student_main_path(path) + if canonical_path == path and "/students/" not in f"/{path}": + continue + + blocks = _family_evidence_blocks( + query, + canonical_path, + chunks, + ) + if blocks: + item["family_evidence_blocks"] = blocks + item["family_expansion"] = { + "strategy": "student_family_evidence", + "applied": True, + "canonical_document_path": canonical_path, + "evidence_count": len(blocks), + } + + title_candidates = _title_anchor_candidates( + query, + documents, + chunks, + ) + has_title_anchor = bool(title_candidates) + multi_document = is_multi_document_query(query) + + global_candidates: list[ + tuple[float, str, dict[str, Any], str] + ] = [] + + # If the query already names an exact document/person, anchor that document. + # Otherwise use a conservative global scan. Multi-document questions always + # receive the scan because one base retrieval result is not sufficient. + if not has_title_anchor or multi_document: + global_candidates = _global_chunk_candidates( + query, + chunks, + multi_document=multi_document, + ) + + all_candidates = sorted( + title_candidates + global_candidates, + key=lambda item: (-item[0], item[1]), + ) + + existing_by_path: dict[str, dict[str, Any]] = { + canonical_student_main_path( + str(item.get("document_path") or "") + ): item + for item in base_results + if str(item.get("document_path") or "").strip() + } + + synthetic: list[dict[str, Any]] = [] + seen_candidate_paths: set[str] = set() + + for score, canonical_path, chunk, reason in all_candidates: + if canonical_path in seen_candidate_paths: + continue + + seen_candidate_paths.add(canonical_path) + existing = existing_by_path.get(canonical_path) + + if existing is not None: + existing["supplemental_expansion"] = { + "strategy": "global_or_entity_anchor", + "applied": True, + "reason": reason, + "score": round(score, 6), + "evidence_document_path": chunk.get("document_path"), + "evidence_chunk_id": chunk.get("chunk_id"), + } + continue + + synthetic.append( + _source_from_chunk( + canonical_path, + chunk, + metadata_by_path, + reason=reason, + score=score, + ) + ) + + if len(synthetic) >= limit: + break + + # Exact title/person anchors are intentionally promoted. For broad + # multi-document queries, supplemental sources are merged by score ahead of + # weaker base-only results, but the final output still respects the tool limit. + if title_candidates: + promoted_paths = { + canonical_path + for _, canonical_path, _, _ in title_candidates + } + promoted_existing = [ + item + for item in base_results + if canonical_student_main_path( + str(item.get("document_path") or "") + ) in promoted_paths + ] + remainder = [ + item + for item in base_results + if item not in promoted_existing + ] + combined = promoted_existing + synthetic + remainder + elif multi_document: + combined = synthetic + base_results + else: + combined = synthetic + base_results + + deduped: list[dict[str, Any]] = [] + seen_paths: set[str] = set() + + for item in combined: + path = canonical_student_main_path( + str(item.get("document_path") or "") + ) + dedupe_key = path or str(item.get("source_url") or "") + + if dedupe_key in seen_paths: + continue + + seen_paths.add(dedupe_key) + deduped.append(item) + + if len(deduped) >= limit: + break + + return deduped