[播报文本净化]:完成TTS表情包过滤,包含emoji清理、上下文净化和回归测试
This commit is contained in:
@@ -27,7 +27,7 @@ from .llm import MockLlmProvider, OpenAICompatibleLlmProvider
|
||||
from .pipeline import PipelineResult, VoicePipeline
|
||||
from .runtime import LiveVoiceRuntime, RuntimeSummary, TerminalRuntimeReporter, TurnResult, build_live_runtime
|
||||
from .simulation import run_simulated_live
|
||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, SineTtsProvider
|
||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, SineTtsProvider, sanitize_tts_text
|
||||
from .assets import validate_pet_assets
|
||||
from .ui import ConsolePetWindow, PetStateController, PetVisualState
|
||||
|
||||
@@ -70,6 +70,7 @@ __all__ = [
|
||||
"MacSayTtsProvider",
|
||||
"SentenceBuffer",
|
||||
"SineTtsProvider",
|
||||
"sanitize_tts_text",
|
||||
"validate_pet_assets",
|
||||
"ConsolePetWindow",
|
||||
"PetStateController",
|
||||
|
||||
@@ -47,7 +47,7 @@ from .protocols import (
|
||||
WakeWordProvider,
|
||||
)
|
||||
from .stt import is_valid_transcript_text
|
||||
from .tts import SentenceBuffer
|
||||
from .tts import SentenceBuffer, sanitize_tts_text
|
||||
from .vad import VadRecorder
|
||||
|
||||
|
||||
@@ -389,7 +389,7 @@ class TurnController:
|
||||
spoken_parts.append(speak_result.spoken_text)
|
||||
except ProviderError as exc:
|
||||
return self._recover(exc, turn_id, completed_turns=completed_turns)
|
||||
if not assistant_text.strip() and not "".join(spoken_parts).strip():
|
||||
if not assistant_text.strip():
|
||||
return self._recover(
|
||||
ProviderError(
|
||||
ErrorCode.LLM_EMPTY_REPLY,
|
||||
@@ -402,6 +402,18 @@ class TurnController:
|
||||
completed_turns=completed_turns,
|
||||
)
|
||||
spoken_text = "".join(spoken_parts)
|
||||
if not spoken_text.strip() and not interrupted:
|
||||
return self._recover(
|
||||
ProviderError(
|
||||
ErrorCode.TTS_EMPTY_AUDIO,
|
||||
"LLM reply contained no speakable text after TTS sanitization",
|
||||
True,
|
||||
"voice-assistant-pipeline",
|
||||
"tts",
|
||||
),
|
||||
turn_id,
|
||||
completed_turns=completed_turns,
|
||||
)
|
||||
if spoken_text.strip():
|
||||
self.context.append_assistant(spoken_text)
|
||||
states = list(self._states)
|
||||
@@ -410,7 +422,7 @@ class TurnController:
|
||||
return TurnResult(
|
||||
True,
|
||||
user_text,
|
||||
spoken_text or assistant_text,
|
||||
spoken_text,
|
||||
states=states,
|
||||
completed_turns=1,
|
||||
)
|
||||
@@ -465,20 +477,23 @@ class TurnController:
|
||||
return user_text
|
||||
|
||||
def _speak(self, sentence: str, turn_id: int) -> SpeakResult:
|
||||
spoken_sentence = sanitize_tts_text(sentence)
|
||||
if not spoken_sentence:
|
||||
return SpeakResult("")
|
||||
self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
|
||||
segment = self.tts.synthesize(sentence)
|
||||
segment = self.tts.synthesize(spoken_sentence)
|
||||
if not self._can_interrupt_playback(segment):
|
||||
playback = self.transport.play_pcm(segment)
|
||||
if playback.error:
|
||||
raise playback.error
|
||||
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id)
|
||||
self._drain_input_after_playback()
|
||||
return SpeakResult(sentence)
|
||||
return SpeakResult(spoken_sentence)
|
||||
if self._play_interruptible(segment, turn_id=turn_id):
|
||||
return SpeakResult("", interrupted=True)
|
||||
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id)
|
||||
self._drain_input_after_playback()
|
||||
return SpeakResult(sentence)
|
||||
return SpeakResult(spoken_sentence)
|
||||
|
||||
def _can_interrupt_playback(self, segment: AudioSegment) -> bool:
|
||||
return (
|
||||
|
||||
@@ -9,7 +9,7 @@ from .models import Message
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ConversationContext:
|
||||
system_prompt: str = "你是一个中文桌宠助手,回答要简洁、自然、适合语音播报。"
|
||||
system_prompt: str = "你是一个中文桌宠助手,回答要简洁、自然、适合语音播报。不要输出 emoji、表情包、Markdown 图片。"
|
||||
max_messages: int = 12
|
||||
max_chars: int = 12000
|
||||
_messages: list[Message] = field(default_factory=list)
|
||||
|
||||
@@ -5,7 +5,7 @@ from dataclasses import dataclass, field
|
||||
from .conversation import ConversationContext
|
||||
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
|
||||
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
|
||||
from .tts import SentenceBuffer
|
||||
from .tts import SentenceBuffer, sanitize_tts_text
|
||||
from .vad import VadRecorder
|
||||
|
||||
|
||||
@@ -98,27 +98,36 @@ class VoicePipeline:
|
||||
self.context.append_user(text)
|
||||
self._state(PipelineState.THINKING)
|
||||
assistant_text = ""
|
||||
spoken_parts: list[str] = []
|
||||
played = 0
|
||||
try:
|
||||
for delta in self.llm.stream_reply(self.context.build_llm_messages()):
|
||||
assistant_text += delta.text_delta
|
||||
for sentence in self.sentence_buffer.feed(delta.text_delta, bool(delta.finish_reason)):
|
||||
spoken_sentence = sanitize_tts_text(sentence)
|
||||
if not spoken_sentence:
|
||||
continue
|
||||
self._state(PipelineState.SPEAKING)
|
||||
self.suppress_input = True
|
||||
audio = self.tts.synthesize(sentence)
|
||||
audio = self.tts.synthesize(spoken_sentence)
|
||||
playback = self.transport.play_pcm(audio)
|
||||
self.suppress_input = False
|
||||
if playback.error:
|
||||
return self._recover(playback.error)
|
||||
spoken_parts.append(spoken_sentence)
|
||||
played += 1
|
||||
for sentence in self.sentence_buffer.flush():
|
||||
spoken_sentence = sanitize_tts_text(sentence)
|
||||
if not spoken_sentence:
|
||||
continue
|
||||
self._state(PipelineState.SPEAKING)
|
||||
self.suppress_input = True
|
||||
audio = self.tts.synthesize(sentence)
|
||||
audio = self.tts.synthesize(spoken_sentence)
|
||||
playback = self.transport.play_pcm(audio)
|
||||
self.suppress_input = False
|
||||
if playback.error:
|
||||
return self._recover(playback.error)
|
||||
spoken_parts.append(spoken_sentence)
|
||||
played += 1
|
||||
except ProviderError as exc:
|
||||
self.suppress_input = False
|
||||
@@ -133,9 +142,20 @@ class VoicePipeline:
|
||||
"llm",
|
||||
)
|
||||
)
|
||||
self.context.append_assistant(assistant_text)
|
||||
spoken_text = "".join(spoken_parts)
|
||||
if not spoken_text.strip():
|
||||
return self._recover(
|
||||
ProviderError(
|
||||
ErrorCode.TTS_EMPTY_AUDIO,
|
||||
"LLM reply contained no speakable text after TTS sanitization",
|
||||
True,
|
||||
"voice-pipeline",
|
||||
"tts",
|
||||
)
|
||||
)
|
||||
self.context.append_assistant(spoken_text)
|
||||
self._state(PipelineState.WAKE_LISTENING)
|
||||
return PipelineResult(True, text, assistant_text, list(self.states), played_segments=played)
|
||||
return PipelineResult(True, text, spoken_text, list(self.states), played_segments=played)
|
||||
|
||||
def _recover(self, error: ProviderError) -> PipelineResult:
|
||||
self._state(PipelineState.ERROR_RECOVERING)
|
||||
|
||||
@@ -33,7 +33,7 @@ from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
|
||||
from .protocols import AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider, WakeWordProvider
|
||||
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from .transport import SoundDeviceAudioTransport
|
||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
|
||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, sanitize_tts_text
|
||||
from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, SherpaOnnxVadProvider, VadRecorder
|
||||
from .wakeword import SherpaOnnxKeywordWakeWordProvider
|
||||
|
||||
@@ -267,13 +267,18 @@ class LiveVoiceRuntime:
|
||||
self.context.append_user(user_text)
|
||||
self._event(LLM_STARTED, PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
|
||||
assistant_text = ""
|
||||
spoken_parts: list[str] = []
|
||||
try:
|
||||
for delta in self.llm.stream_reply(self.context.build_llm_messages()):
|
||||
assistant_text += delta.text_delta
|
||||
for sentence in self.sentence_buffer.feed(delta.text_delta, bool(delta.finish_reason)):
|
||||
self._speak(sentence, turn_id)
|
||||
spoken = self._speak(sentence, turn_id)
|
||||
if spoken:
|
||||
spoken_parts.append(spoken)
|
||||
for sentence in self.sentence_buffer.flush():
|
||||
self._speak(sentence, turn_id)
|
||||
spoken = self._speak(sentence, turn_id)
|
||||
if spoken:
|
||||
spoken_parts.append(spoken)
|
||||
except ProviderError as exc:
|
||||
return self._recover(exc, turn_id)
|
||||
if not assistant_text.strip():
|
||||
@@ -287,18 +292,34 @@ class LiveVoiceRuntime:
|
||||
),
|
||||
turn_id,
|
||||
)
|
||||
self.context.append_assistant(assistant_text)
|
||||
spoken_text = "".join(spoken_parts)
|
||||
if not spoken_text.strip():
|
||||
return self._recover(
|
||||
ProviderError(
|
||||
ErrorCode.TTS_EMPTY_AUDIO,
|
||||
"LLM reply contained no speakable text after TTS sanitization",
|
||||
True,
|
||||
"live-runtime",
|
||||
"tts",
|
||||
),
|
||||
turn_id,
|
||||
)
|
||||
self.context.append_assistant(spoken_text)
|
||||
self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
|
||||
return TurnResult(True, user_text, assistant_text, states=list(self._states))
|
||||
return TurnResult(True, user_text, spoken_text, states=list(self._states))
|
||||
|
||||
def _speak(self, sentence: str, turn_id: int) -> None:
|
||||
def _speak(self, sentence: str, turn_id: int) -> str:
|
||||
spoken_sentence = sanitize_tts_text(sentence)
|
||||
if not spoken_sentence:
|
||||
return ""
|
||||
self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
|
||||
segment = self.tts.synthesize(sentence)
|
||||
segment = self.tts.synthesize(spoken_sentence)
|
||||
playback = self.transport.play_pcm(segment)
|
||||
if playback.error:
|
||||
raise playback.error
|
||||
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "播放完成", turn_id=turn_id)
|
||||
self._drain_input_after_playback()
|
||||
return spoken_sentence
|
||||
|
||||
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
|
||||
self._event(STAGE_ERROR, PipelineState.ERROR_RECOVERING, error.message, turn_id=turn_id, payload={"error": error})
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import math
|
||||
import base64
|
||||
import json
|
||||
import re
|
||||
import socket
|
||||
import struct
|
||||
import subprocess
|
||||
@@ -18,6 +19,131 @@ from .config import AppConfig
|
||||
from .models import AudioSegment, ErrorCode, ProviderError
|
||||
|
||||
|
||||
_MARKDOWN_IMAGE_RE = re.compile(r"!\[[^\]]*]\([^)]*\)")
|
||||
_SHORTCODE_EMOJI_RE = re.compile(r":(?:[A-Za-z][A-Za-z0-9_+\-]{1,31}):")
|
||||
_ASCII_KAOMOJI_RE = re.compile(r"(?:\^_?\^|T_T|QAQ|QwQ|qwq|orz)")
|
||||
_SYMBOL_KAOMOJI_RE = re.compile(r"[\((][^()()]{0,20}[\u00b0\u2500-\u2bff][^()()]{0,20}[\))][^\s,。!?,.!?]{0,10}")
|
||||
_EMOJI_RANGES = (
|
||||
(0x1F000, 0x1FAFF),
|
||||
(0x2600, 0x27BF),
|
||||
(0x2300, 0x23FF),
|
||||
)
|
||||
_EMOJI_JOINERS = {0x200D, 0x20E3}
|
||||
_BRACKET_PAIRS = {"[": "]", "【": "】", "(": ")", "(": ")"}
|
||||
_BRACKET_EMOTE_WORDS = {
|
||||
"ok",
|
||||
"doge",
|
||||
"emoji",
|
||||
"一笑",
|
||||
"偷笑",
|
||||
"傻笑",
|
||||
"加油",
|
||||
"发呆",
|
||||
"吐舌",
|
||||
"呲牙",
|
||||
"哭",
|
||||
"哭泣",
|
||||
"困",
|
||||
"大哭",
|
||||
"大笑",
|
||||
"委屈",
|
||||
"害羞",
|
||||
"尴尬",
|
||||
"开心",
|
||||
"微笑",
|
||||
"怒",
|
||||
"惊讶",
|
||||
"惊喜",
|
||||
"惊恐",
|
||||
"捂脸",
|
||||
"抱拳",
|
||||
"擦汗",
|
||||
"晕",
|
||||
"汗",
|
||||
"流汗",
|
||||
"流泪",
|
||||
"滑稽",
|
||||
"爱心",
|
||||
"玫瑰",
|
||||
"生气",
|
||||
"白眼",
|
||||
"点赞",
|
||||
"破涕为笑",
|
||||
"笑",
|
||||
"笑哭",
|
||||
"鼓掌",
|
||||
"比心",
|
||||
"亲亲",
|
||||
"调皮",
|
||||
"难过",
|
||||
"高兴",
|
||||
"鼓励",
|
||||
"狗头",
|
||||
}
|
||||
|
||||
|
||||
def sanitize_tts_text(text: str) -> str:
|
||||
clean = text.strip()
|
||||
if not clean:
|
||||
return ""
|
||||
clean = _MARKDOWN_IMAGE_RE.sub("", clean)
|
||||
clean = _SHORTCODE_EMOJI_RE.sub("", clean)
|
||||
clean = _strip_bracket_emotes(clean)
|
||||
clean = _SYMBOL_KAOMOJI_RE.sub("", clean)
|
||||
clean = _ASCII_KAOMOJI_RE.sub("", clean)
|
||||
clean = "".join(ch for ch in clean if not _is_emoji_char(ch))
|
||||
return _normalize_spoken_text(clean)
|
||||
|
||||
|
||||
def _strip_bracket_emotes(text: str) -> str:
|
||||
result: list[str] = []
|
||||
index = 0
|
||||
while index < len(text):
|
||||
ch = text[index]
|
||||
close = _BRACKET_PAIRS.get(ch)
|
||||
if close is None:
|
||||
result.append(ch)
|
||||
index += 1
|
||||
continue
|
||||
close_index = text.find(close, index + 1)
|
||||
if close_index == -1 or close_index - index > 12:
|
||||
result.append(ch)
|
||||
index += 1
|
||||
continue
|
||||
content = text[index + 1 : close_index]
|
||||
if _is_bracket_emote(content):
|
||||
index = close_index + 1
|
||||
continue
|
||||
result.append(ch)
|
||||
index += 1
|
||||
return "".join(result)
|
||||
|
||||
|
||||
def _is_bracket_emote(content: str) -> bool:
|
||||
token = re.sub(r"\s+", "", content.strip()).lower()
|
||||
if not token or len(token) > 8:
|
||||
return False
|
||||
if token in _BRACKET_EMOTE_WORDS or token.endswith("表情"):
|
||||
return True
|
||||
return all(_is_emoji_char(ch) for ch in token)
|
||||
|
||||
|
||||
def _is_emoji_char(ch: str) -> bool:
|
||||
codepoint = ord(ch)
|
||||
if codepoint in _EMOJI_JOINERS or 0xFE00 <= codepoint <= 0xFE0F:
|
||||
return True
|
||||
return any(start <= codepoint <= end for start, end in _EMOJI_RANGES)
|
||||
|
||||
|
||||
def _normalize_spoken_text(text: str) -> str:
|
||||
clean = re.sub(r"[ \t]+", " ", text)
|
||||
clean = re.sub(r"\s+([,。!?,.!?;;::])", r"\1", clean)
|
||||
clean = re.sub(r"([,,])\s*([。!?!?])", r"\2", clean)
|
||||
clean = re.sub(r"([。!?!?]){2,}", r"\1", clean)
|
||||
clean = re.sub(r"\s{2,}", " ", clean)
|
||||
return clean.strip(" \t\r\n,,;;::")
|
||||
|
||||
|
||||
class SentenceBuffer:
|
||||
def __init__(self, max_chars: int = 80) -> None:
|
||||
self.max_chars = max_chars
|
||||
|
||||
@@ -26,7 +26,7 @@ from owner_voice_pet.events import (
|
||||
PipelineEventBus,
|
||||
)
|
||||
from owner_voice_pet.llm import MockLlmProvider
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, PlaybackResult, ReplyDelta, Transcript, TransportHealth
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, PlaybackResult, ReplyDelta, Transcript, TransportHealth
|
||||
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
|
||||
from owner_voice_pet.runtime import build_live_runtime
|
||||
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider
|
||||
@@ -409,6 +409,64 @@ class LiveRuntimeTests(unittest.TestCase):
|
||||
self.assertEqual(llm.calls[0][-1].content, "第一问")
|
||||
self.assertNotIn("小杰小杰", llm.calls[0][-1].content)
|
||||
|
||||
def test_assistant_reply_sanitizes_tts_text_and_context(self) -> None:
|
||||
frames = [wake_frame(0, 0)]
|
||||
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
||||
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
||||
stt = QueueSttProvider(["第一问"])
|
||||
llm = QueueLlmProvider([["你好 😊。没问题[捂脸],我来帮你。"]])
|
||||
context = ConversationContext()
|
||||
runtime = VoiceAssistantPipeline(
|
||||
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""),
|
||||
transport=transport,
|
||||
wakeword=KeywordWakeWordProvider(),
|
||||
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
||||
stt=stt,
|
||||
realtime_stt=MetadataSttProvider(),
|
||||
llm=llm,
|
||||
tts=SineTtsProvider(),
|
||||
context=context,
|
||||
reporter=RecordingReporter(),
|
||||
event_bus=PipelineEventBus(),
|
||||
)
|
||||
|
||||
summary = runtime.run(max_turns=1)
|
||||
context_texts = [message.content for message in context.messages()]
|
||||
|
||||
self.assertEqual(summary.completed_turns, 1)
|
||||
self.assertEqual(transport.played_segments[0].metadata["text"], "你好。没问题,我来帮你。")
|
||||
self.assertIn("你好。没问题,我来帮你。", context_texts)
|
||||
self.assertFalse(any("😊" in text or "[捂脸]" in text for text in context_texts))
|
||||
|
||||
def test_emoji_only_reply_recovers_without_tts_playback(self) -> None:
|
||||
frames = [wake_frame(0, 0)]
|
||||
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
||||
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
||||
runtime = VoiceAssistantPipeline(
|
||||
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""),
|
||||
transport=transport,
|
||||
wakeword=KeywordWakeWordProvider(),
|
||||
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
||||
stt=QueueSttProvider(["第一问"]),
|
||||
realtime_stt=MetadataSttProvider(),
|
||||
llm=QueueLlmProvider([["😂😂"]]),
|
||||
tts=SineTtsProvider(),
|
||||
context=ConversationContext(),
|
||||
reporter=RecordingReporter(),
|
||||
event_bus=PipelineEventBus(),
|
||||
)
|
||||
|
||||
summary = runtime.run(once=True)
|
||||
event_types = [event.type for event in runtime.event_bus.events]
|
||||
|
||||
self.assertEqual(summary.completed_turns, 0)
|
||||
self.assertEqual(summary.failed_turns, 1)
|
||||
self.assertIsNotNone(summary.last_error)
|
||||
self.assertEqual(summary.last_error.code, ErrorCode.TTS_EMPTY_AUDIO)
|
||||
self.assertEqual(transport.played_segments, [])
|
||||
self.assertNotIn(TTS_STARTED, event_types)
|
||||
self.assertEqual(event_types[-1], STANDBY_RESUMED)
|
||||
|
||||
def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None:
|
||||
runtime = build_live_runtime(AppConfig(llm_api_key="secret"))
|
||||
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
|
||||
|
||||
@@ -12,7 +12,7 @@ from owner_voice_pet.models import AudioFrame, ErrorCode, Message, PipelineState
|
||||
from owner_voice_pet.pipeline import VoicePipeline
|
||||
from owner_voice_pet.stt import MetadataSttProvider
|
||||
from owner_voice_pet.transport import MemoryAudioTransport
|
||||
from owner_voice_pet.tts import CloudTtsProvider, SentenceBuffer, SineTtsProvider
|
||||
from owner_voice_pet.tts import CloudTtsProvider, SentenceBuffer, SineTtsProvider, sanitize_tts_text
|
||||
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
|
||||
from owner_voice_pet.wakeword import KeywordWakeWordProvider
|
||||
|
||||
@@ -60,6 +60,12 @@ class PipelineLlmTtsTests(unittest.TestCase):
|
||||
self.assertEqual(buffer.feed("剩余").copy(), [])
|
||||
self.assertEqual(buffer.flush(), ["剩余"])
|
||||
|
||||
def test_sanitize_tts_text_removes_unspeakable_expression_tokens(self) -> None:
|
||||
self.assertEqual(sanitize_tts_text("你好 😊"), "你好")
|
||||
self.assertEqual(sanitize_tts_text("好的"), "好的")
|
||||
self.assertEqual(sanitize_tts_text("没问题[捂脸],我来帮你。"), "没问题,我来帮你。")
|
||||
self.assertEqual(sanitize_tts_text("😂😂"), "")
|
||||
|
||||
def test_sine_tts_generates_non_empty_audio(self) -> None:
|
||||
provider = SineTtsProvider()
|
||||
provider.load()
|
||||
@@ -126,6 +132,22 @@ class PipelineLlmTtsTests(unittest.TestCase):
|
||||
self.assertEqual(len(transport.played_segments), 1)
|
||||
self.assertEqual(pipeline.context.messages()[-1].role, "assistant")
|
||||
|
||||
def test_pipeline_sanitizes_tts_text_and_assistant_context(self) -> None:
|
||||
frames = [
|
||||
frame(0, 0, {"wake_word": "小杰小杰", "wake_confidence": 0.95}),
|
||||
frame(1, 20, {"speech": True, "transcript": "你好"}),
|
||||
frame(2, 40, {"speech": True}),
|
||||
frame(3, 60, {"speech": False}),
|
||||
frame(4, 80, {"speech": False}),
|
||||
]
|
||||
pipeline, transport = make_pipeline(frames, MockLlmProvider(["你好 😊。没问题[捂脸],我来帮你。"]))
|
||||
result = pipeline.run_once()
|
||||
|
||||
self.assertTrue(result.success)
|
||||
self.assertEqual(result.assistant_text, "你好。没问题,我来帮你。")
|
||||
self.assertEqual(transport.played_segments[0].metadata["text"], "你好。没问题,我来帮你。")
|
||||
self.assertEqual(pipeline.context.messages()[-1].content, "你好。没问题,我来帮你。")
|
||||
|
||||
def test_pipeline_skips_llm_on_empty_transcript(self) -> None:
|
||||
llm = MockLlmProvider(["不应调用"])
|
||||
frames = [
|
||||
|
||||
Reference in New Issue
Block a user