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