Files
2026-linux-sumka/asr_eventizer.py
2026-09-17 01:44:49 +03:00

128 lines
4.5 KiB
Python

from typing import Any
import traceback
import shutil
import os
import json
from pydantic import BaseModel
from agent import Agent, AgentMessage
from windowizer import Windowizer, Window
from asr_filter import AsrFilterResult, AsrFilterSegment
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 = None
if os.path.isdir("debug"):
debug_dir = f"debug/AsrEventizer/{AsrEventizer.DEBUG_ID}"
AsrEventizer.DEBUG_ID += 1
os.makedirs(debug_dir, exist_ok=True)
if debug_dir:
with open(f"{debug_dir}/request.txt", "w") as f:
f.write(messages[1].content)
retries_left = 5
while retries_left > 0:
retries_left -= 1
response = self._agent.completion(messages=messages)
if debug_dir:
with open(f"{debug_dir}/{retries_left}-retries-left.txt", "w") as f:
f.write(response)
# 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
with open("prompts/asr_eventizer.md", "r") as f:
self._system_prompt = AgentMessage(
content=f.read(),
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)