/
/
/
1"""Store, gate, and API tests for audio_analysis_failures against a real temp DB."""
2
3from __future__ import annotations
4
5import pathlib
6from datetime import UTC, datetime, timedelta
7from typing import TYPE_CHECKING
8from unittest.mock import MagicMock
9
10import pytest
11
12from music_assistant.constants import (
13 DB_TABLE_AUDIO_ANALYSIS,
14 DB_TABLE_AUDIO_ANALYSIS_FAILURES,
15 DB_TABLE_PROVIDER_MAPPINGS,
16)
17from music_assistant.controllers.streams.audio_analysis import AudioAnalysisController
18from music_assistant.helpers.database import DatabaseConnection
19from music_assistant.models.audio_analysis import AudioAnalysisData
20from music_assistant.models.music_provider import MusicProvider
21
22if TYPE_CHECKING:
23 from collections.abc import AsyncGenerator
24
25
26@pytest.fixture
27async def real_db(tmp_path: pathlib.Path) -> AsyncGenerator[DatabaseConnection]:
28 """Create a real on-disk sqlite DB with the minimal tables the gate/store touch."""
29 db = DatabaseConnection(str(tmp_path / "test.db"))
30 await db.setup()
31 await db.execute(
32 f"CREATE TABLE {DB_TABLE_PROVIDER_MAPPINGS}("
33 "provider_item_id TEXT, provider_instance TEXT, provider_domain TEXT, media_type TEXT)"
34 )
35 await db.execute(
36 f"CREATE TABLE {DB_TABLE_AUDIO_ANALYSIS}("
37 "id INTEGER PRIMARY KEY AUTOINCREMENT, media_type TEXT, item_id TEXT, provider TEXT, "
38 "aa_provider_domain TEXT, analysis_data json, analysis_version INTEGER, "
39 "timestamp_created INTEGER DEFAULT (cast(strftime('%s','now') as int)), "
40 "UNIQUE(item_id,provider,aa_provider_domain,media_type))"
41 )
42 await db.execute(
43 f"CREATE TABLE {DB_TABLE_AUDIO_ANALYSIS_FAILURES}("
44 "id INTEGER PRIMARY KEY AUTOINCREMENT, media_type TEXT, item_id TEXT, provider TEXT, "
45 "aa_provider_domain TEXT, reason TEXT, analysis_version INTEGER NOT NULL DEFAULT 1, "
46 "next_retry INTEGER, "
47 "timestamp_created INTEGER DEFAULT (cast(strftime('%s','now') as int)), "
48 "UNIQUE(item_id,provider,aa_provider_domain,media_type))"
49 )
50 await db.commit()
51 yield db
52 await db.close()
53
54
55def _make_fs_music_provider() -> MagicMock:
56 """Return a fake filesystem (non-streaming) MusicProvider keyed by instance_id."""
57 prov = MagicMock(spec=MusicProvider)
58 prov.is_streaming_provider = False
59 prov.domain = "filesystem_local"
60 prov.instance_id = "filesystem_local--abc"
61 prov.available = True
62 return prov
63
64
65def _make_controller(real_db: DatabaseConnection, music_prov: MagicMock) -> AudioAnalysisController:
66 """Return an AudioAnalysisController whose mass.music.database is the real temp DB."""
67 streams = MagicMock()
68 mass = MagicMock()
69 streams.mass = mass
70 mass.music.database = real_db
71 mass.get_provider = MagicMock(return_value=music_prov)
72 mass.get_providers = MagicMock(return_value=[music_prov])
73 return AudioAnalysisController(streams)
74
75
76@pytest.mark.asyncio
77async def test_table_roundtrip(real_db: DatabaseConnection) -> None:
78 """A row inserted into the failures table reads back with next_retry NULL preserved."""
79 await real_db.insert_or_replace(
80 DB_TABLE_AUDIO_ANALYSIS_FAILURES,
81 {
82 "media_type": "track",
83 "item_id": "t1",
84 "provider": "filesystem_local--abc",
85 "aa_provider_domain": "sonic_analysis",
86 "reason": "boom",
87 "analysis_version": 1,
88 "next_retry": None,
89 },
90 )
91 rows = await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)
92 assert len(rows) == 1
93 assert rows[0]["reason"] == "boom"
94 assert rows[0]["next_retry"] is None
95
96
97@pytest.mark.asyncio
98async def test_record_and_clear_failure_roundtrip(real_db: DatabaseConnection) -> None:
99 """record_analysis_failure writes prov_key + NULL retry; clear deletes the row."""
100 music_prov = _make_fs_music_provider()
101 controller = _make_controller(real_db, music_prov)
102
103 await controller.record_analysis_failure(
104 item_id="t1",
105 provider_instance_id_or_domain="filesystem_local--abc",
106 aa_provider_domain="sonic_analysis",
107 reason="no usable audio frames extracted",
108 )
109 rows = await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)
110 assert len(rows) == 1
111 assert rows[0]["provider"] == "filesystem_local--abc" # instance_id for non-streaming
112 assert rows[0]["next_retry"] is None
113 assert rows[0]["reason"] == "no usable audio frames extracted"
114
115 await controller.clear_analysis_failure(
116 item_id="t1",
117 provider_instance_id_or_domain="filesystem_local--abc",
118 aa_provider_domain="sonic_analysis",
119 )
120 rows = await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)
121 assert rows == []
122
123
124@pytest.mark.asyncio
125async def test_record_failure_converts_retry_at_to_epoch(real_db: DatabaseConnection) -> None:
126 """A datetime retry_at is stored as an integer epoch in next_retry."""
127 music_prov = _make_fs_music_provider()
128 controller = _make_controller(real_db, music_prov)
129 when = datetime(2030, 6, 1, 12, 0, tzinfo=UTC)
130
131 await controller.record_analysis_failure(
132 item_id="t2",
133 provider_instance_id_or_domain="filesystem_local--abc",
134 aa_provider_domain="sonic_analysis",
135 reason="offline",
136 retry_at=when,
137 )
138 rows = await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)
139 assert rows[0]["next_retry"] == int(when.timestamp())
140
141
142@pytest.mark.asyncio
143async def test_record_failure_skips_when_not_music_provider(real_db: DatabaseConnection) -> None:
144 """No row is written when the provider lookup is not a MusicProvider."""
145 controller = _make_controller(real_db, music_prov=MagicMock()) # not a MusicProvider spec
146
147 await controller.record_analysis_failure(
148 item_id="t3",
149 provider_instance_id_or_domain="whatever",
150 aa_provider_domain="sonic_analysis",
151 reason="x",
152 )
153 rows = await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)
154 assert rows == []
155
156
157@pytest.mark.asyncio
158async def test_set_audio_analysis_clears_existing_failure(real_db: DatabaseConnection) -> None:
159 """A successful set_audio_analysis removes any prior failure row for the same key."""
160 music_prov = _make_fs_music_provider()
161 controller = _make_controller(real_db, music_prov)
162
163 await controller.record_analysis_failure(
164 item_id="t1",
165 provider_instance_id_or_domain="filesystem_local--abc",
166 aa_provider_domain="sonic_analysis",
167 reason="boom",
168 )
169 assert len(await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)) == 1
170
171 await controller.set_audio_analysis(
172 item_id="t1",
173 provider_instance_id_or_domain="filesystem_local--abc",
174 aa_provider_domain="sonic_analysis",
175 analysis=AudioAnalysisData(energy=0.5),
176 )
177 assert await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0) == []
178
179
180async def _insert_pm(db: DatabaseConnection, item_id: str) -> None:
181 await db.insert(
182 DB_TABLE_PROVIDER_MAPPINGS,
183 {
184 "provider_item_id": item_id,
185 "provider_instance": "filesystem_local--abc",
186 "provider_domain": "filesystem_local",
187 "media_type": "track",
188 },
189 )
190
191
192async def _insert_failure(
193 db: DatabaseConnection, item_id: str, *, next_retry: int | None, version: int = 1
194) -> None:
195 await db.insert_or_replace(
196 DB_TABLE_AUDIO_ANALYSIS_FAILURES,
197 {
198 "media_type": "track",
199 "item_id": item_id,
200 "provider": "filesystem_local--abc",
201 "aa_provider_domain": "sonic_analysis",
202 "reason": "x",
203 "analysis_version": version,
204 "next_retry": next_retry,
205 },
206 )
207
208
209async def _insert_analysis(
210 db: DatabaseConnection,
211 item_id: str,
212 *,
213 version: int | None,
214 aa_domain: str = "sonic_analysis",
215) -> None:
216 await db.insert_or_replace(
217 DB_TABLE_AUDIO_ANALYSIS,
218 {
219 "media_type": "track",
220 "item_id": item_id,
221 "provider": "filesystem_local--abc",
222 "aa_provider_domain": aa_domain,
223 "analysis_data": "{}",
224 "analysis_version": version,
225 },
226 )
227
228
229@pytest.mark.asyncio
230async def test_candidate_gate_excludes_blocked_includes_eligible(
231 real_db: DatabaseConnection,
232) -> None:
233 """Blocked (NULL / future, current version) failures are excluded; others are candidates."""
234 music_prov = _make_fs_music_provider()
235 controller = _make_controller(real_db, music_prov)
236
237 # never-retry, current version -> excluded
238 await _insert_pm(real_db, "blocked_null")
239 await _insert_failure(real_db, "blocked_null", next_retry=None, version=2)
240 # future retry, current version -> excluded
241 future = int((datetime.now(UTC) + timedelta(days=1)).timestamp())
242 await _insert_pm(real_db, "blocked_future")
243 await _insert_failure(real_db, "blocked_future", next_retry=future, version=2)
244 # past-due retry -> included
245 past = int((datetime.now(UTC) - timedelta(days=1)).timestamp())
246 await _insert_pm(real_db, "due")
247 await _insert_failure(real_db, "due", next_retry=past, version=2)
248 # stale version (1 < current 2) -> included
249 await _insert_pm(real_db, "stale")
250 await _insert_failure(real_db, "stale", next_retry=None, version=1)
251 # no failure at all -> included
252 await _insert_pm(real_db, "clean")
253
254 candidates = await controller._find_candidates_missing_analysis({"sonic_analysis": 2}, limit=0)
255 found = {c["item_id"] for c in candidates}
256 assert found == {"due", "stale", "clean"}
257
258
259@pytest.mark.asyncio
260async def test_candidate_gate_resurfaces_stale_analysis_versions(
261 real_db: DatabaseConnection,
262) -> None:
263 """Analysis rows below the current version (or NULL) surface; current-or-newer do not."""
264 music_prov = _make_fs_music_provider()
265 controller = _make_controller(real_db, music_prov)
266
267 # row at current version -> excluded
268 await _insert_pm(real_db, "current")
269 await _insert_analysis(real_db, "current", version=2)
270 # row at newer version -> excluded
271 await _insert_pm(real_db, "newer")
272 await _insert_analysis(real_db, "newer", version=3)
273 # row at older version -> included
274 await _insert_pm(real_db, "stale")
275 await _insert_analysis(real_db, "stale", version=1)
276 # pre-versioning row (NULL version) -> included
277 await _insert_pm(real_db, "nullver")
278 await _insert_analysis(real_db, "nullver", version=None)
279
280 candidates = await controller._find_candidates_missing_analysis({"sonic_analysis": 2}, limit=0)
281 found = {c["item_id"] for c in candidates}
282 assert found == {"stale", "nullver"}
283
284
285@pytest.mark.asyncio
286async def test_candidate_gate_tracks_versions_per_domain(real_db: DatabaseConnection) -> None:
287 """With multiple AA domains, each domain is gated by its own current version."""
288 music_prov = _make_fs_music_provider()
289 controller = _make_controller(real_db, music_prov)
290
291 # analyzed at v1 for both domains; only sonic_analysis bumped to v2
292 await _insert_pm(real_db, "t1")
293 await _insert_analysis(real_db, "t1", version=1, aa_domain="sonic_analysis")
294 await _insert_analysis(real_db, "t1", version=1, aa_domain="loudness_analysis")
295
296 candidates = await controller._find_candidates_missing_analysis(
297 {"sonic_analysis": 2, "loudness_analysis": 1}, limit=0
298 )
299 assert len(candidates) == 1
300 assert candidates[0]["item_id"] == "t1"
301 assert candidates[0]["missing_domains"] == ["sonic_analysis"]
302
303
304@pytest.mark.asyncio
305async def test_get_failures_returns_rows_and_filters_by_domain(
306 real_db: DatabaseConnection,
307) -> None:
308 """get_failures returns the stored shape, optionally filtered by aa_domain."""
309 music_prov = _make_fs_music_provider()
310 controller = _make_controller(real_db, music_prov)
311 await _insert_failure(real_db, "t1", next_retry=None)
312 await real_db.insert_or_replace(
313 DB_TABLE_AUDIO_ANALYSIS_FAILURES,
314 {
315 "media_type": "track",
316 "item_id": "t2",
317 "provider": "filesystem_local--abc",
318 "aa_provider_domain": "loudness_analysis",
319 "reason": "y",
320 "analysis_version": 1,
321 "next_retry": None,
322 },
323 )
324
325 all_rows = await controller.get_failures()
326 assert {r["item_id"] for r in all_rows} == {"t1", "t2"}
327 assert set(all_rows[0]) == {
328 "item_id",
329 "provider",
330 "aa_provider_domain",
331 "reason",
332 "next_retry",
333 "timestamp_created",
334 }
335
336 sonic_rows = await controller.get_failures(aa_domain="sonic_analysis")
337 assert {r["item_id"] for r in sonic_rows} == {"t1"}
338
339
340@pytest.mark.asyncio
341async def test_clear_failures_requires_a_filter(real_db: DatabaseConnection) -> None:
342 """clear_failures with no filter deletes nothing and returns 0."""
343 music_prov = _make_fs_music_provider()
344 controller = _make_controller(real_db, music_prov)
345 await _insert_failure(real_db, "t1", next_retry=None)
346
347 deleted = await controller.clear_failures()
348 assert deleted == 0
349 assert len(await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0)) == 1
350
351
352@pytest.mark.asyncio
353async def test_clear_failures_by_domain(real_db: DatabaseConnection) -> None:
354 """clear_failures(aa_domain=...) deletes all rows for that domain and returns the count."""
355 music_prov = _make_fs_music_provider()
356 controller = _make_controller(real_db, music_prov)
357 await _insert_failure(real_db, "t1", next_retry=None)
358 await _insert_failure(real_db, "t2", next_retry=None)
359
360 deleted = await controller.clear_failures(aa_domain="sonic_analysis")
361 assert deleted == 2
362 assert await real_db.get_rows(DB_TABLE_AUDIO_ANALYSIS_FAILURES, limit=0) == []
363