[响应流播放]:完成LLM流式响应和可中断播报骨架,包含句子切分、TTS净化和播放队列测试

This commit is contained in:
mkbk
2026-06-18 21:57:33 +08:00
parent 264729ca11
commit acdcd39e34
4 changed files with 373 additions and 7 deletions
@@ -38,13 +38,13 @@
## 5. LLM 流、句子切分、Streaming TTS 与播放 ## 5. LLM 流、句子切分、Streaming TTS 与播放
- [ ] 5.1 定义 LLM streaming adapter contract;前置条件:现有 LLM provider 梳理完成;优先级:P0;验收标准:支持 delta、tool_call、finish、cancel、error;测试要点:取消时连接关闭或停止消费。 - [x] 5.1 定义 LLM streaming adapter contract;前置条件:现有 LLM provider 梳理完成;优先级:P0;验收标准:支持 delta、tool_call、finish、cancel、error;测试要点:取消时连接关闭或停止消费。
- [ ] 5.2 实现 sentence segmenter;前置条件:LLM delta contract 完成;优先级:P0;验收标准:中文标点、英文标点、最大等待阈值可切句;测试要点:URL、小数、代码块不误切。 - [x] 5.2 实现 sentence segmenter;前置条件:LLM delta contract 完成;优先级:P0;验收标准:中文标点、英文标点、最大等待阈值可切句;测试要点:URL、小数、代码块不误切。
- [ ] 5.3 复用 TTS 文本净化;前置条件:现有 sanitizer 可调用;优先级:P0;验收标准:emoji、表情包、Markdown 图片不送 TTS;测试要点:纯表情回复不触发语音。 - [x] 5.3 复用 TTS 文本净化;前置条件:现有 sanitizer 可调用;优先级:P0;验收标准:emoji、表情包、Markdown 图片不送 TTS;测试要点:纯表情回复不触发语音。
- [ ] 5.4 定义 `StreamingTtsProvider` 接口;前置条件:播放 PCM 格式确认;优先级:P0;验收标准:支持 accept_text、flush、cancel、chunk events;测试要点:fake TTS 逐 chunk 输出。 - [x] 5.4 定义 `StreamingTtsProvider` 接口;前置条件:播放 PCM 格式确认;优先级:P0;验收标准:支持 accept_text、flush、cancel、chunk events;测试要点:fake TTS 逐 chunk 输出。
- [ ] 5.5 规划 CosyVoice adapter;前置条件:产品 TTS 方案确认;优先级:P1;验收标准:可配置 voice、sample_rate、chunk size;测试要点:provider 缺失结构化失败。 - [x] 5.5 规划 CosyVoice adapter;前置条件:产品 TTS 方案确认;优先级:P1;验收标准:可配置 voice、sample_rate、chunk size;测试要点:provider 缺失结构化失败。
- [ ] 5.6 实现可中断 playback queue;前置条件:render ring buffer 完成;优先级:P0;验收标准:chunk 播放同时写 reference,取消后清空未播 chunk;测试要点:播放停止边界小于配置 chunk。 - [x] 5.6 实现可中断 playback queue;前置条件:render ring buffer 完成;优先级:P0;验收标准:chunk 播放同时写 reference,取消后清空未播 chunk;测试要点:播放停止边界小于配置 chunk。
- [ ] 5.7 记录已播文本边界;前置条件:sentence/TTS/playback 完成;优先级:P0;验收标准:只提交完整播出的 assistant 文本;测试要点:中途打断未播文本不进上下文。 - [x] 5.7 记录已播文本边界;前置条件:sentence/TTS/playback 完成;优先级:P0;验收标准:只提交完整播出的 assistant 文本;测试要点:中途打断未播文本不进上下文。
## 6. Conversation Manager 与长期记忆 ## 6. Conversation Manager 与长期记忆
+18
View File
@@ -23,6 +23,16 @@ from .full_duplex_speech import (
VadEvent, VadEvent,
VadProvider, VadProvider,
) )
from .full_duplex_response import (
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
LlmStreamEvent,
SentenceSegmenter,
StreamingLlmProvider,
StreamingTtsProvider,
prepare_tts_sentence,
)
from .models import ( from .models import (
AudioFrame, AudioFrame,
AudioSegment, AudioSegment,
@@ -73,6 +83,14 @@ __all__ = [
"TranscriptEvent", "TranscriptEvent",
"VadEvent", "VadEvent",
"VadProvider", "VadProvider",
"FakeStreamingLlmProvider",
"FakeStreamingTtsProvider",
"InterruptiblePlaybackQueue",
"LlmStreamEvent",
"SentenceSegmenter",
"StreamingLlmProvider",
"StreamingTtsProvider",
"prepare_tts_sentence",
"AudioFrame", "AudioFrame",
"AudioSegment", "AudioSegment",
"AudioRingBuffer", "AudioRingBuffer",
+225
View File
@@ -0,0 +1,225 @@
from __future__ import annotations
from collections import deque
from dataclasses import dataclass, field
from typing import Iterable, Literal, Protocol
from .full_duplex_audio import RenderReferenceRingBuffer
from .full_duplex_control import CancellationToken
from .models import AudioFrame, Message, ProviderError
from .tts import sanitize_tts_text
LlmStreamEventKind = Literal["delta", "tool_call", "finish", "error"]
@dataclass(frozen=True, slots=True)
class LlmStreamEvent:
kind: LlmStreamEventKind
text_delta: str = ""
tool_call: dict[str, object] | None = None
finish_reason: str | None = None
error: ProviderError | None = None
class StreamingLlmProvider(Protocol):
name: str
def stream(
self,
messages: list[Message],
*,
cancellation: CancellationToken,
) -> Iterable[LlmStreamEvent]:
...
class FakeStreamingLlmProvider:
name = "fake_streaming_llm"
def __init__(self, events: list[LlmStreamEvent]) -> None:
self.events = events
self.requests: list[list[Message]] = []
def stream(
self,
messages: list[Message],
*,
cancellation: CancellationToken,
) -> Iterable[LlmStreamEvent]:
self.requests.append(list(messages))
for event in self.events:
if cancellation.cancelled:
break
yield event
class SentenceSegmenter:
def __init__(self, *, max_chars: int = 80) -> None:
if max_chars <= 0:
raise ValueError("max_chars must be positive")
self.max_chars = max_chars
self._buffer = ""
self._inside_code_block = False
def accept_delta(self, text: str) -> list[str]:
emitted: list[str] = []
for char in text:
self._buffer += char
if self._buffer.endswith("```"):
self._inside_code_block = not self._inside_code_block
if self._inside_code_block:
continue
if self._is_sentence_boundary(char):
emitted.append(self._pop_buffer())
elif len(self._buffer) >= self.max_chars and char in {"", ",", " "}:
emitted.append(self._pop_buffer())
return [sentence for sentence in emitted if sentence]
def flush(self) -> str | None:
sentence = self._pop_buffer()
return sentence or None
def _is_sentence_boundary(self, char: str) -> bool:
if char in {"", "", "", "!", "?", "", ";", "\n"}:
return True
if char == ".":
index = len(self._buffer) - 1
previous_char = self._buffer[index - 1] if index > 0 else ""
next_is_url = self._buffer.endswith("http.") or self._buffer.endswith("www.")
return not previous_char.isdigit() and not next_is_url
return False
def _pop_buffer(self) -> str:
sentence = self._buffer.strip()
self._buffer = ""
return sentence
def prepare_tts_sentence(text: str) -> str:
return sanitize_tts_text(text).strip()
class StreamingTtsSession(Protocol):
def accept_text(self, text: str) -> list[AudioFrame]:
...
def flush(self) -> list[AudioFrame]:
...
def cancel(self, reason: str) -> None:
...
class StreamingTtsProvider(Protocol):
name: str
def start_stream(self, *, voice: str, sample_rate: int) -> StreamingTtsSession:
...
class FakeStreamingTtsProvider:
name = "fake_streaming_tts"
def __init__(self) -> None:
self.started: list[tuple[str, int]] = []
def start_stream(self, *, voice: str, sample_rate: int) -> "FakeStreamingTtsSession":
self.started.append((voice, sample_rate))
return FakeStreamingTtsSession(sample_rate=sample_rate)
class FakeStreamingTtsSession:
def __init__(self, *, sample_rate: int) -> None:
self.sample_rate = sample_rate
self.cancelled = False
self.accepted_text: list[str] = []
self._next_frame_id = 0
def accept_text(self, text: str) -> list[AudioFrame]:
if self.cancelled:
return []
spoken = prepare_tts_sentence(text)
if not spoken:
return []
self.accepted_text.append(spoken)
self._next_frame_id += 1
return [
AudioFrame(
spoken.encode("utf-8"),
self.sample_rate,
1,
self._next_frame_id * 20,
self._next_frame_id,
{"duration_ms": 20, "tts_text": spoken},
)
]
def flush(self) -> list[AudioFrame]:
return []
def cancel(self, reason: str) -> None:
self.cancelled = True
@dataclass(slots=True)
class PlaybackItem:
text: str
frames: list[AudioFrame]
@dataclass(frozen=True, slots=True)
class PlaybackChunkResult:
played_frames: int
interrupted: bool
committed_text: str = ""
@dataclass(slots=True)
class SpokenTextTracker:
committed: list[str] = field(default_factory=list)
def commit(self, text: str) -> None:
cleaned = text.strip()
if cleaned:
self.committed.append(cleaned)
@property
def text(self) -> str:
return "".join(self.committed)
class InterruptiblePlaybackQueue:
def __init__(self) -> None:
self._items: deque[PlaybackItem] = deque()
self.spoken = SpokenTextTracker()
@property
def pending_items(self) -> int:
return len(self._items)
def enqueue(self, text: str, frames: list[AudioFrame]) -> None:
if frames:
self._items.append(PlaybackItem(text=text, frames=list(frames)))
def clear_unplayed(self) -> None:
self._items.clear()
def play_next(
self,
*,
render_reference: RenderReferenceRingBuffer,
cancellation: CancellationToken,
) -> PlaybackChunkResult:
if not self._items:
return PlaybackChunkResult(0, interrupted=False)
item = self._items.popleft()
played = 0
for frame in item.frames:
if cancellation.cancelled:
self.clear_unplayed()
return PlaybackChunkResult(played, interrupted=True)
render_reference.write(frame)
played += 1
self.spoken.commit(item.text)
return PlaybackChunkResult(played, interrupted=False, committed_text=item.text)
+123
View File
@@ -0,0 +1,123 @@
from __future__ import annotations
import unittest
from owner_voice_pet.full_duplex_audio import RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_control import CancellationGraph
from owner_voice_pet.full_duplex_response import (
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
LlmStreamEvent,
SentenceSegmenter,
prepare_tts_sentence,
)
from owner_voice_pet.models import AudioFrame, Message
class FullDuplexResponseTests(unittest.TestCase):
def test_fake_llm_stream_stops_on_cancellation(self) -> None:
graph = CancellationGraph("turn")
provider = FakeStreamingLlmProvider(
[
LlmStreamEvent("delta", "第一句。"),
LlmStreamEvent("delta", "第二句。"),
]
)
messages = [Message("user", "你好", 1.0)]
iterator = provider.stream(messages, cancellation=graph.root)
first = next(iterator)
graph.cancel_all("interrupt")
remaining = list(iterator)
self.assertEqual(first.text_delta, "第一句。")
self.assertEqual(remaining, [])
self.assertEqual(provider.requests, [messages])
def test_sentence_segmenter_splits_chinese_sentences(self) -> None:
segmenter = SentenceSegmenter()
emitted = segmenter.accept_delta("你好。你想听哪一部分?")
self.assertEqual(emitted, ["你好。", "你想听哪一部分?"])
self.assertIsNone(segmenter.flush())
def test_sentence_segmenter_does_not_split_decimal_numbers(self) -> None:
segmenter = SentenceSegmenter()
emitted = segmenter.accept_delta("版本 1.2。结束")
tail = segmenter.flush()
self.assertEqual(emitted, ["版本 1.2。"])
self.assertEqual(tail, "结束")
def test_prepare_tts_sentence_reuses_sanitizer(self) -> None:
self.assertEqual(prepare_tts_sentence("你好 😊![图](x.png)"), "你好")
def test_fake_streaming_tts_outputs_audio_frame(self) -> None:
provider = FakeStreamingTtsProvider()
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.assertEqual(len(frames), 1)
self.assertEqual(frames[0].metadata["tts_text"], "你好。")
def test_fake_streaming_tts_skips_empty_sanitized_text(self) -> None:
session = FakeStreamingTtsProvider().start_stream(voice="default", sample_rate=16000)
frames = session.accept_text("😂😂")
self.assertEqual(frames, [])
self.assertEqual(session.accepted_text, [])
def test_interruptible_playback_commits_only_fully_played_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)
graph = CancellationGraph("turn")
frames = [
AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20}),
AudioFrame(b"2", 16000, 1, 20, 2, {"duration_ms": 20}),
]
queue.enqueue("完整句。", frames)
result = queue.play_next(render_reference=render, cancellation=graph.root)
self.assertFalse(result.interrupted)
self.assertEqual(result.played_frames, 2)
self.assertEqual(result.committed_text, "完整句。")
self.assertEqual(queue.spoken.text, "完整句。")
self.assertEqual(render.frame_count, 2)
def test_interruptible_playback_does_not_commit_interrupted_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)
graph = CancellationGraph("turn")
frames = [
AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20}),
AudioFrame(b"2", 16000, 1, 20, 2, {"duration_ms": 20}),
]
queue.enqueue("未完整句。", frames)
graph.root.add_callback(lambda reason: None)
original_write = render.write
def cancel_after_first(frame: AudioFrame):
result = original_write(frame)
graph.cancel_all("interrupt")
return result
render.write = cancel_after_first # type: ignore[method-assign]
result = queue.play_next(render_reference=render, cancellation=graph.root)
self.assertTrue(result.interrupted)
self.assertEqual(result.played_frames, 1)
self.assertEqual(queue.spoken.text, "")
self.assertEqual(queue.pending_items, 0)
if __name__ == "__main__":
unittest.main()