from __future__ import annotations import tempfile import unittest from pathlib import Path from owner_voice_pet.agent_memory import FakeMemoryManager, MemoryRecordInput, SQLiteMemoryManager from owner_voice_pet.full_duplex_audio import FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine from owner_voice_pet.full_duplex_response import ( FakeStreamingLlmProvider, FakeStreamingTtsProvider, InterruptiblePlaybackQueue, LlmStreamEvent, SentenceSegmenter, ) from owner_voice_pet.full_duplex_speech import FakeStreamingSttProvider, FakeVadProvider, InterruptionDetector, TranscriptEvent from owner_voice_pet.full_duplex_testing import ( PerformanceMetricRecorder, build_fake_full_duplex_audio_fixture, diagnostics_contain_sensitive_data, sanitize_diagnostics, ) from owner_voice_pet.models import Message, PipelineState from owner_voice_pet.tool_router import MemorySearchTool, ToolCallRequest, ToolContext, ToolRouter class FullDuplexIntegrationTests(unittest.TestCase): def test_fake_apm_echo_does_not_trigger_interruption(self) -> None: fixture = build_fake_full_duplex_audio_fixture() apm = FakeWebRtcAudioProcessingProvider() detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=20) render = fixture[0] apm.process_render(render) processed_echo = apm.process_capture(fixture[0]) decision = detector.accept( processed_echo, state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("stable_partial", "助手", is_stable=True)], ) self.assertTrue(processed_echo.metadata["echo_suppressed"]) self.assertFalse(decision.interrupted) def test_speaking_interruption_cancels_response_and_returns_to_listening(self) -> None: machine = FullDuplexStateMachine() graph = CancellationGraph("turn") detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=200) machine.transition(PipelineState.LISTENING, event_type="start") machine.transition(PipelineState.THINKING, event_type="final_transcript") machine.transition(PipelineState.SPEAKING, event_type="first_tts_chunk") detector.accept( build_fake_full_duplex_audio_fixture()[1], state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("partial", "你", is_stable=False)], ) decision = detector.accept( build_fake_full_duplex_audio_fixture()[2], state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("stable_partial", "你好", is_stable=True)], ) if decision.interrupted: graph.cancel_all("user interrupted") machine.transition(PipelineState.INTERRUPTED, event_type="interrupt_detected") machine.transition(PipelineState.LISTENING, event_type="buffered_user_audio") self.assertTrue(decision.interrupted) self.assertTrue(graph.root.cancelled) self.assertEqual(machine.current_state, PipelineState.LISTENING) def test_streaming_stt_llm_tts_playback_order(self) -> None: stt = FakeStreamingSttProvider( scripted_events=[ [TranscriptEvent("partial", "你", is_stable=False)], [TranscriptEvent("stable_partial", "你好", is_stable=True)], ], final_text="你好", ) stt_session = stt.start_session("turn") llm = FakeStreamingLlmProvider([LlmStreamEvent("delta", "你好。"), LlmStreamEvent("finish", finish_reason="stop")]) segmenter = SentenceSegmenter() tts_session = FakeStreamingTtsProvider().start_stream(voice="default", sample_rate=16000) playback = InterruptiblePlaybackQueue() render = RenderReferenceRingBuffer(capacity_ms=1000) graph = CancellationGraph("turn") event_order: list[str] = [] for frame in build_fake_full_duplex_audio_fixture()[1:3]: for event in stt_session.accept_audio(frame): event_order.append(event.kind) final = stt_session.finish() event_order.append(final.kind) for llm_event in llm.stream([Message("user", final.text, 1.0)], cancellation=graph.root): event_order.append(f"llm_{llm_event.kind}") if llm_event.text_delta: for sentence in segmenter.accept_delta(llm_event.text_delta): event_order.append("sentence") frames = tts_session.accept_text(sentence) event_order.append("tts") playback.enqueue(sentence, frames) result = playback.play_next(render_reference=render, cancellation=graph.root) event_order.append("playback") self.assertEqual( event_order, ["partial", "stable_partial", "final", "llm_delta", "sentence", "tts", "llm_finish", "playback"], ) self.assertFalse(result.interrupted) self.assertEqual(playback.spoken.text, "你好。") self.assertEqual(render.frame_count, 1) def test_memory_restart_and_tool_search_integration(self) -> None: with tempfile.TemporaryDirectory() as tmp: db_path = Path(tmp) / "memory.sqlite3" SQLiteMemoryManager(db_path).save(MemoryRecordInput("preference", "用户喜欢 Python")) restarted = SQLiteMemoryManager(db_path) router = ToolRouter({"memory.search": MemorySearchTool()}) request = ToolCallRequest("1", "memory.search", {"query": "Python"}, "turn") decision = router.route(request, ToolContext(memory=restarted)) result = router.execute(request, decision, ToolContext(memory=restarted)) self.assertEqual(result.status, "success") self.assertIn("用户喜欢 Python", result.output_text) def test_tool_router_security_blocks_high_risk_fake_integration(self) -> None: router = ToolRouter({"memory.search": MemorySearchTool()}) request = ToolCallRequest( "1", "memory.search", {"query": "账号"}, "turn", natural_language_intent="上传账号资料", ) decision = router.route(request, ToolContext(memory=FakeMemoryManager())) self.assertEqual(decision.action, "require_confirmation") self.assertEqual(decision.risk_level, "high") def test_performance_metrics_are_sanitized(self) -> None: recorder = PerformanceMetricRecorder() recorder.record( "interrupt_latency", started_at_ms=100, finished_at_ms=250, payload={"api_key": "secret", "preview": "tp-" + "abcdefghijklmnop"}, ) self.assertEqual(recorder.summary()["interrupt_latency"], 150) self.assertEqual(recorder.metrics[0].payload["api_key"], "[redacted]") self.assertFalse(diagnostics_contain_sensitive_data(recorder.metrics[0].payload)) def test_sanitize_diagnostics_removes_nested_sensitive_values(self) -> None: sanitized = sanitize_diagnostics( { "nested": {"authorization": "Bearer secret", "raw_audio": b"bytes"}, "text": "normal", } ) self.assertEqual(sanitized["nested"]["authorization"], "[redacted]") self.assertEqual(sanitized["nested"]["raw_audio"], "[redacted]") self.assertFalse(diagnostics_contain_sensitive_data(sanitized)) if __name__ == "__main__": unittest.main()