/
/
/
1"""Unit tests for the soxr resampling path in SonicAnalysisProvider (T1.2)."""
2
3from __future__ import annotations
4
5from unittest.mock import AsyncMock, MagicMock, patch
6
7import numpy as np
8import pytest
9import soxr
10from music_assistant_models.enums import ContentType, MediaType
11from music_assistant_models.media_items import AudioFormat
12
13import music_assistant.providers.sonic_analysis as sonic_mod
14from music_assistant.providers.sonic_analysis import (
15 ANALYSIS_SAMPLE_RATE,
16 SonicAnalysisProvider,
17 SonicSessionData,
18)
19from music_assistant.providers.sonic_analysis.helpers import extract_block_features as _real_extract
20
21
22def _make_provider() -> SonicAnalysisProvider:
23 """Construct a SonicAnalysisProvider with mocked MA infrastructure, bypassing __init__."""
24 mass = MagicMock()
25 mass.streams.audio_analysis.set_audio_analysis = AsyncMock()
26 mass.streams.audio_analysis.get_audio_analysis_version = AsyncMock(return_value=None)
27 mass.create_task = MagicMock(side_effect=lambda coro: coro.close() or MagicMock())
28
29 manifest = MagicMock()
30 manifest.domain = "sonic_analysis"
31
32 p = SonicAnalysisProvider.__new__(SonicAnalysisProvider)
33 p.logger = MagicMock()
34 p.mass = mass
35 p.manifest = manifest
36 p._sessions = {}
37 p._clap_model = MagicMock()
38 p._clap_prompt_order = []
39 p._clap_text_embeddings = None
40 p.analysis_version = 1
41 p.post_analysis = AsyncMock() # type: ignore[method-assign]
42 p.config = MagicMock()
43 p.config.get_value = MagicMock(return_value="fast")
44 return p
45
46
47def _make_streamdetails(duration: float | None = 60.0) -> MagicMock:
48 """Return a minimal streamdetails mock."""
49 sd = MagicMock()
50 sd.item_id = "track-1"
51 sd.provider = "test_provider"
52 sd.media_type = MediaType.TRACK
53 sd.duration = duration
54 return sd
55
56
57def _make_audio_format(sample_rate: int) -> AudioFormat:
58 """Return a 16-bit mono PCM AudioFormat at the given sample rate."""
59 return AudioFormat(
60 content_type=ContentType.PCM_S16LE,
61 sample_rate=sample_rate,
62 bit_depth=16,
63 channels=1,
64 )
65
66
67async def _start_session(
68 provider: SonicAnalysisProvider,
69 session_id: str,
70 sample_rate: int,
71) -> None:
72 """
73 Seed _sessions with a base AnalysisSessionData then call _start_analysis.
74
75 :param provider: The provider instance to register the session on.
76 :param session_id: The session ID to register.
77 :param sample_rate: PCM sample rate in Hz.
78 """
79 from music_assistant.models.audio_analysis_provider import AnalysisSessionData # noqa: PLC0415
80
81 af = _make_audio_format(sample_rate)
82 sd = _make_streamdetails()
83 provider._sessions[session_id] = AnalysisSessionData(streamdetails=sd, audio_format=af)
84 await provider._start_analysis(session_id, sd, af)
85
86
87# ---------------------------------------------------------------------------
88# Test 1: no resampler at 22050 Hz
89# ---------------------------------------------------------------------------
90
91
92@pytest.mark.asyncio
93async def test_start_analysis_no_resampler_at_native_rate() -> None:
94 """_start_analysis with 22050 Hz must leave resampler as None."""
95 provider = _make_provider()
96 await _start_session(provider, "sess-22k", 22050)
97 session = provider._sessions["sess-22k"]
98 assert isinstance(session, SonicSessionData)
99 assert session.resampler is None
100
101
102# ---------------------------------------------------------------------------
103# Test 2: resampler instantiated at 44100 Hz
104# ---------------------------------------------------------------------------
105
106
107@pytest.mark.asyncio
108async def test_start_analysis_creates_resampler_at_44100() -> None:
109 """_start_analysis with 44100 Hz must instantiate a soxr.ResampleStream."""
110 provider = _make_provider()
111 await _start_session(provider, "sess-44k", 44100)
112 session = provider._sessions["sess-44k"]
113 assert isinstance(session, SonicSessionData)
114 assert isinstance(session.resampler, soxr.ResampleStream)
115
116
117# ---------------------------------------------------------------------------
118# Test 3: process_pcm_chunk passes resampled audio (22050 Hz length) to extract_block_features
119# ---------------------------------------------------------------------------
120
121
122@pytest.mark.asyncio
123async def test_process_pcm_chunk_resamples_before_feature_extraction() -> None:
124 """
125 process_pcm_chunk must resample 44100 Hz audio to 22050 Hz before calling extract_block_features.
126
127 Feeds exactly one 10-second block at 44100 Hz (block_bytes = 44100*2*1*10 bytes).
128 extract_block_features must be called with audio whose length is approximately
129 ANALYSIS_SAMPLE_RATE * 10 samples, not 44100 * 10 samples.
130 """
131 provider = _make_provider()
132 session_id = "sess-resample-check"
133 await _start_session(provider, session_id, 44100)
134
135 # One full 10s block at 44100 Hz, 16-bit mono
136 n_samples = 44100 * 10
137 audio_f32 = (np.sin(2 * np.pi * 440 * np.arange(n_samples) / 44100) * 0.5).astype(np.float32)
138 pcm_bytes = (audio_f32 * 32767).astype(np.int16).tobytes()
139
140 captured_lengths: list[int] = []
141
142 def _capture(audio: np.ndarray, sample_rate: int) -> object:
143 captured_lengths.append(len(audio))
144 return _real_extract(audio, sample_rate)
145
146 with patch.object(sonic_mod, "extract_block_features", side_effect=_capture):
147 await provider.process_pcm_chunk(session_id, pcm_bytes)
148
149 assert len(captured_lengths) == 1, (
150 f"Expected exactly one extract_block_features call, got {len(captured_lengths)}"
151 )
152 # Allow 5% tolerance for soxr's internal delay buffering at block boundaries
153 expected = ANALYSIS_SAMPLE_RATE * 10
154 tolerance = int(expected * 0.05)
155 assert abs(captured_lengths[0] - expected) <= tolerance, (
156 f"Expected audio length ~{expected} (22050 Hz), got {captured_lengths[0]} (44100 Hz input)"
157 )
158
159
160# ---------------------------------------------------------------------------
161# Test 4: extract_block_features assertion guard rejects non-22050 input
162# ---------------------------------------------------------------------------
163
164
165def test_extract_block_features_rejects_wrong_sample_rate() -> None:
166 """extract_block_features must raise AssertionError when called with sample_rate != 22050."""
167 audio = np.zeros(8192, dtype=np.float32)
168 with pytest.raises(AssertionError):
169 _real_extract(audio, 44100)
170