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

125 lines
4.5 KiB
Python

from typing import Any
import traceback
import json
from pydantic import BaseModel
from agent import Agent, AgentMessage
from windowizer import Windowizer, Window
from asr_filter import AsrFilterResult, AsrFilterSegment
from paths import debug_path, prompt_path
from utils import Timeline, Event
class _PreEvent(BaseModel):
"""Objects of this schema are returned by AI"""
ids: list[int]
"""IDs of merged segment"""
text: str
"""Text of the event after segments merging"""
class _EventizeResult(BaseModel):
"""Result of a single window eventizing, as returned by AI"""
preevents: list[_PreEvent]
"""PreEvents, as returned by AI"""
context: dict[str, Any]
"""Context, as returned by AI"""
class AsrEventizer:
"""This class creates a list of events from AsrFilterResult"""
DEBUG_ID = 0
def _eventize_window(self, window: Window[AsrFilterSegment]) -> _EventizeResult:
"""Eventize a single window"""
messages = [
self._system_prompt,
AgentMessage(
content=json.dumps(window.model_dump(mode="json"), indent=2, ensure_ascii=False),
role="user"
)
]
debug_dir = debug_path("AsrEventizer", AsrEventizer.DEBUG_ID)
if debug_dir:
AsrEventizer.DEBUG_ID += 1
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
retries_left = 5
while retries_left > 0:
retries_left -= 1
response = self._agent.completion(messages=messages, response_format=_EventizeResult)
if debug_dir:
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
response,
encoding="utf-8",
)
# validate data
try:
response = json.loads(response)
obj = _EventizeResult(**response)
past_ids = [e.id for e in window.past]
present_ids = [e.id for e in window.present]
usable_ids = past_ids + present_ids
# find invalid IDs
for p in obj.preevents:
for id in p.ids:
if id not in usable_ids:
raise RuntimeError("Model has tried to use ID that was not provided")
return obj
except:
traceback.print_exc()
continue
raise RuntimeError(
"Agent has failed to provide valid schema too many times"
)
def __init__(self, agent: Agent, windowizer: Windowizer[AsrFilterSegment]) -> None:
"""Create the eventizer.
Args:
- agent - agent that will be used
- windowizer - windowizer to use
"""
self._agent = agent
self._windowizer = windowizer
self._system_prompt = AgentMessage(
content=prompt_path("asr_eventizer.md").read_text(encoding="utf-8"),
role="system",
)
def eventize(self, asr_filter_result: AsrFilterResult) -> Timeline:
"""Builds event timeline from `AsrFilterResult`. Resulting timeline
consists only of `voice` events.
"""
events = []
windows = self._windowizer.windowize(asr_filter_result.segments)
context = {}
output_id = 0
last_processed_ids: list[int] = []
for window in windows:
# call the agent
window.context = context
window.past = [e for e in window.past if e.id not in last_processed_ids]
intermediate = self._eventize_window(window)
context = intermediate.context
last_processed_ids = []
# process preevents
for preevent in intermediate.preevents:
related_segments = [
ev for ev in window.past + window.present if ev.id in preevent.ids
]
last_processed_ids += [ev.id for ev in related_segments]
min_time = min(s.start for s in related_segments)
max_time = max(s.end for s in related_segments)
events.append(Event(
id=output_id,
type="voice",
timestamp=min_time,
duration=max_time - min_time,
text=preevent.text,
payload=None
))
output_id += 1
# final timeline
return Timeline(events=events)