"""MoE post-correction pipeline: two VLM experts + one VLM judge (full page).""" from __future__ import annotations import asyncio import json import re from provider import LLMClient, image_part, text_part from prompts import ( render_expert_system, render_expert_user, render_judge_system, render_judge_user, ) def _strip_code_fence(text: str) -> str: t = (text or "").strip() if t.startswith("```"): t = t.strip("`") if t.lower().startswith("json"): t = t[4:] t = t.strip() return t def _try_parse_json(raw: str) -> dict | None: t = _strip_code_fence(raw) if not t: return None try: obj = json.loads(t) return obj if isinstance(obj, dict) else None except json.JSONDecodeError: pass # Fallback: find the first balanced {...} block m = re.search(r"\{.*\}", t, flags=re.DOTALL) if m: try: obj = json.loads(m.group(0)) return obj if isinstance(obj, dict) else None except json.JSONDecodeError: return None return None def _coerce_confidence(v) -> float: try: f = float(v) except (TypeError, ValueError): return 0.5 return max(0.0, min(1.0, f)) def _coerce_corrections(v) -> list[str]: if isinstance(v, list): return [str(x)[:200] for x in v[:10]] if isinstance(v, str): return [v[:200]] return [] async def _run_expert( *, client: LLMClient, model: str, image_b64: str, ocr_text: str, language: str, guidelines: str, temperature: float, system_override: str | None = None, user_override: str | None = None, ) -> dict: system = render_expert_system(language, guidelines, override=system_override) user_text = render_expert_user(ocr_text=ocr_text, override=user_override) user_parts = [text_part(user_text), image_part(image_b64)] res = await client.chat( model=model, system=system, user_parts=user_parts, temperature=temperature, max_tokens=4096, force_json=True, ) out = { "model": model, "ok": res.ok, "latency_s": res.latency_s, "error": res.error, "raw": res.content, "usage": res.usage, "corrected_text": "", "confidence": 0.0, "corrections": [], } if not res.ok: return out parsed = _try_parse_json(res.content) if parsed is None: out["ok"] = False out["error"] = "Could not parse JSON from model response." return out out["corrected_text"] = str(parsed.get("corrected_text") or "").rstrip() out["confidence"] = _coerce_confidence(parsed.get("confidence", 0.5)) out["corrections"] = _coerce_corrections(parsed.get("corrections")) if not out["corrected_text"]: out["ok"] = False out["error"] = "Expert returned empty corrected_text." return out async def _run_judge( *, client: LLMClient, model: str, image_b64: str, ocr_text: str, expert_a: dict, expert_b: dict, language: str, guidelines: str, temperature: float, system_override: str | None = None, user_override: str | None = None, ) -> dict: system = render_judge_system(language, guidelines, override=system_override) user_text = render_judge_user( ocr_text=ocr_text, expert_a_text=expert_a.get("corrected_text", ""), expert_a_conf=expert_a.get("confidence", 0.0), expert_a_corrections=expert_a.get("corrections", []), expert_b_text=expert_b.get("corrected_text", ""), expert_b_conf=expert_b.get("confidence", 0.0), expert_b_corrections=expert_b.get("corrections", []), override=user_override, ) user_parts = [text_part(user_text), image_part(image_b64)] res = await client.chat( model=model, system=system, user_parts=user_parts, temperature=temperature, max_tokens=4096, force_json=True, ) out = { "model": model, "ok": res.ok, "latency_s": res.latency_s, "error": res.error, "raw": res.content, "usage": res.usage, "final_text": "", "confidence": 0.0, "source": "synthesis", "rationale": "", } if not res.ok: return out parsed = _try_parse_json(res.content) if parsed is None: out["ok"] = False out["error"] = "Could not parse JSON from judge response." return out out["final_text"] = str(parsed.get("final_text") or "").rstrip() out["confidence"] = _coerce_confidence(parsed.get("confidence", 0.5)) out["source"] = str(parsed.get("source") or "synthesis") out["rationale"] = str(parsed.get("rationale") or "")[:500] if not out["final_text"]: out["ok"] = False out["error"] = "Judge returned empty final_text." return out async def run_moe_correction( *, api_key: str, image_b64: str, ocr_text: str, expert_a_model: str, expert_b_model: str, judge_model: str, language: str, guidelines: str, temperature: float = 0.1, expert_system_override: str | None = None, expert_user_override: str | None = None, judge_system_override: str | None = None, judge_user_override: str | None = None, ) -> dict: """Top-level MoE entry point. Returns a dict with: - final_text: the consensus / judged correction - source: "expert_consensus" | "judge" | "fallback_expert" | "ocr_unchanged" - confidence: float - rationale: short text - expert_a, expert_b: per-model dicts - judge: judge dict (may be None if skipped) """ async with LLMClient(api_key=api_key) as client: ea, eb = await asyncio.gather( _run_expert( client=client, model=expert_a_model, image_b64=image_b64, ocr_text=ocr_text, language=language, guidelines=guidelines, temperature=temperature, system_override=expert_system_override, user_override=expert_user_override, ), _run_expert( client=client, model=expert_b_model, image_b64=image_b64, ocr_text=ocr_text, language=language, guidelines=guidelines, temperature=temperature, system_override=expert_system_override, user_override=expert_user_override, ), ) ocr_norm = ocr_text.strip() # Both experts agree among themselves if ea["ok"] and eb["ok"] and ea["corrected_text"].strip() == eb["corrected_text"].strip(): text = ea["corrected_text"] source = "expert_consensus" if text.strip() != ocr_norm else "ocr_unchanged" return { "final_text": text, "confidence": max(ea["confidence"], eb["confidence"]), "source": source, "rationale": "Both experts produced the same output; judge skipped.", "expert_a": ea, "expert_b": eb, "judge": None, } # Both experts failed — fall back to OCR if not ea["ok"] and not eb["ok"]: return { "final_text": ocr_text, "confidence": 0.0, "source": "fallback_ocr", "rationale": f"Both experts failed: A={ea['error']}; B={eb['error']}", "expert_a": ea, "expert_b": eb, "judge": None, } # One expert failed — use the other one (no judge needed for arbitration) if not ea["ok"] or not eb["ok"]: winner = ea if ea["ok"] else eb return { "final_text": winner["corrected_text"], "confidence": winner["confidence"] * 0.8, "source": "fallback_expert", "rationale": f"Only {winner['model']} succeeded; the other expert failed.", "expert_a": ea, "expert_b": eb, "judge": None, } # Disagreement → invoke judge judge = await _run_judge( client=client, model=judge_model, image_b64=image_b64, ocr_text=ocr_text, expert_a=ea, expert_b=eb, language=language, guidelines=guidelines, temperature=temperature, system_override=judge_system_override, user_override=judge_user_override, ) if judge["ok"]: return { "final_text": judge["final_text"], "confidence": judge["confidence"], "source": "judge", "rationale": judge["rationale"], "expert_a": ea, "expert_b": eb, "judge": judge, } # Judge failed → pick the higher-confidence expert winner = ea if ea["confidence"] >= eb["confidence"] else eb return { "final_text": winner["corrected_text"], "confidence": winner["confidence"] * 0.7, "source": "fallback_expert", "rationale": f"Judge failed ({judge['error']}); using highest-confidence expert {winner['model']}.", "expert_a": ea, "expert_b": eb, "judge": judge, }