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 crop_box: list[int] | 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.crop_box) ): raise ValueError("An unmatched result must not contain image data") return self class _OCRResult(_StrictModel): matched: bool ocr_text: str | None @model_validator(mode="after") def validate_match(self) -> "_OCRResult": if self.matched and (self.ocr_text is None or not self.ocr_text.strip()): raise ValueError("A matched OCR result must contain ocr_text") if not self.matched and self.ocr_text is not None: raise ValueError("An unmatched OCR result must not contain ocr_text") return self @dataclass(frozen=True) class _Frame: timestamp: float path: str class ReferenceResolver: """Resolve visual references against sampled video frames.""" DEBUG_ID = 0 OCR_DEBUG_ID = 0 MAX_RETRIES = 3 CROP_PADDING = 30 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") with open("prompts/reference_resolver_ocr.md", "r", encoding="utf-8") as f: self._ocr_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.crop_box is not None: raise ValueError("A visual match must not contain crop_box") else: if result.description is not None: raise ValueError("An OCR match must not contain description") if result.crop_box is None or len(result.crop_box) != 4: raise ValueError("An OCR match must contain crop_box") y_min, x_min, y_max, x_max = result.crop_box if not ( 0 <= y_min < y_max <= 1000 and 0 <= x_min < x_max <= 1000 ): raise ValueError("crop_box must be within 0..1000") 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="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 _crop_frame( self, frame: _Frame, crop_box: list[int], ) -> str | None: y_min, x_min, y_max, x_max = crop_box padding = self.CROP_PADDING y_min = max(0, y_min - padding) x_min = max(0, x_min - padding) y_max = min(1000, y_max + padding) x_max = min(1000, x_max + padding) width = (x_max - x_min) / 1000 height = (y_max - y_min) / 1000 x = x_min / 1000 y = y_min / 1000 crop_path = str(Path(frame.path).with_name(f"{Path(frame.path).stem}_crop.jpg")) crop_filter = ( f"crop=trunc(iw*{width:.6f}/2)*2:" f"trunc(ih*{height:.6f}/2)*2:" f"iw*{x:.6f}:ih*{y:.6f}," "scale=1600:1600:force_original_aspect_ratio=decrease:" "force_divisible_by=2" ) command = [ "ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-ss", f"{frame.timestamp:.3f}", "-i", self._video_path, "-frames:v", "1", "-vf", crop_filter, "-q:v", "2", crop_path, ] try: subprocess.run(command, check=True) except subprocess.CalledProcessError: logging.warning("Failed to crop frame at %.3f", frame.timestamp) return None return crop_path if os.path.isfile(crop_path) else None def _recognize_ocr( self, reference: UnresolvedReference, frame: _Frame, crop_path: str, ) -> str | None: request = { "reference": reference.model_dump(mode="json"), "timestamp": frame.timestamp, } messages = [ self._ocr_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/ReferenceResolverOCR/{ReferenceResolver.OCR_DEBUG_ID}" ReferenceResolver.OCR_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) shutil.copyfile(crop_path, f"{debug_dir}/crop.jpg") retries_left = self.MAX_RETRIES while retries_left > 0: retries_left -= 1 response = self._agent.completion_with_images( messages, [crop_path], detail="high", response_format=_OCRResult, ) if debug_dir: with open( f"{debug_dir}/{retries_left}-retries-left.txt", "w", encoding="utf-8", ) as f: f.write(response) try: result = _OCRResult.model_validate_json(response) if not result.matched: return None assert result.ocr_text is not None return result.ocr_text.strip() except (ValidationError, ValueError, TypeError): traceback.print_exc() raise RuntimeError("Agent has failed to provide valid OCR schema too many times") def _find_frame( self, reference: UnresolvedReference, frames: list[_Frame], ) -> tuple[_Frame, _FrameResult, str | None] | None: for level in self._search_levels(len(frames)): batch = [frames[index] for index in level] result = self._check_frames(reference, batch) if not result.matched: continue assert result.image_number is not None frame = batch[result.image_number - 1] if reference.type == "vis": return frame, result, None assert result.crop_box is not None crop_path = self._crop_frame(frame, result.crop_box) if crop_path is None: continue ocr_text = self._recognize_ocr(reference, frame, crop_path) if ocr_text is not None: return frame, result, ocr_text logging.info( "OCR crop at %.3f is unreadable, continuing search", frame.timestamp, ) return None def _resolved_event( self, reference_index: int, reference: UnresolvedReference, frame: _Frame, result: _FrameResult, ocr_text: str | None, ) -> Event: if reference.type == "ocr": assert ocr_text is not None return Event( id=-1, type="ocr", timestamp=frame.timestamp, duration=0.0, text=reference.text, payload=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, ocr_text = match resolved.append( self._resolved_event( reference_index, reference, frame, result, ocr_text, ) ) 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, )