[流式回复播放]:完成LLM到TTS流式播报,包含句子切分和可取消播放队列
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user