[响应流播放]:完成LLM流式响应和可中断播报骨架,包含句子切分、TTS净化和播放队列测试
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user