/
/
1"""Tests for the Smart Fades audio analysis provider."""
2
3from __future__ import annotations
4
5import asyncio
6import math
7from collections.abc import Callable, Generator
8from datetime import UTC, datetime
9from pathlib import Path
10from unittest.mock import AsyncMock, Mock, patch
11
12import numpy as np
13import pytest
14import torch
15from music_assistant_models.enums import ContentType, MediaType
16from music_assistant_models.errors import SetupFailedError
17from music_assistant_models.media_items import AudioFormat
18from torchaudio.transforms import SpectralCentroid
19
20from music_assistant.models.audio_analysis import AudioAnalysisData, AudioAnalysisError
21from music_assistant.providers.smart_fades.provider import ANALYSIS_SAMPLE_RATE, SmartFadesProvider
22from music_assistant.providers.smart_fades.vocal_activity import (
23 FIRERED_MEL_BINS,
24 infer_firered_chunk,
25)
26
27FIXTURE_DIR = Path(__file__).parent / "fixtures"
28# Synthetic 120 BPM drum pattern (kick-hat-snare-hat): 44100 Hz, stereo, float32, ~15.7s
29# 8 bars of 4 beats = 32 beats, downbeats on every 4th beat (kick).
30# Clip ends 0.2s after the last beat to avoid end-of-track hallucination.
31FIXTURE_PCM = FIXTURE_DIR / "test_120bpm_44100_2_32.pcm"
32
33# Expected beat grid for the deterministic 120 BPM fixture.
34EXPECTED_BEATS = [
35 0.000,
36 0.500,
37 1.000,
38 1.500,
39 2.020,
40 2.500,
41 3.000,
42 3.500,
43 4.020,
44 4.500,
45 5.020,
46 5.500,
47 6.020,
48 6.500,
49 7.000,
50 7.520,
51 8.020,
52 8.500,
53 9.000,
54 9.520,
55 10.020,
56 10.500,
57 11.000,
58 11.500,
59 12.020,
60 12.500,
61 13.020,
62 13.520,
63 14.020,
64 14.500,
65 15.020,
66 15.520,
67]
68
69EXPECTED_DOWNBEATS = [
70 0.000,
71 2.020,
72 4.020,
73 6.020,
74 8.020,
75 10.020,
76 12.020,
77 14.020,
78]
79
80
81@pytest.fixture
82def mass_mock() -> Mock:
83 """Return a mock MusicAssistant instance."""
84 mass = Mock()
85 mass.cache = Mock()
86 mass.cache.get = AsyncMock(return_value=None)
87 mass.cache.set = AsyncMock()
88 mass.music = Mock()
89 mass.streams = Mock()
90 mass.streams.audio_analysis = Mock()
91 mass.streams.audio_analysis.get_audio_analysis_version = AsyncMock(return_value=None)
92 mass.streams.audio_analysis.set_audio_analysis = AsyncMock()
93 mass.config = Mock()
94 mass.config.get = Mock(return_value={})
95 return mass
96
97
98@pytest.fixture
99def manifest_mock() -> Mock:
100 """Return a mock provider manifest."""
101 manifest = Mock()
102 manifest.domain = "smart_fades"
103 manifest.name = "Smart Fades"
104 return manifest
105
106
107@pytest.fixture
108def config_mock() -> Mock:
109 """Return a mock provider config."""
110 config = Mock()
111 config.instance_id = "smart_fades_test"
112 config.name = "Smart Fades Test"
113 config.enabled = True
114 config.get_value = Mock(return_value="GLOBAL")
115 config.values = {}
116 return config
117
118
119@pytest.fixture
120async def provider(mass_mock: Mock, manifest_mock: Mock, config_mock: Mock) -> SmartFadesProvider:
121 """Return a SmartFadesProvider without loading external model assets."""
122 prov = SmartFadesProvider(mass_mock, manifest_mock, config_mock, set())
123 with patch.object(prov, "_load_models", new=AsyncMock()):
124 await prov.handle_async_init()
125 prov._spectral_centroid = SpectralCentroid(
126 sample_rate=ANALYSIS_SAMPLE_RATE,
127 hop_length=512,
128 )
129 prov._firered_model = Mock()
130 prov._firered_cmvn_means = np.zeros(FIRERED_MEL_BINS, dtype=np.float64)
131 prov._firered_cmvn_inverse_std = np.ones(FIRERED_MEL_BINS, dtype=np.float64)
132 return prov
133
134
135@pytest.fixture
136def feature_provider(provider: SmartFadesProvider) -> Generator[SmartFadesProvider]:
137 """Return a provider with model-backed feature extraction disabled."""
138 with (
139 patch.object(provider, "_compute_musical_key_features"),
140 patch.object(provider, "_compute_vocal_features"),
141 ):
142 yield provider
143
144
145@pytest.fixture
146def deterministic_provider(
147 feature_provider: SmartFadesProvider,
148) -> Generator[SmartFadesProvider]:
149 """Return a provider with deterministic model inference."""
150 with (
151 patch.object(
152 feature_provider,
153 "_infer_beat_timings",
154 return_value=(
155 np.asarray(EXPECTED_BEATS, dtype=np.float32),
156 np.asarray(EXPECTED_DOWNBEATS, dtype=np.float32),
157 4,
158 ),
159 ),
160 patch.object(
161 feature_provider,
162 "_infer_musical_key",
163 return_value=("C", "major"),
164 ),
165 patch.object(
166 feature_provider,
167 "_infer_vocal_activity",
168 new=AsyncMock(return_value=np.zeros(2, dtype=np.float32)),
169 ),
170 ):
171 yield feature_provider
172
173
174async def test_beat_analysis_pipeline(
175 deterministic_provider: SmartFadesProvider,
176 mass_mock: Mock,
177) -> None:
178 """Test that the provider stores beats and downbeats from a PCM fixture."""
179 audio_format = AudioFormat(
180 content_type=ContentType.PCM_F32LE,
181 bit_depth=32,
182 sample_rate=44100,
183 channels=2,
184 )
185
186 stream_details = Mock()
187 stream_details.item_id = "test_120bpm"
188 stream_details.provider = "test"
189 stream_details.queue_id = "test"
190 stream_details.uri = "test://120bpm"
191 stream_details.media_type = MediaType.TRACK
192 stream_details.duration = 120
193
194 session_id = "test:test:test_120bpm"
195 await deterministic_provider.start_analysis(session_id, stream_details, audio_format)
196
197 # Feed PCM in 1-second chunks (matching the real streaming pipeline)
198 pcm_data = FIXTURE_PCM.read_bytes()
199 chunk_size = 44100 * 2 * 4 # 1 second at 44100 Hz, stereo, float32
200 offset = 0
201 while offset < len(pcm_data):
202 chunk = pcm_data[offset : offset + chunk_size]
203 await deterministic_provider.process_pcm_chunk(session_id, chunk)
204 offset += chunk_size
205
206 await deterministic_provider.finalize(session_id)
207
208 # Verify set_audio_analysis was called with correct data
209 set_aa_mock = mass_mock.streams.audio_analysis.set_audio_analysis
210 set_aa_mock.assert_awaited_once()
211 analysis = set_aa_mock.call_args.kwargs["analysis"]
212
213 beats = analysis.beats
214 downbeats = analysis.downbeats
215
216 assert len(beats) == len(EXPECTED_BEATS), (
217 f"Expected {len(EXPECTED_BEATS)} beats, got {len(beats)}"
218 )
219 assert len(downbeats) == len(EXPECTED_DOWNBEATS), (
220 f"Expected {len(EXPECTED_DOWNBEATS)} downbeats, got {len(downbeats)}"
221 )
222
223 # All beats must be within 20ms (1 frame at 50fps) of expected values
224 for i, (actual, expected) in enumerate(zip(beats, EXPECTED_BEATS, strict=True)):
225 assert abs(float(actual) - expected) < 0.021, (
226 f"Beat {i}: expected {expected:.3f}s, got {float(actual):.3f}s"
227 )
228
229 for i, (actual, expected) in enumerate(zip(downbeats, EXPECTED_DOWNBEATS, strict=True)):
230 assert abs(float(actual) - expected) < 0.021, (
231 f"Downbeat {i}: expected {expected:.3f}s, got {float(actual):.3f}s"
232 )
233
234 # Verify BPM is close to 120
235 assert analysis.bpm is not None
236 assert 115 < analysis.bpm < 125, f"Expected BPM ~120, got {analysis.bpm:.1f}"
237
238 # the 120 BPM fixture is common time
239 assert analysis.beats_per_bar == 4
240
241
242async def test_extended_analysis_fields(
243 deterministic_provider: SmartFadesProvider,
244 mass_mock: Mock,
245) -> None:
246 """Test that extended analysis fields (energy, centroid, key) are populated."""
247 audio_format = AudioFormat(
248 content_type=ContentType.PCM_F32LE,
249 bit_depth=32,
250 sample_rate=44100,
251 channels=2,
252 )
253
254 stream_details = Mock()
255 stream_details.item_id = "test_120bpm"
256 stream_details.provider = "test"
257 stream_details.queue_id = "test"
258 stream_details.uri = "test://120bpm"
259 stream_details.media_type = MediaType.TRACK
260 stream_details.duration = 120
261
262 session_id = "test:test:test_120bpm_extended"
263 await deterministic_provider.start_analysis(session_id, stream_details, audio_format)
264
265 pcm_data = FIXTURE_PCM.read_bytes()
266 chunk_size = 44100 * 2 * 4 # 1 second at 44100 Hz, stereo, float32
267 offset = 0
268 while offset < len(pcm_data):
269 chunk = pcm_data[offset : offset + chunk_size]
270 await deterministic_provider.process_pcm_chunk(session_id, chunk)
271 offset += chunk_size
272
273 await deterministic_provider.finalize(session_id)
274
275 set_aa_mock = mass_mock.streams.audio_analysis.set_audio_analysis
276 analysis = set_aa_mock.call_args.kwargs["analysis"]
277 assert set_aa_mock.call_args.kwargs["analysis_version"] == 3
278
279 # Energy curve should be 1800 bins, normalized to [0, 1]
280 assert analysis.rms_energy is not None
281 assert len(analysis.rms_energy) == 1800
282 assert max(analysis.rms_energy) <= 1.0
283 assert min(analysis.rms_energy) >= 0.0
284
285 # Spectral centroid should be 1800 bins with positive Hz values
286 assert analysis.spectral_centroid is not None
287 assert len(analysis.spectral_centroid) == 1800
288 assert all(v >= 0 for v in analysis.spectral_centroid)
289
290 # Musical key should be detected
291 assert analysis.key is not None
292 assert analysis.key in [
293 "C",
294 "C#",
295 "D",
296 "D#",
297 "E",
298 "F",
299 "F#",
300 "G",
301 "G#",
302 "A",
303 "A#",
304 "B",
305 "Bb",
306 ]
307 assert analysis.mode in ["major", "minor"]
308
309 # BPM and beats should still be correct
310 assert analysis.bpm is not None
311 assert 115 < analysis.bpm < 125
312
313 # v2: per-band RMS envelopes in extra_data, normalized by the full-band peak
314 assert analysis.extra_data is not None
315 band_rms = analysis.extra_data["band_rms"]
316 assert set(band_rms) == {"low", "low_mid", "mid", "high"}
317 for band in band_rms.values():
318 assert len(band) == 1800
319 assert all(v >= 0.0 for v in band)
320 # music with drums has real low-band content; bands are lists (JSON-safe)
321 assert isinstance(band_rms["low"], list)
322 assert max(band_rms["low"]) > 0.05
323
324 vocal_probabilities = analysis.extra_data["vocal_activity"]
325 assert len(vocal_probabilities) == 1800
326 assert all(math.isfinite(value) and 0.0 <= value <= 1.0 for value in vocal_probabilities)
327
328
329async def test_finalize_returns_audio_analysis_data(
330 deterministic_provider: SmartFadesProvider,
331) -> None:
332 """Test that _finalize returns an AudioAnalysisData on success."""
333 audio_format = AudioFormat(
334 content_type=ContentType.PCM_F32LE,
335 bit_depth=32,
336 sample_rate=44100,
337 channels=2,
338 )
339
340 stream_details = Mock()
341 stream_details.item_id = "test_finalize_return"
342 stream_details.provider = "test"
343 stream_details.queue_id = "test"
344 stream_details.uri = "test://finalize_return"
345 stream_details.media_type = MediaType.TRACK
346 stream_details.duration = 120
347
348 session_id = "test:test:test_finalize_return"
349 await deterministic_provider.start_analysis(session_id, stream_details, audio_format)
350
351 pcm_data = FIXTURE_PCM.read_bytes()
352 chunk_size = 44100 * 2 * 4
353 offset = 0
354 while offset < len(pcm_data):
355 chunk = pcm_data[offset : offset + chunk_size]
356 await deterministic_provider.process_pcm_chunk(session_id, chunk)
357 offset += chunk_size
358
359 result = await deterministic_provider._finalize(session_id)
360
361 assert isinstance(result, AudioAnalysisData)
362
363
364async def test_finalize_raises_when_not_enough_beats(
365 feature_provider: SmartFadesProvider,
366) -> None:
367 """Test that _finalize raises AudioAnalysisError when not enough beats are detected."""
368 audio_format = AudioFormat(
369 content_type=ContentType.PCM_F32LE,
370 bit_depth=32,
371 sample_rate=44100,
372 channels=2,
373 )
374
375 stream_details = Mock()
376 stream_details.item_id = "test_finalize_none"
377 stream_details.provider = "test"
378 stream_details.queue_id = "test"
379 stream_details.uri = "test://finalize_none"
380 stream_details.media_type = MediaType.TRACK
381 stream_details.duration = 120
382
383 session_id = "test:test:test_finalize_none"
384 await feature_provider.start_analysis(session_id, stream_details, audio_format)
385
386 pcm_data = FIXTURE_PCM.read_bytes()
387 chunk_size = 44100 * 2 * 4
388 offset = 0
389 while offset < len(pcm_data):
390 chunk = pcm_data[offset : offset + chunk_size]
391 await feature_provider.process_pcm_chunk(session_id, chunk)
392 offset += chunk_size
393
394 # Patch _infer_beat_timings to return fewer than 2 beats → deterministic skip
395 with (
396 patch.object(
397 feature_provider,
398 "_infer_beat_timings",
399 return_value=(np.array([0.5]), np.array([]), 4),
400 ),
401 pytest.raises(AudioAnalysisError, match="beat"),
402 ):
403 await feature_provider._finalize(session_id)
404
405
406async def test_digital_silence_yields_finite_spectral_centroid(
407 provider: SmartFadesProvider,
408) -> None:
409 """Digitally-silent audio yields 0 Hz centroid frames instead of non-finite values."""
410 sample_rate = 22050
411 tone = np.sin(2 * np.pi * 440 * np.arange(sample_rate, dtype=np.float32) / sample_rate)
412 pcm = np.concatenate([tone.astype(np.float32), np.zeros(sample_rate, dtype=np.float32)])
413 data = Mock()
414 data.energy_chunks = []
415 data.frequency_band_chunks = {}
416 data.centroid_chunks = []
417
418 provider._compute_energy_and_spectral_centroids(pcm, data)
419
420 assert data.centroid_chunks
421 assert np.isfinite(np.concatenate(data.centroid_chunks)).all()
422
423
424async def test_setup_raises_when_requirements_not_met(
425 mass_mock: Mock, manifest_mock: Mock, config_mock: Mock
426) -> None:
427 """setup() fails (before importing the heavy provider module) when requirements aren't met."""
428 from music_assistant.providers import smart_fades # noqa: PLC0415
429
430 with (
431 patch(
432 "music_assistant.providers.smart_fades.verify_system_meets_requirements",
433 side_effect=SetupFailedError("unsupported system"),
434 ),
435 pytest.raises(SetupFailedError),
436 ):
437 await smart_fades.setup(mass_mock, manifest_mock, config_mock)
438
439
440def test_initialize_models_uses_expected_components(provider: SmartFadesProvider) -> None:
441 """Model initialization wires the expected local and third-party components."""
442 beat_model = Mock()
443 original_beat_module = beat_model.model
444 quantized_beat_module = Mock()
445 beat_post_processor = Mock()
446 skey_components = (Mock(), Mock(), Mock())
447 spectral_centroid = Mock()
448 firered_components = (
449 Mock(),
450 np.zeros(FIRERED_MEL_BINS, dtype=np.float64),
451 np.ones(FIRERED_MEL_BINS, dtype=np.float64),
452 )
453
454 with (
455 patch(
456 "music_assistant.providers.smart_fades.provider.Spect2Frames",
457 return_value=beat_model,
458 ) as spect2frames,
459 patch(
460 "music_assistant.providers.smart_fades.provider.torch.ao.quantization.quantize_dynamic",
461 return_value=quantized_beat_module,
462 ) as quantize_dynamic,
463 patch(
464 "music_assistant.providers.smart_fades.provider.DBNDownBeatTracker",
465 return_value=beat_post_processor,
466 ),
467 patch(
468 "music_assistant.providers.smart_fades.provider.load_skey_components",
469 return_value=skey_components,
470 ),
471 patch(
472 "music_assistant.providers.smart_fades.provider.SpectralCentroid",
473 return_value=spectral_centroid,
474 ),
475 patch(
476 "music_assistant.providers.smart_fades.provider.load_firered_components",
477 return_value=firered_components,
478 ),
479 ):
480 components = provider._initialize_models()
481
482 spect2frames.assert_called_once_with(checkpoint_path="small0", device="cpu")
483 quantize_dynamic.assert_called_once_with(
484 original_beat_module,
485 {torch.nn.Linear},
486 dtype=torch.qint8,
487 )
488 assert beat_model.model is quantized_beat_module
489 assert components == (
490 beat_model,
491 beat_post_processor,
492 *skey_components,
493 spectral_centroid,
494 *firered_components,
495 )
496
497
498async def test_analysis_version_is_3(provider: SmartFadesProvider) -> None:
499 """v3 adds FireRed AED vocal activity."""
500 assert provider.analysis_version == 3
501
502
503async def test_cancel_clears_all_session_state(provider: SmartFadesProvider) -> None:
504 """Cancellation releases every retained PCM and feature block."""
505 audio_format = AudioFormat(
506 content_type=ContentType.PCM_F32LE,
507 bit_depth=32,
508 sample_rate=44100,
509 channels=2,
510 )
511 stream_details = Mock(
512 item_id="cancel",
513 provider="test",
514 media_type=MediaType.TRACK,
515 duration=120,
516 )
517 session_id = "test:test:cancel"
518 await provider.start_analysis(session_id, stream_details, audio_format)
519 data = provider._data[session_id]
520 data.pcm_buffer.append(np.ones(10, dtype=np.float32))
521 data.beats_feature_blocks.append(np.ones((1, 128), dtype=np.float32))
522 data.energy_chunks.append(np.ones(1, dtype=np.float32))
523 data.centroid_chunks.append(np.ones(1, dtype=np.float32))
524 data.frequency_band_chunks["low"] = [np.ones(1, dtype=np.float32)]
525 data.musical_key_feature_blocks.append(torch.ones((1, 1, 84, 1)))
526 data.vocal_feature_blocks.append(np.ones((1, 80), dtype=np.float32))
527
528 await provider.cancel(session_id)
529
530 assert session_id not in provider._data
531 assert not data.pcm_buffer
532 assert not data.beats_feature_blocks
533 assert not data.energy_chunks
534 assert not data.centroid_chunks
535 assert not data.frequency_band_chunks
536 assert not data.musical_key_feature_blocks
537 assert not data.vocal_feature_blocks
538 assert data.resampler is None
539 assert data.vocal_resampler is None
540 assert data.vocal_fbank is None
541
542
543async def test_finalize_error_clears_all_session_state(
544 provider: SmartFadesProvider,
545) -> None:
546 """Finalization errors release the popped session state."""
547 audio_format = AudioFormat(
548 content_type=ContentType.PCM_F32LE,
549 bit_depth=32,
550 sample_rate=22050,
551 channels=1,
552 )
553 stream_details = Mock(
554 item_id="finalize_error",
555 provider="test",
556 media_type=MediaType.TRACK,
557 duration=120,
558 )
559 session_id = "test:test:finalize_error"
560 await provider.start_analysis(session_id, stream_details, audio_format)
561 data = provider._data[session_id]
562 data.beats_feature_blocks.append(np.ones((1, 128), dtype=np.float32))
563
564 with (
565 patch.object(provider, "_process_block", new=AsyncMock()),
566 patch.object(data.features, "finalize", new=AsyncMock(return_value=np.empty((0, 128)))),
567 patch.object(
568 provider,
569 "_run_final_inference",
570 new=AsyncMock(side_effect=RuntimeError("inference failed")),
571 ),
572 pytest.raises(RuntimeError, match="inference failed"),
573 ):
574 await provider._finalize(session_id)
575
576 assert session_id not in provider._data
577 assert not data.beats_feature_blocks
578 assert data.vocal_fbank is None
579
580
581@pytest.mark.parametrize(("frame_count", "model_calls"), [(100, 1), (30_001, 2)])
582async def test_vocal_inference_uses_bounded_model_calls(
583 provider: SmartFadesProvider,
584 frame_count: int,
585 model_calls: int,
586) -> None:
587 """Normal inputs use one model call while long inputs use bounded chunks."""
588 features = np.zeros((frame_count, 80), dtype=np.float32)
589
590 with patch(
591 "music_assistant.providers.smart_fades.provider.infer_firered_chunk",
592 side_effect=lambda _model, chunk, _device: np.zeros((len(chunk), 3), dtype=np.float32),
593 ) as infer:
594 probabilities = await provider._infer_vocal_activity(features, frame_count / 100)
595
596 assert infer.call_count == model_calls
597 assert len(probabilities) == math.ceil((frame_count / 100) / 0.1)
598
599
600async def test_vocal_inference_starts_before_beat_finishes(
601 provider: SmartFadesProvider,
602) -> None:
603 """FireRed starts concurrently while key inference remains behind beat inference."""
604 beat_started = asyncio.Event()
605 vocal_started = asyncio.Event()
606 beat_finished = False
607
608 async def run_offloaded(func: Callable[..., object], *args: object) -> object:
609 nonlocal beat_finished
610 if func.__name__ == "_infer_beat_timings":
611 beat_started.set()
612 await asyncio.wait_for(vocal_started.wait(), timeout=1)
613 beat_finished = True
614 return np.array([0.0, 0.5]), np.array([0.0]), 4
615 if func is infer_firered_chunk:
616 assert beat_started.is_set()
617 vocal_started.set()
618 features = args[1]
619 assert isinstance(features, np.ndarray)
620 return np.zeros((len(features), 3), dtype=np.float32)
621 if func.__name__ == "_infer_musical_key":
622 assert beat_finished
623 return "C", "major"
624 raise AssertionError(f"unexpected offload: {func}")
625
626 async def run_offloaded_timed(func: Callable[..., object], *args: object) -> object:
627 return await run_offloaded(func, *args), 0.0
628
629 with (
630 patch.object(provider, "_run_offloaded", side_effect=run_offloaded),
631 patch.object(provider, "_run_offloaded_timed", side_effect=run_offloaded_timed),
632 ):
633 beat_key, vocal = await provider._run_final_inference(
634 np.zeros((10, 128), dtype=np.float32),
635 None,
636 np.zeros((10, 80), dtype=np.float32),
637 0.1,
638 )
639
640 assert vocal_started.is_set()
641 assert beat_key[3:] == ("C", "major")
642 assert vocal.shape == (1,)
643
644
645async def test_final_inference_cancellation_stops_long_vocal_loop(
646 provider: SmartFadesProvider,
647) -> None:
648 """Cancelling finalization prevents dispatch of later FireRed chunks."""
649 vocal_started = asyncio.Event()
650 never_finish = asyncio.Event()
651 vocal_calls = 0
652
653 async def run_offloaded(func: Callable[..., object], *_args: object) -> object:
654 nonlocal vocal_calls
655 if func is infer_firered_chunk:
656 vocal_calls += 1
657 vocal_started.set()
658 await never_finish.wait()
659 raise AssertionError("unreachable")
660
661 async def run_offloaded_timed(func: Callable[..., object], *args: object) -> object:
662 return await run_offloaded(func, *args), 0.0
663
664 with (
665 patch.object(provider, "_run_offloaded", side_effect=run_offloaded),
666 patch.object(provider, "_run_offloaded_timed", side_effect=run_offloaded_timed),
667 ):
668 task = asyncio.create_task(
669 provider._run_final_inference(
670 np.zeros((10, 128), dtype=np.float32),
671 None,
672 np.zeros((60_001, 80), dtype=np.float32),
673 600.01,
674 )
675 )
676 await asyncio.wait_for(vocal_started.wait(), timeout=1)
677 task.cancel()
678 with pytest.raises(asyncio.CancelledError):
679 await task
680
681 assert vocal_calls == 1
682
683
684async def test_beat_failure_cancels_vocal_inference(
685 provider: SmartFadesProvider,
686) -> None:
687 """An early beat failure cancels the FireRed branch instead of awaiting it."""
688 vocal_started = asyncio.Event()
689 vocal_cancelled = asyncio.Event()
690 never_finish = asyncio.Event()
691
692 async def run_offloaded(func: Callable[..., object], *_args: object) -> object:
693 if func.__name__ == "_infer_beat_timings":
694 # Fail only once the vocal branch is in flight, so cancellation is observable.
695 await asyncio.wait_for(vocal_started.wait(), timeout=1)
696 return np.array([0.0]), np.array([]), 0
697 if func is infer_firered_chunk:
698 vocal_started.set()
699 try:
700 await never_finish.wait()
701 except asyncio.CancelledError:
702 vocal_cancelled.set()
703 raise
704 raise AssertionError(f"unexpected offload: {func}")
705
706 async def run_offloaded_timed(func: Callable[..., object], *args: object) -> object:
707 return await run_offloaded(func, *args), 0.0
708
709 with (
710 patch.object(provider, "_run_offloaded", side_effect=run_offloaded),
711 patch.object(provider, "_run_offloaded_timed", side_effect=run_offloaded_timed),
712 pytest.raises(AudioAnalysisError, match="no rhythmic beat"),
713 ):
714 await asyncio.wait_for(
715 provider._run_final_inference(
716 np.zeros((10, 128), dtype=np.float32),
717 None,
718 np.zeros((10, 80), dtype=np.float32),
719 0.1,
720 ),
721 timeout=1,
722 )
723
724 assert vocal_cancelled.is_set()
725
726
727async def test_vocal_worker_keeps_local_state_during_cancel_cleanup(
728 provider: SmartFadesProvider,
729) -> None:
730 """A worker keeps valid resampler and fbank references while session fields clear."""
731 data = Mock()
732 data.vocal_feature_blocks = []
733 fbank = Mock()
734 fbank.process.return_value = np.ones((1, 80), dtype=np.float32)
735 fbank.finalize.return_value = np.ones((1, 80), dtype=np.float32)
736 resampler = Mock()
737
738 def resample_chunk(pcm: np.ndarray, _last: bool) -> np.ndarray:
739 data.vocal_resampler = None
740 data.vocal_fbank = None
741 return pcm
742
743 resampler.resample_chunk.side_effect = resample_chunk
744 data.vocal_resampler = resampler
745 data.vocal_fbank = fbank
746
747 provider._compute_vocal_features(
748 np.ones(1600, dtype=np.float32),
749 data,
750 True,
751 )
752
753 assert len(data.vocal_feature_blocks) == 2
754 fbank.process.assert_called_once()
755 fbank.finalize.assert_called_once()
756
757
758async def test_vocal_inference_failure_is_retryable(provider: SmartFadesProvider) -> None:
759 """A torch/hardware FireRed failure records a retryable error, never a permanent row."""
760 features = np.zeros((100, 80), dtype=np.float32)
761
762 with (
763 patch.object(
764 provider,
765 "_run_offloaded",
766 new=AsyncMock(side_effect=RuntimeError("torch kernel crashed")),
767 ),
768 pytest.raises(AudioAnalysisError, match="FireRed vocal inference failed") as excinfo,
769 ):
770 await provider._infer_vocal_activity(features, 1.0)
771
772 assert excinfo.value.retry_at is not None
773 assert excinfo.value.retry_at > datetime.now(UTC)
774
775
776async def test_vocal_inference_with_unloaded_model_is_retryable(
777 provider: SmartFadesProvider,
778) -> None:
779 """The idle-unload race (models freed mid-finalize) must not poison the track forever."""
780 provider._firered_model = None
781
782 with pytest.raises(AudioAnalysisError) as excinfo:
783 await provider._infer_vocal_activity(np.zeros((10, 80), dtype=np.float32), 0.1)
784
785 assert excinfo.value.retry_at is not None
786