first commit
This commit is contained in:
@@ -0,0 +1,200 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import numpy as np
|
||||
|
||||
from app.danmu_queue import Broadcaster, TTSEngine, _BoundedPriorityQueue, _StreamingAudioBuffer
|
||||
|
||||
|
||||
class BoundedPriorityQueueTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_returns_high_priority_first(self):
|
||||
queue = _BoundedPriorityQueue(4)
|
||||
await queue.put({"priority": 2, "sequence": 1, "name": "normal"})
|
||||
await queue.put({"priority": 0, "sequence": 2, "name": "urgent"})
|
||||
|
||||
self.assertEqual((await queue.get())["name"], "urgent")
|
||||
self.assertEqual((await queue.get())["name"], "normal")
|
||||
|
||||
async def test_replaces_worst_pending_item_when_full(self):
|
||||
queue = _BoundedPriorityQueue(2)
|
||||
await queue.put({"priority": 3, "sequence": 1, "name": "low-old"})
|
||||
await queue.put({"priority": 3, "sequence": 2, "name": "low-new"})
|
||||
|
||||
accepted, dropped = await queue.put(
|
||||
{"priority": 0, "sequence": 3, "name": "urgent"}
|
||||
)
|
||||
|
||||
self.assertTrue(accepted)
|
||||
self.assertEqual(dropped["name"], "low-new")
|
||||
self.assertEqual((await queue.get())["name"], "urgent")
|
||||
self.assertEqual((await queue.get())["name"], "low-old")
|
||||
|
||||
async def test_rejects_new_item_when_it_is_not_more_valuable(self):
|
||||
queue = _BoundedPriorityQueue(1)
|
||||
await queue.put({"priority": 0, "sequence": 1, "name": "urgent"})
|
||||
|
||||
accepted, dropped = await queue.put(
|
||||
{"priority": 3, "sequence": 2, "name": "low"}
|
||||
)
|
||||
|
||||
self.assertFalse(accepted)
|
||||
self.assertIsNone(dropped)
|
||||
self.assertEqual((await queue.get())["name"], "urgent")
|
||||
|
||||
|
||||
class StreamingAudioBufferTests(unittest.TestCase):
|
||||
def test_error_unblocks_consumer(self):
|
||||
buffer = _StreamingAudioBuffer()
|
||||
error = RuntimeError("failed")
|
||||
buffer.finish(error)
|
||||
|
||||
self.assertTrue(buffer.ready.is_set())
|
||||
self.assertIs(buffer.error, error)
|
||||
self.assertIs(buffer.chunks.get_nowait(), _StreamingAudioBuffer.END)
|
||||
|
||||
|
||||
class StreamingPreparationTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_prepare_returns_after_first_chunk_before_generation_finishes(self):
|
||||
generation_gate = asyncio.Event()
|
||||
|
||||
class FakeStreamingEngine:
|
||||
streaming = True
|
||||
|
||||
async def _synthesize_to_buffer(self, text, buffer):
|
||||
buffer.put(np.zeros(128, dtype=np.float32), 24000)
|
||||
await generation_gate.wait()
|
||||
buffer.finish()
|
||||
|
||||
tts = TTSEngine.__new__(TTSEngine)
|
||||
tts._engine = FakeStreamingEngine()
|
||||
tts._synthesis_lock = asyncio.Lock()
|
||||
tts._log_event = lambda _message: None
|
||||
|
||||
prepared = await tts.prepare_request({"request_id": "test", "text": "测试"})
|
||||
|
||||
self.assertEqual(prepared["kind"], "stream")
|
||||
self.assertFalse(prepared["synth_task"].done())
|
||||
generation_gate.set()
|
||||
await prepared["synth_task"]
|
||||
|
||||
|
||||
class PipelineOverlapTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_next_item_synthesizes_while_previous_item_is_playing(self):
|
||||
first_play_started = asyncio.Event()
|
||||
release_first_play = asyncio.Event()
|
||||
second_synthesis_started = asyncio.Event()
|
||||
|
||||
class FakeTTS:
|
||||
async def prepare_request(self, request):
|
||||
if request["request_id"] == "second":
|
||||
second_synthesis_started.set()
|
||||
return {"kind": "audio", "audio": request["request_id"].encode()}
|
||||
|
||||
def mark_playing(self, request):
|
||||
return None
|
||||
|
||||
async def play_prepared(self, prepared):
|
||||
if prepared["audio"] == b"first":
|
||||
first_play_started.set()
|
||||
await release_first_play.wait()
|
||||
return True
|
||||
|
||||
def complete_request(self, request):
|
||||
return None
|
||||
|
||||
def fail_request(self, request, error, **kwargs):
|
||||
return None
|
||||
|
||||
broadcaster = Broadcaster.__new__(Broadcaster)
|
||||
broadcaster._stop = False
|
||||
broadcaster._tts_warmup_done = asyncio.Event()
|
||||
broadcaster._tts_warmup_done.set()
|
||||
broadcaster._tts_pending = _BoundedPriorityQueue(4)
|
||||
broadcaster._tts_playback_queue = asyncio.Queue(maxsize=2)
|
||||
broadcaster.tts = FakeTTS()
|
||||
broadcaster.logger = logging.getLogger("tts-pipeline-test")
|
||||
broadcaster._finish_tts_job = lambda *_args, **_kwargs: None
|
||||
|
||||
synth_worker = asyncio.create_task(broadcaster._tts_synthesis_loop())
|
||||
play_worker = asyncio.create_task(broadcaster._tts_playback_loop())
|
||||
expiry = time.monotonic() + 30
|
||||
await broadcaster._tts_pending.put({
|
||||
"priority": 1,
|
||||
"sequence": 1,
|
||||
"expires_at": expiry,
|
||||
"tts_request": {"request_id": "first"},
|
||||
})
|
||||
await broadcaster._tts_pending.put({
|
||||
"priority": 1,
|
||||
"sequence": 2,
|
||||
"expires_at": expiry,
|
||||
"tts_request": {"request_id": "second"},
|
||||
})
|
||||
|
||||
await asyncio.wait_for(first_play_started.wait(), timeout=1)
|
||||
await asyncio.wait_for(second_synthesis_started.wait(), timeout=1)
|
||||
release_first_play.set()
|
||||
|
||||
broadcaster._stop = True
|
||||
synth_worker.cancel()
|
||||
play_worker.cancel()
|
||||
await asyncio.gather(synth_worker, play_worker, return_exceptions=True)
|
||||
|
||||
|
||||
class TTSExpiryPolicyTests(unittest.TestCase):
|
||||
def test_default_expiry_windows_match_configured_policy(self):
|
||||
broadcaster = Broadcaster.__new__(Broadcaster)
|
||||
broadcaster._tts_queue_cfg = {}
|
||||
|
||||
self.assertEqual(broadcaster._tts_policy("login"), (0, 60.0))
|
||||
self.assertEqual(broadcaster._tts_policy("system"), (1, 40.0))
|
||||
self.assertEqual(broadcaster._tts_policy("queue"), (2, 40.0))
|
||||
self.assertEqual(broadcaster._tts_policy("points"), (3, 30.0))
|
||||
|
||||
|
||||
class AudioPlaybackTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_pygame_mixer_is_initialized_once_and_reused(self):
|
||||
class FakeChannel:
|
||||
def get_busy(self):
|
||||
return False
|
||||
|
||||
class FakeSound:
|
||||
def __init__(self, *, file):
|
||||
self.file = file
|
||||
|
||||
def play(self):
|
||||
return FakeChannel()
|
||||
|
||||
class FakeMixer:
|
||||
def __init__(self):
|
||||
self.initialized = False
|
||||
self.init_calls = 0
|
||||
self.Sound = FakeSound
|
||||
|
||||
def get_init(self):
|
||||
return self.initialized
|
||||
|
||||
def init(self):
|
||||
self.initialized = True
|
||||
self.init_calls += 1
|
||||
|
||||
def quit(self):
|
||||
self.initialized = False
|
||||
|
||||
class FakePygame:
|
||||
mixer = FakeMixer()
|
||||
|
||||
tts = TTSEngine.__new__(TTSEngine)
|
||||
tts._pygame = None
|
||||
with patch.dict("sys.modules", {"pygame": FakePygame}):
|
||||
await tts._play_audio(b"first")
|
||||
await tts._play_audio(b"second")
|
||||
|
||||
self.assertEqual(FakePygame.mixer.init_calls, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user