/
/
/
1"""Tests for the _decode_resample_extract helper and its asyncio.to_thread offload (T2.3)."""
2
3from __future__ import annotations
4
5import inspect
6from unittest.mock import AsyncMock, MagicMock, patch
7
8import numpy as np
9import pytest
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)
18
19# ---------------------------------------------------------------------------
20# Helpers shared across tests
21# ---------------------------------------------------------------------------
22
23
24def _make_provider() -> SonicAnalysisProvider:
25 """Construct a SonicAnalysisProvider with mocked MA infrastructure."""
26 mass = MagicMock()
27 mass.streams.audio_analysis.set_audio_analysis = AsyncMock()
28 mass.streams.audio_analysis.get_audio_analysis_version = AsyncMock(return_value=None)
29 mass.create_task = MagicMock(side_effect=lambda coro: coro.close() or MagicMock())
30
31 manifest = MagicMock()
32 manifest.domain = "sonic_analysis"
33
34 p = SonicAnalysisProvider.__new__(SonicAnalysisProvider)
35 p.logger = MagicMock()
36 p.mass = mass
37 p.manifest = manifest
38 p._sessions = {}
39 p._clap_model = MagicMock()
40 p._clap_prompt_order = []
41 p._clap_text_embeddings = None
42 p.analysis_version = 1
43 p.post_analysis = AsyncMock() # type: ignore[method-assign]
44 p.config = MagicMock()
45 p.config.get_value = MagicMock(return_value="fast")
46 return p
47
48
49def _make_audio_format(sample_rate: int = 22050) -> AudioFormat:
50 """Return a 16-bit mono PCM AudioFormat at the given sample rate."""
51 return AudioFormat(
52 content_type=ContentType.PCM_S16LE,
53 sample_rate=sample_rate,
54 bit_depth=16,
55 channels=1,
56 )
57
58
59async def _start_session(
60 provider: SonicAnalysisProvider,
61 session_id: str,
62 sample_rate: int = 22050,
63) -> None:
64 """
65 Seed _sessions and call _start_analysis.
66
67 :param provider: The provider instance to register the session on.
68 :param session_id: The session ID to register.
69 :param sample_rate: PCM sample rate in Hz.
70 """
71 from music_assistant.models.audio_analysis_provider import AnalysisSessionData # noqa: PLC0415
72
73 af = _make_audio_format(sample_rate)
74 sd = MagicMock()
75 sd.item_id = "track-t23"
76 sd.provider = "test_provider"
77 sd.media_type = MediaType.TRACK
78 sd.duration = 60.0
79 provider._sessions[session_id] = AnalysisSessionData(streamdetails=sd, audio_format=af)
80 await provider._start_analysis(session_id, sd, af)
81
82
83def _make_one_block_pcm(sample_rate: int = 22050) -> bytes:
84 """Build exactly one 10-second block of 16-bit mono PCM at the given sample rate."""
85 n_samples = sample_rate * 10
86 audio_f32 = (np.sin(2 * np.pi * 440 * np.arange(n_samples) / sample_rate) * 0.5).astype(
87 np.float32
88 )
89 return (audio_f32 * 32767).astype(np.int16).tobytes()
90
91
92# ---------------------------------------------------------------------------
93# Test 1: _decode_resample_extract exists and is callable
94# ---------------------------------------------------------------------------
95
96
97def test_decode_resample_extract_is_callable() -> None:
98 """_decode_resample_extract must be a callable exported from the module."""
99 fn = getattr(sonic_mod, "_decode_resample_extract", None)
100 assert callable(fn), "_decode_resample_extract must be a module-level callable"
101
102
103# ---------------------------------------------------------------------------
104# Test 2: _decode_resample_extract returns the expected 3-tuple
105# ---------------------------------------------------------------------------
106
107
108def test_decode_resample_extract_returns_three_tuple() -> None:
109 """_decode_resample_extract must return (pre_resample, post_resample, BlockFeatures|None)."""
110 from music_assistant.providers.sonic_analysis import _decode_resample_extract # noqa: PLC0415
111
112 n_samples = ANALYSIS_SAMPLE_RATE * 10
113 audio_f32 = np.zeros(n_samples, dtype=np.float32)
114 pcm_bytes = (audio_f32 * 32767).astype(np.int16).tobytes()
115 af = _make_audio_format(ANALYSIS_SAMPLE_RATE)
116
117 result = _decode_resample_extract(af, pcm_bytes, None, ANALYSIS_SAMPLE_RATE, None)
118
119 assert isinstance(result, tuple), "Return value must be a tuple"
120 assert len(result) == 3, "Return tuple must have exactly 3 elements"
121 pre, post, bf = result
122 assert isinstance(pre, np.ndarray), "pre_resample must be np.ndarray"
123 assert isinstance(post, np.ndarray), "post_resample must be np.ndarray"
124 # bf may be None for silent/short audio, but type must be correct
125 from music_assistant.providers.sonic_analysis.helpers import BlockFeatures # noqa: PLC0415
126
127 assert bf is None or isinstance(bf, BlockFeatures)
128
129
130# ---------------------------------------------------------------------------
131# Test 3: process_pcm_chunk uses asyncio.to_thread for decode+extract
132# ---------------------------------------------------------------------------
133
134
135@pytest.mark.asyncio
136async def test_process_pcm_chunk_offloads_via_run_offloaded() -> None:
137 """
138 process_pcm_chunk must call _decode_resample_extract off the event loop.
139
140 Verified by:
141 1. Patching _decode_resample_extract with a MagicMock that returns a valid 3-tuple.
142 2. Asserting the mock was called once per block.
143 3. Asserting the source of process_pcm_chunk routes through the offload seam.
144 """
145 provider = _make_provider()
146 session_id = "sess-t23-offload"
147 await _start_session(provider, session_id, ANALYSIS_SAMPLE_RATE)
148
149 pcm_bytes = _make_one_block_pcm(ANALYSIS_SAMPLE_RATE)
150 n_samples = ANALYSIS_SAMPLE_RATE * 10
151 dummy_audio = np.zeros(n_samples, dtype=np.float32)
152
153 mock_helper = MagicMock(return_value=(dummy_audio, dummy_audio, None))
154
155 with patch.object(sonic_mod, "_decode_resample_extract", mock_helper):
156 await provider.process_pcm_chunk(session_id, pcm_bytes)
157
158 assert mock_helper.call_count == 1, (
159 f"_decode_resample_extract must be called once per block; got {mock_helper.call_count}"
160 )
161
162 # Structural check: process_pcm_chunk must offload decode+extract off the event loop
163 # (via the concurrency-bounded _run_offloaded seam rather than running it inline).
164 source = inspect.getsource(provider.process_pcm_chunk)
165 assert "_run_offloaded" in source, (
166 "process_pcm_chunk must offload decode+extract via _run_offloaded (CPU offload)"
167 )
168 assert "_decode_resample_extract" in source, (
169 "process_pcm_chunk must reference _decode_resample_extract"
170 )
171