Files
2026-linux-sumka/structure_refiner.py

133 lines
4.7 KiB
Python

import json
import logging
import traceback
from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator
from agent import Agent, AgentMessage
from paths import debug_path, prompt_path
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
self._system_prompt = AgentMessage(
content=prompt_path("structure_refiner.md").read_text(encoding="utf-8"),
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")
@staticmethod
def _keep_present_images(
result: _RefineResult,
window: Window[StructureElement],
) -> _RefineResult:
"""Discard image elements invented from context windows."""
input_image_ids = {
element.event_id
for element in window.present
if isinstance(element, ImageElement)
}
elements: list[StructureElement] = []
for element in result.elements:
if isinstance(element, ImageElement) and element.event_id not in input_image_ids:
logging.warning("Ignoring image outside present window: id=%s", element.event_id)
continue
elements.append(element)
return result.model_copy(update={"elements": elements})
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 = debug_path("StructureRefiner", StructureRefiner.DEBUG_ID)
if debug_dir:
StructureRefiner.DEBUG_ID += 1
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
retries_left = self.MAX_RETRIES
while retries_left > 0:
retries_left -= 1
response = self._agent.completion(
messages=messages,
response_format=_RefineResult,
)
if debug_dir:
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
response,
encoding="utf-8",
)
try:
result = _RefineResult.model_validate_json(response)
result = self._keep_present_images(result, window)
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)