music-assistant-server

32.9 KBPY
test_provider.py
32.9 KB961 lines • python
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