import json import os import traceback from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator from agent import Agent, AgentMessage from utils import Event, Timeline from windowizer import Window, Windowizer class _StrictModel(BaseModel): """Base model for schemas returned by the agent.""" model_config = ConfigDict(extra="forbid", strict=True) class HeadingElement(_StrictModel): """A document section heading.""" type: Literal["heading"] level: Literal[2, 3, 4] text: str = Field(min_length=1) class ParagraphElement(_StrictModel): """A paragraph containing one complete thought.""" type: Literal["paragraph"] text: str = Field(min_length=1) class UnorderedListElement(_StrictModel): """A list whose item order is not significant.""" type: Literal["unordered"] items: list[str] = Field(min_length=1) @model_validator(mode="after") def validate_items(self) -> "UnorderedListElement": if any(not item.strip() for item in self.items): raise ValueError("List items must not be empty") return self class OrderedListElement(_StrictModel): """A list whose item order is significant.""" type: Literal["ordered"] items: list[str] = Field(min_length=1) @model_validator(mode="after") def validate_items(self) -> "OrderedListElement": if any(not item.strip() for item in self.items): raise ValueError("List items must not be empty") return self class DefinitionElement(_StrictModel): """A definition of a term or concept.""" type: Literal["definition"] term: str = Field(min_length=1) text: str = Field(min_length=1) class ImportantElement(_StrictModel): """An important statement that should stand out in the document.""" type: Literal["important"] text: str = Field(min_length=1) class ImageElement(_StrictModel): """A reference to an image stored in a visual timeline event.""" type: Literal["image"] event_id: int StructureElement = ( HeadingElement | ParagraphElement | UnorderedListElement | OrderedListElement | DefinitionElement | ImportantElement | ImageElement ) class Structure(_StrictModel): """Structure of the future Markdown document.""" elements: list[StructureElement] class _Pending(_StrictModel): """An unfinished thought carried over to the next window.""" event_ids: list[int] = Field(min_length=1) summary: str = Field(min_length=1) @model_validator(mode="after") def validate_event_ids(self) -> "_Pending": if len(self.event_ids) != len(set(self.event_ids)): raise ValueError("pending.event_ids must not contain duplicates") return self class _BuilderContext(_StrictModel): """Small amount of document state preserved between agent calls.""" 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) pending: _Pending | None = None @model_validator(mode="after") def validate_recent_headings(self) -> "_BuilderContext": 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 _BuildResult(_StrictModel): """Result returned by the agent for a single window.""" elements: list[StructureElement] context: _BuilderContext class StructureBuilder: """Build a document structure from a timeline of lecture events.""" DEBUG_ID = 0 MAX_RETRIES = 5 def __init__(self, agent: Agent, windowizer: Windowizer[Event]) -> None: """Create a structure builder. Args: agent: Agent used to transform timeline windows. windowizer: Windowizer used to split timeline events. """ self._agent = agent self._windowizer = windowizer with open("prompts/structure_builder.md", "r", encoding="utf-8") as f: self._system_prompt = AgentMessage(content=f.read(), role="system") @staticmethod def _validate_result(result: _BuildResult, window: Window[Event]) -> None: """Validate constraints that depend on the current input window.""" usable_visual_ids = { event.id for event in window.past + window.present if event.type == "vis" } for element in result.elements: if isinstance(element, ImageElement) and element.event_id not in usable_visual_ids: raise ValueError( "image.event_id must reference a vis event from past or present" ) if result.context.pending is None: return window_ids = { event.id for event in window.past + window.present + window.future } invalid_pending_ids = ( set(result.context.pending.event_ids) - window_ids ) if invalid_pending_ids: raise ValueError( "pending.event_ids must reference events from the current window" ) def _build_window(self, window: Window[Event]) -> _BuildResult: """Build document elements from a single timeline window.""" 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/StructureBuilder/{StructureBuilder.DEBUG_ID}" StructureBuilder.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=_BuildResult) if debug_dir: with open( f"{debug_dir}/{retries_left}-retries-left.txt", "w", encoding="utf-8", ) as f: f.write(response) try: result = _BuildResult.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 build(self, timeline: Timeline) -> Structure: """Build the ordered structure of a future Markdown document.""" elements: list[StructureElement] = [] context = _BuilderContext() for window in self._windowizer.windowize(timeline.events): window.context = context.model_dump(mode="json") intermediate = self._build_window(window) elements.extend(intermediate.elements) context = intermediate.context return Structure(elements=elements)