""" conflict_checker.py - Stage 3: cross-discipline conflict reasoning. One reasoning call per co-located cluster (batch-size-1 per the IronBid token-budget lesson: small, focused calls don't truncate). Each call gets the cluster's assertions as evidence plus the relevant sheet images, and returns zero or more validated conflicts. Calls run in parallel. """ import threading from concurrent.futures import ThreadPoolExecutor, as_completed from typing import List, Dict, Optional from backend import config from backend.llm import call_json from backend.prompts import ( CONFLICT_SYSTEM_PROMPT, CONFLICT_USER_INSTRUCTION, CONFLICT_CATEGORIES, ) # Cap images per cluster to bound tokens/cost. Most clusters touch 2-4 sheets. MAX_IMAGES_PER_CLUSTER = 6 def _evidence_block(cluster: Dict) -> str: lines = [] for a in cluster["assertions"]: sheet = a.get("sheet_number") or "?" lines.append( f"- [{a.get('discipline','Unknown')}] {sheet} | " f"{a.get('attribute','')} = {a.get('value','')} | " f"\"{a.get('source_text','')}\"" ) disputes = cluster.get("disputed_attributes") or [] if disputes: lines.append("") for d in disputes: lines.append( "DISPUTED VALUE (possible extraction misread): " f"attribute={d.get('attribute','')} " f"values={' | '.join(d.get('values') or [])} " f"(assertions {', '.join(d.get('assertion_ids') or [])})" ) return "\n".join(lines) def _images_for(cluster: Dict, page_to_b64: Dict[int, str]) -> List[str]: imgs = [] for pn in cluster.get("page_numbers", [])[:MAX_IMAGES_PER_CLUSTER]: b64 = page_to_b64.get(pn) if b64: imgs.append(b64) return imgs def _valid_conflict(c: Dict, cluster: Dict) -> Optional[Dict]: """Coerce + validate one model-returned conflict; drop if malformed.""" if not isinstance(c, dict): return None category = (c.get("category") or "").strip() if category not in CONFLICT_CATEGORIES: # Unknown category -> keep but mark, so tuning can see what the model wanted. category = category or "uncategorized" severity = (c.get("severity") or "medium").strip().lower() if severity not in ("high", "medium", "low"): severity = "medium" description = (c.get("description") or "").strip() if not description: return None # a conflict with no description is noise return { "category": category, "severity": severity, "disciplines": c.get("disciplines") or cluster.get("disciplines", []), "location": c.get("location") or cluster.get("location", ""), "sheets": c.get("sheets") or cluster.get("sheets", []), "description": description, "evidence": c.get("evidence") or [], "recommended_resolution": (c.get("recommended_resolution") or "").strip(), "confidence": (c.get("confidence") or "medium").strip().lower(), "cluster_key": cluster.get("key"), } def _check_one(cluster: Dict, page_to_b64: Dict[int, str]) -> List[Dict]: user_text = ( CONFLICT_USER_INSTRUCTION .replace("{location}", cluster.get("location", "")) .replace("{evidence}", _evidence_block(cluster)) ) parsed = call_json( system_prompt=CONFLICT_SYSTEM_PROMPT, user_text=user_text, images_b64=_images_for(cluster, page_to_b64), max_tokens=config.REASON_MAX_TOKENS, reasoning_effort=config.EXTRACT_REASONING_EFFORT or None, reasoning_max_tokens=config.EXTRACT_REASONING_MAX_TOKENS or None, ) if isinstance(parsed, list): candidates = parsed elif isinstance(parsed, dict): candidates = parsed.get("conflicts") or [] else: return [] out = [] for c in candidates: v = _valid_conflict(c, cluster) if v: out.append(v) return out def check_conflicts(clusters: List[Dict], pages: List[Dict]) -> List[Dict]: """ Reason over every cluster and return a flat list of validated conflicts. pages: the Stage-0 page dicts (need 'page_number' and 'base64') so each cluster can be shown its own sheets. """ page_to_b64 = {p["page_number"]: p["base64"] for p in pages} print(f"[Conflicts] Reasoning over {len(clusters)} clusters...") conflicts: List[Dict] = [] completed = 0 lock = threading.Lock() def work(cluster: Dict) -> List[Dict]: nonlocal completed found = _check_one(cluster, page_to_b64) with lock: completed += 1 if completed % 10 == 0 or completed == len(clusters): print(f"[Conflicts] {completed}/{len(clusters)} clusters checked") return found with ThreadPoolExecutor(max_workers=config.REASON_CONCURRENCY) as ex: futures = [ex.submit(work, c) for c in clusters] for fut in as_completed(futures): try: conflicts.extend(fut.result()) except Exception as e: print(f"[Conflicts] cluster error: {e}") # Sort by severity then category for a stable, readable report. sev_rank = {"high": 0, "medium": 1, "low": 2} conflicts.sort(key=lambda c: (sev_rank.get(c["severity"], 1), c["category"])) print(f"[Conflicts] Found {len(conflicts)} conflict(s)") return conflicts