285 lines
10 KiB
Python
285 lines
10 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, flush_clears_input=False)
|
|
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,
|
|
barge_in_enabled=False,
|
|
)
|
|
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})
|