""" Provenance registry for tracking all sources the LLM sees during a chat session. Every tool result (search hits, fetched documents, etc.) registers its sources here. After the LLM produces an answer with [src:ID] tags, the registry validates cited IDs, renumbers them to [1], [2], and generates the "Källor" section deterministically. """ from __future__ import annotations import re from dataclasses import dataclass, field from typing import Any, Dict, List, Optional, Tuple BODY_CAP_CHARS = 3000 @dataclass class SourceRecord: """A single source (talk) that the LLM saw during research.""" source_id: str # bare talk ID, e.g. "H40911" tool: str # which tool produced it speaker: str | None = None party: str | None = None date: str | None = None heading: str | None = None debateurl: str | None = None snippet: str = "" intressent_id: str | None = None score: float = 0.0 body: str = "" # Grounding text (chunk + neighbours, summary, or capped full text). Capped at BODY_CAP_CHARS. class ProvenanceRegistry: """ Collects and deduplicates sources across all tool calls in a single chat session. Keyed by bare talk ID (e.g. "H40911"). Multiple chunks from the same talk update the existing record (keeping the best snippet) rather than creating duplicate entries. """ def __init__(self) -> None: self._sources: Dict[str, SourceRecord] = {} self._order: List[str] = [] # insertion order def register(self, record: SourceRecord) -> str: """Register a source. Deduplicates by source_id, keeps best snippet/body.""" sid = record.source_id # Cap body on the way in. if record.body and len(record.body) > BODY_CAP_CHARS: record.body = record.body[:BODY_CAP_CHARS].rstrip() + "…" if sid in self._sources: existing = self._sources[sid] # Keep the longer/better snippet if len(record.snippet) > len(existing.snippet): existing.snippet = record.snippet if len(record.body) > len(existing.body): existing.body = record.body # Fill in missing metadata if record.speaker and not existing.speaker: existing.speaker = record.speaker if record.party and not existing.party: existing.party = record.party if record.date and not existing.date: existing.date = record.date if record.heading and not existing.heading: existing.heading = record.heading if record.debateurl and not existing.debateurl: existing.debateurl = record.debateurl if record.intressent_id and not existing.intressent_id: existing.intressent_id = record.intressent_id if record.score > existing.score: existing.score = record.score else: self._sources[sid] = record self._order.append(sid) return sid def get(self, source_id: str) -> Optional[SourceRecord]: return self._sources.get(source_id) def size(self) -> int: return len(self._sources) def all_sources(self) -> List[SourceRecord]: """Return all sources in registration order.""" return [self._sources[sid] for sid in self._order if sid in self._sources] def get_persons(self) -> Dict[str, Dict]: """Return {intressent_id: {name, party}} for person link injection.""" persons: Dict[str, Dict] = {} for src in self._sources.values(): if src.intressent_id and src.speaker and src.intressent_id not in persons: persons[src.intressent_id] = { "name": src.speaker, "party": src.party or "", } return persons def to_cited_sources(self, cited_ids: List[str]) -> List[Dict[str, Any]]: """Convert a list of cited source IDs to ChatSource dicts for the frontend.""" result = [] for sid in cited_ids: src = self._sources.get(sid) if not src: continue result.append( { "_id": f"talks/{sid}", "heading": src.heading, "snippet": _trim_snippet(src.snippet), "chunk_index": -1, "debateurl": src.debateurl, "speaker": src.speaker, "party": src.party, "intressent_id": src.intressent_id, "date": src.date, } ) return result # --------------------------------------------------------------------------- # Citation parsing and renumbering # --------------------------------------------------------------------------- # Tag format: [src:ID] or [src:ID | Speaker (Party) | date]. Also tolerates # [source:ID], [[src:ID]], and mixed case — all produced by LLMs in practice. # Double brackets ([[...]]) are matched by \[{1,2} / \]{1,2}. # Capture group 1 is always the bare ID; metadata after "|" is ignored. _SRC_PATTERN = re.compile( r"\[{1,2}(?:src|source):([A-Za-z0-9_-]+)(?:\s*\|[^\]]*?)?\]{1,2}", re.IGNORECASE, ) _KALLOR_SPLIT = re.compile(r"\n#+\s*K[äa]ll[ao]r", re.IGNORECASE) def parse_and_renumber_citations( answer_text: str, registry: ProvenanceRegistry, max_fallback: int = 5, ) -> Tuple[str, List[Dict[str, Any]], List[str], List[str]]: """ Parse [src:ID] tags from the model's answer, validate against the registry, replace with [1], [2], and generate a "Källor" section. Returns (validated_answer, cited_sources_list). """ cited_ids_raw = _SRC_PATTERN.findall(answer_text) # Deduplicate preserving first-appearance order, validate against registry seen: set[str] = set() unique_cited_ids: List[str] = [] invalid_ids: List[str] = [] for cid in cited_ids_raw: if cid in seen: continue if registry.get(cid): seen.add(cid) unique_cited_ids.append(cid) else: invalid_ids.append(cid) seen.add(cid) # don't report same invalid ID twice # Build replacement map id_to_number = {cid: i + 1 for i, cid in enumerate(unique_cited_ids)} def _replace_src(m: re.Match) -> str: src_id = m.group(1) num = id_to_number.get(src_id) return f"[{num}]" if num else "" validated_answer = _SRC_PATTERN.sub(_replace_src, answer_text) # Strip any leftover malformed tags the main pattern didn't catch # (e.g. mismatched brackets, extra punctuation around the tag). validated_answer = re.sub( r"\[{1,2}(?:src|source):[A-Za-z0-9_:.\-]{1,80}(?:\s*\|[^\]\n]{0,150})?\]{1,2}", "", validated_answer, flags=re.IGNORECASE, ) # Strip any model-generated "Källor" section validated_answer = _KALLOR_SPLIT.split(validated_answer)[0].rstrip() # Build cited sources from registry if unique_cited_ids: cited_sources = registry.to_cited_sources(unique_cited_ids) else: cited_sources = [] # Generate "Källor" section server-side if cited_sources: kallor_lines = [] for i, src in enumerate(cited_sources, 1): speaker = src.get("speaker") or "Okänd" src_date = src.get("date") or "" line = f"[{i}] {speaker} – {src_date}" kallor_lines.append(line) validated_answer += "\n\n### Källor\n\n" + "\n\n".join(kallor_lines) return validated_answer, cited_sources, unique_cited_ids, invalid_ids def normalize_talk_id(raw: str | None) -> str | None: """Strip 'talks/' prefix to get bare talk ID.""" if not raw: return None if "/" in raw: return raw.split("/", 1)[1] return raw def _trim_snippet(text: str, length: int = 400) -> str: cleaned = text.strip() if len(cleaned) <= length: return cleaned return f"{cleaned[:length].rstrip()}…"