[响应流播放]:完成LLM流式响应和可中断播报骨架,包含句子切分、TTS净化和播放队列测试
This commit is contained in:
@@ -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 与长期记忆
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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