修复直播测试与 TTS 预热
This commit is contained in:
@@ -2,6 +2,7 @@ import asyncio
|
||||
import logging
|
||||
import time
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import numpy as np
|
||||
@@ -144,6 +145,49 @@ class PipelineOverlapTests(unittest.IsolatedAsyncioTestCase):
|
||||
await asyncio.gather(synth_worker, play_worker, return_exceptions=True)
|
||||
|
||||
|
||||
class BackgroundWarmupTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_start_returns_before_tts_warmup_finishes(self):
|
||||
warmup_started = asyncio.Event()
|
||||
release_warmup = asyncio.Event()
|
||||
|
||||
class FakeTTS:
|
||||
enabled = True
|
||||
|
||||
async def warmup(self):
|
||||
warmup_started.set()
|
||||
await release_warmup.wait()
|
||||
|
||||
async def idle_worker():
|
||||
await asyncio.Event().wait()
|
||||
|
||||
broadcaster = Broadcaster.__new__(Broadcaster)
|
||||
broadcaster._tts_started = False
|
||||
broadcaster._tts_warmup_done = asyncio.Event()
|
||||
broadcaster._tts_warmup_task = None
|
||||
broadcaster._tts_worker_tasks = set()
|
||||
broadcaster._tts_queue_cfg = {"warmup_on_start": True}
|
||||
broadcaster._tts_pending = SimpleNamespace(maxsize=8)
|
||||
broadcaster._tts_playback_queue = SimpleNamespace(maxsize=2)
|
||||
broadcaster._tts_synthesis_loop = idle_worker
|
||||
broadcaster._tts_playback_loop = idle_worker
|
||||
broadcaster.tts = FakeTTS()
|
||||
broadcaster.logger = logging.getLogger("tts-background-warmup-test")
|
||||
|
||||
await asyncio.wait_for(broadcaster.start(), timeout=0.2)
|
||||
await asyncio.wait_for(warmup_started.wait(), timeout=0.2)
|
||||
|
||||
self.assertFalse(broadcaster._tts_warmup_done.is_set())
|
||||
self.assertIsNotNone(broadcaster._tts_warmup_task)
|
||||
|
||||
release_warmup.set()
|
||||
await asyncio.wait_for(broadcaster._tts_warmup_done.wait(), timeout=0.2)
|
||||
|
||||
tasks = list(broadcaster._tts_worker_tasks)
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
|
||||
class TTSExpiryPolicyTests(unittest.TestCase):
|
||||
def test_default_expiry_windows_match_configured_policy(self):
|
||||
broadcaster = Broadcaster.__new__(Broadcaster)
|
||||
|
||||
Reference in New Issue
Block a user