| from __future__ import annotations |
|
|
| import logging |
| import math |
| import re |
| from collections import Counter, defaultdict |
| from dataclasses import dataclass |
| from typing import Any, Dict, Iterable, List, Optional, Set, Tuple |
|
|
| import chromadb |
| from sentence_transformers import SentenceTransformer |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class NormReference: |
| """Structured representation of a legal reference found in a user question.""" |
|
|
| paragraph: str |
| subsection: str | None = None |
| sentence: str | None = None |
| number: str | None = None |
| letter: str | None = None |
|
|
| @property |
| def section_id(self) -> str: |
| return self.paragraph |
|
|
| @property |
| def canonical_ref(self) -> str: |
| parts = [self.paragraph] |
| if self.subsection: |
| parts.append(f"Abs. {self.subsection}") |
| if self.sentence: |
| parts.append(f"Satz {self.sentence}") |
| if self.number: |
| parts.append(f"Nr. {self.number}") |
| if self.letter: |
| parts.append(f"Buchst. {self.letter}") |
| return " ".join(parts) |
|
|
|
|
| class LegalRetriever: |
| """ |
| Standalone production-oriented retriever for legal RAG applications on Chroma. |
| |
| This class intentionally does not import from an ingestion package. It can live in a |
| separate backend/chatbot folder and only depends on Chroma metadata produced during |
| ingestion. |
| |
| It is backward-compatible with the previous retriever API while adding support for |
| richer legal metadata generated by a parent-child legal chunker. |
| |
| Core capabilities: |
| - default filtering to the main contract, |
| - semantic/dense retrieval, |
| - exact norm lookup for § / Abs. / Satz / Nr. / Buchst., |
| - explicit section lookup for old collections, |
| - definition lookup via is_definition / defined_terms / § 2 fallback, |
| - BM25-like lexical fallback, |
| - optional neighbor chunk expansion, |
| - parent-context expansion via parent_unit_id, |
| - result fusion, deduplication and ranking, |
| - negative-answer verification before saying something is not regulated, |
| - source and context formatting for downstream LLM calls. |
| |
| Expected legacy metadata per chunk: |
| - container_id, e.g. "Vertrag", "Anlage 4", "Anhang zu Anlage 11" |
| - container_type, e.g. "vertrag", "anlage", "anhang" |
| - section_id, e.g. "§ 6" |
| - section_path |
| - page_start, page_end |
| - chunk_index_in_section |
| - text_hash optional |
| |
| Additional metadata supported from the optimized ingestion/chunking layer: |
| - chunk_kind: "parent" | "child" | ... |
| - legal_unit_id |
| - parent_unit_id |
| - canonical_ref |
| - paragraph |
| - subsection |
| - sentence |
| - number |
| - letter |
| - unit_type |
| - is_definition |
| - defined_terms |
| """ |
|
|
| SECTION_REF_RE = re.compile( |
| r"§{1,2}\s*(\d+[a-zA-Z]?)" |
| r"(?:\s*(?:-|–|bis)\s*(\d+[a-zA-Z]?))?", |
| re.IGNORECASE, |
| ) |
|
|
| NORM_REF_RE = re.compile( |
| r"§{1,2}\s*(?P<para>\d+[a-zA-Z]?)" |
| r"(?:\s*(?:Abs\.?|Absatz)\s*(?P<abs>\d+[a-zA-Z]?))?" |
| r"(?:\s*Satz\s*(?P<satz>\d+[a-zA-Z]?))?" |
| r"(?:\s*(?:Nr\.?|Nummer)\s*(?P<nr>\d+[a-zA-Z]?))?" |
| r"(?:\s*(?:Buchst\.?|Buchstabe|lit\.?)\s*(?P<letter>[a-zA-Z]))?", |
| re.IGNORECASE, |
| ) |
|
|
| QUOTED_TERM_RE = re.compile(r"[„\"']([^„“\"']{2,120})[“\"']") |
|
|
| STOPWORDS = { |
| "aber", |
| "alle", |
| "alles", |
| "als", |
| "also", |
| "am", |
| "an", |
| "auch", |
| "auf", |
| "aus", |
| "bei", |
| "bis", |
| "da", |
| "das", |
| "dass", |
| "dem", |
| "den", |
| "der", |
| "des", |
| "die", |
| "dies", |
| "diese", |
| "dieser", |
| "dieses", |
| "ein", |
| "eine", |
| "einem", |
| "einen", |
| "einer", |
| "eines", |
| "er", |
| "es", |
| "für", |
| "gilt", |
| "hat", |
| "im", |
| "in", |
| "ist", |
| "kann", |
| "mit", |
| "nach", |
| "oder", |
| "rahmenvertrag", |
| "regelt", |
| "sagt", |
| "sind", |
| "unter", |
| "und", |
| "von", |
| "wann", |
| "was", |
| "welche", |
| "welchen", |
| "welcher", |
| "welches", |
| "wenn", |
| "wer", |
| "wie", |
| "wird", |
| "wo", |
| "zum", |
| "zur", |
| } |
|
|
| QUERY_EXPANSIONS: Dict[str, List[str]] = { |
| "nicht verfügbar": [ |
| "nicht lieferbar", |
| "nicht vorrätig", |
| "Lieferengpass", |
| "Verfügbarkeit", |
| "lieferfähig", |
| "lieferbar", |
| ], |
| "nicht lieferbar": ["nicht verfügbar", "Lieferengpass", "lieferfähig", "Verfügbarkeit"], |
| "lieferengpass": ["nicht verfügbar", "nicht lieferbar", "Verfügbarkeit", "lieferfähig"], |
| "auseinzelung": ["Teilmenge", "Auseinzelung", "aus Packungen entnehmen", "§ 16"], |
| "teilmenge": ["Auseinzelung", "Teilmenge", "§ 16"], |
| "beitritt": ["beitreten", "teilnehmen", "Mitgliedsverband", "DAV", "Erklärung", "§ 4"], |
| "teilnahme": ["Beitritt", "teilnehmen", "Mitgliedsverband", "DAV", "§ 4"], |
| "wunscharzneimittel": ["Wunscharzneimittel", "Kostenerstattung", "anderes Fertigarzneimittel", "§ 15"], |
| "pharmazeutische dienstleistungen": ["pharmazeutische Dienstleistung", "Anlage 11", "§ 33"], |
| "biosimilar": ["Biosimilar", "biotechnologisch", "Referenzarzneimittel"], |
| "bioidentical": ["Bioidentical", "Ausgangsstoff", "Herstellungsprozess"], |
| "importarzneimittel": ["Importarzneimittel", "Parallelimport", "Reimport", "Referenzarzneimittel"], |
| "rabattvertrag": ["Rabattvertrag", "rabattbegünstigt", "§ 11", "§ 130a"], |
| } |
|
|
| def __init__( |
| self, |
| persist_dir: str, |
| collection: str, |
| model_name: str = "auto", |
| *, |
| default_container_id: str = "Vertrag", |
| normalize_embeddings: bool = True, |
| default_to_contract: bool = False, |
| min_score: float = 0.20, |
| verbose_startup: bool = True, |
| lexical_scan_limit: int = 5000, |
| query_prefix: str | None = None, |
| enable_reranker: bool = False, |
| reranker_model: str = "cross-encoder/mmarco-mMiniLMv2-L12-H384-v1", |
| reranker_candidates: int = 20, |
| ): |
| self.persist_dir = persist_dir |
| self.collection_name = collection |
| self.default_container_id = default_container_id |
| self.normalize_embeddings = normalize_embeddings |
| self.default_to_contract = default_to_contract |
| self.min_score = float(min_score) |
| self.lexical_scan_limit = int(max(100, lexical_scan_limit)) |
|
|
| self.client = chromadb.PersistentClient(path=persist_dir) |
|
|
| try: |
| self.col = self.client.get_collection(name=collection) |
| except Exception as exc: |
| available = [c.name for c in self.client.list_collections()] |
| raise RuntimeError( |
| "Chroma Collection wurde nicht gefunden. " |
| f"Pfad={persist_dir!r}, gesuchte Collection={collection!r}, " |
| f"verfügbare Collections={available!r}" |
| ) from exc |
|
|
| |
| |
| |
| |
| |
| |
| |
| collection_meta = dict(getattr(self.col, "metadata", None) or {}) |
| indexed_model = str(collection_meta.get("embedding_model") or "").strip() |
|
|
| requested = (model_name or "").strip() |
| if not requested or requested.lower() in {"auto", "collection"}: |
| if not indexed_model: |
| raise RuntimeError( |
| "EMBEDDING_MODEL='auto' verlangt, dass die Collection-Metadata " |
| "'embedding_model' enthält. Diese Collection wurde offenbar mit " |
| "einer älteren Ingest-Version gebaut; setze EMBEDDING_MODEL explizit." |
| ) |
| model_name = indexed_model |
| elif indexed_model and indexed_model != requested: |
| raise RuntimeError( |
| "Embedding-Modell passt nicht zur Collection: " |
| f"konfiguriert={requested!r}, Collection wurde indexiert mit {indexed_model!r}. " |
| "Entweder EMBEDDING_MODEL='auto' setzen oder die Collection neu ingestieren." |
| ) |
|
|
| if verbose_startup: |
| self._log_startup_diagnostics() |
|
|
| self.model_name = model_name |
| self.embedder = SentenceTransformer(model_name) |
|
|
| expected_dim = collection_meta.get("embedding_dim") |
| actual_dim = self.embedder.get_sentence_embedding_dimension() |
| if expected_dim and actual_dim and int(expected_dim) > 0 and int(expected_dim) != int(actual_dim): |
| raise RuntimeError( |
| "Embedding-Dimension passt nicht zur Collection: " |
| f"Modell {model_name!r} liefert {actual_dim}, Collection erwartet {expected_dim}." |
| ) |
|
|
| |
| |
| |
| |
| if query_prefix is None: |
| query_prefix = "query: " if "e5" in model_name.lower() else "" |
| self.query_prefix = query_prefix |
|
|
| self.enable_reranker = bool(enable_reranker) |
| self.reranker_model = reranker_model |
| self.reranker_candidates = int(max(1, reranker_candidates)) |
| self._reranker: Any = None |
|
|
| |
| |
| |
| def _log_startup_diagnostics(self) -> None: |
| count = self.col.count() |
| msg = ( |
| f"USING CHROMA PATH: {self.persist_dir}\n" |
| f"USING CHROMA COLLECTION: {self.collection_name}\n" |
| f"COLLECTION COUNT: {count}" |
| ) |
| print(msg) |
| logger.info( |
| "retriever initialized", |
| extra={ |
| "persist_dir": self.persist_dir, |
| "collection": self.collection_name, |
| "count": count, |
| }, |
| ) |
|
|
| def diagnostics(self, *, sample: int = 3) -> Dict[str, Any]: |
| """Return compact runtime diagnostics to catch stale paths or collections.""" |
| collections = [c.name for c in self.client.list_collections()] |
| result: Dict[str, Any] = { |
| "persist_dir": self.persist_dir, |
| "collection": self.collection_name, |
| "available_collections": collections, |
| "count": self.col.count(), |
| } |
|
|
| try: |
| res = self.col.get(limit=sample, include=["metadatas"]) |
| result["metadata_sample"] = res.get("metadatas", []) |
| result["metadata_keys"] = sorted( |
| {key for meta in result["metadata_sample"] if isinstance(meta, dict) for key in meta.keys()} |
| ) |
| except Exception as exc: |
| result["metadata_sample_error"] = repr(exc) |
|
|
| return result |
|
|
| |
| |
| |
| @staticmethod |
| def _build_where(filters: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: |
| """ |
| Build a Chroma-compatible where filter. |
| |
| Chroma permits only one top-level logical operator for compound filters. |
| """ |
| if not filters: |
| return None |
|
|
| clean = {k: v for k, v in filters.items() if v is not None} |
| if not clean: |
| return None |
|
|
| if len(clean) == 1 and next(iter(clean)).startswith("$"): |
| return clean |
|
|
| if len(clean) == 1: |
| return clean |
|
|
| return {"$and": [{k: v} for k, v in clean.items()]} |
|
|
| @staticmethod |
| def _where_and(*conditions: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: |
| parts: List[Dict[str, Any]] = [] |
|
|
| for condition in conditions: |
| if not condition: |
| continue |
| if "$and" in condition and len(condition) == 1: |
| parts.extend(condition["$and"]) |
| else: |
| parts.append(condition) |
|
|
| if not parts: |
| return None |
| if len(parts) == 1: |
| return parts[0] |
| return {"$and": parts} |
|
|
| @staticmethod |
| def _where_or(*conditions: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: |
| parts: List[Dict[str, Any]] = [condition for condition in conditions if condition] |
| if not parts: |
| return None |
| if len(parts) == 1: |
| return parts[0] |
| return {"$or": parts} |
|
|
| def _effective_where( |
| self, |
| where: Optional[Dict[str, Any]], |
| *, |
| restrict_to_default_container: Optional[bool], |
| ) -> Optional[Dict[str, Any]]: |
| """ |
| By default, search the main contract. To search all containers intentionally, |
| call query(..., restrict_to_default_container=False). |
| """ |
| if where: |
| return where |
|
|
| restrict = self.default_to_contract if restrict_to_default_container is None else restrict_to_default_container |
| if restrict: |
| return {"container_id": self.default_container_id} |
|
|
| return None |
|
|
| |
| |
| |
| @classmethod |
| def _parse_section_refs(cls, text: str, *, max_range: int = 30) -> List[str]: |
| """ |
| Recognize paragraph references such as § 6, §§ 7-14, §§ 7 bis 14. |
| Returns normalized section_ids such as ["§ 6", "§ 7", ...]. |
| """ |
| if not text: |
| return [] |
|
|
| found: List[str] = [] |
|
|
| for match in cls.SECTION_REF_RE.finditer(text): |
| start_raw, end_raw = match.group(1), match.group(2) |
|
|
| if end_raw and start_raw.isdigit() and end_raw.isdigit(): |
| start, end = int(start_raw), int(end_raw) |
| if start <= end and (end - start) <= max_range: |
| found.extend([f"§ {i}" for i in range(start, end + 1)]) |
| else: |
| found.append(f"§ {start_raw}") |
| found.append(f"§ {end_raw}") |
| else: |
| found.append(f"§ {start_raw}") |
| if end_raw: |
| found.append(f"§ {end_raw}") |
|
|
| return list(dict.fromkeys(found)) |
|
|
| @classmethod |
| def _parse_norm_refs(cls, text: str) -> List[NormReference]: |
| """Recognize structured references including Abs., Satz, Nr. and Buchst.""" |
| if not text: |
| return [] |
|
|
| refs: List[NormReference] = [] |
| for match in cls.NORM_REF_RE.finditer(text): |
| para = match.group("para") |
| if not para: |
| continue |
| ref = NormReference( |
| paragraph=f"§ {para}", |
| subsection=match.group("abs"), |
| sentence=match.group("satz"), |
| number=match.group("nr"), |
| letter=(match.group("letter") or "").lower() or None, |
| ) |
| refs.append(ref) |
|
|
| seen: Set[str] = set() |
| unique: List[NormReference] = [] |
| for ref in refs: |
| key = ref.canonical_ref.lower() |
| if key not in seen: |
| unique.append(ref) |
| seen.add(key) |
| return unique |
|
|
| @staticmethod |
| def _section_variants(section: str) -> List[str]: |
| if not section: |
| return [] |
|
|
| s = str(section).strip() |
| num = s.replace("§", "").strip() |
| variants = [s] |
|
|
| if num: |
| variants.extend([f"§ {num}", f"§{num}", num]) |
|
|
| return list(dict.fromkeys(variants)) |
|
|
| @staticmethod |
| def _paragraph_variants(paragraph: str) -> List[str]: |
| return LegalRetriever._section_variants(paragraph) |
|
|
| @staticmethod |
| def _subsection_variants(subsection: str | None) -> List[str]: |
| if not subsection: |
| return [] |
| s = str(subsection).strip() |
| return list(dict.fromkeys([s, f"Abs. {s}", f"Absatz {s}", f"({s})"])) |
|
|
| @staticmethod |
| def _sentence_variants(sentence: str | None) -> List[str]: |
| if not sentence: |
| return [] |
| s = str(sentence).strip() |
| return list(dict.fromkeys([s, f"Satz {s}"])) |
|
|
| @staticmethod |
| def _number_variants(number: str | None) -> List[str]: |
| if not number: |
| return [] |
| s = str(number).strip() |
| return list(dict.fromkeys([s, f"Nr. {s}", f"Nummer {s}"])) |
|
|
| @staticmethod |
| def _letter_variants(letter: str | None) -> List[str]: |
| if not letter: |
| return [] |
| s = str(letter).strip().lower() |
| return list(dict.fromkeys([s, f"Buchst. {s}", f"Buchstabe {s}", f"{s})"])) |
|
|
| |
| |
| |
| def query( |
| self, |
| question: str, |
| top_k: int = 8, |
| where: Optional[Dict[str, Any]] = None, |
| *, |
| fetch_k: Optional[int] = None, |
| include_explicit_sections: bool = True, |
| explicit_sections: Optional[List[str]] = None, |
| explicit_section_container_ids: Optional[List[str]] = None, |
| max_chunks_per_explicit_section: int = 6, |
| include_definitions: bool = True, |
| definition_k: int = 8, |
| include_lexical: bool = True, |
| lexical_k: int = 12, |
| lexical_scan_limit: Optional[int] = None, |
| include_parent_context: bool = True, |
| max_parent_contexts: int = 8, |
| include_neighbors: bool = True, |
| neighbor_window: int = 1, |
| max_final_results: Optional[int] = 12, |
| min_score: Optional[float] = None, |
| restrict_to_default_container: Optional[bool] = None, |
| max_chunks_per_section: int = 5, |
| expand_query: bool = True, |
| verify_negative_answer: bool = False, |
| rerank: Optional[bool] = None, |
| ) -> List[Dict[str, Any]]: |
| """ |
| Hybrid legal retrieval. |
| |
| The old behavior is preserved and expanded: |
| - semantic query, |
| - exact norm/section lookup, |
| - optional neighbors, |
| - section limits. |
| |
| New paths: |
| - definitions, |
| - lexical fallback, |
| - parent context expansion, |
| - negative answer verification. |
| """ |
| if not question or not question.strip(): |
| return [] |
|
|
| effective_where = self._effective_where( |
| where, |
| restrict_to_default_container=restrict_to_default_container, |
| ) |
| threshold = self.min_score if min_score is None else float(min_score) |
| intent = self.classify_query_intent(question) |
| expanded_terms = self.expand_query_terms(question) if expand_query else [] |
|
|
| all_results: List[Dict[str, Any]] = [] |
|
|
| |
| semantic_results = self._semantic_query( |
| question=question, |
| top_k=top_k, |
| where=effective_where, |
| fetch_k=fetch_k, |
| ) |
| semantic_results = [hit for hit in semantic_results if float(hit.get("score", 0.0)) >= threshold] |
| all_results.extend(semantic_results) |
|
|
| |
| if include_explicit_sections: |
| norm_refs = self._parse_norm_refs(question) |
| sections = list(explicit_sections or []) |
| sections.extend(self._parse_section_refs(question)) |
| sections = list(dict.fromkeys(sections)) |
| containers = explicit_section_container_ids or [self.default_container_id] |
|
|
| if norm_refs: |
| all_results.extend( |
| self._get_explicit_norms( |
| refs=norm_refs, |
| container_ids=containers, |
| base_where=effective_where, |
| max_chunks_per_ref=max_chunks_per_explicit_section, |
| ) |
| ) |
| elif sections: |
| all_results.extend( |
| self._get_explicit_sections( |
| sections=sections, |
| container_ids=containers, |
| base_where=effective_where, |
| max_chunks_per_section=max_chunks_per_explicit_section, |
| ) |
| ) |
|
|
| |
| if include_definitions and (intent == "definition" or self._extract_definition_terms(question)): |
| all_results.extend( |
| self._definition_search( |
| question=question, |
| base_where=effective_where, |
| top_k=definition_k, |
| expanded_terms=expanded_terms, |
| ) |
| ) |
|
|
| |
| if include_lexical: |
| all_results.extend( |
| self._lexical_search( |
| question=question, |
| base_where=effective_where, |
| top_k=lexical_k, |
| expanded_terms=expanded_terms, |
| scan_limit=lexical_scan_limit or self.lexical_scan_limit, |
| ) |
| ) |
|
|
| |
| if include_neighbors and neighbor_window > 0: |
| neighbor_results: List[Dict[str, Any]] = [] |
| seed_hits = self._dedupe_and_rank(all_results)[: max(top_k, 1)] |
| for hit in seed_hits: |
| neighbor_results.extend( |
| self._get_neighbors(hit, window=neighbor_window, base_where=effective_where) |
| ) |
| all_results.extend(neighbor_results) |
|
|
| |
| if include_parent_context: |
| parent_results = self._expand_parent_context( |
| all_results, |
| base_where=effective_where, |
| max_parent_contexts=max_parent_contexts, |
| ) |
| all_results.extend(parent_results) |
|
|
| merged = self._dedupe_and_rank(all_results, question=question, intent=intent) |
|
|
| use_reranker = self.enable_reranker if rerank is None else bool(rerank) |
| if use_reranker: |
| merged = self._rerank(question, merged) |
|
|
| merged = self._limit_chunks_per_section(merged, max_chunks_per_section=max_chunks_per_section) |
|
|
| if max_final_results is not None: |
| merged = merged[:max_final_results] |
|
|
| |
| |
| if verify_negative_answer and self._looks_like_negative_risk(question, merged): |
| check = self.verify_negative_result(question, max_results=max_final_results or 12) |
| if not check.get("safe_to_answer_negative") and check.get("results"): |
| merged = check["results"] |
|
|
| return merged |
|
|
| def query_paragraph( |
| self, |
| question: str, |
| paragraph: str, |
| container: str = "Vertrag", |
| top_k: int = 6, |
| **kwargs: Any, |
| ) -> List[Dict[str, Any]]: |
| return self.query( |
| question, |
| top_k=top_k, |
| where={ |
| "container_id": container, |
| "section_id": paragraph, |
| }, |
| include_explicit_sections=False, |
| **kwargs, |
| ) |
|
|
| def query_anlage( |
| self, |
| question: str, |
| anlage_nr: int, |
| paragraph: Optional[str] = None, |
| top_k: int = 6, |
| **kwargs: Any, |
| ) -> List[Dict[str, Any]]: |
| where: Dict[str, Any] = {"container_id": f"Anlage {anlage_nr}"} |
| if paragraph: |
| where["section_id"] = paragraph |
|
|
| return self.query(question, top_k=top_k, where=where, **kwargs) |
|
|
| def query_anhang( |
| self, |
| question: str, |
| anlage_nr: int, |
| top_k: int = 6, |
| **kwargs: Any, |
| ) -> List[Dict[str, Any]]: |
| return self.query( |
| question, |
| top_k=top_k, |
| where={ |
| "container_type": "anhang", |
| "container_id": f"Anhang zu Anlage {anlage_nr}", |
| }, |
| **kwargs, |
| ) |
|
|
| def get_section( |
| self, |
| section: str, |
| *, |
| container: str = "Vertrag", |
| max_chunks: int = 20, |
| include_parent_context: bool = False, |
| ) -> List[Dict[str, Any]]: |
| variants = self._section_variants(section) |
| where = {"container_id": container, "section_id": {"$in": variants}} |
|
|
| try: |
| res = self.col.get(where=self._build_where(where), include=["documents", "metadatas"]) |
| except Exception: |
| return [] |
|
|
| formatted = self._format_get(res, retrieval_kind="section_lookup", score=1.0) |
| formatted.sort(key=lambda x: x.get("chunk_index", 0)) |
| out = formatted[:max_chunks] |
|
|
| if include_parent_context: |
| out.extend(self._expand_parent_context(out, base_where={"container_id": container})) |
| out = self._dedupe_and_rank(out) |
|
|
| return out[:max_chunks] |
|
|
| def get_units( |
| self, |
| section: str, |
| *, |
| container: Optional[str] = None, |
| subsection: Optional[str] = None, |
| max_chunks: int = 2, |
| retrieval_kind: str = "norm_anchor", |
| ) -> List[Dict[str, Any]]: |
| """Fetch one norm by its address, independent of any question. |
| |
| Two things `get_section` cannot do, and both are needed for a mandatory |
| fetch driven by a curated register: |
| |
| * No container is assumed. A statute's §§ live in Kapitel containers, |
| and `get_section` would filter on "Vertrag" and return nothing. |
| * The Absatz can be narrowed. "§ 129 SGB V" is eighty chunks; |
| "§ 129 Abs. 1" is two, and both of them carry. |
| |
| Parent units are preferred over their children because the parent chunk |
| holds the full text of the Absatz while the children repeat fragments of |
| it — one chunk is then usually the whole answer. |
| """ |
| where = { |
| "container_id": container or None, |
| "section_id": {"$in": self._section_variants(section)}, |
| "subsection": str(subsection) if subsection else None, |
| } |
|
|
| try: |
| res = self.col.get(where=self._build_where(where), include=["documents", "metadatas"]) |
| except Exception as exc: |
| logger.debug("norm unit lookup failed", exc_info=exc) |
| return [] |
|
|
| formatted = self._format_get(res, retrieval_kind=retrieval_kind, score=1.0) |
| formatted.sort( |
| key=lambda hit: ( |
| 0 if (hit.get("metadata") or {}).get("chunk_kind") == "parent" else 1, |
| hit.get("chunk_index", 0), |
| ) |
| ) |
| return formatted[:max_chunks] |
|
|
| def verify_negative_result( |
| self, |
| question: str, |
| *, |
| max_results: int = 12, |
| strong_score: float = 0.55, |
| ) -> Dict[str, Any]: |
| """ |
| Run a broad second-pass search before the answer layer says that something |
| is not regulated or not found. |
| |
| Returns: |
| { |
| "safe_to_answer_negative": bool, |
| "reason": str, |
| "results": list[dict], |
| "strong_result_count": int, |
| } |
| """ |
| results = self.query( |
| question, |
| top_k=max(10, max_results), |
| fetch_k=max(40, max_results * 4), |
| include_explicit_sections=True, |
| include_definitions=True, |
| include_lexical=True, |
| lexical_k=max(20, max_results * 2), |
| include_parent_context=True, |
| include_neighbors=True, |
| neighbor_window=1, |
| restrict_to_default_container=False, |
| min_score=0.0, |
| max_final_results=max_results, |
| max_chunks_per_section=6, |
| verify_negative_answer=False, |
| ) |
|
|
| strong_kinds = {"exact_norm", "explicit_section", "definition", "lexical", "parent_context"} |
| strong = [ |
| hit |
| for hit in results |
| if float(hit.get("rank_score", hit.get("score", 0.0))) >= strong_score |
| or bool(strong_kinds.intersection(set(hit.get("retrieval_kinds", [])))) |
| ] |
|
|
| if strong: |
| return { |
| "safe_to_answer_negative": False, |
| "reason": "Der breite Kontrollabruf hat potenziell relevante Regelungen gefunden.", |
| "results": results, |
| "strong_result_count": len(strong), |
| } |
|
|
| return { |
| "safe_to_answer_negative": True, |
| "reason": "Auch der breite Kontrollabruf hat keine belastbaren Treffer gefunden.", |
| "results": results, |
| "strong_result_count": 0, |
| } |
|
|
| |
| |
| |
| @classmethod |
| def classify_query_intent(cls, question: str) -> str: |
| q = cls._norm_text(question) |
| if re.search(r"\b(was\s+versteht|wie\s+definiert|definition|legaldefinition|begriff|bedeutet)\b", q): |
| return "definition" |
| if cls._parse_norm_refs(question): |
| return "norm_lookup" |
| if re.search(r"\b(welche|voraussetzungen|kriterien|tatbestandsmerkmale|nennt|liste|auflistung)\b", q): |
| return "enumeration" |
| if re.search(r"\b(nicht\s+geregelt|keine\s+regelung|steht\s+nicht|nicht\s+enthalten)\b", q): |
| return "negative_check" |
| return "semantic" |
|
|
| @classmethod |
| def expand_query_terms(cls, question: str) -> List[str]: |
| q = cls._norm_text(question) |
| expansions: List[str] = [] |
| for key, values in cls.QUERY_EXPANSIONS.items(): |
| if key in q: |
| expansions.extend(values) |
| |
| expansions.extend(cls._extract_definition_terms(question)) |
| return list(dict.fromkeys([x.strip() for x in expansions if x and x.strip()])) |
|
|
| @classmethod |
| def _extract_definition_terms(cls, question: str) -> List[str]: |
| terms: List[str] = [] |
| for match in cls.QUOTED_TERM_RE.finditer(question or ""): |
| term = match.group(1).strip() |
| if term: |
| terms.append(term) |
|
|
| q = (question or "").strip() |
| patterns = [ |
| r"unter\s+(.+?)(?:\?|$)", |
| r"begriff\s+(.+?)(?:\?|$)", |
| r"bedeutet\s+(.+?)(?:\?|$)", |
| r"definiert\s+(?:der\s+vertrag\s+|der\s+rahmenvertrag\s+)?(?:ein\s+|eine\s+|einen\s+|das\s+|den\s+|die\s+)?(.+?)(?:\?|$)", |
| ] |
| for pattern in patterns: |
| m = re.search(pattern, q, flags=re.I) |
| if not m: |
| continue |
| raw = m.group(1) |
| raw = re.sub(r"\b(im|in|nach|des|der|die|das|ein|eine|einen|rahmenvertrag|vertrag)\b", " ", raw, flags=re.I) |
| raw = re.sub(r"\s+", " ", raw).strip(" .,:;!?\"'„“") |
| if 2 <= len(raw) <= 80: |
| terms.append(raw) |
|
|
| return list(dict.fromkeys(terms)) |
|
|
| |
| |
| |
| def _semantic_query( |
| self, |
| *, |
| question: str, |
| top_k: int, |
| where: Optional[Dict[str, Any]], |
| fetch_k: Optional[int], |
| ) -> List[Dict[str, Any]]: |
| n_results = fetch_k if fetch_k is not None else max(top_k * 4, top_k) |
|
|
| query_text = question |
| if self.query_prefix and not question.startswith(self.query_prefix): |
| query_text = f"{self.query_prefix}{question}" |
|
|
| if hasattr(self.embedder, "encode_query"): |
| q_emb = self.embedder.encode_query([query_text], normalize_embeddings=self.normalize_embeddings)[0].tolist() |
| else: |
| q_emb = self.embedder.encode([query_text], normalize_embeddings=self.normalize_embeddings)[0].tolist() |
|
|
| chroma_where = self._build_where(where) |
|
|
| res = self.col.query( |
| query_embeddings=[q_emb], |
| n_results=n_results, |
| where=chroma_where, |
| include=["documents", "metadatas", "distances"], |
| ) |
|
|
| return self._format(res, retrieval_kind="semantic")[:n_results] |
|
|
| def _get_explicit_norms( |
| self, |
| *, |
| refs: List[NormReference], |
| container_ids: List[str], |
| base_where: Optional[Dict[str, Any]] = None, |
| max_chunks_per_ref: int = 6, |
| ) -> List[Dict[str, Any]]: |
| out: List[Dict[str, Any]] = [] |
|
|
| for container_id in container_ids: |
| for ref in refs: |
| |
| metadata_filter: Dict[str, Any] = {"container_id": container_id} |
| metadata_filter["paragraph"] = {"$in": self._paragraph_variants(ref.paragraph)} |
| if ref.subsection: |
| metadata_filter["subsection"] = {"$in": self._subsection_variants(ref.subsection)} |
| if ref.sentence: |
| metadata_filter["sentence"] = {"$in": self._sentence_variants(ref.sentence)} |
| if ref.number: |
| metadata_filter["number"] = {"$in": self._number_variants(ref.number)} |
| if ref.letter: |
| metadata_filter["letter"] = {"$in": self._letter_variants(ref.letter)} |
|
|
| where = self._where_and(self._build_where(base_where), metadata_filter) |
| formatted: List[Dict[str, Any]] = [] |
|
|
| try: |
| res = self.col.get(where=where, include=["documents", "metadatas"]) |
| formatted = self._format_get(res, retrieval_kind="exact_norm", score=1.0) |
| except Exception as exc: |
| logger.debug("structured exact norm lookup failed", exc_info=exc) |
|
|
| |
| if not formatted: |
| formatted = self._get_explicit_sections( |
| sections=[ref.section_id], |
| container_ids=[container_id], |
| base_where=base_where, |
| max_chunks_per_section=max_chunks_per_ref * 2, |
| retrieval_kind="exact_norm", |
| ) |
| if ref.subsection or ref.sentence or ref.number or ref.letter: |
| narrowed = [hit for hit in formatted if self._hit_matches_norm_ref(hit, ref)] |
| if narrowed: |
| formatted = narrowed |
|
|
| formatted.sort(key=lambda hit: self._exact_norm_sort_key(hit, ref), reverse=True) |
| out.extend(formatted[:max_chunks_per_ref]) |
|
|
| return out |
|
|
| def _get_explicit_sections( |
| self, |
| *, |
| sections: List[str], |
| container_ids: List[str], |
| base_where: Optional[Dict[str, Any]] = None, |
| max_chunks_per_section: int = 4, |
| retrieval_kind: str = "explicit_section", |
| ) -> List[Dict[str, Any]]: |
| out: List[Dict[str, Any]] = [] |
|
|
| for container_id in container_ids: |
| for section in sections: |
| variants = self._section_variants(section) |
|
|
| where = self._where_and( |
| self._build_where(base_where), |
| {"container_id": container_id, "section_id": {"$in": variants}}, |
| ) |
|
|
| try: |
| res = self.col.get(where=where, include=["documents", "metadatas"]) |
| except Exception as exc: |
| logger.debug("explicit section lookup failed", exc_info=exc) |
| continue |
|
|
| formatted = self._format_get(res, retrieval_kind=retrieval_kind, score=1.0) |
| formatted.sort(key=lambda x: x.get("chunk_index", 0)) |
| out.extend(formatted[:max_chunks_per_section]) |
|
|
| return out |
|
|
| def _definition_search( |
| self, |
| *, |
| question: str, |
| base_where: Optional[Dict[str, Any]], |
| top_k: int, |
| expanded_terms: Optional[List[str]] = None, |
| ) -> List[Dict[str, Any]]: |
| terms = self._extract_definition_terms(question) |
| terms.extend(expanded_terms or []) |
| if not terms: |
| terms = self._content_terms(question)[:5] |
|
|
| terms = list(dict.fromkeys([term for term in terms if term])) |
| candidates: List[Dict[str, Any]] = [] |
|
|
| |
| where = self._where_and(self._build_where(base_where), {"is_definition": True}) |
| try: |
| res = self.col.get(where=where, include=["documents", "metadatas"], limit=self.lexical_scan_limit) |
| candidates.extend(self._format_get(res, retrieval_kind="definition", score=0.85)) |
| except Exception as exc: |
| logger.debug("definition metadata lookup failed", exc_info=exc) |
|
|
| |
| if not candidates: |
| containers = [self.default_container_id] |
| candidates.extend( |
| self._get_explicit_sections( |
| sections=["§ 2"], |
| container_ids=containers, |
| base_where=base_where, |
| max_chunks_per_section=80, |
| retrieval_kind="definition", |
| ) |
| ) |
|
|
| scored: List[Dict[str, Any]] = [] |
| for hit in candidates: |
| score = self._definition_score(hit, terms) |
| if score <= 0: |
| continue |
| new_hit = dict(hit) |
| new_hit["score"] = round(min(score, 1.0), 4) |
| new_hit["rank_score"] = new_hit["score"] |
| kinds = set(new_hit.get("retrieval_kinds", [])) |
| kinds.add("definition") |
| new_hit["retrieval_kinds"] = sorted(kinds) |
| scored.append(new_hit) |
|
|
| scored.sort(key=lambda h: (h.get("rank_score", 0.0), h.get("score", 0.0)), reverse=True) |
| return scored[:top_k] |
|
|
| def _lexical_search( |
| self, |
| *, |
| question: str, |
| base_where: Optional[Dict[str, Any]], |
| top_k: int, |
| expanded_terms: Optional[List[str]] = None, |
| scan_limit: int, |
| ) -> List[Dict[str, Any]]: |
| tokens = self._content_terms(question) |
| phrases = list(expanded_terms or []) |
| for term in self._extract_definition_terms(question): |
| if term not in phrases: |
| phrases.append(term) |
|
|
| if not tokens and not phrases: |
| return [] |
|
|
| try: |
| res = self.col.get( |
| where=self._build_where(base_where), |
| include=["documents", "metadatas"], |
| limit=max(scan_limit, top_k), |
| ) |
| except Exception as exc: |
| logger.debug("lexical scan failed", exc_info=exc) |
| return [] |
|
|
| candidates = self._format_get(res, retrieval_kind="lexical", score=0.0) |
| scored: List[Dict[str, Any]] = [] |
| for hit in candidates: |
| score = self._lexical_score(hit, tokens=tokens, phrases=phrases) |
| if score <= 0: |
| continue |
| new_hit = dict(hit) |
| new_hit["score"] = round(score, 4) |
| new_hit["rank_score"] = round(score, 4) |
| new_hit["retrieval_kinds"] = sorted(set(new_hit.get("retrieval_kinds", [])) | {"lexical"}) |
| scored.append(new_hit) |
|
|
| scored.sort( |
| key=lambda hit: ( |
| hit.get("rank_score", 0.0), |
| hit.get("score", 0.0), |
| hit.get("metadata", {}).get("is_definition") is True, |
| ), |
| reverse=True, |
| ) |
| return scored[:top_k] |
|
|
| def _get_neighbors( |
| self, |
| hit: Dict[str, Any], |
| *, |
| window: int = 1, |
| base_where: Optional[Dict[str, Any]] = None, |
| ) -> List[Dict[str, Any]]: |
| container = hit.get("container") |
| section = hit.get("section") |
| chunk_index = hit.get("chunk_index") |
|
|
| if container is None or section is None or chunk_index is None: |
| return [] |
|
|
| try: |
| idx = int(chunk_index) |
| except (TypeError, ValueError): |
| return [] |
|
|
| neighbor_indices = [i for i in range(idx - window, idx + window + 1) if i >= 0 and i != idx] |
| if not neighbor_indices: |
| return [] |
|
|
| where = self._where_and( |
| self._build_where(base_where), |
| { |
| "container_id": container, |
| "section_id": section, |
| "chunk_index_in_section": {"$in": neighbor_indices}, |
| }, |
| ) |
|
|
| try: |
| res = self.col.get(where=where, include=["documents", "metadatas"]) |
| except Exception as exc: |
| logger.debug("neighbor lookup failed", exc_info=exc) |
| return [] |
|
|
| formatted = self._format_get( |
| res, |
| retrieval_kind="neighbor", |
| score=max(float(hit.get("score", 0.0)) - 0.05, 0.0), |
| ) |
| formatted.sort(key=lambda x: x.get("chunk_index", 0)) |
| return formatted |
|
|
| def _expand_parent_context( |
| self, |
| hits: List[Dict[str, Any]], |
| *, |
| base_where: Optional[Dict[str, Any]] = None, |
| max_parent_contexts: int = 8, |
| ) -> List[Dict[str, Any]]: |
| parent_ids: List[str] = [] |
| score_by_parent: Dict[str, float] = defaultdict(float) |
|
|
| for hit in hits: |
| meta = hit.get("metadata") or {} |
| parent_id = meta.get("parent_unit_id") or meta.get("parent_id") |
| legal_id = meta.get("legal_unit_id") |
| chunk_kind = str(meta.get("chunk_kind") or "").lower() |
|
|
| if not parent_id: |
| continue |
| if legal_id and str(parent_id) == str(legal_id): |
| continue |
| if chunk_kind == "parent": |
| continue |
|
|
| parent = str(parent_id) |
| parent_ids.append(parent) |
| score_by_parent[parent] = max(score_by_parent[parent], float(hit.get("rank_score", hit.get("score", 0.0)))) |
|
|
| parent_ids = list(dict.fromkeys(parent_ids))[:max_parent_contexts] |
| out: List[Dict[str, Any]] = [] |
|
|
| for parent_id in parent_ids: |
| where = self._where_and( |
| self._build_where(base_where), |
| {"legal_unit_id": parent_id}, |
| ) |
| try: |
| res = self.col.get(where=where, include=["documents", "metadatas"], limit=4) |
| except Exception as exc: |
| logger.debug("parent context lookup failed", exc_info=exc) |
| continue |
|
|
| score = min(score_by_parent.get(parent_id, 0.75) + 0.05, 1.0) |
| formatted = self._format_get(res, retrieval_kind="parent_context", score=score) |
| for hit in formatted: |
| hit["retrieval_kinds"] = sorted(set(hit.get("retrieval_kinds", [])) | {"parent_context"}) |
| out.extend(formatted) |
|
|
| return out |
|
|
| |
| |
| |
| def _get_reranker(self) -> Any: |
| if self._reranker is None: |
| from sentence_transformers import CrossEncoder |
|
|
| logger.info("loading cross-encoder reranker", extra={"model": self.reranker_model}) |
| self._reranker = CrossEncoder(self.reranker_model, max_length=512) |
| return self._reranker |
|
|
| def _rerank(self, question: str, hits: List[Dict[str, Any]]) -> List[Dict[str, Any]]: |
| """Re-order the fused candidate pool with a cross-encoder. |
| |
| The bi-encoder retrieves candidates cheaply; the cross-encoder reads |
| query and passage together and orders them far more precisely. In the |
| golden evaluation this lifted hybrid retrieval from 17/18 to 18/18 and |
| stabilizes which hit becomes [Quelle 1] for the LLM. |
| """ |
| if not hits: |
| return hits |
|
|
| candidates = hits[: self.reranker_candidates] |
| rest = hits[self.reranker_candidates :] |
| pairs = [(question, str(hit.get("text") or "")) for hit in candidates] |
|
|
| try: |
| scores = self._get_reranker().predict(pairs, batch_size=16, show_progress_bar=False) |
| except Exception as exc: |
| logger.warning("reranker failed; keeping fused order", exc_info=exc) |
| return hits |
|
|
| reranked: List[Dict[str, Any]] = [] |
| for hit, score in zip(candidates, scores): |
| enriched = dict(hit) |
| enriched["rerank_score"] = round(float(score), 4) |
| |
| |
| |
| |
| enriched["rank_score"] = round(1.0 / (1.0 + math.exp(-float(score))), 4) |
| reranked.append(enriched) |
|
|
| reranked.sort(key=lambda h: h["rerank_score"], reverse=True) |
| return reranked + rest |
|
|
| |
| |
| |
| @staticmethod |
| def build_rag_context( |
| results: List[Dict[str, Any]], |
| *, |
| max_chars: int = 12000, |
| include_neighbors: bool = True, |
| include_metadata: bool = True, |
| ) -> str: |
| """Build stable RAG context blocks for the downstream LLM client.""" |
| if not results: |
| return "" |
|
|
| parts: List[str] = [] |
| total = 0 |
|
|
| for i, hit in enumerate(results, start=1): |
| kinds = set(hit.get("retrieval_kinds", [])) |
| if not include_neighbors and kinds == {"neighbor"}: |
| continue |
|
|
| meta = hit.get("metadata") or {} |
| canonical = meta.get("canonical_ref") or hit.get("canonical_ref") or hit.get("section", "ohne Abschnitt") |
| chunk_kind = meta.get("chunk_kind") or hit.get("chunk_kind") or "" |
| unit_type = meta.get("unit_type") or "" |
| definition = "Definition" if meta.get("is_definition") is True else "" |
| details = ", ".join([x for x in [str(chunk_kind), str(unit_type), definition] if x]) |
|
|
| source = ( |
| f"[Quelle {i}: {hit.get('container', 'Unbekannt')}::" |
| f"{canonical}, Seiten {hit.get('page_range', '?')}, " |
| f"Chunk {hit.get('chunk_index', '?')}, " |
| f"Typ {','.join(hit.get('retrieval_kinds', []))}" |
| ) |
| if include_metadata and details: |
| source += f", Metadaten {details}" |
| source += "]" |
|
|
| text = (hit.get("text") or "").strip() |
| block = f"{source}\n{text}" |
|
|
| if max_chars and total + len(block) + 2 > max_chars: |
| break |
|
|
| parts.append(block) |
| total += len(block) + 2 |
|
|
| return "\n\n".join(parts) |
|
|
| @staticmethod |
| def build_sources( |
| results: List[Dict[str, Any]], |
| *, |
| max_sources: int = 5, |
| include_neighbor_only: bool = False, |
| ) -> List[Dict[str, Any]]: |
| """ |
| Build a clean, deduplicated source list. |
| |
| Neighbor-only hits are excluded by default because they are usually context, |
| not the primary legal basis. |
| """ |
| seen: Set[Tuple[Any, Any, Any, Any]] = set() |
| sources: List[Dict[str, Any]] = [] |
|
|
| for hit in results: |
| kinds = set(hit.get("retrieval_kinds", [])) |
| if kinds == {"neighbor"} and not include_neighbor_only: |
| continue |
|
|
| meta = hit.get("metadata") or {} |
| canonical = meta.get("canonical_ref") or hit.get("canonical_ref") or hit.get("section") |
| key = (hit.get("container"), canonical, hit.get("page_range"), meta.get("legal_unit_id")) |
|
|
| if key in seen: |
| continue |
|
|
| seen.add(key) |
|
|
| sources.append( |
| { |
| "container": hit.get("container"), |
| "section": hit.get("section"), |
| "canonical_ref": canonical, |
| "path": hit.get("path"), |
| "page_range": hit.get("page_range"), |
| "page_start": hit.get("page_start"), |
| "page_end": hit.get("page_end"), |
| "score": hit.get("score"), |
| "rank_score": hit.get("rank_score"), |
| "retrieval_kinds": hit.get("retrieval_kinds", []), |
| "section_title": meta.get("section_title"), |
| "chunk_kind": meta.get("chunk_kind"), |
| "unit_type": meta.get("unit_type"), |
| "is_definition": meta.get("is_definition"), |
| "defined_terms": meta.get("defined_terms"), |
| } |
| ) |
|
|
| if len(sources) >= max_sources: |
| break |
|
|
| return sources |
|
|
| @staticmethod |
| def format_sources_markdown(sources: List[Dict[str, Any]], *, title: str = "Quellen") -> str: |
| if not sources: |
| return f"{title}: Keine Quellen gefunden." |
|
|
| lines = [f"{title}:"] |
| for source in sources: |
| container = source.get("container") or "Unbekannt" |
| section = source.get("section") or "ohne Abschnitt" |
| pages = source.get("page_range") or "?" |
| canonical = source.get("canonical_ref") |
| label = canonical or section |
| if canonical and canonical != section: |
| lines.append(f"- {container}::{section} ({canonical}), Seiten {pages}") |
| else: |
| lines.append(f"- {container}::{label}, Seiten {pages}") |
|
|
| return "\n".join(lines) |
|
|
| |
| |
| |
| @staticmethod |
| def _format(res: Dict[str, Any], *, retrieval_kind: str) -> List[Dict[str, Any]]: |
| docs = (res.get("documents") or [[]])[0] |
| metas = (res.get("metadatas") or [[]])[0] |
| dists = (res.get("distances") or [[]])[0] |
|
|
| out: List[Dict[str, Any]] = [] |
| for doc, meta, dist in zip(docs, metas, dists): |
| similarity = LegalRetriever._similarity_from_distance(dist) |
| out.append( |
| LegalRetriever._normalize_hit( |
| doc=doc, |
| meta=meta or {}, |
| score=similarity, |
| retrieval_kind=retrieval_kind, |
| ) |
| ) |
| return out |
|
|
| @staticmethod |
| def _format_get(res: Dict[str, Any], *, retrieval_kind: str, score: float) -> List[Dict[str, Any]]: |
| docs = res.get("documents") or [] |
| metas = res.get("metadatas") or [] |
|
|
| out: List[Dict[str, Any]] = [] |
| for doc, meta in zip(docs, metas): |
| out.append( |
| LegalRetriever._normalize_hit( |
| doc=doc, |
| meta=meta or {}, |
| score=score, |
| retrieval_kind=retrieval_kind, |
| ) |
| ) |
| return out |
|
|
| @staticmethod |
| def _similarity_from_distance(dist: Any) -> float: |
| """ |
| For cosine space in Chroma, distance is typically 1 - cosine_similarity. |
| For other metrics this remains only an approximate rank signal. |
| """ |
| try: |
| similarity = 1.0 - float(dist) |
| except (TypeError, ValueError): |
| similarity = 0.0 |
|
|
| return round(max(min(similarity, 1.0), -1.0), 4) |
|
|
| @staticmethod |
| def _normalize_hit(*, doc: str, meta: Dict[str, Any], score: float, retrieval_kind: str) -> Dict[str, Any]: |
| page_start = meta.get("page_start") |
| page_end = meta.get("page_end", page_start) |
| page_range = "?" |
|
|
| if page_start is not None and page_end is not None: |
| page_range = f"{page_start}–{page_end}" |
| elif page_start is not None: |
| page_range = str(page_start) |
|
|
| chunk_index = meta.get("chunk_index_in_section") |
| try: |
| chunk_index = int(chunk_index) |
| except (TypeError, ValueError): |
| chunk_index = 0 |
|
|
| canonical_ref = meta.get("canonical_ref") or LegalRetriever._canonical_from_metadata(meta) or meta.get("section_id") |
|
|
| hit = { |
| "score": round(float(score), 4), |
| "rank_score": round(float(score), 4), |
| "text": doc or "", |
| "container": meta.get("container_id", "Unbekannt"), |
| "container_type": meta.get("container_type"), |
| "section": meta.get("section_id") or meta.get("paragraph") or "ohne Abschnitt", |
| "canonical_ref": canonical_ref, |
| "path": meta.get("section_path", ""), |
| "page_range": page_range, |
| "page_start": page_start, |
| "page_end": page_end, |
| "chunk_index": chunk_index, |
| "retrieval_kinds": [retrieval_kind], |
| "metadata": dict(meta), |
| } |
| hit["source_key"] = LegalRetriever._source_key(hit) |
| return hit |
|
|
| @staticmethod |
| def _canonical_from_metadata(meta: Dict[str, Any]) -> str | None: |
| paragraph = meta.get("paragraph") or meta.get("section_id") |
| if not paragraph: |
| return None |
| parts = [str(paragraph)] |
| if meta.get("subsection"): |
| sub = str(meta["subsection"]) |
| parts.append(sub if sub.lower().startswith("abs") else f"Abs. {sub}") |
| if meta.get("sentence"): |
| sent = str(meta["sentence"]) |
| parts.append(sent if sent.lower().startswith("satz") else f"Satz {sent}") |
| if meta.get("number"): |
| num = str(meta["number"]) |
| parts.append(num if num.lower().startswith(("nr", "nummer")) else f"Nr. {num}") |
| if meta.get("letter"): |
| letter = str(meta["letter"]).lower().replace(")", "") |
| parts.append(letter if letter.lower().startswith("buchst") else f"Buchst. {letter}") |
| return " ".join(parts) |
|
|
| @staticmethod |
| def _source_key(hit: Dict[str, Any]) -> Tuple[Any, Any, Any, Any]: |
| meta = hit.get("metadata") or {} |
| legal_unit_id = meta.get("legal_unit_id") |
| if legal_unit_id: |
| return (hit.get("container"), legal_unit_id, meta.get("chunk_kind"), meta.get("text_hash")) |
|
|
| text_hash = meta.get("text_hash") |
| if text_hash: |
| return (hit.get("container"), hit.get("section"), hit.get("chunk_index"), text_hash) |
|
|
| return (hit.get("container"), hit.get("section"), hit.get("chunk_index"), (hit.get("text") or "")[:120]) |
|
|
| @staticmethod |
| def _dedupe_and_rank( |
| results: List[Dict[str, Any]], |
| *, |
| question: str | None = None, |
| intent: str | None = None, |
| ) -> List[Dict[str, Any]]: |
| """Deduplicate hits and rank legal-specific retrieval kinds above generic context.""" |
| merged: Dict[Tuple[Any, Any, Any, Any], Dict[str, Any]] = {} |
|
|
| for hit in results: |
| key = hit.get("source_key") or LegalRetriever._source_key(hit) |
| existing = merged.get(key) |
|
|
| if existing is None: |
| merged[key] = dict(hit) |
| continue |
|
|
| existing["score"] = max(float(existing.get("score", 0.0)), float(hit.get("score", 0.0))) |
| existing["rank_score"] = max(float(existing.get("rank_score", 0.0)), float(hit.get("rank_score", 0.0))) |
|
|
| kinds: Set[str] = set(existing.get("retrieval_kinds", [])) |
| kinds.update(hit.get("retrieval_kinds", [])) |
| existing["retrieval_kinds"] = sorted(kinds) |
|
|
| ranked = list(merged.values()) |
|
|
| target_refs = LegalRetriever._parse_norm_refs(question or "") |
| target_terms = LegalRetriever._extract_definition_terms(question or "") if question else [] |
|
|
| for hit in ranked: |
| meta = hit.get("metadata") or {} |
| kinds = set(hit.get("retrieval_kinds", [])) |
| boost = 0.0 |
|
|
| if "exact_norm" in kinds: |
| boost += 0.16 |
| if "explicit_section" in kinds or "section_lookup" in kinds: |
| boost += 0.10 |
| if "definition" in kinds: |
| boost += 0.13 |
| if "parent_context" in kinds: |
| boost += 0.09 |
| if "lexical" in kinds: |
| boost += 0.06 |
| if "semantic" in kinds: |
| boost += 0.04 |
| if kinds == {"neighbor"}: |
| boost -= 0.05 |
|
|
| if hit.get("container") == "Vertrag": |
| boost += 0.01 |
|
|
| if meta.get("chunk_kind") == "parent": |
| boost += 0.03 |
| if meta.get("is_definition") is True and intent == "definition": |
| boost += 0.10 |
|
|
| canonical = str(meta.get("canonical_ref") or hit.get("canonical_ref") or "").lower() |
| for ref in target_refs: |
| if ref.canonical_ref.lower() in canonical or ref.section_id.lower() in canonical: |
| boost += 0.12 |
| break |
|
|
| meta_text = LegalRetriever._metadata_text(meta).lower() |
| for term in target_terms: |
| if term.lower() in meta_text: |
| boost += 0.08 |
| break |
|
|
| hit["rank_score"] = round(float(hit.get("score", 0.0)) + boost, 4) |
|
|
| ranked.sort( |
| key=lambda x: ( |
| x.get("rank_score", 0.0), |
| x.get("score", 0.0), |
| 1 if (x.get("metadata") or {}).get("chunk_kind") == "parent" else 0, |
| -int(x.get("chunk_index", 0)), |
| ), |
| reverse=True, |
| ) |
|
|
| return ranked |
|
|
| @staticmethod |
| def _limit_chunks_per_section(results: List[Dict[str, Any]], *, max_chunks_per_section: int) -> List[Dict[str, Any]]: |
| if max_chunks_per_section <= 0: |
| return results |
|
|
| counts: Dict[Tuple[Any, Any], int] = defaultdict(int) |
| limited: List[Dict[str, Any]] = [] |
|
|
| for hit in results: |
| meta = hit.get("metadata") or {} |
| |
| key = (hit.get("container"), hit.get("section"), meta.get("chunk_kind")) |
| if counts[key] >= max_chunks_per_section: |
| continue |
| counts[key] += 1 |
| limited.append(hit) |
|
|
| return limited |
|
|
| |
| |
| |
| @staticmethod |
| def _hit_matches_norm_ref(hit: Dict[str, Any], ref: NormReference) -> bool: |
| meta = hit.get("metadata") or {} |
| haystack = " ".join( |
| [ |
| str(meta.get("canonical_ref") or ""), |
| str(hit.get("canonical_ref") or ""), |
| str(meta.get("paragraph") or ""), |
| str(meta.get("subsection") or ""), |
| str(meta.get("sentence") or ""), |
| str(meta.get("number") or ""), |
| str(meta.get("letter") or ""), |
| hit.get("text") or "", |
| ] |
| ).lower() |
|
|
| if ref.section_id.lower() not in haystack and ref.section_id.replace(" ", "").lower() not in haystack.replace(" ", ""): |
| return False |
| if ref.subsection and f"abs. {ref.subsection}".lower() not in haystack and f"({ref.subsection})" not in haystack: |
| return False |
| if ref.sentence and f"satz {ref.sentence}".lower() not in haystack: |
| return False |
| if ref.number and f"nr. {ref.number}".lower() not in haystack and f"nummer {ref.number}".lower() not in haystack: |
| return False |
| if ref.letter and f"buchst. {ref.letter}".lower() not in haystack and f"{ref.letter})" not in haystack: |
| return False |
| return True |
|
|
| @staticmethod |
| def _exact_norm_sort_key(hit: Dict[str, Any], ref: NormReference) -> Tuple[int, int, float, int]: |
| meta = hit.get("metadata") or {} |
| canonical = str(meta.get("canonical_ref") or hit.get("canonical_ref") or "").lower() |
| exact = 1 if ref.canonical_ref.lower() in canonical else 0 |
| parent = 1 if str(meta.get("chunk_kind") or "").lower() == "parent" else 0 |
| score = float(hit.get("score", 0.0)) |
| |
| idx = -int(hit.get("chunk_index", 0)) |
| return exact, parent, score, idx |
|
|
| @classmethod |
| def _definition_score(cls, hit: Dict[str, Any], terms: List[str]) -> float: |
| meta = hit.get("metadata") or {} |
| text = f"{hit.get('text') or ''}\n{cls._metadata_text(meta)}" |
| haystack = cls._norm_text(text) |
| score = 0.0 |
|
|
| if meta.get("is_definition") is True: |
| score += 0.45 |
| if str(meta.get("unit_type") or "").lower() == "definition": |
| score += 0.35 |
| if str(meta.get("section_id") or meta.get("paragraph") or "").replace(" ", "") in {"§2", "2"}: |
| score += 0.12 |
|
|
| for term in terms: |
| t = cls._norm_text(term) |
| if not t: |
| continue |
| if t in haystack: |
| score += 0.35 |
| else: |
| token_hits = sum(1 for token in cls._tokenize(t) if token in haystack) |
| if token_hits: |
| score += min(0.18, 0.06 * token_hits) |
|
|
| return min(score, 1.0) |
|
|
| @classmethod |
| def _lexical_score(cls, hit: Dict[str, Any], *, tokens: List[str], phrases: List[str]) -> float: |
| meta = hit.get("metadata") or {} |
| haystack_raw = f"{hit.get('text') or ''}\n{cls._metadata_text(meta)}" |
| haystack = cls._norm_text(haystack_raw) |
| if not haystack: |
| return 0.0 |
|
|
| raw = 0.0 |
| token_counts = Counter(cls._tokenize(haystack)) |
| for token in tokens: |
| freq = token_counts.get(token, 0) |
| if freq: |
| raw += 1.0 + min(freq - 1, 3) * 0.2 |
|
|
| for phrase in phrases: |
| p = cls._norm_text(phrase) |
| if p and p in haystack: |
| raw += 3.0 if " " in p else 1.4 |
|
|
| |
| if meta.get("is_definition") is True: |
| raw += 0.5 |
| if meta.get("canonical_ref") and any(token in cls._norm_text(str(meta.get("canonical_ref"))) for token in tokens): |
| raw += 0.5 |
|
|
| denom = max(len(tokens) + len(phrases) * 1.5, 4.0) |
| score = raw / denom |
| return round(min(score, 0.99), 4) |
|
|
| @classmethod |
| def _content_terms(cls, text: str) -> List[str]: |
| tokens = cls._tokenize(cls._norm_text(text)) |
| return list(dict.fromkeys([t for t in tokens if len(t) >= 3 and t not in cls.STOPWORDS])) |
|
|
| @staticmethod |
| def _metadata_text(meta: Dict[str, Any]) -> str: |
| keys = [ |
| "canonical_ref", |
| "paragraph", |
| "subsection", |
| "sentence", |
| "number", |
| "letter", |
| "unit_type", |
| "defined_terms", |
| "section_title", |
| "section_id", |
| "section_path", |
| "container_id", |
| ] |
| values: List[str] = [] |
| for key in keys: |
| value = meta.get(key) |
| if value is None: |
| continue |
| if isinstance(value, (list, tuple, set)): |
| values.extend(str(v) for v in value) |
| else: |
| values.append(str(value)) |
| return " ".join(values) |
|
|
| @staticmethod |
| def _looks_like_negative_risk(question: str, results: List[Dict[str, Any]]) -> bool: |
| if not results: |
| return True |
| q = LegalRetriever._norm_text(question) |
| if re.search(r"\b(nicht\s+geregelt|keine\s+regelung|steht\s+nicht|nicht\s+enthalten|nicht\s+gefunden)\b", q): |
| return True |
| strongest = max(float(hit.get("rank_score", hit.get("score", 0.0))) for hit in results) |
| return strongest < 0.35 |
|
|
| @staticmethod |
| def _norm_text(text: str) -> str: |
| s = str(text or "").lower() |
| s = s.replace("§§", "§") |
| s = re.sub(r"[\u00a0\t\r\n]+", " ", s) |
| s = re.sub(r"\s+", " ", s) |
| return s.strip() |
|
|
| @staticmethod |
| def _tokenize(text: str) -> List[str]: |
| return re.findall(r"[a-zäöüß0-9]{2,}", text.lower()) |
|
|
|
|