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