[响应流播放]:完成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
+18
View File
@@ -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",
+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)