feat(voice): add Yandex streaming ASR
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from services.ai_voice_runtime_service.audiosocket import pcm16le_to_wav_bytes
|
||||
from services.ai_voice_runtime_service.providers import asr as asr_module
|
||||
|
||||
@@ -198,10 +200,133 @@ def test_yandex_buffered_streaming_provider_buffers_pcm_until_finalize():
|
||||
assert calls == [{"audio_bytes": b"\x10\x00\x20\x00", "language_hint": "ru"}]
|
||||
|
||||
|
||||
def test_yandex_grpc_streaming_provider_returns_partial_and_final():
|
||||
calls: list[dict] = []
|
||||
|
||||
class _Message:
|
||||
def __init__(self, **kwargs) -> None:
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
def HasField(self, name: str) -> bool:
|
||||
return getattr(self, name, None) is not None
|
||||
|
||||
class _RawAudio(_Message):
|
||||
LINEAR16_PCM = 1
|
||||
|
||||
class _TextNormalizationOptions(_Message):
|
||||
TEXT_NORMALIZATION_ENABLED = 1
|
||||
PHONE_FORMATTING_MODE_DISABLED = 1
|
||||
|
||||
class _LanguageRestrictionOptions(_Message):
|
||||
WHITELIST = 1
|
||||
|
||||
class _RecognitionModelOptions(_Message):
|
||||
REAL_TIME = 1
|
||||
|
||||
class _DefaultEouClassifier(_Message):
|
||||
HIGH = 2
|
||||
|
||||
class _FakeSttPb2:
|
||||
RawAudio = _RawAudio
|
||||
AudioFormatOptions = _Message
|
||||
TextNormalizationOptions = _TextNormalizationOptions
|
||||
LanguageRestrictionOptions = _LanguageRestrictionOptions
|
||||
RecognitionModelOptions = _RecognitionModelOptions
|
||||
DefaultEouClassifier = _DefaultEouClassifier
|
||||
EouClassifierOptions = _Message
|
||||
StreamingOptions = _Message
|
||||
StreamingRequest = _Message
|
||||
AudioChunk = _Message
|
||||
Eou = _Message
|
||||
|
||||
class _FakeChannel:
|
||||
def close(self) -> None:
|
||||
calls.append({"method": "close"})
|
||||
|
||||
class _FakeGrpc:
|
||||
@staticmethod
|
||||
def ssl_channel_credentials():
|
||||
return "ssl"
|
||||
|
||||
@staticmethod
|
||||
def secure_channel(target: str, credentials):
|
||||
calls.append({"method": "secure_channel", "target": target, "credentials": credentials})
|
||||
return _FakeChannel()
|
||||
|
||||
class _RecognizerStub:
|
||||
def __init__(self, channel) -> None:
|
||||
calls.append({"method": "stub_init", "channel": channel})
|
||||
|
||||
def RecognizeStreaming(self, request_iterator, *, metadata, timeout):
|
||||
calls.append({"method": "recognize", "metadata": metadata, "timeout": timeout})
|
||||
for request in request_iterator:
|
||||
calls.append({"method": "request", "request": request})
|
||||
if getattr(request, "chunk", None) is not None:
|
||||
yield _Message(
|
||||
partial=_Message(
|
||||
alternatives=[
|
||||
_Message(text="need schedule", confidence=0.7, languages=[]),
|
||||
]
|
||||
)
|
||||
)
|
||||
continue
|
||||
if getattr(request, "eou", None) is not None:
|
||||
yield _Message(
|
||||
final_refinement=_Message(
|
||||
normalized_text=_Message(
|
||||
alternatives=[
|
||||
_Message(text="need schedule in almaty", confidence=0.91, languages=[]),
|
||||
]
|
||||
)
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
class _FakeGrpcPb2:
|
||||
RecognizerStub = _RecognizerStub
|
||||
|
||||
provider = asr_module.YandexSpeechKitGrpcStreamingASRProvider(
|
||||
api_key="asr-key",
|
||||
grpc_target="stt.test:443",
|
||||
timeout_seconds=3,
|
||||
grpc_module=_FakeGrpc,
|
||||
stt_pb2_module=_FakeSttPb2,
|
||||
stt_service_pb2_grpc_module=_FakeGrpcPb2,
|
||||
)
|
||||
stream_id = provider.open_stream("session-1", language_hint="ru")
|
||||
provider.push_pcm(stream_id, b"\x10\x00\x20\x00")
|
||||
|
||||
partial = None
|
||||
for _ in range(20):
|
||||
partial = provider.poll_partial(stream_id)
|
||||
if partial is not None:
|
||||
break
|
||||
time.sleep(0.02)
|
||||
|
||||
assert partial is not None
|
||||
assert partial.text == "need schedule"
|
||||
assert partial.is_final is False
|
||||
|
||||
final = provider.finalize(stream_id)
|
||||
provider.close_stream(stream_id)
|
||||
|
||||
assert final.text == "need schedule in almaty"
|
||||
recognize_call = next(call for call in calls if call.get("method") == "recognize")
|
||||
assert ("authorization", "Api-Key asr-key") in recognize_call["metadata"]
|
||||
requests = [call["request"] for call in calls if call.get("method") == "request"]
|
||||
assert getattr(requests[0], "session_options", None) is not None
|
||||
assert getattr(requests[1], "chunk", None).data == b"\x10\x00\x20\x00"
|
||||
assert getattr(requests[-1], "eou", None) is not None
|
||||
|
||||
|
||||
def test_yandex_asr_builders():
|
||||
assert isinstance(asr_module.build_asr_provider("yandex"), asr_module.YandexSpeechKitASRProvider)
|
||||
assert isinstance(asr_module.build_asr_provider("speechkit"), asr_module.YandexSpeechKitASRProvider)
|
||||
assert isinstance(
|
||||
asr_module.build_streaming_asr_provider("yandex_speechkit"),
|
||||
asr_module.YandexSpeechKitGrpcStreamingASRProvider,
|
||||
)
|
||||
assert isinstance(
|
||||
asr_module.build_streaming_asr_provider("yandex_speechkit_buffered"),
|
||||
asr_module.YandexSpeechKitBufferedStreamingASRProvider,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user