M1+M2: RAG-Pipeline mit verbindlichem Grounding
agent/-Paket: Ingest (601 Layer-2-Eintraege -> 3005 Chunks, FTS5-BM25 + Vektoren-Cache), Hybrid-Retrieval (RRF, Stand-Boost, cross_ref-Erweiterung), Ollama-Client (embed/chat, think-Flag-Fallback, kurzes Connect-Budget), Systemprompt mit Zitierpflicht, Post-Validierung (zitierte IDs gemaess Retrieved-Set, 1x Regenerierung, dann Verweigerung), FastAPI (/ask, /health, /reindex), CLI, Goldset (31 Fragen, IDs gegen kb.json verifiziert, inkl. ATZ-Konfliktfall + 4 Verweigerungsfaelle), Eval-Suite, Test-Chat. 41 Offline-Tests gruen. Baseline BM25-only: Hit-Rate 0,871 / Recall@8 0,855 / MRR 0,476. Hybrid-Messung, Antwortmodus-Eval und Modell-Bake-off (M3) auf dem Host ausstaendig (Ollama aus der Zed-Sandbox nicht erreichbar). MEMORY.md und planung.md Umsetzungsstand aktualisiert.
This commit is contained in:
@@ -0,0 +1,252 @@
|
||||
"""Hybrid-Retrieval: BM25 (FTS5) + Dense (bge-m3) -> RRF-Fusion,
|
||||
milde Stand-Aktualitätsgewichtung und kontrollierte cross_ref-Erweiterung.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .config import Config
|
||||
from .normalize import fts_query
|
||||
from .ollama_client import OllamaClient
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChunkResult:
|
||||
chunk_id: int
|
||||
entry_id: str
|
||||
section: str
|
||||
text: str
|
||||
title: str
|
||||
stand: str
|
||||
work: str
|
||||
chapter: str
|
||||
topic: str
|
||||
tags: list = field(default_factory=list)
|
||||
legal_bases: list = field(default_factory=list)
|
||||
cross_refs: list = field(default_factory=list)
|
||||
batch: int = 0
|
||||
score: float = 0.0
|
||||
source: str = "fused" # bm25 | dense | fused | cross_ref
|
||||
|
||||
|
||||
class Retriever:
|
||||
def __init__(self, cfg: Config, db_path: str | None = None, client=None):
|
||||
self.cfg = cfg
|
||||
self.db_path = str(db_path or cfg.db_path)
|
||||
if not Path(self.db_path).is_file():
|
||||
raise RuntimeError(
|
||||
f"index fehlt ({self.db_path}) — zuerst 'python -m agent.cli ingest' ausführen"
|
||||
)
|
||||
self._con = sqlite3.connect(self.db_path)
|
||||
self._con.row_factory = sqlite3.Row
|
||||
self._client = client
|
||||
self._owns_client = client is None
|
||||
self._mat: np.ndarray | None = None
|
||||
self._mat_chunk_ids: list[int] | None = None
|
||||
row = self._con.execute("SELECT MIN(stand), MAX(stand) FROM chunks").fetchone()
|
||||
self._stand_min = int((row[0] or "2026-01").replace("-", ""))
|
||||
self._stand_max = int((row[1] or "2026-01").replace("-", ""))
|
||||
|
||||
def close(self) -> None:
|
||||
self._con.close()
|
||||
if self._owns_client and self._client is not None:
|
||||
self._client.close()
|
||||
|
||||
# -- Index-Kennzahlen ---------------------------------------------------
|
||||
|
||||
def stats(self) -> dict:
|
||||
n_chunks = self._con.execute("SELECT COUNT(*) FROM chunks").fetchone()[0]
|
||||
n_entries = self._con.execute(
|
||||
"SELECT COUNT(DISTINCT entry_id) FROM chunks"
|
||||
).fetchone()[0]
|
||||
n_vec = self._con.execute(
|
||||
"SELECT COUNT(*) FROM vectors WHERE model = ?",
|
||||
(self.cfg.embed_model,),
|
||||
).fetchone()[0]
|
||||
meta = dict(self._con.execute("SELECT key, value FROM meta").fetchall())
|
||||
return {
|
||||
"n_entries": n_entries,
|
||||
"n_chunks": n_chunks,
|
||||
"n_vectors": n_vec,
|
||||
"dense_available": n_vec > 0 and not self.cfg.embed_off,
|
||||
"stand_min": str(self._stand_min),
|
||||
"stand_max": str(self._stand_max),
|
||||
"built_at": meta.get("built_at"),
|
||||
}
|
||||
|
||||
# -- Einzelverfahren ----------------------------------------------------
|
||||
|
||||
def _bm25(self, question: str, limit: int) -> dict[int, float]:
|
||||
q = fts_query(question)
|
||||
if not q:
|
||||
return {}
|
||||
rows = self._con.execute(
|
||||
"SELECT rowid, bm25(chunks_fts) AS rank FROM chunks_fts "
|
||||
"WHERE chunks_fts MATCH ? ORDER BY rank LIMIT ?",
|
||||
(q, limit),
|
||||
).fetchall()
|
||||
# bm25(): kleinere Werte = besser -> negieren für "größer = besser"
|
||||
return {r["rowid"]: -float(r["rank"]) for r in rows}
|
||||
|
||||
def _dense(self, question: str, limit: int) -> dict[int, float]:
|
||||
if self.cfg.embed_off:
|
||||
return {}
|
||||
self._ensure_matrix()
|
||||
if self._mat is None or len(self._mat) == 0:
|
||||
return {}
|
||||
if self._client is None:
|
||||
self._client = OllamaClient(
|
||||
self.cfg.ollama_url,
|
||||
embed_timeout_s=self.cfg.embed_timeout_s,
|
||||
chat_timeout_s=self.cfg.chat_timeout_s,
|
||||
)
|
||||
try:
|
||||
qvec = np.asarray(
|
||||
self._client.embed(self.cfg.embed_model, [question])[0],
|
||||
dtype=np.float32,
|
||||
)
|
||||
except Exception:
|
||||
return {} # Ollama nicht erreichbar -> BM25-only weiter
|
||||
qn = np.linalg.norm(qvec)
|
||||
if qn == 0:
|
||||
return {}
|
||||
sims = self._mat_norm @ (qvec / qn)
|
||||
order = np.argsort(-sims)[:limit]
|
||||
return {self._mat_chunk_ids[i]: float(sims[i]) for i in order}
|
||||
|
||||
def _ensure_matrix(self) -> None:
|
||||
if self._mat is not None:
|
||||
return
|
||||
rows = self._con.execute(
|
||||
"SELECT c.chunk_id, v.vec, v.dim FROM chunks c "
|
||||
"JOIN vectors v ON v.content_hash = c.content_hash AND v.model = ?",
|
||||
(self.cfg.embed_model,),
|
||||
).fetchall()
|
||||
if not rows:
|
||||
self._mat = np.zeros((0, 1), dtype=np.float32)
|
||||
self._mat_chunk_ids = []
|
||||
return
|
||||
ids = [r[0] for r in rows]
|
||||
mat = np.vstack(
|
||||
[np.frombuffer(r[1], dtype=np.float32) for r in rows]
|
||||
)
|
||||
norms = np.linalg.norm(mat, axis=1, keepdims=True)
|
||||
self._mat = mat
|
||||
self._mat_norm = mat / np.where(norms == 0, 1.0, norms)
|
||||
self._mat_chunk_ids = ids
|
||||
|
||||
# -- Metadaten & Fusion -------------------------------------------------
|
||||
|
||||
def _stand_factor(self, stand: str) -> float:
|
||||
if self._stand_max <= self._stand_min:
|
||||
return 0.0
|
||||
try:
|
||||
s = int(stand.replace("-", ""))
|
||||
except (ValueError, AttributeError):
|
||||
return 0.0
|
||||
f = (s - self._stand_min) / (self._stand_max - self._stand_min)
|
||||
return min(1.0, max(0.0, f))
|
||||
|
||||
def _fetch_chunks(self, chunk_ids: list[int]) -> dict[int, sqlite3.Row]:
|
||||
out: dict[int, sqlite3.Row] = {}
|
||||
for i in range(0, len(chunk_ids), 500):
|
||||
part = chunk_ids[i:i + 500]
|
||||
qm = ",".join("?" * len(part))
|
||||
for r in self._con.execute(
|
||||
f"SELECT * FROM chunks WHERE chunk_id IN ({qm})", part
|
||||
).fetchall():
|
||||
out[r["chunk_id"]] = r
|
||||
return out
|
||||
|
||||
def _row_to_result(self, row: sqlite3.Row, score: float, source: str) -> ChunkResult:
|
||||
return ChunkResult(
|
||||
chunk_id=row["chunk_id"],
|
||||
entry_id=row["entry_id"],
|
||||
section=row["section"],
|
||||
text=row["text"],
|
||||
title=row["title"],
|
||||
stand=row["stand"],
|
||||
work=row["work"],
|
||||
chapter=row["chapter"],
|
||||
topic=row["topic"],
|
||||
tags=json.loads(row["tags"]),
|
||||
legal_bases=json.loads(row["legal_bases"]),
|
||||
cross_refs=json.loads(row["cross_refs"]),
|
||||
batch=row["batch"],
|
||||
score=score,
|
||||
source=source,
|
||||
)
|
||||
|
||||
def _best_chunk_of_entry(self, entry_id: str) -> ChunkResult | None:
|
||||
rows = self._con.execute(
|
||||
"SELECT * FROM chunks WHERE entry_id = ? "
|
||||
"ORDER BY CASE WHEN section LIKE 'Zusammenfassung%' THEN 0 ELSE 1 END, "
|
||||
"chunk_id LIMIT 1",
|
||||
(entry_id,),
|
||||
).fetchall()
|
||||
if not rows:
|
||||
return None
|
||||
return self._row_to_result(rows[0], 0.0, "cross_ref")
|
||||
|
||||
# -- öffentliche Suche --------------------------------------------------
|
||||
|
||||
def search(self, question: str, n_entries: int | None = None) -> list[ChunkResult]:
|
||||
"""Liefert die Top-Kontextblöcke (Hauptretrieval + cross_ref-Erweiterung)."""
|
||||
n = n_entries or self.cfg.context_blocks
|
||||
pool = self.cfg.candidate_pool
|
||||
bm = self._bm25(question, pool)
|
||||
try:
|
||||
dn = self._dense(question, pool)
|
||||
except Exception:
|
||||
dn = {}
|
||||
fused: dict[int, float] = {}
|
||||
for ranking in (bm, dn):
|
||||
ordered = sorted(ranking.items(), key=lambda kv: -kv[1])
|
||||
for rank, (cid, _) in enumerate(ordered):
|
||||
fused[cid] = fused.get(cid, 0.0) + 1.0 / (self.cfg.rrf_k + rank)
|
||||
if not fused:
|
||||
return []
|
||||
|
||||
rows = self._fetch_chunks(list(fused))
|
||||
results: list[ChunkResult] = []
|
||||
for cid, score in fused.items():
|
||||
if cid not in rows:
|
||||
continue
|
||||
source = "fused" if (cid in bm and cid in dn) else (
|
||||
"bm25" if cid in bm else "dense"
|
||||
)
|
||||
score += self.cfg.recency_boost * self._stand_factor(rows[cid]["stand"])
|
||||
results.append(self._row_to_result(rows[cid], score, source))
|
||||
results.sort(key=lambda r: -r.score)
|
||||
|
||||
# Bester Chunk je Eintrag -> Kontext (Entry-Level-Dedup)
|
||||
main: list[ChunkResult] = []
|
||||
seen: set[str] = set()
|
||||
for r in results:
|
||||
if r.entry_id in seen:
|
||||
continue
|
||||
seen.add(r.entry_id)
|
||||
main.append(r)
|
||||
if len(main) >= n:
|
||||
break
|
||||
|
||||
# cross_ref-Erweiterung (kontrolliert, markiert, begrenzt)
|
||||
extra: list[ChunkResult] = []
|
||||
budget = self.cfg.cross_ref_max_extra
|
||||
for r in main[: self.cfg.cross_ref_expand]:
|
||||
for ref in r.cross_refs:
|
||||
if budget <= 0:
|
||||
break
|
||||
if ref in seen:
|
||||
continue
|
||||
er = self._best_chunk_of_entry(ref)
|
||||
if er is not None:
|
||||
extra.append(er)
|
||||
seen.add(ref)
|
||||
budget -= 1
|
||||
return main + extra
|
||||
Reference in New Issue
Block a user