Files
2026-linux-sumka/reference_resolver.py
2026-09-23 03:26:09 +03:00

500 lines
17 KiB
Python

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 paths import debug_path, prompt_path
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
self._system_prompt = AgentMessage(
content=prompt_path("reference_resolver.md").read_text(encoding="utf-8"),
role="system",
)
self._ocr_prompt = AgentMessage(
content=prompt_path("reference_resolver_ocr.md").read_text(encoding="utf-8"),
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 = debug_path("ReferenceResolver", ReferenceResolver.DEBUG_ID)
if debug_dir:
ReferenceResolver.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_with_images(
messages,
[frame.path for frame in frames],
detail="low",
response_format=_FrameResult,
)
if debug_dir:
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
response,
encoding="utf-8",
)
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 = debug_path("ReferenceResolverOCR", ReferenceResolver.OCR_DEBUG_ID)
if debug_dir:
ReferenceResolver.OCR_DEBUG_ID += 1
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
shutil.copyfile(crop_path, 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:
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
response,
encoding="utf-8",
)
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,
)