| """Moonley agentic TOOL REGISTRY. Loads the corpus once (lean — no BM25) and exposes the tools the |
| ReAct controller can call. Each tool returns a list of compact case dicts (doc_id + the fields the |
| LLM needs to reason) or a small structured result. Most tools are ports of serve.py primitives; |
| the statute tools are new (see themis-statute-layer). keyword_search(BM25) is intentionally omitted |
| here — it is 68s/query on the Mac; it runs on the GPU box in production. |
| |
| Usage: from tools import Corpus ; C = Corpus(DATA, STATUTE_DIR) ; C.vector_search("...", k=8) |
| """ |
| import json, os, re, difflib |
| import numpy as np |
| from collections import Counter, defaultdict |
| from sentence_transformers import SentenceTransformer, CrossEncoder |
| from case_summary import case_summary_record, load_case_summaries |
| from statute_crosswalk import load_default_crosswalk |
| from statute_library import ExactStatuteLibrary |
|
|
| BGE_Q = "Represent this sentence for searching relevant passages: " |
| _NAME_STOP = {"v","vs","of","and","the","ors","anr","etc","state","union","govt","government","in","re", |
| "others","another","ltd","co","pvt","dead","thr","lrs","alias","through","etc","anrs","ms", |
| "shri","smt","sri","mr","mrs","dr","justice","sh","kum","mohd"} |
| BAD_STATUS = {"overruled", "per_incuriam", "doubted"} |
|
|
| def _clean(s): |
| if not s: return "" |
| return re.sub(r"\s+", " ", s).strip() |
|
|
| def _ntok(s): |
| """Tokenize a case name, MERGING runs of single letters so abbreviations match: U.P.->up, A.K.->ak.""" |
| out = []; buf = "" |
| for t in re.findall(r"[a-z]+", (s or "").lower()): |
| if len(t) == 1: buf += t |
| else: |
| if buf: out.append(buf); buf = "" |
| out.append(t) |
| if buf: out.append(buf) |
| return out |
|
|
| class Corpus: |
| def __init__(self, data_dir, statute_dir, device="cpu"): |
| self.device = device |
| print("[tools] loading corpus ...", flush=True) |
| self.texts = []; self.chunk_doc = [] |
| with open(os.path.join(data_dir, "escr_chunks.jsonl"), encoding="utf-8") as f: |
| for l in f: |
| c = json.loads(l); self.texts.append(c["text"]); self.chunk_doc.append(c["doc_id"]) |
| self.M = np.load(os.path.join(data_dir, "escr_vectors.npy"), mmap_mode="r") |
| self.chunk_doc_arr = np.array(self.chunk_doc) |
| self.doc_chunks = defaultdict(list) |
| for i, d in enumerate(self.chunk_doc): self.doc_chunks[d].append(i) |
| self.meta = {} |
| for l in open(os.path.join(data_dir, "escr_meta.jsonl"), encoding="utf-8"): |
| m = json.loads(l); self.meta[m["doc_id"]] = m |
| |
| |
| |
| |
| self.eligible_doc_ids = { |
| d for d, cis in self.doc_chunks.items() |
| if d in self.meta and any(_clean(self.texts[i]) for i in cis) |
| } |
| self.goodlaw = {} |
| |
| _gl = "good_law_v2.jsonl" if os.path.exists(os.path.join(data_dir, "good_law_v2.jsonl")) else "good_law.jsonl" |
| for l in open(os.path.join(data_dir, _gl), encoding="utf-8"): |
| g = json.loads(l); self.goodlaw[g["doc_id"]] = g |
| |
| self.decision_year = {}; self.bench_n = {}; self.canonical = set(); self.cluster_of = {} |
| _lp = os.path.join(data_dir, "corpus_ledger.jsonl") |
| if os.path.exists(_lp): |
| for l in open(_lp, encoding="utf-8"): |
| r = json.loads(l); d = r["doc_id"] |
| if r.get("decision_year"): self.decision_year[d] = r["decision_year"] |
| self.bench_n[d] = r.get("bench_n", 0) |
| if r.get("canonical", True): self.canonical.add(d) |
| if r.get("cluster_id"): self.cluster_of[d] = r["cluster_id"] |
| print(f"[tools] ledger: {len(self.decision_year)} decision-years, " |
| f"{len(self.cluster_of)} sibling-clustered docs", flush=True) |
| |
| self.aliases = {} |
| _ap = os.path.join(data_dir, "case_aliases.json") |
| if os.path.exists(_ap): |
| self.aliases = {k.lower(): v for k, v in json.load(open(_ap, encoding="utf-8")).items()} |
| print(f"[tools] aliases: {len(self.aliases)}", flush=True) |
| self.in_edges = defaultdict(list); self.out_edges = defaultdict(list); self.edge_meta = {} |
| self.cite_indeg = Counter() |
| |
| _ep = "edges_v2.jsonl" if os.path.exists(os.path.join(data_dir, "edges_v2.jsonl")) else "edges.jsonl" |
| for l in open(os.path.join(data_dir, _ep), encoding="utf-8"): |
| e = json.loads(l); f, t = e["from"], e["target"] |
| self.out_edges[f].append(t); self.in_edges[t].append(f) |
| self.edge_meta[(f, t)] = {"treatment": e.get("treatment"), "method": e.get("method")} |
| if e.get("method") in ("cite", "body", "headnote"): self.cite_indeg[t] += 1 |
| if _ep == "edges_v2.jsonl": |
| print(f"[tools] edges_v2: {sum(len(v) for v in self.out_edges.values())} edges", flush=True) |
| |
| _tp = os.path.join(data_dir, "edges_treatment.jsonl") |
| if os.path.exists(_tp): |
| n = 0 |
| for l in open(_tp, encoding="utf-8"): |
| r = json.loads(l); k = (r["from"], r["target"]) |
| if k in self.edge_meta: self.edge_meta[k]["treatment"] = r["treatment"]; n += 1 |
| print(f"[tools] treatments: {n} classified edges", flush=True) |
| |
| self.syn_held = {} |
| _sp = os.path.join(data_dir, "synthetic_headnotes.jsonl") |
| if os.path.exists(_sp): |
| for l in open(_sp, encoding="utf-8"): |
| r = json.loads(l) |
| if r.get("held"): self.syn_held[r["doc_id"]] = r |
| print(f"[tools] synthetic headnotes: {len(self.syn_held)}", flush=True) |
| |
| |
| self.case_summaries, _summary_file = load_case_summaries(data_dir) |
| if _summary_file: |
| print(f"[tools] case summaries: {len(self.case_summaries)} from {_summary_file}", flush=True) |
| self.name_vocab = set(); self.name_postings = defaultdict(set) |
| self.nc2doc = {}; self.cite_resolver = {} |
| for d, m in self.meta.items(): |
| for w in _ntok(m.get("case_name") or ""): |
| if len(w) >= 4: self.name_vocab.add(w) |
| if len(w) > 1: self.name_postings[w].add(d) |
| if m.get("neutral_citation"): self.nc2doc[m["neutral_citation"]] = d |
| for k in [m.get("neutral_citation")] + (m.get("equivalent_citations") or []): |
| if k: self.cite_resolver.setdefault(re.sub(r"\s+", " ", k.replace(".", "")).strip().upper(), d) |
| |
| self.statute_idx = json.load(open(os.path.join(statute_dir, "statute_index.json"))) |
| self.statute_V = np.load(os.path.join(statute_dir, "statute_vectors.npy")) |
| _statute_records = json.load(open(os.path.join(statute_dir, "all_statutes.json"))) |
| self.statute_texts = [s.get("retrieval_text", "") for s in _statute_records] |
| self.statute_library = ExactStatuteLibrary.from_env( |
| fallback_path=os.path.join(statute_dir, "all_statutes.json") |
| ) |
| self.concord = json.load(open(os.path.join(statute_dir, "concordance.json"))) |
| self.crosswalk = load_default_crosswalk(os.environ.get("THEMIS_SECTION_CROSSWALK", "").strip() or None) |
| |
| self.held_V = None |
| hv = os.path.join(data_dir, "held_vectors.npy") |
| if os.path.exists(hv): |
| self.held_V = np.load(hv) |
| self.held_docs = json.load(open(os.path.join(data_dir, "held_docids.json"))) |
| print(f"[tools] HELD vectors: {self.held_V.shape[0]}", flush=True) |
| |
| |
| self.citectx_V = None |
| cv = os.path.join(data_dir, "citectx_vectors.npy") |
| if os.path.exists(cv): |
| self.citectx_V = np.load(cv, mmap_mode="r") |
| self.citectx_docs = json.load(open(os.path.join(data_dir, "citectx_docids.json"))) |
| print(f"[tools] CITECTX vectors: {self.citectx_V.shape[0]}", flush=True) |
| self.st = SentenceTransformer("BAAI/bge-small-en-v1.5", device=device) |
| self.ce = CrossEncoder("cross-encoder/ms-marco-MiniLM-L-6-v2", device=device) |
| |
| |
| |
| self.kw = None; self.fts = None |
| _fts_path = os.path.join(data_dir, "escr_fts.sqlite") |
| if os.environ.get("THEMIS_FTS", "1") == "1" and os.path.exists(_fts_path): |
| import sqlite3, threading |
| self.fts = sqlite3.connect(f"file:{_fts_path}?mode=ro", uri=True, check_same_thread=False) |
| self.fts_lock = threading.Lock() |
| print(f"[tools] FTS5 keyword index: {os.path.getsize(_fts_path)/1e9:.2f} GB (disk-backed)", flush=True) |
| if self.fts is None and os.environ.get("THEMIS_KEYWORD", "1") == "1": |
| import math, time as _t |
| t0 = _t.time(); print("[tools] building keyword index ...", flush=True) |
| self.kw_docs = list(self.doc_chunks.keys()) |
| self.kw_postings = defaultdict(list); df = Counter() |
| self.kw_dl = np.zeros(len(self.kw_docs), dtype=np.float32) |
| for i, d in enumerate(self.kw_docs): |
| toks = [] |
| for ci in self.doc_chunks[d]: toks += re.findall(r"[a-z0-9]+", self.texts[ci].lower()) |
| tf = Counter(toks); self.kw_dl[i] = len(toks) |
| for t, c in tf.items(): self.kw_postings[t].append((i, c)); df[t] += 1 |
| N = len(self.kw_docs); self.kw_avgdl = float(self.kw_dl.mean()) or 1.0 |
| self.kw_idf = {t: math.log(1 + (N - n + 0.5) / (n + 0.5)) for t, n in df.items()} |
| self.kw = True |
| print(f"[tools] keyword index: {len(self.kw_postings)} terms, {N} docs, {_t.time()-t0:.0f}s", flush=True) |
| print( |
| f"[tools] ready — {len(self.eligible_doc_ids)} source-grounded judgments " |
| f"({len(self.meta) - len(self.eligible_doc_ids)} metadata-only excluded), " |
| f"{len(self.statute_idx)} statute sections", |
| flush=True, |
| ) |
|
|
| |
| def _enc(self, q): |
| return self.st.encode(BGE_Q + q, normalize_embeddings=True, convert_to_numpy=True).astype(np.float32) |
| def is_retrieval_eligible(self, doc_id): |
| """True only for judgments whose source text is in the active corpus.""" |
| return str(doc_id) in self.eligible_doc_ids |
| def coverage(self): |
| return { |
| "accepted_judgments": len(self.eligible_doc_ids), |
| "metadata_only_excluded": len(self.meta) - len(self.eligible_doc_ids), |
| "scope": "Supreme Court of India judgments stored in this release", |
| } |
| def _card(self, d, rr=0.0): |
| m = self.meta.get(d, {}); gl = self.goodlaw.get(d, {}) |
| cis = self.doc_chunks.get(d, []) |
| snip = _clean((m.get("held") or m.get("issue") |
| or (self.syn_held.get(d) or {}).get("held") |
| or (self.texts[cis[0]] if cis else "")))[:240] |
| return {"doc_id": d, "judgment_id": str(d), "case_name": m.get("case_name"), |
| "year": self.decision_year.get(d) or m.get("year") or m.get("date"), |
| "neutral_citation": m.get("neutral_citation"), "bench_strength": m.get("bench_strength"), |
| "cited_by": self.cite_indeg.get(d, 0), "good_law": gl.get("good_law_status", "unknown"), |
| "rr": round(float(rr), 2), "snippet": snip} |
| def _dense_pool(self, qv, n=120): |
| sim = self.M @ qv |
| requested = min(max(n * 2, n + 1), len(sim) - 1) |
| top = np.argpartition(-sim, requested)[:requested] |
| return [ |
| int(i) for i in top[np.argsort(-sim[top])] |
| if self.is_retrieval_eligible(self.chunk_doc[int(i)]) |
| ][:n] |
| def _rerank_docs(self, q, cand_idx, topk): |
| cand_idx = [ |
| ci for ci in cand_idx |
| if self.is_retrieval_eligible(self.chunk_doc[ci]) |
| and _clean(self.texts[ci]) |
| ] |
| if not cand_idx: |
| return [] |
| rr = self.ce.predict([(q, self.texts[ci]) for ci in cand_idx]) |
| best = {} |
| for ci, s in zip(cand_idx, rr): |
| d = self.chunk_doc[ci] |
| if not self.is_retrieval_eligible(d): continue |
| if d not in best or s > best[d]: best[d] = float(s) |
| ranked = sorted(best.items(), key=lambda x: -x[1])[:topk] |
| return [self._card(d, s) for d, s in ranked] |
|
|
| |
| def vector_search(self, q, k=8): |
| """Semantic retrieval — doctrine described in the user's words. dense pool -> cross-encoder.""" |
| return self._rerank_docs(q, self._dense_pool(self._enc(q), 120), k) |
|
|
| def authority_search(self, q, k=8, alpha=0.3): |
| """Retrieve, then rank by AUTHORITY (cross-encoder + alpha*log1p(cite_indeg)) — 'the leading case on X'.""" |
| cand = self._dense_pool(self._enc(q), 120) |
| rr = self.ce.predict([(q, self.texts[ci]) for ci in cand]); best = {} |
| for ci, s in zip(cand, rr): |
| d = self.chunk_doc[ci] |
| if not self.is_retrieval_eligible(d): continue |
| if d not in best or s > best[d]: best[d] = float(s) |
| sig = lambda x: 1/(1+np.exp(-x)) |
| scored = sorted(best.items(), key=lambda x: -(sig(x[1]) + alpha*np.log1p(self.cite_indeg.get(x[0], 0))))[:k] |
| return [self._card(d, s) for d, s in scored] |
|
|
| def keyword_search(self, q, k=12, k1=1.5, b=0.75): |
| """BM25 keyword retrieval — exact terms / names / section nums that dense misses. |
| FTS5 path: chunk-level match, best-chunk-per-doc aggregation (a doc with one strong |
| exact-term chunk ranks high), OR semantics to mirror the legacy scorer.""" |
| if self.fts is not None: |
| _stop = {"of","the","and","or","in","to","a","an","is","for","on","by","at", |
| "with","under","was","were","be","has","had","it","that","this"} |
| toks = [t for t in re.findall(r"[a-z0-9]+", q.lower()) if t not in _stop] |
| if not toks: return [] |
| match = " OR ".join(f'"{t}"' for t in toks[:24]) |
| with self.fts_lock: |
| rows = self.fts.execute( |
| "SELECT rowid, bm25(fts) FROM fts WHERE fts MATCH ? ORDER BY bm25(fts) LIMIT ?", |
| (match, max(k * 12, 240))).fetchall() |
| best = {} |
| for ci, s in rows: |
| d = self.chunk_doc[ci] |
| if not self.is_retrieval_eligible(d): continue |
| if d not in best or s < best[d]: best[d] = s |
| top = sorted(best.items(), key=lambda x: x[1])[:k] |
| return [self._card(d) for d, _ in top] |
| if not self.kw: return [] |
| scores = defaultdict(float) |
| for t in set(re.findall(r"[a-z0-9]+", q.lower())): |
| idf = self.kw_idf.get(t) |
| if not idf: continue |
| for i, tf in self.kw_postings[t]: |
| scores[i] += idf * (tf * (k1 + 1)) / (tf + k1 * (1 - b + b * self.kw_dl[i] / self.kw_avgdl)) |
| top = [ |
| item for item in sorted(scores.items(), key=lambda x: -x[1]) |
| if self.is_retrieval_eligible(self.kw_docs[item[0]]) |
| ][:k] |
| return [self._card(self.kw_docs[i]) for i, _ in top] |
|
|
| def dense_docs(self, q, k=60): |
| """Doc-level dense ranking (first chunk-hit per doc).""" |
| qv = self._enc(q); sim = self.M @ qv |
| top = np.argpartition(-sim, 2500)[:2500]; top = top[np.argsort(-sim[top])] |
| seen = []; s = set() |
| for ci in top: |
| d = self.chunk_doc[ci] |
| if not self.is_retrieval_eligible(d): continue |
| if d not in s: s.add(d); seen.append(d) |
| if len(seen) >= k: break |
| return seen |
|
|
| def held_search(self, q, k=12): |
| """Rank judgments by HELD-headnote similarity (doc-level, clean reporter language).""" |
| if self.held_V is None: return [] |
| qv = self._enc(q); sim = self.held_V @ qv |
| top = np.argpartition(-sim, min(k, len(sim) - 1))[:k] |
| return [ |
| self.held_docs[int(i)] for i in top[np.argsort(-sim[top])] |
| if self.is_retrieval_eligible(self.held_docs[int(i)]) |
| ][:k] |
|
|
| def citectx_search(self, q, k=12): |
| """Rank judgments by how LATER courts describe them (citation-context vectors). |
| Multiple contexts per doc -> dedupe keeping best rank.""" |
| if self.citectx_V is None: return [] |
| qv = self._enc(q); sim = np.asarray(self.citectx_V @ qv) |
| n = min(k * 6, len(sim) - 1) |
| top = np.argpartition(-sim, n)[:n] |
| out, seen = [], set() |
| for i in top[np.argsort(-sim[top])]: |
| d = self.citectx_docs[int(i)] |
| if not self.is_retrieval_eligible(d): continue |
| if d not in seen: |
| seen.add(d); out.append(d) |
| if len(out) >= k: break |
| return out |
|
|
| def hybrid_search(self, q, k=8, pool=60): |
| """The strong base primitive (panel + recall ablation: RRF@100=0.96): dense + BM25 -> RRF -> rerank.""" |
| dd = self.dense_docs(q, pool) |
| kd = [c["doc_id"] for c in self.keyword_search(q, pool)] |
| sc = {} |
| for r in (dd, kd): |
| for rank, d in enumerate(r): sc[d] = sc.get(d, 0.0) + 1.0 / (60 + rank + 1) |
| fused = [d for d, _ in sorted(sc.items(), key=lambda x: -x[1])][:max(40, k * 4)] |
| rr = self.score_docs(q, fused) |
| for d in fused: |
| if d not in rr: rr[d] = -9.0 |
| ranked = sorted(fused, key=lambda d: -rr[d]) |
| return [self._card(d, rr.get(d, 0.0)) for d in ranked[:k]] |
|
|
| def statute_search(self, q, k=3): |
| """Find the statute SECTION(S) most relevant to the query (BNS/IPC/CrPC/IEA/...).""" |
| qv = self._enc(q); sim = self.statute_V @ qv |
| out = [] |
| for j in np.argsort(-sim)[:k]: |
| s = self.statute_idx[int(j)] |
| out.append({"act": s.get("act_short"), "section": s.get("section_number"), |
| "title": s.get("title"), "i": int(j)}) |
| return out |
|
|
| def cases_on_section(self, act_section_text, k=8): |
| """Cases discussing a statute section: embed the section text, retrieve nearest judgments.""" |
| qv = self.st.encode(act_section_text, normalize_embeddings=True, convert_to_numpy=True).astype(np.float32) |
| return self._rerank_docs(act_section_text[:300], self._dense_pool(qv, 120), k) |
|
|
| def statute_crosswalk(self, code, section): |
| """Map a section across the new/old codes (BNS<->IPC, BNSS<->CrPC, BSA<->IEA).""" |
| return self.crosswalk.lookup(code, section) |
|
|
| def statute_provision(self, code, section): |
| return self.statute_library.lookup(code, section) |
|
|
| def encode_documents(self, texts): |
| values = [_clean(value) for value in texts if _clean(value)] |
| if not values: |
| return np.empty((0, int(self.M.shape[1])), dtype=np.float32) |
| return np.asarray( |
| self.st.encode(values, normalize_embeddings=True, convert_to_numpy=True), |
| dtype=np.float32, |
| ) |
|
|
| def find_similar_cases(self, doc_id, k=8): |
| """'More like this' — nearest judgments to doc_id by embedding centroid.""" |
| cis = self.doc_chunks.get(doc_id, []) |
| if not cis: return [] |
| centroid = self.M[cis].mean(0); centroid /= (np.linalg.norm(centroid) + 1e-9) |
| pool = self._dense_pool(centroid.astype(np.float32), 60) |
| seen = set([doc_id]); out = [] |
| for ci in pool: |
| d = self.chunk_doc[ci] |
| if self.is_retrieval_eligible(d) and d not in seen: |
| seen.add(d); out.append(self._card(d)) |
| if len(out) >= k: break |
| return out |
|
|
| def cited_authorities(self, doc_id, k=12): |
| """Note-UP: the cases doc_id relies on (its authority chain).""" |
| return [ |
| self._card(d) |
| for d in list(dict.fromkeys(self.out_edges.get(doc_id, []))) |
| if self.is_retrieval_eligible(d) |
| ][:k] |
|
|
| def progeny(self, doc_id, k=12): |
| """Note-DOWN: the cases that cite doc_id (its progeny + treatment).""" |
| out = [] |
| for d in list(dict.fromkeys(self.in_edges.get(doc_id, [])))[:k]: |
| if not self.is_retrieval_eligible(d): continue |
| c = self._card(d); c["treatment"] = self.edge_meta.get((d, doc_id), {}).get("treatment") |
| out.append(c) |
| return out |
|
|
| def co_cited_cases(self, doc_id, k=8): |
| """Cases similar by SHARED AUTHORITIES (bibliographic coupling) — cases that cite what doc_id cites.""" |
| mine = set(self.out_edges.get(doc_id, [])) |
| if not mine: return [] |
| score = Counter() |
| for t in mine: |
| for citer in self.in_edges.get(t, []): |
| if citer != doc_id and self.is_retrieval_eligible(citer): score[citer] += 1 |
| return [self._card(d) for d, _ in score.most_common(k)] |
|
|
| def good_law_check(self, doc_id): |
| """Citator: current status + treatment breakdown + the overruling case if any.""" |
| gl = self.goodlaw.get(doc_id, {}) |
| status = gl.get("good_law_status", "unknown") |
| overruled_by = None |
| if status in BAD_STATUS: |
| for s, t in [(s, t) for (s, t) in self.edge_meta if t == doc_id]: |
| if self.edge_meta[(s, t)].get("treatment") in ("overruled", "overrules"): |
| if self.is_retrieval_eligible(s): |
| overruled_by = self._card(s); break |
| return {"doc_id": doc_id, "good_law": status, "treatment_breakdown": gl.get("treatment_breakdown", {}), |
| "overruled_by": overruled_by} |
|
|
| def metadata_filter(self, cards, min_bench=None, year_from=None, year_to=None): |
| """Filter a candidate list by bench strength (Constitution Bench = 5+), date range.""" |
| out = [] |
| _BN = {"single": 1, "division": 2, "full": 3, "constitution": 5, "larger": 7} |
| for c in cards: |
| d = c["doc_id"]; m = self.meta.get(d, {}) |
| bs = self.bench_n.get(d) or _BN.get(str(m.get("bench_strength") or "").lower(), 0) |
| yr = self.decision_year.get(d) or 0 |
| if not yr: |
| try: yr = int(str(m.get("year") or 0)[:4]) |
| except Exception: yr = 0 |
| if min_bench and bs < min_bench: continue |
| if year_from and yr and yr < year_from: continue |
| if year_to and yr and yr > year_to: continue |
| out.append(c) |
| return out |
|
|
| def read_case(self, doc_id): |
| """Read a case's headnote/held/issue (for the agent to verify relevance + for grounding).""" |
| if not self.is_retrieval_eligible(doc_id): |
| return {} |
| m = self.meta.get(doc_id, {}) |
| held = _clean(m.get("held")); issue = _clean(m.get("issue")) |
| if not held and doc_id in self.syn_held: |
| s = self.syn_held[doc_id] |
| held = _clean(s.get("held")); issue = issue or _clean(s.get("issue")) |
| if not held and not issue: |
| cis = self.doc_chunks.get(doc_id, []) |
| held = _clean(" ".join(self.texts[i] for i in cis[:2])) |
| return {"doc_id": doc_id, "case_name": m.get("case_name"), "neutral_citation": m.get("neutral_citation"), |
| "bench_strength": m.get("bench_strength"), "good_law": self.goodlaw.get(doc_id, {}).get("good_law_status", "unknown"), |
| "issue": issue[:1200], "held": held[:1800]} |
|
|
| def score_docs(self, q, doc_ids, per_doc=3): |
| """Uniformly cross-encoder-score a heterogeneous pool of docs vs q (max over each doc's first |
| chunks). One batched CE pass. Returns {doc_id: rr}. Lets graph/authority/statute additions be |
| ranked on the same scale as dense hits.""" |
| pairs = []; owner = [] |
| for d in doc_ids: |
| if not self.is_retrieval_eligible(d): continue |
| for ci in self.doc_chunks.get(d, [])[:per_doc]: |
| pairs.append((q, self.texts[ci])); owner.append(d) |
| if not pairs: return {d: 0.0 for d in doc_ids} |
| sc = self.ce.predict(pairs, batch_size=256) |
| best = {d: -9e9 for d in doc_ids} |
| for d, s in zip(owner, sc): |
| if s > best[d]: best[d] = float(s) |
| return {d: (best[d] if best[d] > -9e9 else 0.0) for d in doc_ids} |
|
|
| def front_text(self, doc_id, n=1800): |
| """The judgment's FRONT MATTER (reporter headnote lives here in 70-100% of judgments across |
| all decades — more reliable than meta.held, which craters to ~2% in the 1970s-80s).""" |
| held = _clean(self.meta.get(doc_id, {}).get("held") or "") |
| if len(held) > 200: return held[:n] |
| syn = _clean((self.syn_held.get(doc_id) or {}).get("held") or "") |
| if len(syn) > 200: return syn[:n] |
| cis = self.doc_chunks.get(doc_id, []) |
| return _clean(" ".join(self.texts[i] for i in cis[:3]))[:n] |
|
|
| def full_text_for_read(self, q, doc_id, cap_chars=90000): |
| """Layer-2 reading surface: the FULL judgment up to ~cap (≈22k tokens). Above-cap monsters get a |
| tiered pack: HELD + opening + a window around the query's best-matching chunk + the ending — |
| the controlling passage in multi-issue judgments sits mid-text where head/tail packs go blind.""" |
| cis = self.doc_chunks.get(doc_id, []) |
| if not cis: return "" |
| parts = [self.texts[i] for i in cis] |
| full = "\n".join(parts) |
| if len(full) <= cap_chars: return full |
| held = _clean((self.meta.get(doc_id, {}).get("held") or ""))[:6000] |
| probe = cis[:40] |
| sc = self.ce.predict([(q, self.texts[i]) for i in probe]) |
| bi = int(np.argmax(sc)) |
| win = "\n".join(self.texts[i] for i in cis[max(0, bi - 2):bi + 3]) |
| head = "\n".join(parts[:8]); tail = "\n".join(parts[-6:]) |
| pack = (("HELD: " + held + "\n\n") if held else "") + head + "\n[...]\n" + win + "\n[...]\n" + tail |
| return pack[:cap_chars] |
|
|
| def best_chunk_text(self, q, doc_id, limit=1600): |
| """The doc's single chunk most relevant to q (for the grounding gate to quote from).""" |
| cis = self.doc_chunks.get(doc_id, [])[:6] |
| if not cis: return "" |
| sc = self.ce.predict([(q, self.texts[ci]) for ci in cis]) |
| return _clean(self.texts[cis[int(np.argmax(sc))]])[:limit] |
|
|
| def judgment_view(self, doc_id): |
| """Full case view for the pilot UI (metadata + issue/held + good-law + cited-by). issue is in |
| ~3% of metadata, held in ~56% — so fall back to the judgment's opening text when missing.""" |
| if not self.is_retrieval_eligible(doc_id): |
| return {} |
| m = self.meta.get(doc_id, {}); gl = self.goodlaw.get(doc_id, {}) |
| cis = self.doc_chunks.get(doc_id, []) |
| body = _clean(" ".join(self.texts[i] for i in cis[:14])) |
| extracted = self.case_summaries.get(doc_id, {}) |
| issue = _clean(m.get("issue") or extracted.get("issue")) |
| held = _clean(m.get("held") or extracted.get("held")); synthetic = False |
| if not held and doc_id in self.syn_held: |
| s = self.syn_held[doc_id] |
| held = _clean(s.get("held")); issue = issue or _clean(s.get("issue")); synthetic = bool(held) |
| if not held: held = body[:6000] |
| return {"doc_id": doc_id, "synthetic_headnote": synthetic, |
| "summary": case_summary_record(m, self.syn_held.get(doc_id), extracted), |
| "case_name": m.get("case_name"), "neutral_citation": m.get("neutral_citation"), |
| "equivalent_citations": m.get("equivalent_citations"), "court": m.get("court"), "date": m.get("date"), |
| "bench_strength": m.get("bench_strength"), "disposition": m.get("disposition"), |
| "good_law_status": gl.get("good_law_status", "unknown"), "treatment_breakdown": gl.get("treatment_breakdown", {}), |
| "cited_by": self.cite_indeg.get(doc_id, 0), "issue": issue[:4000], |
| "held": held[:6000], "text": body[:60000]} |
|
|
| def case_chat_passages(self, q, doc_id, k=4, limit=1800): |
| """Return only source passages from one eligible opened judgment. |
| |
| The legacy bundle has chunk identity rather than schema-v5 paragraph |
| identity, so its stable fallback IDs are disclosed as chunk anchors. |
| The v5 adapter supplies real paragraph IDs through the same contract. |
| """ |
| if not self.is_retrieval_eligible(doc_id): |
| return [] |
| cis = self.doc_chunks.get(doc_id, []) |
| if not cis: |
| return [] |
| scores = self.ce.predict([(q, self.texts[ci]) for ci in cis[:40]]) |
| ranked = sorted(zip(cis[:40], scores), key=lambda item: -float(item[1]))[:k] |
| return [ |
| { |
| "paragraph_id": f"{doc_id}:chunk:{ci}", |
| "label": f"Indexed passage {rank + 1}", |
| "text": _clean(self.texts[ci])[:limit], |
| "source_kind": "legacy_chunk", |
| } |
| for rank, (ci, _) in enumerate(ranked, 1) |
| if _clean(self.texts[ci]) |
| ] |
|
|
| def identity_hits(self, q): |
| """Known-item route: a citation or a 'X v Y' case-name query resolves to the EXACT case(s) |
| (cite_indeg salience tiebreak), not semantic search. Restores serve.py's identity routing.""" |
| ql = q.strip() |
| m = re.search(r"\[\d{4}\]\s*\d+\s*S\.?C\.?R\.?\s*\d+|\(\d{4}\)\s*\d+\s*SCC\s*\d+|\d{4}\s+INSC\s+\d+|AIR\s+\d{4}\s+SC\s+\d+", ql, re.I) |
| if m: |
| rid = self.cite_resolver.get(re.sub(r"\s+", " ", m.group(0).replace(".", "")).strip().upper()) or self.nc2doc.get(m.group(0)) |
| if rid and self.is_retrieval_eligible(rid): return [rid], "citation" |
| |
| if self.aliases and len(ql) <= 60: |
| qa = re.sub(r"[^a-z0-9 ]", " ", ql.lower()) |
| qa = re.sub(r"\b(the|case|judgment|judgement|in|re|of)\b", " ", qa) |
| qa = re.sub(r"\s+", " ", qa).strip() |
| if qa in self.aliases and self.is_retrieval_eligible(self.aliases[qa]): |
| return [self.aliases[qa]], "case name" |
| |
| |
| |
| hits = [(a, d) for a, d in self.aliases.items() |
| if len(a) >= 8 and a in qa and len(qa) - len(a) <= 10] |
| if hits: |
| best = max(hits, key=lambda ad: self.cite_indeg.get(ad[1], 0)) |
| if self.is_retrieval_eligible(best[1]): |
| return [best[1]], "case name" |
| if re.search(r"\bv[s.]?\b|\bversus\b", ql, re.I) and len(ql) <= 90: |
| hits = [c["doc_id"] for c in self.name_lookup(ql, 6)] |
| if hits: return hits, "case name" |
| return [], None |
|
|
| def name_lookup(self, name, k=4): |
| """Resolve a case NAME to corpus doc(s) — the recall tool for LLM-named authorities.""" |
| raw = [t for t in _ntok(name) if t not in _NAME_STOP and len(t) > 1] |
| if not raw: return [] |
| |
| extra = [] |
| for t in raw: |
| if t not in self.name_vocab and len(t) >= 7: |
| for cut in range(3, len(t) - 2): |
| a, b = t[:cut], t[cut:] |
| if a in self.name_vocab and b in self.name_vocab: extra += [a, b]; break |
| for a, b in zip(raw, raw[1:]): |
| if (a + b) in self.name_vocab: extra.append(a + b) |
| raw += extra |
| qtok = set() |
| for t in raw: |
| if t in self.name_vocab or len(t) <= 3: qtok.add(t) |
| else: qtok.update(difflib.get_close_matches(t, self.name_vocab, n=3, cutoff=0.82) or [t]) |
| cand = set() |
| for t in qtok: cand |= self.name_postings.get(t, set()) |
| qdist = {t for t in qtok if len(t) >= 5} |
| scored = [] |
| for d in cand: |
| if not self.is_retrieval_eligible(d): continue |
| ntok = set(_ntok(self.meta.get(d, {}).get("case_name") or "")) |
| ov = qtok & ntok |
| if qdist and not (qdist & ntok): continue |
| if len(ov) >= 2 or (len(ov) == 1 and any(len(t) >= 5 for t in ov)): |
| |
| |
| scored.append((len(ov), 1 if d in self.canonical else 0, |
| -(len(ntok) - len(ov)), self.cite_indeg.get(d, 0), d)) |
| scored.sort(reverse=True) |
| return [self._card(t[-1]) for t in scored[:k]] |
|
|