From da91be6e5cae7a4e9332c7ea818dc0f6e594e8e9 Mon Sep 17 00:00:00 2001 From: mkbk Date: Wed, 17 Jun 2026 21:33:33 +0800 Subject: [PATCH] =?UTF-8?q?[=E4=B8=BB=E8=AF=B4=E8=AF=9D=E4=BA=BA=E7=AB=AF?= =?UTF-8?q?=E7=82=B9]=EF=BC=9A=E5=AE=8C=E6=88=90=E9=9F=B3=E8=89=B2?= =?UTF-8?q?=E6=B6=88=E5=A4=B1=E7=BB=93=E6=9D=9F=E5=BD=95=E9=9F=B3=EF=BC=8C?= =?UTF-8?q?=E5=8C=85=E5=90=AB=E4=B8=B4=E6=97=B6=E9=9F=B3=E8=89=B2=E7=94=BB?= =?UTF-8?q?=E5=83=8F=E3=80=81=E7=AB=AF=E7=82=B9=E9=85=8D=E7=BD=AE=E5=92=8C?= =?UTF-8?q?=E5=9B=9E=E5=BD=92=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../tasks.md | 2 +- src/owner_voice_pet/__init__.py | 3 +- src/owner_voice_pet/config.py | 80 +++++++++ src/owner_voice_pet/runtime.py | 27 ++- src/owner_voice_pet/vad.py | 170 ++++++++++++++++++ tests/test_live_runtime.py | 7 +- tests/test_models_config.py | 17 ++ tests/test_wake_vad_stt.py | 78 +++++++- 8 files changed, 372 insertions(+), 12 deletions(-) diff --git a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md index dfa0754..96993d7 100644 --- a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md +++ b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md @@ -51,6 +51,6 @@ - [x] 7.1 更新 OpenSpec 以描述 stage 化 pipeline、事件总线、TurnController 和主说话人端点;前置条件:公开参考已确认;验收标准:proposal/design/spec/tasks 覆盖新架构和任务;测试要点:OpenSpec strict;优先级:P0;预计:45 分钟。 - [x] 7.2 实现 pipeline event bus 和终端事件映射;前置条件:7.1 完成;验收标准:所有 live 用户可见状态由事件产生;测试要点:事件顺序和终端文案测试;优先级:P0;预计:60 分钟。 - [x] 7.3 实现 `TurnController` 和 `VoiceAssistantPipeline`;前置条件:7.2 完成;验收标准:`run-live` 使用统一 pipeline,成功/失败 turn 均恢复待机;测试要点:两轮 fake runtime、错误恢复、上下文回归;优先级:P0;预计:60 分钟。 -- [ ] 7.4 实现本轮主说话人端点;前置条件:7.3 完成;验收标准:主说话人音色消失约 300 ms 后结束采集;测试要点:一次提问后背景噪声不拖尾、短暂停顿不断句、画像不足回退;优先级:P0;预计:60 分钟。 +- [x] 7.4 实现本轮主说话人端点;前置条件:7.3 完成;验收标准:主说话人音色消失约 300 ms 后结束采集;测试要点:一次提问后背景噪声不拖尾、短暂停顿不断句、画像不足回退;优先级:P0;预计:60 分钟。 - [ ] 7.5 更新 README、`.env.example`、本地 `.env` 非密钥配置;前置条件:7.2 至 7.4 完成;验收标准:运行说明匹配新 pipeline;测试要点:`--show-config` 不泄露 key;优先级:P0;预计:30 分钟。 - [ ] 7.6 验证并提交“Pipeline 文档验收”模块;前置条件:7.1 至 7.5 完成;验收标准:compileall、unittest、security-check、model-check、device-check、OpenSpec strict 全通过;优先级:P0;预计:30 分钟。 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index 9316c1b..eb87312 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -19,7 +19,7 @@ from .models import ( ) from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider -from .vad import EnergyVadProvider, HybridVadProvider, VadRecorder +from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .conversation import ConversationContext from .llm import MockLlmProvider, OpenAICompatibleLlmProvider @@ -45,6 +45,7 @@ __all__ = [ "SherpaOnnxKeywordWakeWordProvider", "EnergyVadProvider", "HybridVadProvider", + "PrimarySpeakerVadRecorder", "VadRecorder", "CloudAsrSttProvider", "MetadataSttProvider", diff --git a/src/owner_voice_pet/config.py b/src/owner_voice_pet/config.py index 7fc3ee9..71febc8 100644 --- a/src/owner_voice_pet/config.py +++ b/src/owner_voice_pet/config.py @@ -26,6 +26,12 @@ class AppConfig: wake_kws_score: float = 1.0 wake_ack_text: str = "我在" post_playback_drain_ms: int = 50 + pipeline_mode: str = "live_turn_based" + endpoint_mode: str = "primary_speaker" + speaker_profile_ms: int = 600 + speaker_absent_ms: int = 300 + speaker_similarity_threshold: float = 0.70 + speaker_min_rms: float = 0.012 vad_provider: str = "hybrid" vad_threshold: float = 0.5 vad_min_duration_ms: int = 250 @@ -37,6 +43,7 @@ class AppConfig: tts_model: str = "mimo-v2.5-tts" tts_voice: str = "mimo_default" speech_models_dir: Path = Path("models") + context_mode: str = "session_memory" context_max_messages: int = 12 context_max_chars: int = 12000 @@ -67,6 +74,14 @@ class AppConfig: wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"), wake_ack_text=get("WAKE_ACK_TEXT", "我在") or "我在", post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "50") or "50"), + pipeline_mode=(get("PIPELINE_MODE", "live_turn_based") or "live_turn_based").lower(), + endpoint_mode=(get("ENDPOINT_MODE", "primary_speaker") or "primary_speaker").lower(), + speaker_profile_ms=int(get("SPEAKER_PROFILE_MS", "600") or "600"), + speaker_absent_ms=int(get("SPEAKER_ABSENT_MS", "300") or "300"), + speaker_similarity_threshold=float( + get("SPEAKER_SIMILARITY_THRESHOLD", "0.70") or "0.70" + ), + speaker_min_rms=float(get("SPEAKER_MIN_RMS", "0.012") or "0.012"), vad_provider=(get("VAD_PROVIDER", "hybrid") or "hybrid").lower(), vad_threshold=float(get("VAD_THRESHOLD", "0.5") or "0.5"), vad_min_duration_ms=int(get("VAD_MIN_DURATION_MS", "250") or "250"), @@ -78,6 +93,7 @@ class AppConfig: tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts", tts_voice=get("TTS_VOICE", "mimo_default") or "mimo_default", speech_models_dir=Path(get("SPEECH_MODELS_DIR", "models") or "models"), + context_mode=(get("CONTEXT_MODE", "session_memory") or "session_memory").lower(), context_max_messages=int(get("CONTEXT_MAX_MESSAGES", "12") or "12"), context_max_chars=int(get("CONTEXT_MAX_CHARS", "12000") or "12000"), ) @@ -178,6 +194,60 @@ class AppConfig: "startup", ) ) + if self.pipeline_mode not in {"live_turn_based"}: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_PIPELINE_MODE must be live_turn_based", + False, + "config", + "startup", + ) + ) + if self.endpoint_mode not in {"primary_speaker", "vad"}: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_ENDPOINT_MODE must be primary_speaker or vad", + False, + "config", + "startup", + ) + ) + for name, value in { + "OWNER_SPEAKER_PROFILE_MS": self.speaker_profile_ms, + "OWNER_SPEAKER_ABSENT_MS": self.speaker_absent_ms, + }.items(): + if value <= 0: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + f"{name} must be positive", + False, + "config", + "startup", + ) + ) + if not 0 < self.speaker_similarity_threshold <= 1: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_SPEAKER_SIMILARITY_THRESHOLD must be in (0, 1]", + False, + "config", + "startup", + ) + ) + if self.speaker_min_rms <= 0: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_SPEAKER_MIN_RMS must be positive", + False, + "config", + "startup", + ) + ) if self.vad_provider not in {"hybrid", "local", "energy"}: errors.append( ProviderError( @@ -224,6 +294,16 @@ class AppConfig: "startup", ) ) + if self.context_mode not in {"session_memory"}: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_CONTEXT_MODE must be session_memory", + False, + "config", + "startup", + ) + ) return errors def api_url(self, path: str) -> str: diff --git a/src/owner_voice_pet/runtime.py b/src/owner_voice_pet/runtime.py index b61d22c..0ddee7f 100644 --- a/src/owner_voice_pet/runtime.py +++ b/src/owner_voice_pet/runtime.py @@ -33,7 +33,7 @@ from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, Wa from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .transport import SoundDeviceAudioTransport from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer -from .vad import EnergyVadProvider, HybridVadProvider, SherpaOnnxVadProvider, VadRecorder +from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, SherpaOnnxVadProvider, VadRecorder from .wakeword import SherpaOnnxKeywordWakeWordProvider @@ -345,6 +345,23 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold) else: vad_provider = EnergyVadProvider() + recorder_cls = PrimarySpeakerVadRecorder if config.endpoint_mode == "primary_speaker" else VadRecorder + recorder_kwargs = { + "provider": vad_provider, + "min_duration_ms": config.vad_min_duration_ms, + "end_silence_ms": config.vad_end_silence_ms, + "no_speech_timeout_ms": config.vad_no_speech_timeout_ms, + "max_recording_ms": config.vad_max_recording_ms, + } + if recorder_cls is PrimarySpeakerVadRecorder: + recorder_kwargs.update( + { + "speaker_profile_ms": config.speaker_profile_ms, + "speaker_absent_ms": config.speaker_absent_ms, + "similarity_threshold": config.speaker_similarity_threshold, + "min_rms": config.speaker_min_rms, + } + ) return VoiceAssistantPipeline( config=config, transport=SoundDeviceAudioTransport(output_device=config.audio_output_device), @@ -355,13 +372,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non threshold=config.wake_kws_threshold, score=config.wake_kws_score, ), - vad_recorder=VadRecorder( - vad_provider, - min_duration_ms=config.vad_min_duration_ms, - end_silence_ms=config.vad_end_silence_ms, - no_speech_timeout_ms=config.vad_no_speech_timeout_ms, - max_recording_ms=config.vad_max_recording_ms, - ), + vad_recorder=recorder_cls(**recorder_kwargs), stt=stt, llm=OpenAICompatibleLlmProvider(config), tts=tts, diff --git a/src/owner_voice_pet/vad.py b/src/owner_voice_pet/vad.py index 5146139..0b4bd8b 100644 --- a/src/owner_voice_pet/vad.py +++ b/src/owner_voice_pet/vad.py @@ -204,6 +204,13 @@ class HybridVadProvider: self.fallback.reset() +@dataclass(slots=True) +class SpeakerProfile: + vector: tuple[float, ...] | None = None + speaker_id: str | None = None + speech_ms: int = 0 + + @dataclass(slots=True) class VadRecorder: provider: Any @@ -275,3 +282,166 @@ class VadRecorder: self.reset() self.provider.reset() return segment + + +@dataclass(slots=True) +class PrimarySpeakerVadRecorder(VadRecorder): + speaker_profile_ms: int = 600 + speaker_absent_ms: int = 300 + similarity_threshold: float = 0.70 + min_rms: float = 0.012 + profile: SpeakerProfile = field(default_factory=SpeakerProfile, init=False) + profile_vectors: list[tuple[float, ...]] = field(default_factory=list, init=False) + primary_absent_ms: int = field(default=0, init=False) + + def reset(self) -> None: + VadRecorder.reset(self) + self.profile = SpeakerProfile() + self.profile_vectors = [] + self.primary_absent_ms = 0 + + def feed(self, frame: AudioFrame) -> AudioSegment | ProviderError | None: + if self.first_seen_ms is None: + self.first_seen_ms = frame.timestamp_ms + result = self.provider.analyze(frame) + frame_ms = int(frame.metadata.get("duration_ms", 20)) + + if result.is_speech: + if not self.started: + self.started = True + self.start_time_ms = frame.timestamp_ms + self.frames.append(frame) + self._update_profile(frame, frame_ms) + elif self.started: + self.frames.append(frame) + + if not self.started: + elapsed = frame.timestamp_ms - self.first_seen_ms + if elapsed >= self.no_speech_timeout_ms: + return ProviderError( + ErrorCode.VAD_TIMEOUT_NO_SPEECH, + "no speech detected after wakeword", + True, + "primary-speaker-vad", + "vad", + ) + return None + + start_time = self.start_time_ms if self.start_time_ms is not None else frame.timestamp_ms + duration = frame.timestamp_ms - start_time + if duration >= self.max_recording_ms: + return self._build_segment("max_recording") + if self._profile_ready(): + if self._matches_primary(frame): + self.primary_absent_ms = 0 + else: + self.primary_absent_ms += frame_ms + if self.primary_absent_ms >= self.speaker_absent_ms and duration >= self.min_duration_ms: + return self._build_segment("primary_speaker_absent") + if result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms: + return self._build_segment("silence") + return None + + def _update_profile(self, frame: AudioFrame, frame_ms: int) -> None: + if self.profile.speech_ms >= self.speaker_profile_ms: + return + speaker_id = frame.metadata.get("speaker_id") + if speaker_id is not None: + if self.profile.speaker_id is None: + self.profile.speaker_id = str(speaker_id) + if str(speaker_id) == self.profile.speaker_id: + self.profile.speech_ms += frame_ms + return + vector = extract_timbre_vector(frame, min_rms=self.min_rms) + if vector is None: + return + self.profile_vectors.append(vector) + self.profile.speech_ms += frame_ms + if self.profile_vectors: + width = len(self.profile_vectors[0]) + averaged = [] + for index in range(width): + averaged.append(sum(item[index] for item in self.profile_vectors) / len(self.profile_vectors)) + self.profile.vector = tuple(averaged) + + def _profile_ready(self) -> bool: + minimum_ms = min(self.speaker_profile_ms, max(120, self.min_duration_ms)) + return self.profile.speech_ms >= minimum_ms and ( + self.profile.speaker_id is not None or self.profile.vector is not None + ) + + def _matches_primary(self, frame: AudioFrame) -> bool: + speaker_id = frame.metadata.get("speaker_id") + if self.profile.speaker_id is not None: + return speaker_id is not None and str(speaker_id) == self.profile.speaker_id + if self.profile.vector is None: + return True + vector = extract_timbre_vector(frame, min_rms=self.min_rms) + if vector is None: + return False + return cosine_similarity(self.profile.vector, vector) >= self.similarity_threshold + + +def extract_timbre_vector(frame: AudioFrame, *, min_rms: float) -> tuple[float, ...] | None: + if "timbre_vector" in frame.metadata: + raw = frame.metadata["timbre_vector"] + if isinstance(raw, (list, tuple)) and raw: + return tuple(float(item) for item in raw) + if not frame.pcm: + return None + try: + import numpy as np + + samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) + if samples.size == 0: + return None + if frame.channels > 1: + samples = samples.reshape(-1, frame.channels).mean(axis=1) + normalized = samples / 32768.0 + rms = float(np.sqrt(np.mean(normalized * normalized))) + if rms < min_rms: + return None + signs = np.signbit(normalized) + zcr = float(np.mean(signs[1:] != signs[:-1])) if normalized.size > 1 else 0.0 + windowed = normalized * np.hanning(normalized.size) + spectrum = np.abs(np.fft.rfft(windowed)) + total = float(np.sum(spectrum)) + if total <= 1e-9: + return None + freqs = np.fft.rfftfreq(normalized.size, 1.0 / frame.sample_rate) + nyquist = max(frame.sample_rate / 2.0, 1.0) + centroid = float(np.sum(freqs * spectrum) / total) / nyquist + bandwidth = float(np.sqrt(np.sum(((freqs / nyquist - centroid) ** 2) * spectrum) / total)) + cumulative = np.cumsum(spectrum) + rolloff_index = int(np.searchsorted(cumulative, 0.85 * cumulative[-1])) + rolloff = float(freqs[min(rolloff_index, freqs.size - 1)] / nyquist) + flatness = float(np.exp(np.mean(np.log(spectrum + 1e-9))) / (np.mean(spectrum) + 1e-9)) + + def band_ratio(low: float, high: float) -> float: + mask = (freqs >= low) & (freqs < high) + return float(np.sum(spectrum[mask]) / total) + + return ( + rms, + zcr, + centroid, + bandwidth, + rolloff, + flatness, + band_ratio(80, 500), + band_ratio(500, 2000), + band_ratio(2000, nyquist), + ) + except Exception: + return None + + +def cosine_similarity(left: tuple[float, ...], right: tuple[float, ...]) -> float: + if len(left) != len(right): + return 0.0 + numerator = sum(a * b for a, b in zip(left, right)) + left_norm = sum(a * a for a in left) ** 0.5 + right_norm = sum(b * b for b in right) ** 0.5 + if left_norm <= 1e-9 or right_norm <= 1e-9: + return 0.0 + return numerator / (left_norm * right_norm) diff --git a/tests/test_live_runtime.py b/tests/test_live_runtime.py index 182bd2a..7b146d7 100644 --- a/tests/test_live_runtime.py +++ b/tests/test_live_runtime.py @@ -22,9 +22,10 @@ from owner_voice_pet.events import ( from owner_voice_pet.llm import MockLlmProvider from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline +from owner_voice_pet.runtime import build_live_runtime from owner_voice_pet.transport import MemoryAudioTransport from owner_voice_pet.tts import SineTtsProvider -from owner_voice_pet.vad import EnergyVadProvider, VadRecorder +from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder from owner_voice_pet.wakeword import KeywordWakeWordProvider @@ -181,6 +182,10 @@ class LiveRuntimeTests(unittest.TestCase): self.assertEqual(llm.calls[0][-1].content, "第一问") self.assertNotIn("小杰小杰", llm.calls[0][-1].content) + def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None: + runtime = build_live_runtime(AppConfig(llm_api_key="secret")) + self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_models_config.py b/tests/test_models_config.py index 4fa0e05..99dc4a6 100644 --- a/tests/test_models_config.py +++ b/tests/test_models_config.py @@ -74,6 +74,12 @@ class ModelsConfigTests(unittest.TestCase): self.assertEqual(config.wake_kws_score, 1.0) self.assertEqual(config.wake_ack_text, "我在") self.assertEqual(config.post_playback_drain_ms, 50) + self.assertEqual(config.pipeline_mode, "live_turn_based") + self.assertEqual(config.endpoint_mode, "primary_speaker") + self.assertEqual(config.speaker_profile_ms, 600) + self.assertEqual(config.speaker_absent_ms, 300) + self.assertEqual(config.speaker_similarity_threshold, 0.70) + self.assertEqual(config.speaker_min_rms, 0.012) self.assertEqual(config.vad_provider, "hybrid") self.assertEqual(config.vad_threshold, 0.5) self.assertEqual(config.vad_min_duration_ms, 250) @@ -85,6 +91,7 @@ class ModelsConfigTests(unittest.TestCase): self.assertEqual(config.tts_model, "mimo-v2.5-tts") self.assertEqual(config.tts_voice, "mimo_default") self.assertEqual(str(config.speech_models_dir), "models") + self.assertEqual(config.context_mode, "session_memory") self.assertTrue(config.llm_stream) self.assertEqual(config.validate_basic(), []) @@ -103,6 +110,16 @@ class ModelsConfigTests(unittest.TestCase): errors = config.validate_basic() self.assertTrue(any("OWNER_VAD_PROVIDER" in error.message for error in errors)) + def test_endpoint_provider_must_be_primary_speaker_or_vad(self) -> None: + config = AppConfig(endpoint_mode="invalid") + errors = config.validate_basic() + self.assertTrue(any("OWNER_ENDPOINT_MODE" in error.message for error in errors)) + + def test_speaker_similarity_threshold_range_is_validated(self) -> None: + config = AppConfig(speaker_similarity_threshold=1.5) + errors = config.validate_basic() + self.assertTrue(any("OWNER_SPEAKER_SIMILARITY_THRESHOLD" in error.message for error in errors)) + def test_missing_dotenv_uses_non_secret_defaults(self) -> None: config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env") self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") diff --git a/tests/test_wake_vad_stt.py b/tests/test_wake_vad_stt.py index 7e9ba09..4c7d458 100644 --- a/tests/test_wake_vad_stt.py +++ b/tests/test_wake_vad_stt.py @@ -7,7 +7,7 @@ import unittest from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError from owner_voice_pet.config import AppConfig from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text -from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, VadRecorder +from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder from owner_voice_pet.wakeword import ( KeywordWakeWordProvider, MissingWakeWordModelProvider, @@ -111,6 +111,82 @@ class WakeVadSttTests(unittest.TestCase): self.assertTrue(result.is_speech) self.assertEqual(result.speech_ms, 20) + def test_primary_speaker_endpoint_stops_on_background_noise(self) -> None: + provider = EnergyVadProvider() + provider.load() + recorder = PrimarySpeakerVadRecorder( + provider, + min_duration_ms=40, + end_silence_ms=1000, + speaker_profile_ms=40, + speaker_absent_ms=40, + ) + frames = [ + make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁"}), + make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}), + make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}), + make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}), + make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁第二次"}), + ] + segment = None + consumed = 0 + for item in frames: + consumed += 1 + result = recorder.feed(item) + if isinstance(result, AudioSegment): + segment = result + break + self.assertIsNotNone(segment) + assert segment is not None + self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent") + self.assertEqual(consumed, 4) + self.assertEqual(segment.metadata["transcript"], "你是谁") + + def test_primary_speaker_endpoint_allows_short_pause(self) -> None: + provider = EnergyVadProvider() + provider.load() + recorder = PrimarySpeakerVadRecorder( + provider, + min_duration_ms=40, + end_silence_ms=1000, + speaker_profile_ms=40, + speaker_absent_ms=60, + ) + frames = [ + make_frame(1, 0, speech=True, metadata={"speaker_id": "owner"}), + make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}), + make_frame(3, 40, speech=False), + make_frame(4, 60, speech=True, metadata={"speaker_id": "owner"}), + make_frame(5, 80, speech=False), + make_frame(6, 100, speech=False), + make_frame(7, 120, speech=False), + ] + results = [recorder.feed(item) for item in frames] + self.assertIsNone(results[2]) + self.assertIsNone(results[3]) + self.assertIsInstance(results[-1], AudioSegment) + assert isinstance(results[-1], AudioSegment) + self.assertEqual(results[-1].metadata["end_reason"], "primary_speaker_absent") + + def test_primary_speaker_endpoint_falls_back_to_vad_when_profile_missing(self) -> None: + provider = EnergyVadProvider() + provider.load() + recorder = PrimarySpeakerVadRecorder(provider, min_duration_ms=40, end_silence_ms=40) + frames = [ + AudioFrame(b"", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "transcript": "你好"}), + AudioFrame(b"", 16000, 1, 20, 2, {"duration_ms": 20, "speech": True}), + make_frame(3, 40, speech=False), + make_frame(4, 60, speech=False), + ] + segment = None + for item in frames: + result = recorder.feed(item) + if isinstance(result, AudioSegment): + segment = result + self.assertIsNotNone(segment) + assert segment is not None + self.assertEqual(segment.metadata["end_reason"], "silence") + def test_metadata_stt_transcribes_fixture_text(self) -> None: provider = MetadataSttProvider() provider.load()