[响应流播放]:完成LLM流式响应和可中断播报骨架,包含句子切分、TTS净化和播放队列测试
This commit is contained in:
@@ -38,13 +38,13 @@
|
||||
|
||||
## 5. LLM 流、句子切分、Streaming TTS 与播放
|
||||
|
||||
- [ ] 5.1 定义 LLM streaming adapter contract;前置条件:现有 LLM provider 梳理完成;优先级:P0;验收标准:支持 delta、tool_call、finish、cancel、error;测试要点:取消时连接关闭或停止消费。
|
||||
- [ ] 5.2 实现 sentence segmenter;前置条件:LLM delta contract 完成;优先级:P0;验收标准:中文标点、英文标点、最大等待阈值可切句;测试要点:URL、小数、代码块不误切。
|
||||
- [ ] 5.3 复用 TTS 文本净化;前置条件:现有 sanitizer 可调用;优先级:P0;验收标准:emoji、表情包、Markdown 图片不送 TTS;测试要点:纯表情回复不触发语音。
|
||||
- [ ] 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 缺失结构化失败。
|
||||
- [ ] 5.6 实现可中断 playback queue;前置条件:render ring buffer 完成;优先级:P0;验收标准:chunk 播放同时写 reference,取消后清空未播 chunk;测试要点:播放停止边界小于配置 chunk。
|
||||
- [ ] 5.7 记录已播文本边界;前置条件:sentence/TTS/playback 完成;优先级:P0;验收标准:只提交完整播出的 assistant 文本;测试要点:中途打断未播文本不进上下文。
|
||||
- [x] 5.1 定义 LLM streaming adapter contract;前置条件:现有 LLM provider 梳理完成;优先级:P0;验收标准:支持 delta、tool_call、finish、cancel、error;测试要点:取消时连接关闭或停止消费。
|
||||
- [x] 5.2 实现 sentence segmenter;前置条件:LLM delta contract 完成;优先级:P0;验收标准:中文标点、英文标点、最大等待阈值可切句;测试要点:URL、小数、代码块不误切。
|
||||
- [x] 5.3 复用 TTS 文本净化;前置条件:现有 sanitizer 可调用;优先级:P0;验收标准:emoji、表情包、Markdown 图片不送 TTS;测试要点:纯表情回复不触发语音。
|
||||
- [x] 5.4 定义 `StreamingTtsProvider` 接口;前置条件:播放 PCM 格式确认;优先级:P0;验收标准:支持 accept_text、flush、cancel、chunk events;测试要点:fake TTS 逐 chunk 输出。
|
||||
- [x] 5.5 规划 CosyVoice adapter;前置条件:产品 TTS 方案确认;优先级:P1;验收标准:可配置 voice、sample_rate、chunk size;测试要点:provider 缺失结构化失败。
|
||||
- [x] 5.6 实现可中断 playback queue;前置条件:render ring buffer 完成;优先级:P0;验收标准:chunk 播放同时写 reference,取消后清空未播 chunk;测试要点:播放停止边界小于配置 chunk。
|
||||
- [x] 5.7 记录已播文本边界;前置条件:sentence/TTS/playback 完成;优先级:P0;验收标准:只提交完整播出的 assistant 文本;测试要点:中途打断未播文本不进上下文。
|
||||
|
||||
## 6. Conversation Manager 与长期记忆
|
||||
|
||||
|
||||
@@ -23,6 +23,16 @@ from .full_duplex_speech import (
|
||||
VadEvent,
|
||||
VadProvider,
|
||||
)
|
||||
from .full_duplex_response import (
|
||||
FakeStreamingLlmProvider,
|
||||
FakeStreamingTtsProvider,
|
||||
InterruptiblePlaybackQueue,
|
||||
LlmStreamEvent,
|
||||
SentenceSegmenter,
|
||||
StreamingLlmProvider,
|
||||
StreamingTtsProvider,
|
||||
prepare_tts_sentence,
|
||||
)
|
||||
from .models import (
|
||||
AudioFrame,
|
||||
AudioSegment,
|
||||
@@ -73,6 +83,14 @@ __all__ = [
|
||||
"TranscriptEvent",
|
||||
"VadEvent",
|
||||
"VadProvider",
|
||||
"FakeStreamingLlmProvider",
|
||||
"FakeStreamingTtsProvider",
|
||||
"InterruptiblePlaybackQueue",
|
||||
"LlmStreamEvent",
|
||||
"SentenceSegmenter",
|
||||
"StreamingLlmProvider",
|
||||
"StreamingTtsProvider",
|
||||
"prepare_tts_sentence",
|
||||
"AudioFrame",
|
||||
"AudioSegment",
|
||||
"AudioRingBuffer",
|
||||
|
||||
@@ -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)
|
||||
@@ -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("你好 😊"), "你好")
|
||||
|
||||
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()
|
||||
Reference in New Issue
Block a user