/
/
1"""Tests for the live-CLAP finalize + cancel paths."""
2
3from __future__ import annotations
4
5import asyncio
6import math
7from unittest.mock import MagicMock
8
9import numpy as np
10import pytest
11
12from music_assistant.models.audio_analysis import AudioAnalysisData
13from music_assistant.providers.sonic_analysis import (
14 SonicAnalysisProvider,
15 SonicSessionData,
16)
17from music_assistant.providers.sonic_analysis.clap_prompts import (
18 CALIBRATION,
19 SCALAR_PROMPT_PAIRS,
20)
21
22
23def _make_provider() -> tuple[SonicAnalysisProvider, MagicMock]:
24 """Stub provider with the prompt order needed by _run_live_clap_if_eligible."""
25 p = SonicAnalysisProvider.__new__(SonicAnalysisProvider)
26 fake_logger = MagicMock()
27 p.logger = fake_logger
28 p._clap_prompt_order = list(SCALAR_PROMPT_PAIRS.items())
29 return p, fake_logger
30
31
32def _make_session(target_starts: list[int] | None = None) -> SonicSessionData:
33 starts = list(target_starts) if target_starts is not None else []
34 return SonicSessionData(
35 streamdetails=MagicMock(),
36 audio_format=MagicMock(),
37 clap_target_starts=starts,
38 clap_target_buffers=[[] for _ in starts],
39 clap_target_complete=[False] * len(starts),
40 )
41
42
43@pytest.mark.asyncio
44async def test_no_targets_short_circuits_silently() -> None:
45 """A session with no planned targets returns immediately, no log noise."""
46 p, fake_logger = _make_provider()
47 session = _make_session(target_starts=[])
48 analysis = AudioAnalysisData()
49
50 await p._run_live_clap_if_eligible(session, analysis)
51
52 assert analysis.danceability is None
53 assert analysis.valence is None
54 assert analysis.arousal is None
55 assert analysis.instrumentalness is None
56 assert analysis.acousticness is None
57 assert analysis.speechiness is None
58 assert analysis.extra_data is None or "clap_embedding" not in (analysis.extra_data or {})
59 fake_logger.warning.assert_not_called()
60
61
62@pytest.mark.asyncio
63async def test_no_completions_logs_warning() -> None:
64 """Targets planned but zero windows completed â warning, no scalar updates."""
65 p, fake_logger = _make_provider()
66 session = _make_session(target_starts=[0, 100, 200])
67 # No tasks added, no completed_count incremented
68 analysis = AudioAnalysisData()
69
70 await p._run_live_clap_if_eligible(session, analysis)
71
72 assert analysis.danceability is None
73 assert analysis.valence is None
74 assert analysis.arousal is None
75 assert analysis.instrumentalness is None
76 assert analysis.acousticness is None
77 assert analysis.speechiness is None
78 assert analysis.extra_data is None or "clap_embedding" not in (analysis.extra_data or {})
79 fake_logger.warning.assert_called_once()
80
81
82@pytest.mark.asyncio
83async def test_mean_pools_and_calibrates_scalars() -> None:
84 """Three completed windows with known sums produce calibrated scalars and L2-normalized embedding."""
85 p, _ = _make_provider()
86 session = _make_session(target_starts=[0, 100, 200])
87
88 n_pairs = len(SCALAR_PROMPT_PAIRS)
89 # Sums equivalent to mean_emb = 0.5 (pre-norm), mean_sim = [1.0, 0.0, 1.0, 0.0, ...]
90 session.clap_completed_count = 3
91 session.clap_sum_embedding = np.full(1024, 1.5, dtype=np.float32) # mean = 0.5
92 raw_sims = np.zeros(2 * n_pairs, dtype=np.float32)
93 for i in range(n_pairs):
94 raw_sims[i * 2] = 3.0 # pos_logit mean = 1.0
95 raw_sims[i * 2 + 1] = 0.0 # neg_logit mean = 0.0
96 session.clap_sum_similarities = raw_sims
97
98 analysis = AudioAnalysisData()
99 await p._run_live_clap_if_eligible(session, analysis)
100
101 # Embedding: pre-norm = 0.5 across 1024 dims; ||v|| = sqrt(1024 * 0.25) = 16.0; normalized = 1/32
102 assert analysis.extra_data is not None
103 emb = np.asarray(analysis.extra_data["clap_embedding"], dtype=np.float32)
104 expected_norm = math.sqrt(1024 * 0.25)
105 expected_value = 0.5 / expected_norm
106 np.testing.assert_array_almost_equal(emb, np.full(1024, expected_value), decimal=5)
107
108 for scalar_name in SCALAR_PROMPT_PAIRS:
109 a, b = CALIBRATION[scalar_name]
110 expected = 1.0 / (1.0 + math.exp(-(a * 1.0 + b)))
111 assert getattr(analysis, scalar_name) == pytest.approx(expected)
112
113
114@pytest.mark.asyncio
115async def test_awaits_pending_tasks() -> None:
116 """Tasks still in flight are awaited before the mean-pool runs."""
117 p, _ = _make_provider()
118 session = _make_session(target_starts=[0])
119
120 completion_event = asyncio.Event()
121
122 async def slow_inference() -> None:
123 await completion_event.wait()
124 n_pairs = len(SCALAR_PROMPT_PAIRS)
125 session.clap_sum_embedding = np.ones(1024, dtype=np.float32)
126 session.clap_sum_similarities = np.zeros(2 * n_pairs, dtype=np.float32)
127 session.clap_completed_count = 1
128
129 task = asyncio.create_task(slow_inference())
130 session.clap_inference_tasks.append(task)
131
132 finalize_task = asyncio.create_task(p._run_live_clap_if_eligible(session, AudioAnalysisData()))
133 # Yield once: finalize should be blocked on the gather
134 await asyncio.sleep(0)
135 assert not finalize_task.done()
136
137 # Release the inference task
138 completion_event.set()
139 await finalize_task
140 assert task.done()
141
142
143@pytest.mark.asyncio
144async def test_cancel_aborts_pending_tasks_and_clears_buffers() -> None:
145 """Cancel cancels in-flight inferences and resets per-window buffers."""
146 p = SonicAnalysisProvider.__new__(SonicAnalysisProvider)
147 p.logger = MagicMock()
148 p._sessions = {}
149
150 session = _make_session(target_starts=[0, 100])
151 session.clap_target_buffers = [
152 [np.zeros(1024, dtype=np.float32)],
153 [np.zeros(2048, dtype=np.float32)],
154 ]
155
156 started = asyncio.Event()
157 cancelled = asyncio.Event()
158
159 async def long_running() -> None:
160 started.set()
161 try:
162 await asyncio.sleep(60)
163 except asyncio.CancelledError:
164 cancelled.set()
165 raise
166
167 task = asyncio.create_task(long_running())
168 session.clap_inference_tasks.append(task)
169 await started.wait()
170
171 p._sessions["sess"] = session
172 await p.cancel("sess")
173
174 # Give the cancellation a tick to propagate
175 await asyncio.sleep(0)
176 assert task.cancelled() or cancelled.is_set()
177 assert session.clap_target_buffers == []
178