Files
Owner/src/owner_voice_pet/simulation.py
T

284 lines
9.9 KiB
Python

from __future__ import annotations
import math
import struct
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from .audio_preprocess import NoopAudioPreprocessor
from .assistant_pipeline import VoiceAssistantPipeline
from .config import AppConfig
from .conversation import ConversationContext
from .events import (
LLM_STARTED,
PLAYBACK_FINISHED,
SPEECH_ENDED,
SPEECH_STARTED,
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TRANSCRIPT_PARTIAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
)
from .llm import MockLlmProvider
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError
from .stt import MetadataSttProvider
from .transport import FileReplayTransport, MemoryAudioTransport
from .tts import SineTtsProvider
from .vad import EnergyVadProvider, PrimarySpeakerVadRecorder
from .wakeword import KeywordWakeWordProvider
@dataclass(slots=True)
class SimulationReporter:
statuses: list[str] = field(default_factory=list)
partials: list[str] = field(default_factory=list)
finals: list[str] = field(default_factory=list)
errors: list[str] = field(default_factory=list)
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
prefix = f"{turn_id}轮:" if turn_id is not None else ""
self.statuses.append(f"{prefix}{message}")
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
if final:
self.finals.append(text)
else:
self.partials.append(text)
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
self.errors.append(f"{stage}:{code}:{message}")
class BoundedMemoryAudioTransport(MemoryAudioTransport):
def __init__(self, frames: list[AudioFrame], *, max_empty_reads: int = 5) -> None:
super().__init__(frames)
self.max_empty_reads = max_empty_reads
self.empty_reads = 0
def read_frames(self, timeout_ms: int) -> list[AudioFrame]:
frames = super().read_frames(timeout_ms)
if frames:
self.empty_reads = 0
return frames
self.empty_reads += 1
if self.empty_reads > self.max_empty_reads:
raise ProviderError(
ErrorCode.VALIDATION_FAILED,
"simulated microphone frames were exhausted",
False,
"simulated-microphone",
"transport",
)
return []
class SimulatedNoiseFilter(NoopAudioPreprocessor):
def __init__(self) -> None:
self.loaded = False
self.processed_frames = 0
def load(self) -> None:
self.loaded = True
def reset(self) -> None:
return None
def process_frame(self, frame: AudioFrame) -> AudioFrame:
self.processed_frames += 1
metadata = dict(frame.metadata)
metadata["denoised"] = True
metadata["noise_filter_provider"] = "simulated"
return AudioFrame(
pcm=frame.pcm,
sample_rate=frame.sample_rate,
channels=frame.channels,
timestamp_ms=frame.timestamp_ms,
frame_id=frame.frame_id,
metadata=metadata,
)
def run_simulated_live(
*,
turns: int = 2,
fixture_path: str | Path | None = None,
write_fixture: str | Path | None = None,
) -> dict[str, Any]:
if turns <= 0:
raise ValueError("turns must be positive")
frames = FileReplayTransport.from_jsonl(fixture_path)._frames if fixture_path else _simulated_turn_frames(turns)
frame_list = list(frames)
if write_fixture:
FileReplayTransport.write_jsonl(write_fixture, frame_list)
transport = BoundedMemoryAudioTransport(frame_list)
reporter = SimulationReporter()
preprocessor = SimulatedNoiseFilter()
llm = MockLlmProvider(["这是模拟回复。"])
config = AppConfig(
llm_api_key="simulated",
speech_provider="local",
realtime_transcript_enabled=True,
noise_filter_enabled=True,
post_playback_drain_ms=0,
endpoint_mode="primary_speaker",
speaker_profile_ms=120,
speaker_profile_min_ms=120,
speaker_absent_ms=300,
vad_min_duration_ms=250,
vad_end_silence_ms=350,
vad_no_speech_timeout_ms=3000,
vad_max_recording_ms=6000,
)
pipeline = VoiceAssistantPipeline(
config=config,
transport=transport,
wakeword=KeywordWakeWordProvider(threshold=0.5),
vad_recorder=PrimarySpeakerVadRecorder(
EnergyVadProvider(),
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,
speaker_profile_ms=config.speaker_profile_ms,
speaker_profile_min_ms=config.speaker_profile_min_ms,
speaker_absent_ms=config.speaker_absent_ms,
similarity_threshold=config.speaker_similarity_threshold,
min_rms=config.speaker_min_rms,
),
audio_preprocessor=preprocessor,
stt=MetadataSttProvider(),
realtime_stt=MetadataSttProvider(),
llm=llm,
tts=SineTtsProvider(),
ack_tts=SineTtsProvider(),
context=ConversationContext(),
reporter=reporter,
)
completed_turns = 0
failed_turns = 0
pipeline.load()
transport.start_input(sample_rate=config.sample_rate, channels=config.channels)
try:
for turn_id in range(1, turns + 1):
result = pipeline.run_turn(turn_id)
if result.success:
completed_turns += 1
continue
failed_turns += 1
break
finally:
pipeline.shutdown()
event_types = [event.type for event in pipeline.event_bus.events]
expected_transcripts = [f"{index}轮模拟问题" for index in range(1, turns + 1)]
checks = {
"completed_turns": completed_turns == turns,
"no_failed_turns": failed_turns == 0,
"wake_per_turn": event_types.count(WAKE_DETECTED) == turns,
"speech_per_turn": event_types.count(SPEECH_STARTED) == turns and event_types.count(SPEECH_ENDED) == turns,
"stt_per_turn": event_types.count(STT_STARTED) == turns and event_types.count(TRANSCRIPT_FINAL) == turns,
"llm_per_turn": event_types.count(LLM_STARTED) == turns and len(llm.calls) == turns,
"tts_per_turn": event_types.count(TTS_STARTED) == turns and event_types.count(PLAYBACK_FINISHED) >= turns,
"standby_per_turn": event_types.count(STANDBY_RESUMED) == turns,
"transcripts_match": reporter.finals == expected_transcripts,
"partial_noise_filtered": "" not in reporter.partials and "家确" not in reporter.partials,
"denoised_capture": preprocessor.loaded and preprocessor.processed_frames > 0,
"context_in_second_turn": turns < 2 or _second_turn_has_first_history(llm.calls, expected_transcripts[0]),
}
return {
"success": all(checks.values()) and not reporter.errors,
"turns": turns,
"completed_turns": completed_turns,
"failed_turns": failed_turns,
"checks": checks,
"errors": reporter.errors,
"partials": reporter.partials,
"final_transcripts": reporter.finals,
"played_segments": len(transport.played_segments),
"llm_calls": len(llm.calls),
"event_types": event_types,
}
def _second_turn_has_first_history(calls: list[list[Any]], first_user_text: str) -> bool:
if len(calls) < 2:
return False
return any(getattr(message, "content", "") == first_user_text for message in calls[1])
def _simulated_turn_frames(turns: int) -> list[AudioFrame]:
frames: list[AudioFrame] = []
frame_id = 0
timestamp_ms = 0
for index in range(1, turns + 1):
frames.append(
_frame(
frame_id,
timestamp_ms,
frequency=880.0,
metadata={"duration_ms": 20, "wake_word": "小杰小杰", "wake_confidence": 0.99},
)
)
frame_id += 1
timestamp_ms += 20
question = f"{index}轮模拟问题"
for speech_index in range(6):
partial = "" if speech_index == 0 else "家确" if speech_index == 1 else question
metadata: dict[str, object] = {
"duration_ms": 20,
"speech": True,
"speaker_id": "owner",
"partial_transcript": partial,
}
if speech_index == 0:
metadata["transcript"] = question
frames.append(
_frame(
frame_id,
timestamp_ms,
frequency=440.0,
metadata=metadata,
)
)
frame_id += 1
timestamp_ms += 20
for _ in range(15):
frames.append(
_frame(
frame_id,
timestamp_ms,
frequency=180.0,
amplitude=500,
metadata={
"duration_ms": 20,
"speech": True,
"speaker_id": "background",
"partial_transcript": "家确",
},
)
)
frame_id += 1
timestamp_ms += 20
return frames
def _frame(
frame_id: int,
timestamp_ms: int,
*,
frequency: float,
amplitude: int = 8000,
metadata: dict[str, object] | None = None,
) -> AudioFrame:
sample_rate = 16000
samples = int(sample_rate * 0.02)
pcm = bytearray()
for index in range(samples):
value = int(math.sin(2 * math.pi * frequency * index / sample_rate) * amplitude)
pcm.extend(struct.pack("<h", value))
return AudioFrame(bytes(pcm), sample_rate, 1, timestamp_ms, frame_id, metadata or {"duration_ms": 20})