Added video references search (codex)

This commit is contained in:
2026-09-23 00:56:02 +03:00
parent d2bdb37773
commit 0da97ca225
4 changed files with 292 additions and 4 deletions

174
video_references.py Normal file
View File

@@ -0,0 +1,174 @@
import json
import os
import traceback
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator
from agent import Agent, AgentMessage
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
with open("prompts/video_references.md", "r", encoding="utf-8") as f:
self._system_prompt = AgentMessage(content=f.read(), 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)
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 = None
if os.path.isdir("debug"):
debug_dir = f"debug/VideoReferenceBuilder/{self.DEBUG_ID}"
VideoReferenceBuilder.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=_WindowResult,
)
if debug_dir:
with open(
f"{debug_dir}/{retries_left}-retries-left.txt",
"w",
encoding="utf-8",
) as f:
f.write(response)
try:
result = _WindowResult.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) -> 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)