/
/
/
1"""Tests for the per-window CLAP inference + session accumulator logic."""
2
3from __future__ import annotations
4
5from unittest.mock import MagicMock, patch
6
7import numpy as np
8import pytest
9import torch
10
11from music_assistant.providers.sonic_analysis import (
12 CLAP_WINDOW_SECONDS,
13 SonicAnalysisProvider,
14 SonicSessionData,
15)
16
17SR = 22050
18WINDOW_SAMPLES = CLAP_WINDOW_SECONDS * SR
19
20
21def _make_provider(
22 embedding_value: float = 0.5,
23 similarity_row: list[float] | None = None,
24) -> tuple[SonicAnalysisProvider, MagicMock, MagicMock]:
25 """
26 Stub provider with a mocked CLAP model returning predictable tensors.
27
28 :returns: (provider, fake_model_mock, fake_logger_mock).
29 """
30 p = SonicAnalysisProvider.__new__(SonicAnalysisProvider)
31 fake_logger = MagicMock()
32 p.logger = fake_logger
33 # _run_offloaded reads mass.streams.audio_analysis.analysis_semaphore; a plain Mock
34 # attribute is not an asyncio.Semaphore, so offloads fall back to a plain worker thread.
35 p.mass = MagicMock()
36
37 fake_model = MagicMock()
38 sim = similarity_row if similarity_row is not None else [1.0, 2.0, 3.0, 4.0]
39 fake_model.get_audio_embeddings_from_tensor = MagicMock(
40 return_value=torch.full((1, 1024), embedding_value, dtype=torch.float32)
41 )
42 fake_model.compute_similarity = MagicMock(return_value=torch.tensor([sim], dtype=torch.float32))
43
44 p._clap_model = fake_model
45 p._clap_text_embeddings = MagicMock()
46 return p, fake_model, fake_logger
47
48
49def _make_session() -> SonicSessionData:
50 return SonicSessionData(streamdetails=MagicMock(), audio_format=MagicMock())
51
52
53@pytest.mark.asyncio
54async def test_first_completion_initializes_sums() -> None:
55 """First successful inference allocates zero-arrays for the running sums."""
56 p, _, _ = _make_provider(embedding_value=0.5)
57 session = _make_session()
58 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
59
60 await p._run_single_clap_window(session, window, SR)
61
62 assert session.clap_sum_embedding is not None
63 assert session.clap_sum_embedding.shape == (1024,)
64 assert session.clap_sum_similarities is not None
65 assert session.clap_sum_similarities.shape == (4,)
66 assert session.clap_completed_count == 1
67
68
69@pytest.mark.asyncio
70async def test_subsequent_completions_accumulate() -> None:
71 """Three calls produce running sums equal to 3x the per-call output."""
72 p, _, _ = _make_provider(embedding_value=0.5, similarity_row=[1.0, 2.0, 3.0, 4.0])
73 session = _make_session()
74 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
75
76 for _ in range(3):
77 await p._run_single_clap_window(session, window, SR)
78
79 assert session.clap_completed_count == 3
80 assert session.clap_sum_embedding is not None
81 assert session.clap_sum_similarities is not None
82 np.testing.assert_array_almost_equal(
83 session.clap_sum_embedding, np.full(1024, 1.5, dtype=np.float32)
84 )
85 np.testing.assert_array_almost_equal(
86 session.clap_sum_similarities, np.array([3.0, 6.0, 9.0, 12.0], dtype=np.float32)
87 )
88
89
90@pytest.mark.asyncio
91async def test_clap_failure_logs_and_skips_accumulation() -> None:
92 """A CLAP exception logs at debug, leaves the session sums untouched."""
93 p, fake_model, fake_logger = _make_provider()
94 fake_model.get_audio_embeddings_from_tensor = MagicMock(side_effect=RuntimeError("boom"))
95 session = _make_session()
96 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
97
98 await p._run_single_clap_window(session, window, SR)
99
100 assert session.clap_completed_count == 0
101 assert session.clap_sum_embedding is None
102 assert session.clap_sum_similarities is None
103 fake_logger.debug.assert_called()
104
105
106@pytest.mark.asyncio
107async def test_no_clap_model_is_no_op() -> None:
108 """If the CLAP model never loaded, the helper returns silently."""
109 p, _, _ = _make_provider()
110 p._clap_model = None
111 session = _make_session()
112 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
113
114 await p._run_single_clap_window(session, window, SR)
115
116 assert session.clap_completed_count == 0
117 assert session.clap_sum_embedding is None
118
119
120@pytest.mark.asyncio
121async def test_calls_compute_similarity_with_text_embeddings() -> None:
122 """The session's CLAP text embeddings are forwarded to compute_similarity."""
123 p, fake_model, _ = _make_provider()
124 session = _make_session()
125 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
126
127 await p._run_single_clap_window(session, window, SR)
128
129 fake_model.compute_similarity.assert_called_once()
130 _audio_embs, text_embs_arg = fake_model.compute_similarity.call_args.args
131 assert text_embs_arg is p._clap_text_embeddings
132
133
134def test_single_window_inference_sync_returns_none_when_model_is_none() -> None:
135 """
136 _single_window_inference_sync returns None cleanly when the CLAP model is unloaded.
137
138 Simulates the unload() race: model is nulled before the sync thread reads it.
139 """
140 p, _, _ = _make_provider()
141 p._clap_model = None
142 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
143
144 result = p._single_window_inference_sync(window, SR)
145
146 assert result is None
147
148
149def test_single_window_inference_sync_returns_none_when_text_embeddings_are_none() -> None:
150 """_single_window_inference_sync returns None cleanly when text embeddings are unloaded."""
151 p, _, _ = _make_provider()
152 p._clap_text_embeddings = None
153 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
154
155 result = p._single_window_inference_sync(window, SR)
156
157 assert result is None
158
159
160@pytest.mark.asyncio
161async def test_run_single_clap_window_handles_none_return_from_sync() -> None:
162 """
163 _run_single_clap_window does not log a warning when inference returns None (unload race).
164
165 The unload race: model is non-None at the early-return check in _run_single_clap_window,
166 but is nulled before the sync thread reads it, causing _single_window_inference_sync to
167 return None. The caller must handle this without a warning.
168 """
169 p, _, fake_logger = _make_provider()
170 session = _make_session()
171 window = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
172
173 with patch.object(p, "_single_window_inference_sync", return_value=None):
174 await p._run_single_clap_window(session, window, SR)
175
176 assert session.clap_completed_count == 0
177 fake_logger.warning.assert_not_called()
178 # Discriminate the new graceful path (logs "unloaded mid-flight") from the
179 # old bug path (would log a TypeError about unpacking NoneType via the
180 # broad except handler). Without this assertion the test is always-green.
181 fake_logger.debug.assert_called()
182 debug_messages = [str(call.args[0]) for call in fake_logger.debug.call_args_list]
183 assert any("unloaded" in msg for msg in debug_messages), (
184 f"Expected 'unloaded' debug message; got: {debug_messages}"
185 )
186