harrier-semantic-v1 / patch_mcp_server.py
ArneH's picture
Fix: hybrid search + full 962k corpus + semantic fallback
67cff6a verified
Raw
History Blame Contribute Delete
11.3 kB
"""
Swiss Legal β€” MCP Server Patch
================================
Ersetzt Hertner's stock search_fts5 mit einem hybriden System:
1. Harness v21 (FTS5) β€” MRR ~0.87 auf Benchmark, sehr schnell
2. harrier-semantic-v1 β€” Semantisches Fallback fΓΌr konzeptuelle &
cross-linguale Queries (DE/FR/IT)
3. RRF-Kombination β€” Beide Signale verschmelzen zu einem Score
Resultat: Kein Overfitting-Risiko, volle Corpus-Abdeckung, cross-lingual.
Verwendung:
# Statt mcp_server.py direkt:
python3 patch_mcp_server.py
# Claude Desktop config:
{
"mcpServers": {
"swiss-caselaw": {
"command": "/path/to/.venv/bin/python3",
"args": ["/path/to/patch_mcp_server.py"]
}
}
}
"""
from __future__ import annotations
import sys, os, logging, math
from pathlib import Path
log = logging.getLogger("harness-patch")
logging.basicConfig(level=logging.INFO, stream=sys.stderr,
format="%(asctime)s %(levelname)s %(message)s")
# ── Locate caselaw-repo-1 ──────────────────────────────────────────────────────
_REPO_CANDIDATES = [
Path.home() / "caselaw-repo-1",
Path.home() / "swiss-legal" / "caselaw-repo-1",
Path(__file__).parent,
Path(__file__).parent.parent,
Path("/root/caselaw-repo-1"),
]
REPO_DIR = next((p for p in _REPO_CANDIDATES if (p / "mcp_server.py").exists()), None)
if REPO_DIR is None:
sys.exit("ERROR: mcp_server.py nicht gefunden. git clone https://github.com/jonashertner/caselaw-repo-1")
sys.path.insert(0, str(REPO_DIR))
log.info(f"Repo: {REPO_DIR}")
# ── Locate harness_v21.py ──────────────────────────────────────────────────────
_HARNESS_CANDIDATES = [
Path(__file__).parent / "harness_v21.py",
Path.home() / "swiss-legal" / "patch" / "harness_v21.py",
REPO_DIR / "harness_v21.py",
REPO_DIR / "harnesses" / "harness_no_inject_021.py",
]
HARNESS_PATH = next((p for p in _HARNESS_CANDIDATES if p.exists()), None)
if HARNESS_PATH is None:
sys.exit("ERROR: harness_v21.py nicht gefunden. Download von ArneH/harrier-semantic-v1 (HuggingFace).")
import importlib.util
spec = importlib.util.spec_from_file_location("harness_v21", str(HARNESS_PATH))
harness = importlib.util.module_from_spec(spec)
spec.loader.exec_module(harness)
log.info(f"Harness v21: {HARNESS_PATH}")
# ── Semantic engine (harrier-semantic-v1) ──────────────────────────────────────
DATA_DIR = Path(os.environ.get("SWISS_CASELAW_DIR", Path.home() / ".swiss-caselaw"))
SEMANTIC_DIR = DATA_DIR / "semantic"
MODEL_DIR = SEMANTIC_DIR / "harrier-semantic-v1"
EMB_DIR = SEMANTIC_DIR / "embeddings"
HF_MODEL_REPO = "ArneH/harrier-semantic-v1"
HF_EMBS_REPO = "ArneH/swiss-caselaw-embeddings"
import numpy as np
_sem_model = None
_corpus_ids = []
_corpus_mat = None
_sem_ready = False
def _get_device():
try:
import torch
if torch.cuda.is_available(): return "cuda"
if torch.backends.mps.is_available(): return "mps"
except Exception: pass
return "cpu"
def _load_semantic_engine():
global _sem_model, _corpus_ids, _corpus_mat, _sem_ready
if _sem_ready:
return True
# Model
try:
from sentence_transformers import SentenceTransformer
device = _get_device()
if MODEL_DIR.exists() and (MODEL_DIR / "model.safetensors").exists():
_sem_model = SentenceTransformer(str(MODEL_DIR), device=device)
else:
log.info(f"Downloading harrier-semantic-v1 from HuggingFace...")
MODEL_DIR.mkdir(parents=True, exist_ok=True)
_sem_model = SentenceTransformer(HF_MODEL_REPO, device=device,
cache_folder=str(SEMANTIC_DIR))
log.info(f"Semantic model loaded on {device}")
except Exception as e:
log.warning(f"Semantic model unavailable: {e}")
return False
# Embeddings
npz_files = sorted(EMB_DIR.glob("*.npz"))
if not npz_files:
log.info(f"Downloading corpus embeddings from HuggingFace (~740MB)...")
try:
from huggingface_hub import snapshot_download
EMB_DIR.mkdir(parents=True, exist_ok=True)
snapshot_download(repo_id=HF_EMBS_REPO, repo_type="dataset",
local_dir=str(EMB_DIR),
allow_patterns=["harrier_semantic_v1_*.npz"])
npz_files = sorted(EMB_DIR.glob("*.npz"))
except Exception as e:
log.warning(f"Embeddings download failed: {e}")
return False
log.info(f"Loading {len(npz_files)} embedding shards...")
ids, mats = [], []
for f in npz_files:
d = np.load(f)
ids.extend(d["ids"].tolist())
mats.append(d["embeddings"].astype(np.float32))
_corpus_ids = ids
mat = np.concatenate(mats, axis=0)
norms = np.linalg.norm(mat, axis=1, keepdims=True)
_corpus_mat = mat / np.where(norms < 1e-8, 1e-8, norms)
_sem_ready = True
log.info(f"Semantic engine ready: {len(_corpus_ids):,} vectors")
return True
def _semantic_search(query: str, top_k: int = 50) -> list[tuple[str, float]]:
"""Returns list of (decision_id, cosine_score)."""
if not _load_semantic_engine():
return []
try:
q_emb = _sem_model.encode([query], normalize_embeddings=True, show_progress_bar=False)
scores = (q_emb @ _corpus_mat.T)[0]
top_idx = np.argsort(-scores)[:top_k]
return [(_corpus_ids[i], float(scores[i])) for i in top_idx]
except Exception as e:
log.warning(f"Semantic search error: {e}")
return []
# ── Import and patch mcp_server ────────────────────────────────────────────────
import mcp_server
_original_search_fts5 = mcp_server.search_fts5
def _hybrid_search(query: str, limit: int, filters: dict) -> tuple[list[dict], int]:
"""
Hybrid search: harness_v21 FTS5 + harrier-semantic-v1, combined via RRF.
Falls back to original search_fts5 for pure filter queries.
"""
RRF_K = 60
# 1. Harness v21 (FTS5)
try:
fts_raw = harness.search(query, k=limit * 4)
except Exception as e:
log.warning(f"Harness search failed: {e}")
fts_raw = []
# 2. Semantic search (harrier-semantic-v1)
sem_raw = _semantic_search(query, top_k=limit * 3)
# 3. RRF fusion
rrf: dict[str, float] = {}
for rank, r in enumerate(fts_raw, 1):
did = r["decision_id"]
rrf[did] = rrf.get(did, 0.0) + 0.7 / (RRF_K + rank) # FTS5 weight: 0.7
for rank, (did, _) in enumerate(sem_raw, 1):
rrf[did] = rrf.get(did, 0.0) + 0.3 / (RRF_K + rank) # Semantic weight: 0.3
# If semantic has results but FTS5 doesn't β†’ boost semantic weight
if not fts_raw and sem_raw:
rrf = {}
for rank, (did, score) in enumerate(sem_raw, 1):
rrf[did] = 1.0 / (RRF_K + rank)
sorted_ids = sorted(rrf, key=lambda x: -rrf[x])
# 4. Apply filters post-hoc
if any(filters.values()):
try:
db = mcp_server.get_db()
id_list = ",".join(f"'{i.replace(chr(39), '')}'" for i in sorted_ids[:500])
clauses, params = [], []
for col in ("court", "canton", "language"):
if filters.get(col):
clauses.append(f"{col} = ?"); params.append(filters[col])
if filters.get("date_from"):
clauses.append("decision_date >= ?"); params.append(filters["date_from"])
if filters.get("date_to"):
clauses.append("decision_date <= ?"); params.append(filters["date_to"])
where = " AND ".join(clauses)
allowed = {
row[0] for row in db.execute(
f"SELECT decision_id FROM decisions WHERE decision_id IN ({id_list}) AND {where}",
params
).fetchall()
}
sorted_ids = [d for d in sorted_ids if d in allowed]
except Exception as e:
log.warning(f"Filter error: {e}")
total = len(sorted_ids)
page_ids = sorted_ids[:limit]
if not page_ids:
return [], 0
# 5. Fetch full rows from DB
try:
db = mcp_server.get_db()
id_list = ",".join(f"'{i.replace(chr(39), '')}'" for i in page_ids)
rows = {
r["decision_id"]: dict(r)
for r in db.execute(
f"SELECT * FROM decisions WHERE decision_id IN ({id_list})"
).fetchall()
}
except Exception as e:
log.warning(f"DB fetch error: {e}")
rows = {}
results = []
for did in page_ids:
row = rows.get(did)
if not row:
continue
row["relevance_score"] = rrf.get(did, 0.0)
row["snippet"] = (row.get("regeste") or row.get("title") or "")[:400]
row["citation_count"] = row.get("citation_count", 0) or 0
results.append(row)
return results, total
def _patched_search_fts5(query: str = "", limit: int = 50,
court=None, canton=None, language=None,
date_from=None, date_to=None,
chamber=None, decision_type=None,
legal_area=None, offset: int = 0, sort=None,
**kwargs):
q = (query or "").strip()
filters = dict(court=court, canton=canton, language=language,
date_from=date_from, date_to=date_to)
# Pure filter query (no text) β†’ use original
if not q:
return _original_search_fts5(
query=query, limit=limit, court=court, canton=canton,
language=language, date_from=date_from, date_to=date_to,
chamber=chamber, decision_type=decision_type,
legal_area=legal_area, offset=offset, sort=sort, **kwargs
)
results, total = _hybrid_search(q, limit=limit + offset, filters=filters)
return results[offset:offset + limit], total
mcp_server.search_fts5 = _patched_search_fts5
log.info("βœ“ search_fts5 β†’ Hybrid (Harness v21 FTS5 + harrier-semantic-v1)")
# Pre-load semantic engine in background
import threading
threading.Thread(target=_load_semantic_engine, daemon=True).start()
# ── Run MCP server ─────────────────────────────────────────────────────────────
if __name__ == "__main__":
import asyncio
if "--remote" in sys.argv:
mcp_server.REMOTE_MODE = True
host, port = "0.0.0.0", 8000
for i, arg in enumerate(sys.argv):
if arg == "--host" and i + 1 < len(sys.argv): host = sys.argv[i + 1]
if arg == "--port" and i + 1 < len(sys.argv): port = int(sys.argv[i + 1])
mcp_server.main_remote(host, port)
else:
asyncio.run(mcp_server.main_stdio())