import json import logging import traceback from typing import Literal from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator from agent import Agent, AgentMessage from paths import debug_path, prompt_path from utils import Timeline from windowizer import Window, Windowizer class _StrictModel(BaseModel): model_config = ConfigDict(extra="forbid", strict=True) class VideoReferenceEvent(_StrictModel): id: int timestamp: float duration: float text: str class UnresolvedReference(_StrictModel): ids: list[int] = Field(min_length=1) type: Literal["vis", "ocr"] text: str = Field(min_length=1) @model_validator(mode="after") def validate_ids(self) -> "UnresolvedReference": if len(self.ids) != len(set(self.ids)): raise ValueError("ids must not contain duplicates") return self class UnresolvedReferences(_StrictModel): unresolved: list[UnresolvedReference] class _ReferenceContext(_StrictModel): topic: str | None = Field(default=None, max_length=120) recent_references: list[str] = Field(default_factory=list, max_length=6) @model_validator(mode="after") def validate_recent_references(self) -> "_ReferenceContext": if any( not reference.strip() or len(reference) > 160 for reference in self.recent_references ): raise ValueError( "recent_references must contain non-empty strings up to 160 characters" ) return self class _WindowResult(_StrictModel): unresolved: list[UnresolvedReference] context: _ReferenceContext class VideoReferenceBuilder: """Find timeline regions that require visual or OCR processing.""" DEBUG_ID = 0 MAX_RETRIES = 5 def __init__( self, agent: Agent, windowizer: Windowizer[VideoReferenceEvent], ) -> None: self._agent = agent self._windowizer = windowizer self._system_prompt = AgentMessage( content=prompt_path("video_references.md").read_text(encoding="utf-8"), role="system", ) @staticmethod def _validate_result( result: _WindowResult, window: Window[VideoReferenceEvent], ) -> None: present_ids = [event.id for event in window.present] present_positions = { event_id: position for position, event_id in enumerate(present_ids) } seen: set[tuple[tuple[int, ...], str]] = set() for reference in result.unresolved: if any(event_id not in present_positions for event_id in reference.ids): raise ValueError("Reference ids must come from present") positions = [present_positions[event_id] for event_id in reference.ids] if positions != sorted(positions): raise ValueError("Reference ids must be in chronological order") key = (tuple(reference.ids), reference.type) if key in seen: raise ValueError("Duplicate reference in a single window") seen.add(key) @staticmethod def _keep_present_references( result: _WindowResult, window: Window[VideoReferenceEvent], ) -> _WindowResult: """Keep the usable part of references that cross a window boundary.""" present_positions = { event.id: position for position, event in enumerate(window.present) } references: list[UnresolvedReference] = [] seen: set[tuple[tuple[int, ...], str]] = set() for reference in result.unresolved: ids = sorted( {event_id for event_id in reference.ids if event_id in present_positions}, key=present_positions.__getitem__, ) if not ids: logging.warning( "Ignoring reference outside present window: ids=%s", reference.ids, ) continue if ids != reference.ids: logging.warning( "Trimming reference to present window: ids=%s -> %s", reference.ids, ids, ) normalized = reference.model_copy(update={"ids": ids}) key = (tuple(ids), normalized.type) if key in seen: logging.warning("Ignoring duplicate reference: ids=%s", ids) continue seen.add(key) references.append(normalized) return result.model_copy(update={"unresolved": references}) def _build_window( self, window: Window[VideoReferenceEvent], ) -> _WindowResult: messages = [ self._system_prompt, AgentMessage( content=json.dumps( window.model_dump(mode="json"), indent=2, ensure_ascii=False, ), role="user", ), ] debug_dir = debug_path("VideoReferenceBuilder", self.DEBUG_ID) if debug_dir: VideoReferenceBuilder.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=_WindowResult, ) if debug_dir: (debug_dir / f"{retries_left}-retries-left.txt").write_text( response, encoding="utf-8", ) try: result = _WindowResult.model_validate_json(response) result = self._keep_present_references(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 build(self, timeline: Timeline) -> UnresolvedReferences: events: list[VideoReferenceEvent] = [] seen_ids: set[int] = set() for event in timeline.events: if event.type != "voice": raise ValueError("Video references must be built from voice events") if event.id in seen_ids: raise ValueError(f"Duplicate event ID: {event.id}") seen_ids.add(event.id) events.append( VideoReferenceEvent( id=event.id, timestamp=event.timestamp, duration=event.duration, text=event.text, ) ) unresolved: list[UnresolvedReference] = [] context = _ReferenceContext() for window in self._windowizer.windowize(events): window.context = context.model_dump(mode="json") result = self._build_window(window) unresolved.extend(result.unresolved) context = result.context return UnresolvedReferences(unresolved=unresolved)