Files
2026-linux-sumka/asr_eventizer.py
2026-09-17 00:46:04 +03:00

111 lines
3.9 KiB
Python

from dataclasses import asdict
from typing import Any
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"""
def _eventize_window(self, window: Window[AsrFilterSegment]) -> _EventizeResult:
"""Eventize a single window"""
messages = [
self._system_prompt,
AgentMessage(
content=json.dumps(asdict(window), indent=2, ensure_ascii=False),
role="user"
)
]
retries_left = 5
while retries_left > 0:
retries_left -= 1
response = self._agent.completion(messages=messages)
# 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:
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.json", "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
# 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
))
# final timeline
return Timeline(events=events)