215 lines
7.1 KiB
Python
215 lines
7.1 KiB
Python
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)
|