import json import os import traceback from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator from agent import Agent, AgentMessage from structure_builder import ImageElement, Structure, StructureElement from windowizer import Window, Windowizer class _StrictModel(BaseModel): model_config = ConfigDict(extra="forbid", strict=True) class _RefinerContext(_StrictModel): current_section: str | None = Field(default=None, max_length=120) current_subsection: str | None = Field(default=None, max_length=120) recent_headings: list[str] = Field(default_factory=list, max_length=6) @model_validator(mode="after") def validate_recent_headings(self) -> "_RefinerContext": if any(not heading.strip() or len(heading) > 120 for heading in self.recent_headings): raise ValueError("recent_headings must contain non-empty strings up to 120 characters") return self class _RefineResult(_StrictModel): elements: list[StructureElement] context: _RefinerContext class StructureRefiner: """Perform a final editing pass over a document structure.""" DEBUG_ID = 0 MAX_RETRIES = 5 def __init__(self, agent: Agent, windowizer: Windowizer[StructureElement]) -> None: self._agent = agent self._windowizer = windowizer with open("prompts/structure_refiner.md", "r", encoding="utf-8") as f: self._system_prompt = AgentMessage(content=f.read(), role="system") @staticmethod def _validate_result( result: _RefineResult, window: Window[StructureElement], ) -> None: input_image_ids = { element.event_id for element in window.present if isinstance(element, ImageElement) } for element in result.elements: if isinstance(element, ImageElement) and element.event_id not in input_image_ids: raise ValueError("image.event_id must come from present") def _refine_window(self, window: Window[StructureElement]) -> _RefineResult: messages = [ self._system_prompt, AgentMessage( content=json.dumps( window.model_dump(mode="json"), indent=2, ensure_ascii=False, ), role="user", ), ] debug_dir = None if os.path.isdir("debug"): debug_dir = f"debug/StructureRefiner/{StructureRefiner.DEBUG_ID}" StructureRefiner.DEBUG_ID += 1 os.makedirs(debug_dir, exist_ok=True) with open(f"{debug_dir}/request.txt", "w", encoding="utf-8") as f: f.write(messages[1].content) retries_left = self.MAX_RETRIES while retries_left > 0: retries_left -= 1 response = self._agent.completion( messages=messages, response_format=_RefineResult, ) if debug_dir: with open( f"{debug_dir}/{retries_left}-retries-left.txt", "w", encoding="utf-8", ) as f: f.write(response) try: result = _RefineResult.model_validate_json(response) self._validate_result(result, window) return result except (ValidationError, ValueError, TypeError): traceback.print_exc() raise RuntimeError("Agent has failed to provide valid schema too many times") def refine(self, structure: Structure) -> Structure: elements: list[StructureElement] = [] context = _RefinerContext() for window in self._windowizer.windowize(structure.elements): window.context = context.model_dump(mode="json") intermediate = self._refine_window(window) elements.extend(intermediate.elements) context = intermediate.context return Structure(elements=elements)