import json import logging import math import os import shutil import subprocess import tempfile import traceback from dataclasses import dataclass from pathlib import Path from typing import Literal from pydantic import BaseModel, ConfigDict, ValidationError, model_validator from agent import Agent, AgentMessage from utils import Event, Timeline from video_references import UnresolvedReference, UnresolvedReferences class _StrictModel(BaseModel): model_config = ConfigDict(extra="forbid", strict=True) class FailedReference(UnresolvedReference): reason: Literal["no_frames", "not_found"] class ReferenceResolutionResult(_StrictModel): events: list[Event] unresolved: list[FailedReference] class _FrameResult(_StrictModel): matched: bool image_number: int | None description: str | None ocr_text: str | None @model_validator(mode="after") def validate_match(self) -> "_FrameResult": if self.matched and (self.image_number is None or self.image_number < 1): raise ValueError("A matched result must contain image_number") if not self.matched and any( value is not None for value in (self.image_number, self.description, self.ocr_text) ): raise ValueError("An unmatched result must not contain image data") return self @dataclass(frozen=True) class _Frame: timestamp: float path: str class ReferenceResolver: """Resolve visual references against sampled video frames.""" DEBUG_ID = 0 MAX_RETRIES = 3 def __init__( self, agent: Agent, *, video_path: str = "video.mp4", image_dir: str = "images", frame_interval: float = 1.0, max_frames: int = 9, ) -> None: if frame_interval <= 0: raise ValueError("frame_interval must be positive") if max_frames < 1: raise ValueError("max_frames must be positive") self._agent = agent self._video_path = video_path self._image_dir = Path(image_dir) self._frame_interval = frame_interval self._max_frames = max_frames with open("prompts/reference_resolver.md", "r", encoding="utf-8") as f: self._system_prompt = AgentMessage(content=f.read(), role="system") def _sampling_timestamps(self, start: float, end: float) -> list[float]: start = max(0.0, start) end = max(start, end) duration = end - start if duration <= 0.05: return [start] count = min( max(1, math.ceil(duration / self._frame_interval)), self._max_frames, ) step = duration / count return [start + step * (index + 0.5) for index in range(count)] def _extract_frames( self, timestamps: list[float], output_dir: str, ) -> list[_Frame]: frames: list[_Frame] = [] for index, timestamp in enumerate(timestamps): output_path = os.path.join(output_dir, f"frame_{index:04d}.jpg") command = [ "ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-ss", f"{timestamp:.3f}", "-i", self._video_path, "-frames:v", "1", "-vf", "scale=1280:1280:force_original_aspect_ratio=decrease", "-q:v", "3", output_path, ] try: subprocess.run(command, check=True) except subprocess.CalledProcessError: logging.warning("Failed to extract frame at %.3f", timestamp) continue if os.path.isfile(output_path): frames.append(_Frame(timestamp=timestamp, path=output_path)) return frames @staticmethod def _search_levels(frame_count: int) -> list[list[int]]: if frame_count == 0: return [] first = list(dict.fromkeys([0, (frame_count - 1) // 2, frame_count - 1])) levels = [first] selected = set(first) while len(selected) < frame_count: next_level: list[int] = [] ordered = sorted(selected) for left, right in zip(ordered, ordered[1:]): if right - left <= 1: continue midpoint = (left + right) // 2 if midpoint not in selected: next_level.append(midpoint) if not next_level: break levels.append(next_level) selected.update(next_level) return levels @staticmethod def _validate_frame_result( result: _FrameResult, reference: UnresolvedReference, frame_count: int, ) -> None: if not result.matched: return if result.image_number is None or result.image_number > frame_count: raise ValueError("image_number does not exist in the current batch") if reference.type == "vis": if result.description is None or not result.description.strip(): raise ValueError("A visual match must contain description") if result.ocr_text is not None: raise ValueError("A visual match must not contain ocr_text") else: if result.ocr_text is None or not result.ocr_text.strip(): raise ValueError("An OCR match must contain ocr_text") if result.description is not None: raise ValueError("An OCR match must not contain description") def _check_frames( self, reference: UnresolvedReference, frames: list[_Frame], ) -> _FrameResult: request = { "reference": reference.model_dump(mode="json"), "images": [ { "image_number": index, "timestamp": frame.timestamp, } for index, frame in enumerate(frames, start=1) ], } messages = [ self._system_prompt, AgentMessage( content=json.dumps(request, indent=2, ensure_ascii=False), role="user", ), ] debug_dir = None if os.path.isdir("debug"): debug_dir = f"debug/ReferenceResolver/{ReferenceResolver.DEBUG_ID}" ReferenceResolver.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_with_images( messages, [frame.path for frame in frames], detail="high" if reference.type == "ocr" else "low", response_format=_FrameResult, ) if debug_dir: with open( f"{debug_dir}/{retries_left}-retries-left.txt", "w", encoding="utf-8", ) as f: f.write(response) try: result = _FrameResult.model_validate_json(response) self._validate_frame_result(result, reference, len(frames)) return result except (ValidationError, ValueError, TypeError): traceback.print_exc() raise RuntimeError("Agent has failed to provide valid schema too many times") def _find_frame( self, reference: UnresolvedReference, frames: list[_Frame], ) -> tuple[_Frame, _FrameResult] | None: for level in self._search_levels(len(frames)): batch = [frames[index] for index in level] result = self._check_frames(reference, batch) if result.matched: assert result.image_number is not None return batch[result.image_number - 1], result return None def _resolved_event( self, reference_index: int, reference: UnresolvedReference, frame: _Frame, result: _FrameResult, ) -> Event: if reference.type == "ocr": assert result.ocr_text is not None return Event( id=-1, type="ocr", timestamp=frame.timestamp, duration=0.0, text=reference.text, payload=result.ocr_text, ) assert result.description is not None self._image_dir.mkdir(parents=True, exist_ok=True) timestamp_ms = round(frame.timestamp * 1000) destination = self._image_dir / ( f"reference_{reference_index:04d}_{timestamp_ms:010d}.jpg" ) shutil.copyfile(frame.path, destination) return Event( id=-1, type="vis", timestamp=frame.timestamp, duration=0.0, text=result.description, payload=destination.as_posix(), ) def resolve( self, timeline: Timeline, references: UnresolvedReferences, ) -> ReferenceResolutionResult: events_by_id: dict[int, Event] = {} positions: dict[int, int] = {} for position, event in enumerate(timeline.events): if event.type != "voice": raise ValueError("Reference resolver expects voice events") if event.id in events_by_id: raise ValueError(f"Duplicate event ID: {event.id}") events_by_id[event.id] = event positions[event.id] = position resolved: list[Event] = [] failed: list[FailedReference] = [] frame_cache: dict[tuple[int, ...], list[_Frame]] = {} with tempfile.TemporaryDirectory(prefix="sumka-frames-") as temp_dir: for reference_index, reference in enumerate(references.unresolved): if any(event_id not in events_by_id for event_id in reference.ids): raise ValueError("Reference contains an unknown event ID") reference_positions = [positions[event_id] for event_id in reference.ids] if reference_positions != sorted(reference_positions): raise ValueError("Reference ids must be in chronological order") cache_key = tuple(reference.ids) frames = frame_cache.get(cache_key) if frames is None: related = [events_by_id[event_id] for event_id in reference.ids] start = min(event.timestamp for event in related) end = max( event.timestamp + max(0.0, event.duration) for event in related ) timestamps = self._sampling_timestamps(start, end) reference_dir = os.path.join(temp_dir, f"reference_{reference_index:04d}") os.makedirs(reference_dir, exist_ok=True) frames = self._extract_frames(timestamps, reference_dir) frame_cache[cache_key] = frames if not frames: failed.append(FailedReference( **reference.model_dump(mode="json"), reason="no_frames", )) continue match = self._find_frame(reference, frames) if match is None: failed.append(FailedReference( **reference.model_dump(mode="json"), reason="not_found", )) continue frame, result = match resolved.append( self._resolved_event( reference_index, reference, frame, result, ) ) all_events = timeline.events + resolved all_events.sort(key=lambda event: (event.timestamp, event.type != "voice")) voice_ids: dict[int, int] = { event.id: event_id for event_id, event in enumerate(all_events) if event.type == "voice" } final_events = [ event.model_copy(update={"id": event_id}) for event_id, event in enumerate(all_events) ] final_failed = [ reference.model_copy(update={ "ids": [voice_ids[event_id] for event_id in reference.ids] }) for reference in failed ] return ReferenceResolutionResult( events=final_events, unresolved=final_failed, )