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