[流式回复播放]:完成LLM到TTS流式播报,包含句子切分和可取消播放队列

This commit is contained in:
mkbk
2026-06-19 12:32:32 +08:00
parent f31e3d89d6
commit 79b3e89b79
7 changed files with 206 additions and 11 deletions
@@ -22,11 +22,11 @@
## 4. Streaming LLM/TTS/Playback
- [ ] 4.1 默认启用 LLM streaming;前置条件:Phase 3;优先级:P0;验收标准:`run-agent-live --check-config` 显示 streaming true;测试要点:配置测试。
- [ ] 4.2 实现 streaming TTS provider wrapper;前置条件:4.1;优先级:P0;验收标准:句子进入 TTS 后输出 PCM chunks;测试要点:fake/cosyvoice fallback。
- [ ] 4.3 播放队列写 render reference;前置条件:4.2;优先级:P0;验收标准:播放 chunk 同步进入 APM reference;测试要点:render ring frame count。
- [ ] 4.4 只提交已播 assistant 文本;前置条件:4.3;优先级:P0;验收标准:中途打断不写未播文本;测试要点:上下文断言。
- [ ] 4.5 Phase 4 提交;前置条件:4.1-4.4;优先级:P0;验收标准:中文提交 `[流式回复播放]...`;测试要点:compileall、单测。
- [x] 4.1 默认启用 LLM streaming;前置条件:Phase 3;优先级:P0;验收标准:`run-agent-live --check-config` 显示 streaming true;测试要点:配置测试。
- [x] 4.2 实现 streaming TTS provider wrapper;前置条件:4.1;优先级:P0;验收标准:句子进入 TTS 后输出 PCM chunks;测试要点:fake/cosyvoice fallback。
- [x] 4.3 播放队列写 render reference;前置条件:4.2;优先级:P0;验收标准:播放 chunk 同步进入 APM reference;测试要点:render ring frame count。
- [x] 4.4 只提交已播 assistant 文本;前置条件:4.3;优先级:P0;验收标准:中途打断不写未播文本;测试要点:上下文断言。
- [x] 4.5 Phase 4 提交;前置条件:4.1-4.4;优先级:P0;验收标准:中文提交 `[流式回复播放]...`;测试要点:compileall、单测。
## 5. Memory 与 Tool Router
+4
View File
@@ -54,6 +54,8 @@ from .full_duplex_speech import (
VadProvider,
)
from .full_duplex_response import (
AudioSegmentStreamingTtsProvider,
AudioSegmentStreamingTtsSession,
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
@@ -138,6 +140,8 @@ __all__ = [
"webrtc_apm_probe",
"FullDuplexAgentRuntime",
"FullDuplexRuntimeHealth",
"AudioSegmentStreamingTtsProvider",
"AudioSegmentStreamingTtsSession",
"FakeStreamingSttProvider",
"FakeVadProvider",
"InterruptController",
+98 -4
View File
@@ -4,9 +4,9 @@ from collections import deque
from dataclasses import dataclass, field
from typing import Iterable, Literal, Protocol
from .full_duplex_audio import RenderReferenceRingBuffer
from .full_duplex_audio import AudioHub, RenderReferenceRingBuffer
from .full_duplex_control import CancellationToken
from .models import AudioFrame, Message, ProviderError
from .models import AudioFrame, AudioSegment, Message, ProviderError
from .tts import sanitize_tts_text
@@ -162,6 +162,94 @@ class FakeStreamingTtsSession:
self.cancelled = True
class AudioSegmentStreamingTtsProvider:
name = "audio_segment_streaming_tts"
def __init__(self, tts_provider, *, chunk_ms: int = 30) -> None:
if chunk_ms <= 0:
raise ValueError("chunk_ms must be positive")
self.tts_provider = tts_provider
self.chunk_ms = chunk_ms
self.started: list[tuple[str, int]] = []
def start_stream(self, *, voice: str, sample_rate: int) -> "AudioSegmentStreamingTtsSession":
self.started.append((voice, sample_rate))
return AudioSegmentStreamingTtsSession(
tts_provider=self.tts_provider,
sample_rate=sample_rate,
chunk_ms=self.chunk_ms,
)
class AudioSegmentStreamingTtsSession:
def __init__(self, *, tts_provider, sample_rate: int, chunk_ms: int) -> None:
self.tts_provider = tts_provider
self.sample_rate = sample_rate
self.chunk_ms = chunk_ms
self.cancelled = False
self.cancel_reason = ""
self.accepted_text: list[str] = []
self._next_frame_id = 0
if hasattr(self.tts_provider, "load"):
self.tts_provider.load()
def accept_text(self, text: str) -> list[AudioFrame]:
if self.cancelled:
return []
spoken = prepare_tts_sentence(text)
if not spoken:
return []
segment = self.tts_provider.synthesize(spoken)
self.accepted_text.append(spoken)
return self._segment_to_frames(segment, spoken)
def flush(self) -> list[AudioFrame]:
return []
def cancel(self, reason: str) -> None:
self.cancelled = True
self.cancel_reason = reason
def _segment_to_frames(self, segment: AudioSegment, text: str) -> list[AudioFrame]:
metadata = dict(segment.metadata)
if metadata.get("format"):
self._next_frame_id += 1
metadata.update({"duration_ms": segment.duration_ms, "tts_text": text})
return [
AudioFrame(
segment.pcm,
segment.sample_rate,
segment.channels,
segment.start_time_ms,
self._next_frame_id,
metadata,
)
]
bytes_per_ms = max(1, int(segment.sample_rate * segment.channels * 2 / 1000))
chunk_bytes = max(2 * segment.channels, bytes_per_ms * self.chunk_ms)
chunk_bytes -= chunk_bytes % (2 * segment.channels)
frames: list[AudioFrame] = []
offset = 0
timestamp_ms = segment.start_time_ms
while offset < len(segment.pcm):
data = segment.pcm[offset : offset + chunk_bytes]
duration_ms = max(1, int(len(data) / bytes_per_ms))
self._next_frame_id += 1
frames.append(
AudioFrame(
data,
segment.sample_rate,
segment.channels,
timestamp_ms,
self._next_frame_id,
{"duration_ms": duration_ms, "tts_text": text},
)
)
offset += chunk_bytes
timestamp_ms += duration_ms
return frames
@dataclass(slots=True)
class PlaybackItem:
text: str
@@ -208,9 +296,12 @@ class InterruptiblePlaybackQueue:
def play_next(
self,
*,
render_reference: RenderReferenceRingBuffer,
render_reference: RenderReferenceRingBuffer | None = None,
audio_hub: AudioHub | None = None,
cancellation: CancellationToken,
) -> PlaybackChunkResult:
if render_reference is None and audio_hub is None:
raise ValueError("render_reference or audio_hub is required")
if not self._items:
return PlaybackChunkResult(0, interrupted=False)
item = self._items.popleft()
@@ -219,7 +310,10 @@ class InterruptiblePlaybackQueue:
if cancellation.cancelled:
self.clear_unplayed()
return PlaybackChunkResult(played, interrupted=True)
render_reference.write(frame)
if audio_hub is not None:
audio_hub.accept_render(frame)
elif render_reference is not None:
render_reference.write(frame)
played += 1
self.spoken.commit(item.text)
return PlaybackChunkResult(played, interrupted=False, committed_text=item.text)
+54 -1
View File
@@ -5,8 +5,17 @@ from dataclasses import dataclass
from .config import AppConfig
from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider
from .full_duplex_control import CancellationGraph, FullDuplexStateMachine
from .full_duplex_response import (
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
LlmStreamEvent,
SentenceSegmenter,
StreamingLlmProvider,
StreamingTtsProvider,
)
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
from .models import AudioFrame, PipelineState
from .models import AudioFrame, Message, PipelineState
from .runtime import RuntimeSummary
@@ -32,14 +41,19 @@ class FullDuplexAgentRuntime:
config: AppConfig,
processor: AudioProcessingProvider | None = None,
audio_hub: AudioHub | None = None,
llm_provider: StreamingLlmProvider | None = None,
tts_provider: StreamingTtsProvider | None = None,
) -> None:
self.config = config
self.processor = processor
self.audio_hub = audio_hub
self.llm_provider = llm_provider
self.tts_provider = tts_provider
self.health: FullDuplexRuntimeHealth | None = None
self.state_machine = FullDuplexStateMachine()
self.cancellation_graph = CancellationGraph("turn")
self.interrupt_controller: InterruptController | None = None
self.playback_queue = InterruptiblePlaybackQueue()
def load_audio(self) -> FullDuplexRuntimeHealth:
if self.audio_hub is None:
@@ -116,6 +130,45 @@ class FullDuplexAgentRuntime:
return RuntimeSummary(completed_turns=0, failed_turns=0, interrupted=True)
return RuntimeSummary(completed_turns=1, failed_turns=0)
def run_streaming_response_fixture(
self,
messages: list[Message],
*,
llm_events: list[LlmStreamEvent] | None = None,
) -> str:
self.load_audio()
if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded")
self.cancellation_graph = CancellationGraph("turn")
llm = self.llm_provider or FakeStreamingLlmProvider(
llm_events or [LlmStreamEvent("delta", "你好。"), LlmStreamEvent("finish", finish_reason="stop")]
)
tts = self.tts_provider or FakeStreamingTtsProvider()
tts_session = tts.start_stream(voice=self.config.tts_voice, sample_rate=self.config.sample_rate)
segmenter = SentenceSegmenter()
self.playback_queue = InterruptiblePlaybackQueue()
for event in llm.stream(messages, cancellation=self.cancellation_graph.root):
if event.kind == "delta" and event.text_delta:
for sentence in segmenter.accept_delta(event.text_delta):
self._synthesize_and_play_sentence(sentence, tts_session)
elif event.kind == "finish":
break
tail = segmenter.flush()
if tail:
self._synthesize_and_play_sentence(tail, tts_session)
for frame in tts_session.flush():
self.playback_queue.enqueue(str(frame.metadata.get("tts_text", "")), [frame])
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
return self.playback_queue.spoken.text
def _synthesize_and_play_sentence(self, sentence: str, tts_session) -> None:
if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded")
frames = tts_session.accept_text(sentence)
self.playback_queue.enqueue(sentence, frames)
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
def _run_audio_smoke_once(self) -> None:
if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded")
+1
View File
@@ -88,6 +88,7 @@ class CliAcceptanceTests(unittest.TestCase):
self.assertFalse(data["full_duplex_runtime_ready"])
self.assertEqual(data["audio_apm_provider"], "webrtc")
self.assertEqual(data["audio_apm_error_code"], "AUDIO_APM_UNAVAILABLE")
self.assertTrue(data["llm_streaming_enabled"])
self.assertEqual(data["turn_based_entry"], "run-live")
def test_run_agent_live_once_invokes_agent_runtime(self) -> None:
+16
View File
@@ -134,6 +134,22 @@ class FullDuplexIntegrationTests(unittest.TestCase):
self.assertEqual(playback.spoken.text, "你好。")
self.assertEqual(render.frame_count, 1)
def test_full_duplex_runtime_streaming_response_writes_render_reference_and_spoken_text(self) -> None:
runtime = FullDuplexAgentRuntime(
config=AppConfig(audio_apm_provider="fake", audio_apm_required=False),
llm_provider=FakeStreamingLlmProvider(
[LlmStreamEvent("delta", "第一句。第二句。"), LlmStreamEvent("finish", finish_reason="stop")]
),
tts_provider=FakeStreamingTtsProvider(),
)
spoken = runtime.run_streaming_response_fixture([Message("user", "你好", 1.0)])
self.assertEqual(spoken, "第一句。第二句。")
self.assertIsNotNone(runtime.audio_hub)
self.assertEqual(runtime.audio_hub.render_reference.frame_count, 2)
self.assertEqual(runtime.playback_queue.pending_items, 0)
def test_memory_restart_and_tool_search_integration(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
db_path = Path(tmp) / "memory.sqlite3"
+28 -1
View File
@@ -2,9 +2,10 @@ from __future__ import annotations
import unittest
from owner_voice_pet.full_duplex_audio import RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_audio import AudioHub, FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_control import CancellationGraph
from owner_voice_pet.full_duplex_response import (
AudioSegmentStreamingTtsProvider,
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
@@ -13,6 +14,7 @@ from owner_voice_pet.full_duplex_response import (
prepare_tts_sentence,
)
from owner_voice_pet.models import AudioFrame, Message
from owner_voice_pet.tts import SineTtsProvider
class FullDuplexResponseTests(unittest.TestCase):
@@ -74,6 +76,17 @@ class FullDuplexResponseTests(unittest.TestCase):
self.assertEqual(frames, [])
self.assertEqual(session.accepted_text, [])
def test_audio_segment_streaming_tts_wraps_sync_tts_as_pcm_chunks(self) -> None:
provider = AudioSegmentStreamingTtsProvider(SineTtsProvider(), chunk_ms=30)
session = provider.start_stream(voice="default", sample_rate=16000)
frames = session.accept_text("你好。")
self.assertEqual(provider.started, [("default", 16000)])
self.assertEqual(session.accepted_text, ["你好。"])
self.assertGreater(len(frames), 1)
self.assertTrue(all(frame.metadata["tts_text"] == "你好。" for frame in frames))
def test_interruptible_playback_commits_only_fully_played_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)
@@ -92,6 +105,20 @@ class FullDuplexResponseTests(unittest.TestCase):
self.assertEqual(queue.spoken.text, "完整句。")
self.assertEqual(render.frame_count, 2)
def test_interruptible_playback_writes_audio_hub_render_reference(self) -> None:
queue = InterruptiblePlaybackQueue()
graph = CancellationGraph("turn")
processor = FakeWebRtcAudioProcessingProvider()
hub = AudioHub(processor=processor, render_capacity_ms=1000)
frames = [AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20})]
queue.enqueue("一句。", frames)
result = queue.play_next(audio_hub=hub, cancellation=graph.root)
self.assertFalse(result.interrupted)
self.assertEqual(hub.render_reference.frame_count, 1)
self.assertEqual([item.frame_id for item in processor.render_frames], [1])
def test_interruptible_playback_does_not_commit_interrupted_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)