#!/usr/bin/env python3 """Classify citation treatments over edges_v2 contexts -> a REAL good-law citator. The at-scale build hardcoded every edge treatment to "cited"; good_law_status is "unknown" for 99.9% of the corpus, so the bad-law guardrail was theater. Each edge in edges_v2.jsonl carries `para` — the citing court's sentence about the precedent — which is exactly the input a treatment classifier needs. Labels: relied_on | followed | distinguished | doubted | overruled | cited (neutral). Model: deepseek-v4-flash (non-thinking pinned), temperature 0, batched 20 edges/call. Requires DEEPSEEK_API_KEY. Run: python phase1/scripts/classify_treatments.py [data_dir] [--dry-run N] [--limit N] Out: /edges_treatment.jsonl {from, target, treatment} (resumable) /good_law_v2.jsonl per-target rollup {doc_id, good_law_status, treatment_breakdown, cited_by} Rollup rule: any overruled edge -> "overruled"; >=2 doubted/distinguished and no follow/rely since -> "doubted"; else "good" if >=3 positive treatments else "unknown". """ import json, os, sys, time from collections import Counter, defaultdict data_dir = next((a for a in sys.argv[1:] if not a.startswith("--")), None) \ or os.environ.get("THEMIS_DATA", "phase1/data/thor_artifacts") DRY = 0 if "--dry-run" in sys.argv: i = sys.argv.index("--dry-run"); DRY = int(sys.argv[i + 1]) if len(sys.argv) > i + 1 else 3 LIMIT = int(sys.argv[sys.argv.index("--limit") + 1]) if "--limit" in sys.argv else None SYS = ("You classify how an Indian Supreme Court judgment TREATS a precedent it cites, from the " "sentence surrounding the citation. Labels: relied_on (the holding rests on it), followed " "(applied approvingly), distinguished (held inapplicable on facts/law), doubted (correctness " "questioned), overruled (expressly overruled/no longer good law), cited (neutral mention). " "Input: JSON list of {i, context}. Output STRICT JSON: {\"labels\": [{\"i\": n, " "\"treatment\": \"...\"}]} — one per input, label from the list only.") def load_edges(): edges = [] for line in open(os.path.join(data_dir, "edges_v2.jsonl"), encoding="utf-8"): e = json.loads(line) if len((e.get("para") or "")) >= 60: edges.append(e) return edges def rollup(treat_by_edge, edges): per_tgt = defaultdict(Counter); indeg = Counter() for e in edges: indeg[e["target"]] += 1 for (f_, t_), lab in treat_by_edge.items(): per_tgt[t_][lab] += 1 out = os.path.join(data_dir, "good_law_v2.jsonl") with open(out, "w", encoding="utf-8") as f: for tgt, cnt in per_tgt.items(): if cnt.get("overruled"): status = "overruled" elif cnt.get("doubted", 0) + cnt.get("distinguished", 0) >= 2 and \ cnt.get("relied_on", 0) + cnt.get("followed", 0) == 0: status = "doubted" elif cnt.get("relied_on", 0) + cnt.get("followed", 0) >= 3: status = "good" else: status = "unknown" f.write(json.dumps({"doc_id": tgt, "good_law_status": status, "treatment_breakdown": dict(cnt), "cited_by": indeg.get(tgt, 0)}, ensure_ascii=False) + "\n") print(f"[treat] rollup -> {out} ({len(per_tgt)} targets)", flush=True) def main(): edges = load_edges() outp = os.path.join(data_dir, "edges_treatment.jsonl") done = set() if os.path.exists(outp): for line in open(outp, encoding="utf-8"): r = json.loads(line); done.add((r["from"], r["target"])) todo = [e for e in edges if (e["from"], e["target"]) not in done] if LIMIT: todo = todo[:LIMIT] print(f"[treat] edges with context: {len(edges)} | done: {len(done)} | todo: {len(todo)}", flush=True) if DRY: for e in todo[:DRY]: print(f"\n--- DRY {e['from']} -> {e['target']} ---\n {e['para'][:220]}", flush=True) print(f"\n[treat] dry-run only ({DRY} shown); no API calls made.", flush=True) return key = os.environ.get("DEEPSEEK_API_KEY", "") if not key: env = os.path.join(os.path.dirname(os.path.abspath(__file__)), ".env") if os.path.exists(env): for l in open(env): if l.startswith("DEEPSEEK_API_KEY="): key = l.split("=", 1)[1].strip() if not key: sys.exit("[treat] DEEPSEEK_API_KEY missing (env or phase1/scripts/.env) — aborting.") import requests B = 20 with open(outp, "a", encoding="utf-8") as f: for s in range(0, len(todo), B): batch = todo[s:s + B] payload = [{"i": i, "context": e["para"][:400]} for i, e in enumerate(batch)] try: r = requests.post("https://api.deepseek.com/chat/completions", headers={"Authorization": f"Bearer {key}"}, json={"model": os.environ.get("THEMIS_LLM_MODEL", "deepseek-v4-flash"), "temperature": 0, "thinking": {"type": "disabled"}, "response_format": {"type": "json_object"}, "messages": [{"role": "system", "content": SYS}, {"role": "user", "content": json.dumps(payload)}]}, timeout=120) labs = json.loads(r.json()["choices"][0]["message"]["content"]).get("labels", []) VALID = {"relied_on","followed","distinguished","doubted","overruled","cited"} for l in labs: i = l.get("i"); t = l.get("treatment") if isinstance(i, int) and 0 <= i < len(batch) and t in VALID: e = batch[i] f.write(json.dumps({"from": e["from"], "target": e["target"], "treatment": t}, ensure_ascii=False) + "\n") f.flush() except Exception as ex: print(f"[treat] batch {s}: {ex}", flush=True); time.sleep(3) if (s // B) % 25 == 0: print(f"[treat] {min(s+B,len(todo))}/{len(todo)}", flush=True) treat = {} for line in open(outp, encoding="utf-8"): r = json.loads(line); treat[(r["from"], r["target"])] = r["treatment"] rollup(treat, edges) if __name__ == "__main__": main()