diff --git a/services/ai_orchestrator_service/app.py b/services/ai_orchestrator_service/app.py index 3fc1e2b..dc988b3 100644 --- a/services/ai_orchestrator_service/app.py +++ b/services/ai_orchestrator_service/app.py @@ -8,11 +8,12 @@ import math import os import re import time +from threading import Lock from typing import Any import httpx from fastapi import Depends, FastAPI, HTTPException, Query -from sqlalchemy import select +from sqlalchemy import func, select from services.shared.core import Role, new_id, utc_now_iso from services.shared.db import get_session @@ -1909,13 +1910,61 @@ def _article_snippet(article: KBArticleRow, limit: int = 240) -> str: return f"{raw[: limit - 3]}..." -def _kb_search(session, text: str, *, language: str | None = None) -> list[KBArticleRow]: - if not str(text or "").strip(): - return [] +_KB_ROWS_CACHE: dict[str, tuple[tuple[int, int, str, int], list[KBArticleRow]]] = {} +_KB_ROWS_CACHE_LOCK = Lock() + + +def _kb_rows_fingerprint(session, language: str | None) -> tuple[int, int, str, int]: + # A cheap aggregate (row content lengths, computed server-side - no + # title/body/tags_json actually transferred) that changes on any + # insert, update, or delete, so the cache below can never serve stale + # results. count/max_id/max_updated_at alone aren't sufficient: ids can + # be reused after a row is deleted and updated_at has 1-second + # resolution, so two different row sets can otherwise share a + # fingerprint if they happen to be written within the same second. + content_length = func.length(KBArticleRow.title) + func.length(KBArticleRow.body) + func.length( + KBArticleRow.tags_json + ) + stmt = select( + func.count(KBArticleRow.id), + func.max(KBArticleRow.id), + func.max(KBArticleRow.updated_at), + func.sum(content_length), + ) + if language is not None: + stmt = stmt.where(KBArticleRow.language == normalize_kb_language(language)) + count, max_id, max_updated_at, total_length = session.execute(stmt).one() + return (int(count or 0), int(max_id or 0), str(max_updated_at or ""), int(total_length or 0)) + + +def _load_kb_rows_cached(session, language: str | None) -> list[KBArticleRow]: + # Every voice/chat turn re-runs KB search, and this table rarely changes + # mid-call, so avoid re-fetching + re-transferring the whole table from + # the DB (title/body/tags_json for every article) when nothing changed. + cache_key = normalize_kb_language(language) if language is not None else "__all__" + fingerprint = _kb_rows_fingerprint(session, language) + + with _KB_ROWS_CACHE_LOCK: + cached = _KB_ROWS_CACHE.get(cache_key) + if cached is not None and cached[0] == fingerprint: + return cached[1] + stmt = select(KBArticleRow).order_by(KBArticleRow.id.desc()) if language is not None: stmt = stmt.where(KBArticleRow.language == normalize_kb_language(language)) - rows = session.execute(stmt).scalars().all() + rows = list(session.execute(stmt).scalars().all()) + for row in rows: + session.expunge(row) + + with _KB_ROWS_CACHE_LOCK: + _KB_ROWS_CACHE[cache_key] = (fingerprint, rows) + return rows + + +def _kb_search(session, text: str, *, language: str | None = None) -> list[KBArticleRow]: + if not str(text or "").strip(): + return [] + rows = _load_kb_rows_cached(session, language) return search_kb_rows(rows, text, limit=_ai_max_kb_results()) diff --git a/services/ai_voice_runtime_service/media_runtime.py b/services/ai_voice_runtime_service/media_runtime.py index 40cfa6f..01879ce 100644 --- a/services/ai_voice_runtime_service/media_runtime.py +++ b/services/ai_voice_runtime_service/media_runtime.py @@ -193,6 +193,7 @@ class AudioSocketMediaRuntime: self._partial_asr_min_ms = 320 self._immediate_ack_min_ms = 700 self._v2_ack_post_gap_seconds = 0.10 + self._v1_ack_wait_seconds = 0.6 self._partial_poll_interval_seconds = 0.20 self._streaming_asr_partial_poll_enabled = ( str(os.getenv("AI_VOICE_V2_STREAMING_ASR_PARTIAL_POLL_ENABLED", "0")).strip().lower() @@ -1619,7 +1620,18 @@ class AudioSocketMediaRuntime: early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12) decision = early_plan_decision if early_plan_decision is not None else await decision_task else: - decision = await decision_task + try: + decision = await asyncio.wait_for(asyncio.shield(decision_task), timeout=self._v1_ack_wait_seconds) + except asyncio.TimeoutError: + if not actor.early_ack_started: + await self._emit_early_ack( + actor, + language=transcription.language or actor.registration.language, + metadata=metadata, + ack_source="v1_decision_timeout", + ack_kind="unknown", + ) + decision = await decision_task handoff_task: asyncio.Task | None = None if decision.needs_handoff: await self._set_actor_state(actor, "handoff_requested", decision.handoff_reason) diff --git a/tests/test_ai_voice_media_runtime.py b/tests/test_ai_voice_media_runtime.py index 09407bf..2a9e56e 100644 --- a/tests/test_ai_voice_media_runtime.py +++ b/tests/test_ai_voice_media_runtime.py @@ -187,6 +187,103 @@ def test_media_runtime_streams_greeting_and_turn(): assert errors == [] +def test_media_runtime_plays_filler_ack_when_v1_decision_is_slow(): + registrations: dict[str, MediaRegistration] = {} + delivered: list[tuple[str, str, bool]] = [] + + media_uuid = str(uuid.uuid4()) + registrations[media_uuid] = MediaRegistration( + voice_session_id="avs_media_runtime_slow_v1", + call_id="call_media_runtime_slow_v1", + interaction_id="int_media_runtime_slow_v1", + ai_session_id="ais_media_runtime_slow_v1", + language="ru", + media_uuid=media_uuid, + ) + + def _slow_process_turn(session_id, transcript_text, language, barge_in, metadata): + del transcript_text, barge_in, metadata + time.sleep(0.9) + return VoiceAITurnDecisionOut( + language=language or "ru", + intent="answer", + reply_text="reply", + confidence=0.9, + needs_handoff=False, + handoff_reason=None, + case_action="keep_open", + kb_refs=[], + summary_text="reply ready", + model="stub-voice", + latency_ms=1, + status="active", + ) + + runtime = AudioSocketMediaRuntime( + enabled=True, + host="127.0.0.1", + port=0, + frame_ms=20, + idle_timeout_seconds=2.0, + registration_wait_timeout_seconds=0.5, + min_speech_ms=40, + trailing_silence_ms=40, + max_turn_ms=400, + asr_provider=_StubASRProvider(), + tts_provider=_StubTTSProvider(), + load_registration_by_media_uuid=lambda value: registrations.get(value), + mark_media_connected=lambda session_id, value: None, + mark_media_ended=lambda session_id, reason: None, + touch_media_frame=lambda session_id: None, + set_state=lambda session_id, state, handoff_reason, metadata: None, + get_pending_greeting=lambda session_id: "greeting" if session_id == "avs_media_runtime_slow_v1" else None, + mark_reply_delivered=lambda session_id, text, is_greeting: delivered.append((session_id, text, is_greeting)), + plan_reply=lambda session_id, text, metadata, kind: None, + process_turn=_slow_process_turn, + request_handoff=lambda session_id, customer_request_text, decision: None, + handle_media_error=lambda session_id, message, metadata: None, + ) + + async def _scenario() -> None: + await runtime.start() + port = runtime._server.sockets[0].getsockname()[1] + + reader, writer = await asyncio.open_connection("127.0.0.1", port) + writer.write(encode_packet(AUDIO_SOCKET_PACKET_UUID, uuid.UUID(media_uuid).bytes)) + await writer.drain() + + packet_type, _ = await read_packet(reader, timeout=2.0) + assert packet_type == AUDIO_SOCKET_PACKET_PCM16 + await asyncio.sleep(0.15) + + speech_frame = (1000).to_bytes(2, "little", signed=True) * 160 + silence_frame = b"\x00\x00" * 160 + for _ in range(2): + writer.write(encode_audio_packet(speech_frame)) + for _ in range(2): + writer.write(encode_audio_packet(silence_frame)) + await writer.drain() + + packet_type, _ = await read_packet(reader, timeout=2.0) + assert packet_type == AUDIO_SOCKET_PACKET_PCM16 + + for _ in range(60): + if len(delivered) >= 3: + break + await asyncio.sleep(0.05) + + writer.close() + await writer.wait_closed() + await asyncio.sleep(0.2) + await runtime.stop() + + asyncio.run(_scenario()) + + assert delivered[0] == ("avs_media_runtime_slow_v1", "greeting", True) + assert delivered[1] == ("avs_media_runtime_slow_v1", "Секунду.", False) + assert delivered[2] == ("avs_media_runtime_slow_v1", "reply", False) + + def test_media_runtime_speaks_technical_fallback_when_asr_transcribe_fails(): registrations: dict[str, MediaRegistration] = {} delivered: list[tuple[str, str, bool]] = []