/
/
/
1"""Tests for the AudioAnalysisController."""
2
3from __future__ import annotations
4
5import asyncio
6import inspect
7import sqlite3
8from collections.abc import AsyncGenerator, Mapping
9from concurrent.futures import ThreadPoolExecutor
10from contextlib import closing
11from typing import Any, cast
12from unittest.mock import AsyncMock, MagicMock, patch
13
14import numpy as np
15import pytest
16from music_assistant_models.audio_analysis import AudioAnalysisCoverage
17from music_assistant_models.enums import ContentType, MediaType, StreamType
18from music_assistant_models.errors import ProviderUnavailableError
19from music_assistant_models.media_items import AudioFormat, ProviderMapping, Track
20
21import music_assistant.controllers.streams.audio_analysis as audio_analysis_mod
22from music_assistant.constants import (
23 DB_TABLE_AUDIO_ANALYSIS,
24 DEFAULT_BACKGROUND_SCAN_CONCURRENCY,
25 _default_background_scan_concurrency,
26)
27from music_assistant.controllers.streams.audio_analysis import (
28 LOUDNESS_ANALYSIS_DOMAIN,
29 SMART_FADES_ANALYSIS_DOMAIN,
30 SONIC_ANALYSIS_DOMAIN,
31 AudioAnalysisController,
32 _merged_from_rows,
33)
34from music_assistant.controllers.streams.audio_buffer import AudioBufferEOF
35from music_assistant.helpers.json import json_dumps
36from music_assistant.models.audio_analysis import AudioAnalysisData
37from music_assistant.models.audio_analysis_provider import (
38 AudioAnalysisProvider,
39 InstrumentedSemaphore,
40)
41from music_assistant.models.music_provider import MusicProvider
42
43
44@pytest.mark.asyncio
45async def test_distribute_chunk_calls_all_providers() -> None:
46 """_distribute_chunk must invoke process_pcm_chunk on every active provider."""
47 controller = _make_controller()
48 session_key = "track://provider/abc"
49 controller._active_sessions[session_key] = {"prov-1", "prov-2"}
50
51 p1 = _make_aa_provider("prov-1", available=True)
52 p2 = _make_aa_provider("prov-2", available=True)
53 provider_map = {"prov-1": p1, "prov-2": p2}
54 controller.mass.get_provider = MagicMock(side_effect=provider_map.get) # type: ignore[method-assign]
55
56 await controller._distribute_chunk(session_key, b"\x00" * 1024)
57
58 p1.process_pcm_chunk.assert_awaited_once_with(session_key, b"\x00" * 1024)
59 p2.process_pcm_chunk.assert_awaited_once_with(session_key, b"\x00" * 1024)
60
61
62def test_ensure_inference_runtime_configured_is_idempotent() -> None:
63 """The inference runtime (torch thread caps) is configured once per controller."""
64 controller = _make_controller()
65 with (
66 patch("torch.set_num_threads") as set_threads,
67 patch("torch.set_num_interop_threads"),
68 patch("torch.backends.nnpack.set_flags"),
69 ):
70 controller.ensure_inference_runtime_configured()
71 controller.ensure_inference_runtime_configured()
72 set_threads.assert_called_once()
73 if controller.analysis_executor is not None:
74 controller.analysis_executor.shutdown(wait=False)
75
76
77def test_ensure_inference_runtime_creates_solo_lock_and_executor() -> None:
78 """Runtime config creates the playback-priority solo lock and a dedicated worker pool."""
79 controller = _make_controller()
80 with (
81 patch("torch.set_num_threads"),
82 patch("torch.set_num_interop_threads"),
83 patch("torch.backends.nnpack.set_flags"),
84 ):
85 controller.ensure_inference_runtime_configured()
86 try:
87 assert isinstance(controller.analysis_solo_lock, asyncio.Lock)
88 assert isinstance(controller.analysis_executor, ThreadPoolExecutor)
89 finally:
90 if controller.analysis_executor is not None:
91 controller.analysis_executor.shutdown(wait=False)
92
93
94def test_playback_active_delegates_to_streams() -> None:
95 """playback_active reflects the streams controller's active-output-stream gauge."""
96 controller = _make_controller()
97 controller.streams.output_stream_active = MagicMock(return_value=True) # type: ignore[method-assign]
98 assert controller.playback_active() is True
99 controller.streams.output_stream_active = MagicMock(return_value=False) # type: ignore[method-assign]
100 assert controller.playback_active() is False
101
102
103@pytest.mark.parametrize(
104 ("cpu_count", "expected_permits"),
105 [(2, 1), (4, 2), (8, 4), (16, 8)],
106)
107@pytest.mark.asyncio
108async def test_analysis_concurrency_capped_at_half_cores(
109 cpu_count: int, expected_permits: int
110) -> None:
111 """The analysis concurrency cap is half the cores (min 1) on every host."""
112 controller = _make_controller()
113 with (
114 patch(
115 "music_assistant.controllers.streams.audio_analysis.os.process_cpu_count",
116 return_value=cpu_count,
117 ),
118 patch("torch.set_num_threads"),
119 patch("torch.set_num_interop_threads"),
120 patch("torch.backends.nnpack.set_flags"),
121 ):
122 controller.ensure_inference_runtime_configured()
123 semaphore = controller.analysis_semaphore
124 assert isinstance(semaphore, asyncio.Semaphore)
125 # Exactly `expected_permits` acquires exhaust the cap.
126 for _ in range(expected_permits):
127 await semaphore.acquire()
128 assert semaphore.locked()
129
130
131@pytest.mark.asyncio
132async def test_instrumented_semaphore_tracks_in_flight_and_waiters() -> None:
133 """InstrumentedSemaphore exposes live permit-in-use and queued-acquirer counts."""
134 sem = InstrumentedSemaphore(2)
135 assert (sem.capacity, sem.in_flight, sem.waiters) == (2, 0, 0)
136
137 await sem.acquire()
138 await sem.acquire()
139 assert sem.in_flight == 2
140 assert sem.locked()
141
142 # A third acquire blocks behind the cap and registers as a waiter.
143 blocked = asyncio.ensure_future(sem.acquire())
144 await asyncio.sleep(0)
145 assert sem.waiters == 1
146 assert sem.in_flight == 2
147
148 # Freeing a permit lets the queued acquirer through; the queue drains.
149 sem.release()
150 await blocked
151 assert sem.waiters == 0
152 assert sem.in_flight == 2
153
154 sem.release()
155 sem.release()
156 assert sem.in_flight == 0
157
158
159@pytest.mark.asyncio
160async def test_distribute_chunk_evicts_provider_on_timeout() -> None:
161 """A provider whose process_pcm_chunk exceeds max_interval is evicted."""
162 controller = _make_controller()
163 session_key = "track://provider/abc"
164 controller._active_sessions[session_key] = {"slow", "fast"}
165
166 async def _hang(*_args: object, **_kwargs: object) -> None:
167 await asyncio.sleep(10)
168
169 slow = _make_aa_provider("slow", available=True, process_pcm_chunk=AsyncMock(side_effect=_hang))
170 fast = _make_aa_provider("fast", available=True)
171 provider_map = {"slow": slow, "fast": fast}
172 controller.mass.get_provider = MagicMock(side_effect=provider_map.get) # type: ignore[method-assign]
173
174 await controller._distribute_chunk(session_key, b"\x00" * 1024, max_interval=0.05)
175
176 assert "slow" not in controller._active_sessions[session_key]
177 assert "fast" in controller._active_sessions[session_key]
178
179
180@pytest.mark.asyncio
181async def test_distribute_chunk_evicts_provider_on_exception() -> None:
182 """A provider that raises in process_pcm_chunk is evicted; others continue."""
183 controller = _make_controller()
184 session_key = "track://provider/abc"
185 controller._active_sessions[session_key] = {"raises", "ok"}
186
187 raises = _make_aa_provider(
188 "raises",
189 available=True,
190 process_pcm_chunk=AsyncMock(side_effect=RuntimeError("boom")),
191 )
192 ok = _make_aa_provider("ok", available=True)
193 provider_map = {"raises": raises, "ok": ok}
194 controller.mass.get_provider = MagicMock(side_effect=provider_map.get) # type: ignore[method-assign]
195
196 await controller._distribute_chunk(session_key, b"\x00" * 1024)
197
198 assert "raises" not in controller._active_sessions[session_key]
199 assert "ok" in controller._active_sessions[session_key]
200
201
202def test_get_scan_concurrency_returns_default_on_unset() -> None:
203 """When the config value is unset/None, fall back to DEFAULT_BACKGROUND_SCAN_CONCURRENCY."""
204 controller = _make_controller()
205 controller.mass.config.get_raw_core_config_value = MagicMock(return_value=None) # type: ignore[method-assign]
206 assert controller._get_scan_concurrency() == DEFAULT_BACKGROUND_SCAN_CONCURRENCY
207
208
209def test_get_scan_concurrency_clamps_to_max() -> None:
210 """Values above 16 are clamped to 16."""
211 controller = _make_controller()
212 controller.mass.config.get_raw_core_config_value = MagicMock(return_value=99) # type: ignore[method-assign]
213 assert controller._get_scan_concurrency() == 16
214
215
216def test_get_scan_concurrency_clamps_to_min() -> None:
217 """Values below 1 are clamped to 1."""
218 controller = _make_controller()
219 # Use a truthy negative value so the controller's `value or DEFAULT` fallback
220 # doesn't swap us out for the default before the min-clamp runs.
221 controller.mass.config.get_raw_core_config_value = MagicMock(return_value=-1) # type: ignore[method-assign]
222 assert controller._get_scan_concurrency() == 1
223
224
225@pytest.mark.parametrize(
226 ("cpu_count", "expected"),
227 [(1, 1), (2, 1), (3, 1), (4, 2), (8, 2), (16, 2)],
228)
229def test_default_background_scan_concurrency(cpu_count: int, expected: int) -> None:
230 """Background scan defaults to 1 below 4 cores, 2 at/above (never more than 2)."""
231 with patch("music_assistant.constants.os.process_cpu_count", return_value=cpu_count):
232 assert _default_background_scan_concurrency() == expected
233
234
235def _make_stream_mock(chunks: list[bytes]) -> object:
236 """Return a get_media_stream mock that yields the given chunks."""
237
238 async def _stream(
239 _streamdetails: object, _pcm_format: object, **_kwargs: object
240 ) -> AsyncGenerator[bytes]:
241 for chunk in chunks:
242 yield chunk
243
244 return _stream
245
246
247@pytest.mark.asyncio
248async def test_background_streaming_happy_path(monkeypatch: pytest.MonkeyPatch) -> None:
249 """PCM chunks reach providers; session is cleaned up on clean EOF."""
250 controller = _make_controller()
251 streamdetails = _make_streamdetails(path="/music/test.flac")
252 p = _make_aa_provider("p1", available=True)
253 p.start_analysis = AsyncMock(return_value=True)
254 p.finalize = AsyncMock(return_value=None)
255 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
256
257 fake_chunks = [b"\x00\x01" * 512 for _ in range(5)]
258 controller.mass.streams.audio.get_media_stream = _make_stream_mock(fake_chunks) # type: ignore[method-assign,assignment]
259 monkeypatch.setattr(audio_analysis_mod, "BACKGROUND_PACE_INTERVAL_SECONDS_FLOOR", 0.0)
260
261 await controller._run_background_streaming_for_track(streamdetails, [p])
262
263 assert p.start_analysis.await_count == 1
264 assert p.process_pcm_chunk.await_count == len(fake_chunks)
265 # _finalize_providers pops the session key before dispatching — key must be gone
266 assert streamdetails.uri not in controller._active_sessions
267
268
269@pytest.mark.asyncio
270async def test_background_streaming_paces_chunk_dispatch(
271 monkeypatch: pytest.MonkeyPatch,
272) -> None:
273 """Background dispatch enforces the pacing floor between consecutive chunks."""
274 controller = _make_controller()
275 streamdetails = _make_streamdetails(path="/music/test.flac")
276 p = _make_aa_provider("p1", available=True)
277 p.start_analysis = AsyncMock(return_value=True)
278 p.finalize = AsyncMock(return_value=None)
279 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
280
281 fake_chunks = [b"\x00\x01" * 512 for _ in range(4)]
282 controller.mass.streams.audio.get_media_stream = _make_stream_mock(fake_chunks) # type: ignore[method-assign,assignment]
283 monkeypatch.setattr(audio_analysis_mod, "BACKGROUND_PACE_INTERVAL_SECONDS_FLOOR", 0.05)
284
285 started = asyncio.get_running_loop().time()
286 await controller._run_background_streaming_for_track(streamdetails, [p])
287 elapsed = asyncio.get_running_loop().time() - started
288
289 assert p.process_pcm_chunk.await_count == len(fake_chunks)
290 # First chunk dispatches immediately; each of the remaining 3 waits out the floor.
291 assert elapsed >= 3 * 0.05
292
293
294@pytest.mark.asyncio
295async def test_background_streaming_per_track_timeout(monkeypatch: pytest.MonkeyPatch) -> None:
296 """Per-track timeout cancels providers and cleans up the session."""
297 controller = _make_controller()
298 streamdetails = _make_streamdetails(path="/music/test.flac")
299 p = _make_aa_provider("p1", available=True)
300 p.start_analysis = AsyncMock(return_value=True)
301
302 async def _hang_chunk(*_args: object, **_kwargs: object) -> None:
303 await asyncio.sleep(10)
304
305 p.process_pcm_chunk = AsyncMock(side_effect=_hang_chunk)
306 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
307
308 controller.mass.streams.audio.get_media_stream = _make_stream_mock([b"\x00" * 1024] * 50) # type: ignore[method-assign,assignment]
309 monkeypatch.setattr(audio_analysis_mod, "BACKGROUND_PER_TRACK_TIMEOUT_SECONDS", 0.2)
310
311 await controller._run_background_streaming_for_track(streamdetails, [p])
312
313 assert streamdetails.uri not in controller._active_sessions
314 # Per-track timeout must be surfaced to the TasksController so the run ends
315 # as PARTIAL_SUCCESS with a retryable status.
316 controller.mass.tasks.add_task_failure.assert_called_once() # type: ignore[attr-defined]
317 failure_args = controller.mass.tasks.add_task_failure.call_args.args # type: ignore[attr-defined]
318 assert failure_args[0] == audio_analysis_mod.BACKGROUND_SCAN_TASK_ID
319 assert "Timed out" in failure_args[1]
320 assert streamdetails.uri in failure_args[1]
321
322
323@pytest.mark.asyncio
324async def test_background_streaming_timeout_scales_with_track_duration(
325 monkeypatch: pytest.MonkeyPatch,
326) -> None:
327 """Per-track timeout is derived from track duration when duration is known."""
328 controller = _make_controller()
329 streamdetails = _make_streamdetails(path="/music/long_mix.flac", duration=3600)
330 p = _make_aa_provider("p1", available=True)
331 p.start_analysis = AsyncMock(return_value=True)
332 p.finalize = AsyncMock(return_value=None)
333 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
334 controller.mass.streams.audio.get_media_stream = _make_stream_mock([]) # type: ignore[method-assign,assignment]
335
336 captured_timeouts: list[float | None] = []
337 real_wait_for = asyncio.wait_for
338
339 async def _spy_wait_for(coro: Any, timeout: float | None) -> Any:
340 captured_timeouts.append(timeout)
341 return await real_wait_for(coro, timeout)
342
343 monkeypatch.setattr("asyncio.wait_for", _spy_wait_for)
344
345 await controller._run_background_streaming_for_track(streamdetails, [p])
346
347 expected = int(3600 * audio_analysis_mod.BACKGROUND_PER_TRACK_TIMEOUT_DURATION_MULTIPLIER)
348 assert captured_timeouts[0] == expected
349
350
351@pytest.mark.asyncio
352async def test_background_streaming_ffmpeg_startup_failure() -> None:
353 """get_media_stream failure cancels providers cleanly without raising."""
354 controller = _make_controller()
355 streamdetails = _make_streamdetails(path="/nonexistent.flac")
356 p = _make_aa_provider("p1", available=True)
357 p.start_analysis = AsyncMock(return_value=True)
358 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
359
360 def _failing_stream(*_args: object, **_kwargs: object) -> AsyncGenerator[bytes]:
361 raise RuntimeError("ffmpeg startup failed")
362
363 controller.mass.streams.audio.get_media_stream = _failing_stream # type: ignore[method-assign]
364
365 # Should not raise
366 await controller._run_background_streaming_for_track(streamdetails, [p])
367 assert streamdetails.uri not in controller._active_sessions
368 # Per-track exception must be surfaced to the TasksController.
369 controller.mass.tasks.add_task_failure.assert_called_once() # type: ignore[attr-defined]
370 failure_args = controller.mass.tasks.add_task_failure.call_args.args # type: ignore[attr-defined]
371 assert failure_args[0] == audio_analysis_mod.BACKGROUND_SCAN_TASK_ID
372 assert "Failed" in failure_args[1]
373 assert "ffmpeg startup failed" in failure_args[1]
374
375
376def _make_streamdetails(
377 *, path: str, item_id: str = "test-item", duration: int | None = None
378) -> MagicMock:
379 sd = MagicMock()
380 sd.path = path
381 sd.uri = f"track://test/{path}"
382 sd.audio_format = AudioFormat(
383 content_type=ContentType.FLAC,
384 sample_rate=44100,
385 bit_depth=16,
386 channels=2,
387 )
388 sd.item_id = item_id
389 sd.provider = "test-provider"
390 sd.media_type = MagicMock()
391 sd.duration = duration
392 return sd
393
394
395def _make_controller() -> AudioAnalysisController:
396 streams = MagicMock()
397 streams.mass = MagicMock()
398 streams.mass.logger.getChild.return_value = MagicMock()
399 return AudioAnalysisController(streams)
400
401
402def _make_aa_provider(
403 instance_id: str,
404 *,
405 available: bool = True,
406 process_pcm_chunk: AsyncMock | None = None,
407) -> MagicMock:
408 provider = MagicMock(spec=AudioAnalysisProvider)
409 provider.instance_id = instance_id
410 provider.available = available
411 provider.process_pcm_chunk = process_pcm_chunk or AsyncMock(return_value=None)
412 provider.cancel = AsyncMock(return_value=None)
413 return provider
414
415
416@pytest.mark.asyncio
417async def test_run_background_scan_uses_union_candidate_query(
418 monkeypatch: pytest.MonkeyPatch,
419) -> None:
420 """The new scan loop drives _run_background_streaming_for_track per candidate."""
421 controller = _make_controller()
422 p1 = _make_aa_provider("prov-1", available=True)
423 p1.domain = "p1"
424 p1.start_analysis = AsyncMock(return_value=True)
425 monkeypatch.setattr(
426 controller.__class__,
427 "providers",
428 property(lambda _self: [p1]),
429 )
430
431 candidates = [
432 {"item_id": "track-1", "provider_instance": "filesystem_local", "missing_domains": ["p1"]},
433 {"item_id": "track-2", "provider_instance": "filesystem_local", "missing_domains": ["p1"]},
434 ]
435 monkeypatch.setattr(
436 controller, "_find_candidates_missing_analysis", AsyncMock(return_value=candidates)
437 )
438
439 streamdetails_list = [
440 _make_streamdetails(path=f"/music/{c['item_id']}.flac", item_id=str(c["item_id"]))
441 for c in candidates
442 ]
443 for sd in streamdetails_list:
444 sd.stream_type = StreamType.LOCAL_FILE
445
446 music_prov = MagicMock()
447 music_prov.available = True
448 music_prov.get_stream_details = AsyncMock(side_effect=streamdetails_list)
449 music_prov.instance_id = "filesystem_local"
450 controller.mass.get_provider = MagicMock(return_value=music_prov) # type: ignore[method-assign]
451
452 streaming_calls: list[str] = []
453
454 async def _track_streaming(
455 streamdetails: MagicMock, _providers: object, **_kwargs: object
456 ) -> None:
457 streaming_calls.append(streamdetails.item_id)
458
459 monkeypatch.setattr(controller, "_run_background_streaming_for_track", _track_streaming)
460
461 await controller._run_background_scan()
462
463 assert sorted(streaming_calls) == ["track-1", "track-2"]
464
465
466@pytest.mark.asyncio
467async def test_find_candidates_handles_sqlite_row_without_get(
468 monkeypatch: pytest.MonkeyPatch,
469) -> None:
470 """
471 _find_candidates_missing_analysis must use __getitem__ not .get() on rows.
472
473 sqlite3.Row supports only __getitem__, not .get(). This regression test
474 uses a row class that lacks .get() to ensure we never reintroduce the bug.
475 """
476 controller = _make_controller()
477 p1 = _make_aa_provider("prov-1", available=True)
478 p1.domain = "loudness_analysis"
479 p1.analysis_version = 1
480 p1.available = True
481 monkeypatch.setattr(
482 controller.__class__,
483 "providers",
484 property(lambda _self: [p1]),
485 )
486
487 # Make the filesystem-providers gate succeed
488 fs_prov = MagicMock()
489 fs_prov.domain = "filesystem_local"
490 fs_prov.available = True
491 controller.mass.get_providers = MagicMock(return_value=[fs_prov]) # type: ignore[method-assign]
492
493 class _RowNoGet:
494 """Mimics sqlite3.Row: __getitem__ only, no .get()."""
495
496 def __init__(self, data: dict[str, object]) -> None:
497 self._d = data
498
499 def __getitem__(self, key: str) -> object:
500 return self._d[key]
501
502 # SQL filters out fully-covered tracks via NOT EXISTS + GROUP BY, so the
503 # rows we receive from the database only contain missing-domain pairs.
504 rows = [
505 _RowNoGet(
506 {
507 "item_id": "track-1",
508 "provider_instance": "filesystem_local",
509 "missing_domains": "loudness_analysis",
510 }
511 ),
512 ]
513 controller.mass.music.database.get_rows_from_query = AsyncMock(return_value=rows) # type: ignore[method-assign]
514
515 result = await controller._find_candidates_missing_analysis({"loudness_analysis": 1}, 100)
516
517 assert len(result) == 1
518 assert result[0]["item_id"] == "track-1"
519 assert result[0]["missing_domains"] == ["loudness_analysis"]
520
521
522@pytest.mark.asyncio
523async def test_find_candidates_query_gates_on_current_version(
524 monkeypatch: pytest.MonkeyPatch,
525) -> None:
526 """
527 The candidate query must treat stale-version rows as needing re-analysis.
528
529 The NOT EXISTS gate may only count a stored analysis row as up-to-date when
530 its analysis_version is non-NULL and >= the provider's current version, so a
531 provider bumping analysis_version re-surfaces previously analyzed tracks.
532 """
533 controller = _make_controller()
534 p1 = _make_aa_provider("prov-1", available=True)
535 p1.domain = "sonic_analysis"
536 monkeypatch.setattr(
537 controller.__class__,
538 "providers",
539 property(lambda _self: [p1]),
540 )
541
542 fs_prov = MagicMock()
543 fs_prov.domain = "filesystem_local"
544 fs_prov.available = True
545 controller.mass.get_providers = MagicMock(return_value=[fs_prov]) # type: ignore[method-assign]
546
547 captured: dict[str, Any] = {}
548
549 async def _capture(query: str, params: dict[str, Any], limit: int) -> list[Any]: # noqa: ARG001
550 captured["query"] = query
551 captured["params"] = params
552 return []
553
554 controller.mass.music.database.get_rows_from_query = AsyncMock(side_effect=_capture) # type: ignore[method-assign]
555
556 await controller._find_candidates_missing_analysis({"sonic_analysis": 3}, 0)
557
558 sql = captured["query"]
559 assert "aa.analysis_version IS NOT NULL" in sql
560 assert "aa.analysis_version >= possible.current_version" in sql
561 assert captured["params"]["ver_0"] == 3
562 assert captured["params"]["aa_0"] == "sonic_analysis"
563
564
565@pytest.mark.asyncio
566async def test_run_background_scan_concurrency_semaphore(
567 monkeypatch: pytest.MonkeyPatch,
568) -> None:
569 """At most CONF_BACKGROUND_SCAN_CONCURRENCY tracks run concurrently."""
570 controller = _make_controller()
571 monkeypatch.setattr(controller, "_get_scan_concurrency", lambda: 2)
572
573 p1 = _make_aa_provider("prov-1", available=True)
574 p1.domain = "p1"
575 p1.start_analysis = AsyncMock(return_value=True)
576 monkeypatch.setattr(
577 controller.__class__,
578 "providers",
579 property(lambda _self: [p1]),
580 )
581
582 candidates = [
583 {
584 "item_id": f"track-{i}",
585 "provider_instance": "filesystem_local",
586 "missing_domains": ["p1"],
587 }
588 for i in range(4)
589 ]
590 monkeypatch.setattr(
591 controller, "_find_candidates_missing_analysis", AsyncMock(return_value=candidates)
592 )
593
594 streamdetails_list = [
595 _make_streamdetails(path=f"/music/{c['item_id']}.flac") for c in candidates
596 ]
597 for sd in streamdetails_list:
598 sd.stream_type = StreamType.LOCAL_FILE
599 music_prov = MagicMock()
600 music_prov.available = True
601 music_prov.get_stream_details = AsyncMock(side_effect=streamdetails_list)
602 music_prov.instance_id = "filesystem_local"
603 controller.mass.get_provider = MagicMock(return_value=music_prov) # type: ignore[method-assign]
604
605 in_flight = 0
606 max_in_flight = 0
607 barrier = asyncio.Barrier(2)
608
609 async def _track_streaming(
610 _streamdetails: MagicMock, _providers: object, **_kwargs: object
611 ) -> None:
612 nonlocal in_flight, max_in_flight
613 in_flight += 1
614 max_in_flight = max(max_in_flight, in_flight)
615 await barrier.wait()
616 in_flight -= 1
617
618 monkeypatch.setattr(controller, "_run_background_streaming_for_track", _track_streaming)
619
620 await controller._run_background_scan()
621
622 assert max_in_flight == 2
623
624
625@pytest.mark.asyncio
626async def test_background_streaming_cancellation_cleans_up(
627 monkeypatch: pytest.MonkeyPatch,
628) -> None:
629 """CancelledError mid-track must trigger _cancel_providers and re-raise."""
630 controller = _make_controller()
631 streamdetails = _make_streamdetails(path="/music/test.flac")
632 p = _make_aa_provider("p1", available=True)
633 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
634
635 session_key = streamdetails.uri
636
637 async def _inner_cancelled(
638 _session_key: str, _sd: object, _providers: object, **_kwargs: object
639 ) -> None:
640 # Simulate the inner having registered the session before being cancelled.
641 controller._active_sessions[session_key] = {"p1"}
642 raise asyncio.CancelledError
643
644 monkeypatch.setattr(controller, "_run_background_streaming_inner", _inner_cancelled)
645
646 with pytest.raises(asyncio.CancelledError):
647 await controller._run_background_streaming_for_track(streamdetails, [p])
648
649 # Session must be popped and provider.cancel scheduled.
650 assert session_key not in controller._active_sessions
651 p.cancel.assert_called_once_with(session_key)
652
653
654@pytest.mark.asyncio
655async def test_run_background_scan_defers_past_run_budget(
656 monkeypatch: pytest.MonkeyPatch,
657) -> None:
658 """Tracks past the run-budget deadline are deferred to the next run."""
659 controller = _make_controller()
660
661 p1 = _make_aa_provider("prov-1", available=True)
662 p1.domain = "p1"
663 monkeypatch.setattr(
664 controller.__class__,
665 "providers",
666 property(lambda _self: [p1]),
667 )
668
669 candidates = [
670 {
671 "item_id": f"track-{i}",
672 "provider_instance": "filesystem_local",
673 "missing_domains": ["p1"],
674 }
675 for i in range(3)
676 ]
677 monkeypatch.setattr(
678 controller, "_find_candidates_missing_analysis", AsyncMock(return_value=candidates)
679 )
680
681 # Force budget to negative so every candidate is past deadline.
682 monkeypatch.setattr(audio_analysis_mod, "BACKGROUND_SCAN_RUN_BUDGET_SECONDS", -1)
683
684 streaming_called = False
685
686 async def _track_streaming(_sd: object, _providers: object, **_kwargs: object) -> None:
687 nonlocal streaming_called
688 streaming_called = True
689
690 monkeypatch.setattr(controller, "_run_background_streaming_for_track", _track_streaming)
691
692 await controller._run_background_scan()
693
694 assert not streaming_called
695
696
697@pytest.mark.asyncio
698async def test_close_drains_sessions_and_workers() -> None:
699 """close() cancels in-flight chunk workers and dispatches provider cancels."""
700 controller = _make_controller()
701
702 # Real asyncio task that swallows cancellation cleanly.
703 async def _busy_worker() -> None:
704 try:
705 await asyncio.sleep(60)
706 except asyncio.CancelledError:
707 return
708
709 worker_task = asyncio.create_task(_busy_worker())
710 controller._workers["track://test/a"] = worker_task
711
712 p = _make_aa_provider("p1", available=True)
713 controller.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
714 controller._active_sessions["track://test/a"] = {"p1"}
715
716 await controller.close()
717
718 # Worker awaited to completion; both dicts drained.
719 assert worker_task.done()
720 assert controller._workers == {}
721 assert controller._active_sessions == {}
722 # Provider cancel scheduled with the session key.
723 p.cancel.assert_called_once_with("track://test/a")
724
725
726async def _run_buffer_reader_worker(
727 chunk_count: int, expected_duration: float | None
728) -> tuple[MagicMock, MagicMock]:
729 """Run the reader worker against a buffer yielding chunk_count 1s chunks, then EOF."""
730 controller = _make_controller()
731 session_key = "track://test/worker"
732 controller._active_sessions[session_key] = {"prov-1"}
733 controller._distribute_chunk = AsyncMock() # type: ignore[method-assign]
734 controller._finalize_providers = MagicMock() # type: ignore[method-assign]
735 controller._cancel_providers = MagicMock() # type: ignore[method-assign]
736 audio_buffer = MagicMock()
737 audio_buffer.first_buffered_chunk = 0
738 audio_buffer.read_chunk_for_analysis = AsyncMock(
739 side_effect=[b"pcm"] * chunk_count + [AudioBufferEOF()]
740 )
741 await controller._buffer_reader_worker(session_key, audio_buffer, expected_duration)
742 return controller._finalize_providers, controller._cancel_providers
743
744
745@pytest.mark.asyncio
746async def test_buffer_reader_worker_finalizes_when_near_expected_duration() -> None:
747 """An EOF within the completeness tolerance finalizes the session."""
748 finalize, cancel = await _run_buffer_reader_worker(9, expected_duration=10)
749 finalize.assert_called_once_with("track://test/worker")
750 cancel.assert_not_called()
751
752
753@pytest.mark.asyncio
754async def test_buffer_reader_worker_discards_incomplete_stream() -> None:
755 """A source that ends far short of the expected duration is cancelled, not finalized."""
756 finalize, cancel = await _run_buffer_reader_worker(8, expected_duration=10)
757 finalize.assert_not_called()
758 cancel.assert_called_once_with("track://test/worker")
759
760
761@pytest.mark.asyncio
762async def test_buffer_reader_worker_finalizes_when_duration_unknown() -> None:
763 """Without an expected duration (e.g. radio), any clean EOF finalizes the session."""
764 finalize, cancel = await _run_buffer_reader_worker(1, expected_duration=None)
765 finalize.assert_called_once_with("track://test/worker")
766 cancel.assert_not_called()
767
768
769def _stub_controller(
770 count_result: int = 0,
771 iter_rows: list[dict[str, Any]] | None = None,
772) -> tuple[AudioAnalysisController, MagicMock]:
773 """Build a bare AudioAnalysisController whose database is mocked."""
774 c = AudioAnalysisController.__new__(AudioAnalysisController)
775 c.logger = MagicMock()
776 db = MagicMock()
777 db.get_count_from_query = AsyncMock(return_value=count_result)
778 db.delete = AsyncMock()
779 rows_to_yield = list(iter_rows or [])
780
781 async def _iter_stub(*_args: Any, **_kwargs: Any) -> AsyncGenerator[Mapping[str, Any]]:
782 for row in rows_to_yield:
783 yield row
784
785 db.iter_rows_from_query = MagicMock(side_effect=_iter_stub)
786 c.mass = MagicMock()
787 c.mass.music = MagicMock()
788 c.mass.music.database = db
789 c.mass.get_providers = MagicMock(return_value=[])
790 return c, db
791
792
793# --- provider-scoped merge (loudness regression fix) ---
794
795_ALL_AA_DOMAINS = {LOUDNESS_ANALYSIS_DOMAIN, SMART_FADES_ANALYSIS_DOMAIN, SONIC_ANALYSIS_DOMAIN}
796
797
798def _aa_row(domain: str, row_id: int, **fields: Any) -> dict[str, Any]:
799 """Build one audio_analysis db row (rows are passed oldest-first / ascending row_id)."""
800 return {
801 "id": row_id,
802 "item_id": "track-1",
803 "provider": "test-provider",
804 "media_type": MediaType.TRACK.value,
805 "aa_provider_domain": domain,
806 "analysis_data": json_dumps(AudioAnalysisData(**fields).to_dict()),
807 }
808
809
810def test_merged_from_rows_priority_none_is_last_write_wins() -> None:
811 """Without priority, the newest (last) row wins each non-None field (legacy behaviour)."""
812 rows = [
813 _aa_row(LOUDNESS_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5),
814 _aa_row(SONIC_ANALYSIS_DOMAIN, 2, loudness_integrated=-12.0),
815 ]
816 merged = _merged_from_rows(rows, _ALL_AA_DOMAINS)
817 assert merged is not None
818 assert merged.loudness_integrated == -12.0
819
820
821def test_merged_from_rows_single_priority_uses_only_that_provider() -> None:
822 """A single-domain priority returns only that provider's values; others are ignored."""
823 rows = [
824 _aa_row(LOUDNESS_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5),
825 _aa_row(SONIC_ANALYSIS_DOMAIN, 2, loudness_integrated=-12.0, bpm=120),
826 ]
827 merged = _merged_from_rows(rows, _ALL_AA_DOMAINS, priority=(LOUDNESS_ANALYSIS_DOMAIN,))
828 assert merged is not None
829 assert merged.loudness_integrated == -7.5
830 assert merged.bpm is None # sonic_analysis excluded entirely
831
832
833def test_merged_from_rows_multi_priority_first_listed_wins_and_merges() -> None:
834 """Multi-domain priority merges all listed domains; the first-listed wins conflicts."""
835 rows = [
836 # sonic newer than loudness, but loudness is listed first -> wins loudness_integrated
837 _aa_row(SONIC_ANALYSIS_DOMAIN, 1, loudness_integrated=-12.0, energy=0.5),
838 _aa_row(LOUDNESS_ANALYSIS_DOMAIN, 2, loudness_integrated=-7.5),
839 _aa_row(SMART_FADES_ANALYSIS_DOMAIN, 3, bpm=120),
840 ]
841 merged = _merged_from_rows(
842 rows,
843 _ALL_AA_DOMAINS,
844 priority=(LOUDNESS_ANALYSIS_DOMAIN, SONIC_ANALYSIS_DOMAIN, SMART_FADES_ANALYSIS_DOMAIN),
845 )
846 assert merged is not None
847 assert merged.loudness_integrated == -7.5 # first-listed wins the conflict
848 assert merged.energy == 0.5 # non-conflicting field from sonic still merged in
849 assert merged.bpm == 120 # and from smart_fades
850
851
852def test_merged_from_rows_priority_domain_not_available_is_excluded() -> None:
853 """A priority domain that is not currently available is dropped (can yield None)."""
854 rows = [_aa_row(LOUDNESS_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5)]
855 merged = _merged_from_rows(rows, {SONIC_ANALYSIS_DOMAIN}, priority=(LOUDNESS_ANALYSIS_DOMAIN,))
856 assert merged is None
857
858
859def test_merged_from_rows_regression_sonic_does_not_clobber_loudness() -> None:
860 """
861 Regression: sonic_analysis' RMS loudness must not overwrite the EBU R128 value.
862
863 Reproduces the volume-jump bug: a newer sonic_analysis row carries an RMS-proxy
864 loudness_integrated that wins under last-write-wins, but scoping to loudness_analysis
865 returns the authoritative value.
866 """
867 rows = [
868 _aa_row(LOUDNESS_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5),
869 _aa_row(SONIC_ANALYSIS_DOMAIN, 2, loudness_integrated=-12.0),
870 ]
871 legacy = _merged_from_rows(rows, _ALL_AA_DOMAINS)
872 assert legacy is not None
873 assert legacy.loudness_integrated == -12.0 # old/buggy: sonic clobbers
874 scoped = _merged_from_rows(rows, _ALL_AA_DOMAINS, priority=(LOUDNESS_ANALYSIS_DOMAIN,))
875 assert scoped is not None
876 assert scoped.loudness_integrated == -7.5 # fixed
877
878
879@pytest.mark.asyncio
880async def test_get_audio_analysis_priority_threads_through_to_merge() -> None:
881 """get_audio_analysis forwards priority so the loudness call gets the EBU R128 value."""
882 c, db = _stub_controller()
883 db.get_rows = AsyncMock(
884 return_value=[
885 _aa_row(LOUDNESS_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5),
886 _aa_row(SONIC_ANALYSIS_DOMAIN, 2, loudness_integrated=-12.0),
887 ]
888 )
889 music_prov = MagicMock(spec=MusicProvider)
890 music_prov.is_streaming_provider = True
891 music_prov.domain = "test-provider"
892 c.mass.get_provider = MagicMock(return_value=music_prov) # type: ignore[method-assign]
893 aa_loud = MagicMock()
894 aa_loud.domain = LOUDNESS_ANALYSIS_DOMAIN
895 aa_loud.available = True
896 aa_sonic = MagicMock()
897 aa_sonic.domain = SONIC_ANALYSIS_DOMAIN
898 aa_sonic.available = True
899 c.mass.get_providers = MagicMock(return_value=[aa_loud, aa_sonic]) # type: ignore[method-assign]
900
901 result = await c.get_audio_analysis(
902 "track-1", "test-provider", priority=(LOUDNESS_ANALYSIS_DOMAIN,)
903 )
904 assert result is not None
905 assert result.loudness_integrated == -7.5
906
907
908@pytest.mark.asyncio
909async def test_get_audio_analysis_count_returns_helper_result() -> None:
910 """The controller forwards whatever get_count_from_query returns."""
911 c, _ = _stub_controller(count_result=42)
912 assert await c.get_audio_analysis_count("sonic_analysis") == 42
913
914
915@pytest.mark.asyncio
916async def test_get_audio_analysis_count_filters_by_domain_and_track_media_type() -> None:
917 """Default count filters on aa_provider_domain AND media_type=track."""
918 c, db = _stub_controller(count_result=0)
919 await c.get_audio_analysis_count("sonic_analysis")
920 sql, params = db.get_count_from_query.await_args.args
921 assert "aa_provider_domain = :aa_provider_domain" in sql
922 assert "media_type = :media_type" in sql
923 assert params == {"aa_provider_domain": "sonic_analysis", "media_type": MediaType.TRACK.value}
924
925
926@pytest.mark.asyncio
927async def test_get_audio_analysis_count_respects_media_type_override() -> None:
928 """Caller can count rows for a non-track media type."""
929 c, db = _stub_controller(count_result=7)
930 result = await c.get_audio_analysis_count(
931 "sonic_analysis", media_type=MediaType.PODCAST_EPISODE
932 )
933 assert result == 7
934 params = db.get_count_from_query.await_args.args[1]
935 assert params["media_type"] == MediaType.PODCAST_EPISODE.value
936
937
938@pytest.mark.asyncio
939async def test_iter_audio_analysis_rows_yields_all_rows() -> None:
940 """iter_audio_analysis_rows yields each DB row in order; no filtering or parsing."""
941 rows: list[dict[str, Any]] = [
942 {"item_id": "a", "provider": "filesystem_local", "analysis_data": "{}"},
943 {"item_id": "b", "provider": "filesystem_local", "analysis_data": "{}"},
944 ]
945 c, _ = _stub_controller(iter_rows=rows)
946 result = [r async for r in c.iter_audio_analysis_rows("sonic_analysis")]
947 assert result == rows
948
949
950@pytest.mark.asyncio
951async def test_iter_audio_analysis_rows_filters_by_domain_and_track_media_type() -> None:
952 """Default query filters on aa_provider_domain + media_type=track."""
953 c, db = _stub_controller(iter_rows=[])
954 [r async for r in c.iter_audio_analysis_rows("sonic_analysis")]
955 sql, params = db.iter_rows_from_query.call_args.args
956 assert "aa_provider_domain = :aa_provider_domain" in sql
957 assert "media_type = :media_type" in sql
958 assert params == {
959 "aa_provider_domain": "sonic_analysis",
960 "media_type": MediaType.TRACK.value,
961 }
962
963
964@pytest.mark.asyncio
965async def test_iter_audio_analysis_rows_respects_media_type_override() -> None:
966 """Caller can stream rows for a non-track media type."""
967 c, db = _stub_controller(iter_rows=[])
968 [
969 r
970 async for r in c.iter_audio_analysis_rows(
971 "sonic_analysis", media_type=MediaType.PODCAST_EPISODE
972 )
973 ]
974 params = db.iter_rows_from_query.call_args.args[1]
975 assert params["media_type"] == MediaType.PODCAST_EPISODE.value
976
977
978def _aa_provider_stub(domain: str, available: bool = True) -> MagicMock:
979 """Build a provider stub that satisfies the get_providers().available filter."""
980 p = MagicMock()
981 p.domain = domain
982 p.available = available
983 return p
984
985
986@pytest.mark.asyncio
987async def test_iter_merged_audio_analysis_rows_merges_within_group() -> None:
988 """Two rows for the same (item_id, provider) merge in timestamp order."""
989 rows: list[dict[str, Any]] = [
990 {
991 "item_id": "t1",
992 "provider": "filesystem_local",
993 "aa_provider_domain": "sonic_analysis",
994 "analysis_data": '{"bpm": 100.0, "energy": 0.5}',
995 },
996 {
997 "item_id": "t1",
998 "provider": "filesystem_local",
999 "aa_provider_domain": "smart_fades",
1000 "analysis_data": '{"bpm": 120.0, "key": "C"}',
1001 },
1002 ]
1003 c, _ = _stub_controller(iter_rows=rows)
1004 c.mass.get_providers = MagicMock( # type: ignore[method-assign]
1005 return_value=[
1006 _aa_provider_stub("sonic_analysis"),
1007 _aa_provider_stub("smart_fades"),
1008 ]
1009 )
1010
1011 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1012 assert len(result) == 1
1013 item_id, provider, merged = result[0]
1014 assert (item_id, provider) == ("t1", "filesystem_local")
1015 assert merged.bpm == 120.0 # smart_fades wins on bpm (later row)
1016 assert merged.energy == 0.5 # sonic_analysis still wins where smart_fades is None
1017 assert merged.key == "C"
1018
1019
1020@pytest.mark.asyncio
1021async def test_iter_merged_audio_analysis_rows_skips_unavailable_providers() -> None:
1022 """Rows from unavailable AA providers are skipped during merge."""
1023 rows: list[dict[str, Any]] = [
1024 {
1025 "item_id": "t1",
1026 "provider": "filesystem_local",
1027 "aa_provider_domain": "sonic_analysis",
1028 "analysis_data": '{"bpm": 100.0}',
1029 },
1030 {
1031 "item_id": "t1",
1032 "provider": "filesystem_local",
1033 "aa_provider_domain": "disabled_provider",
1034 "analysis_data": '{"bpm": 999.0}',
1035 },
1036 ]
1037 c, _ = _stub_controller(iter_rows=rows)
1038 c.mass.get_providers = MagicMock(return_value=[_aa_provider_stub("sonic_analysis")]) # type: ignore[method-assign]
1039
1040 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1041 assert len(result) == 1
1042 assert result[0][2].bpm == 100.0 # disabled_provider's row ignored
1043
1044
1045@pytest.mark.asyncio
1046async def test_iter_merged_audio_analysis_rows_groups_by_item_provider() -> None:
1047 """Rows from different (item_id, provider) pairs are emitted as separate entries."""
1048 rows: list[dict[str, Any]] = [
1049 {
1050 "item_id": "t1",
1051 "provider": "filesystem_local",
1052 "aa_provider_domain": "sonic_analysis",
1053 "analysis_data": '{"bpm": 100.0}',
1054 },
1055 {
1056 "item_id": "t2",
1057 "provider": "filesystem_local",
1058 "aa_provider_domain": "sonic_analysis",
1059 "analysis_data": '{"bpm": 200.0}',
1060 },
1061 ]
1062 c, _ = _stub_controller(iter_rows=rows)
1063 c.mass.get_providers = MagicMock(return_value=[_aa_provider_stub("sonic_analysis")]) # type: ignore[method-assign]
1064
1065 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1066 assert len(result) == 2
1067 assert {r[0] for r in result} == {"t1", "t2"}
1068
1069
1070@pytest.mark.asyncio
1071async def test_iter_merged_audio_analysis_rows_skips_unparsable_rows() -> None:
1072 """A row with corrupt JSON is silently skipped without aborting the merge."""
1073 rows: list[dict[str, Any]] = [
1074 {
1075 "item_id": "t1",
1076 "provider": "filesystem_local",
1077 "aa_provider_domain": "sonic_analysis",
1078 "analysis_data": "not-json",
1079 },
1080 {
1081 "item_id": "t1",
1082 "provider": "filesystem_local",
1083 "aa_provider_domain": "smart_fades",
1084 "analysis_data": '{"bpm": 120.0}',
1085 },
1086 ]
1087 c, _ = _stub_controller(iter_rows=rows)
1088 c.mass.get_providers = MagicMock( # type: ignore[method-assign]
1089 return_value=[
1090 _aa_provider_stub("sonic_analysis"),
1091 _aa_provider_stub("smart_fades"),
1092 ]
1093 )
1094
1095 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1096 assert len(result) == 1
1097 assert result[0][2].bpm == 120.0
1098
1099
1100@pytest.mark.asyncio
1101async def test_iter_merged_audio_analysis_rows_empty_db_yields_nothing() -> None:
1102 """An empty DB result yields no entries without flushing a sentinel group."""
1103 c, _ = _stub_controller(iter_rows=[])
1104 c.mass.get_providers = MagicMock(return_value=[_aa_provider_stub("sonic_analysis")]) # type: ignore[method-assign]
1105
1106 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1107 assert result == []
1108
1109
1110@pytest.mark.asyncio
1111async def test_iter_merged_audio_analysis_rows_logs_warning_for_unparsable_rows(
1112 caplog: pytest.LogCaptureFixture,
1113) -> None:
1114 """Unparsable rows must surface a WARNING so storage corruption is observable."""
1115 rows: list[dict[str, Any]] = [
1116 {
1117 "id": 42,
1118 "item_id": "t1",
1119 "provider": "filesystem_local",
1120 "aa_provider_domain": "sonic_analysis",
1121 "analysis_data": "not-json",
1122 },
1123 {
1124 "id": 43,
1125 "item_id": "t1",
1126 "provider": "filesystem_local",
1127 "aa_provider_domain": "smart_fades",
1128 "analysis_data": '{"bpm": 120.0}',
1129 },
1130 ]
1131 c, _ = _stub_controller(iter_rows=rows)
1132 c.mass.get_providers = MagicMock( # type: ignore[method-assign]
1133 return_value=[
1134 _aa_provider_stub("sonic_analysis"),
1135 _aa_provider_stub("smart_fades"),
1136 ]
1137 )
1138
1139 with caplog.at_level("WARNING", logger=audio_analysis_mod.LOGGER.name):
1140 [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1141
1142 assert any(
1143 "Skipping unparsable audio_analysis row" in r.message
1144 and "id=42" in r.message
1145 and "sonic_analysis" in r.message
1146 for r in caplog.records
1147 )
1148
1149
1150@pytest.mark.asyncio
1151async def test_iter_merged_audio_analysis_rows_drops_groups_with_only_corrupt_rows() -> None:
1152 """A group whose only row has corrupt JSON is not emitted at all."""
1153 rows: list[dict[str, Any]] = [
1154 {
1155 "item_id": "broken",
1156 "provider": "filesystem_local",
1157 "aa_provider_domain": "sonic_analysis",
1158 "analysis_data": "not-json",
1159 },
1160 {
1161 "item_id": "good",
1162 "provider": "filesystem_local",
1163 "aa_provider_domain": "sonic_analysis",
1164 "analysis_data": '{"bpm": 100.0}',
1165 },
1166 ]
1167 c, _ = _stub_controller(iter_rows=rows)
1168 c.mass.get_providers = MagicMock(return_value=[_aa_provider_stub("sonic_analysis")]) # type: ignore[method-assign]
1169
1170 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1171 assert len(result) == 1
1172 assert result[0][0] == "good"
1173
1174
1175@pytest.mark.asyncio
1176async def test_iter_merged_audio_analysis_rows_warns_when_primary_domain_offline(
1177 caplog: pytest.LogCaptureFixture,
1178) -> None:
1179 """Offline primary AA domain yields nothing and surfaces a WARNING."""
1180 c, db = _stub_controller(iter_rows=[])
1181 # Only smart_fades is available; sonic_analysis (the queried domain) isn't.
1182 c.mass.get_providers = MagicMock(return_value=[_aa_provider_stub("smart_fades")]) # type: ignore[method-assign]
1183
1184 with caplog.at_level("WARNING", logger=audio_analysis_mod.LOGGER.name):
1185 result = [x async for x in c.iter_merged_audio_analysis_rows("sonic_analysis")]
1186
1187 assert result == []
1188 # Early return must short-circuit before any DB work.
1189 assert not db.iter_rows_from_query.called
1190 assert any(
1191 "offline primary AA domain" in r.message and "sonic_analysis" in r.message
1192 for r in caplog.records
1193 )
1194
1195
1196def _make_aa_provider_with_domain(
1197 domain: str,
1198 *,
1199 available: bool = True,
1200 analysis_version: int = 1,
1201) -> MagicMock:
1202 """AA provider mock with domain and analysis_version set."""
1203 provider = MagicMock(spec=AudioAnalysisProvider)
1204 provider.domain = domain
1205 provider.available = available
1206 provider.analysis_version = analysis_version
1207 return provider
1208
1209
1210@pytest.mark.asyncio
1211async def test_coverage_returns_three_counts_and_version() -> None:
1212 """get_coverage() reports analyzed, pending, stale_version, analysis_version."""
1213 c, _ = _stub_controller()
1214 p = _make_aa_provider_with_domain("sonic_analysis", analysis_version=3)
1215 c.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
1216 c.get_audio_analysis_count = AsyncMock(return_value=100) # type: ignore[method-assign]
1217 c._count_candidates_missing_analysis = AsyncMock(return_value=20) # type: ignore[method-assign]
1218 c.mass.music.database.get_count_from_query = AsyncMock( # type: ignore[method-assign]
1219 return_value=5
1220 )
1221
1222 result = await c.get_coverage(aa_domain="sonic_analysis")
1223
1224 assert result == AudioAnalysisCoverage(
1225 analyzed=100,
1226 pending=20,
1227 stale_version=5,
1228 analysis_version=3,
1229 )
1230
1231
1232@pytest.mark.asyncio
1233async def test_coverage_raises_for_unknown_aa_domain() -> None:
1234 """Unloaded AA provider raises ProviderUnavailableError."""
1235 c, _ = _stub_controller()
1236 c.mass.get_provider = MagicMock(return_value=None) # type: ignore[method-assign]
1237
1238 with pytest.raises(ProviderUnavailableError):
1239 await c.get_coverage(aa_domain="nope")
1240
1241
1242@pytest.mark.asyncio
1243async def test_coverage_stale_query_counts_null_analysis_version_as_stale() -> None:
1244 """Rows with NULL analysis_version must be counted as stale (SQLite `NULL < N` is NULL)."""
1245 c, db = _stub_controller(count_result=0)
1246 p = _make_aa_provider_with_domain("sonic_analysis", analysis_version=3)
1247 c.mass.get_provider = MagicMock(return_value=p) # type: ignore[method-assign]
1248 c.get_audio_analysis_count = AsyncMock(return_value=0) # type: ignore[method-assign]
1249 c._count_candidates_missing_analysis = AsyncMock(return_value=0) # type: ignore[method-assign]
1250
1251 await c.get_coverage(aa_domain="sonic_analysis")
1252
1253 sql, params = db.get_count_from_query.await_args.args
1254 assert "analysis_version IS NULL" in sql
1255 assert "analysis_version < :current_version" in sql
1256 assert params == {
1257 "aa_domain": "sonic_analysis",
1258 "media_type": MediaType.TRACK.value,
1259 "current_version": 3,
1260 }
1261
1262
1263@pytest.mark.asyncio
1264async def test_count_candidates_missing_analysis_zero_without_filesystem() -> None:
1265 """No available filesystem music providers -> 0 pending (no DB query)."""
1266 c, _ = _stub_controller()
1267 c.mass.get_providers = MagicMock(return_value=[]) # type: ignore[method-assign]
1268
1269 assert await c._count_candidates_missing_analysis("sonic_analysis", 1) == 0
1270
1271
1272@pytest.mark.asyncio
1273async def test_count_candidates_missing_analysis_queries_with_available_filesystem() -> None:
1274 """With an available filesystem provider, the NOT EXISTS count query runs with bound params."""
1275 c, db = _stub_controller(count_result=7)
1276 domain = next(iter(audio_analysis_mod.FILESYSTEM_PROVIDER_DOMAINS))
1277 fs_prov = MagicMock()
1278 fs_prov.domain = domain
1279 fs_prov.available = True
1280 c.mass.get_providers = MagicMock(return_value=[fs_prov]) # type: ignore[method-assign]
1281
1282 result = await c._count_candidates_missing_analysis("sonic_analysis", 2)
1283
1284 assert result == 7
1285 db.get_count_from_query.assert_awaited_once()
1286 sql, params = db.get_count_from_query.await_args.args
1287 assert "NOT EXISTS" in sql
1288 assert "aa.analysis_version IS NOT NULL" in sql
1289 assert "aa.analysis_version >= :current_version" in sql
1290 assert f"'{domain}'" in sql
1291 assert params["media_type"] == MediaType.TRACK.value
1292 assert params["aa_domain"] == "sonic_analysis"
1293 assert params["current_version"] == 2
1294 assert "now" in params
1295 assert "aa.analysis_version IS NOT NULL" in sql
1296 assert "aa.analysis_version >= :current_version" in sql
1297
1298
1299def test_controller_has_no_provider_specific_extra_data_keys() -> None:
1300 """Generic controller must not reference any provider extra_data key names."""
1301 source = inspect.getsource(audio_analysis_mod)
1302 # Deliberately a raw source-substring guard: the generic controller must never
1303 # name provider specifics, even in a comment. The brittleness is intentional --
1304 # do not weaken this to an import/attribute check.
1305 assert "_EXPORT_STRIP_EXTRA_DATA_KEYS" not in source
1306 assert "clap_embedding" not in source
1307
1308
1309# --- track audio metadata & waveform export ---
1310
1311
1312def _track_with_mapping(item_id: str = "track-1", provider: str = "test-provider") -> Track:
1313 """Build a Track with a single provider mapping."""
1314 return Track(
1315 item_id=item_id,
1316 provider="library",
1317 name="Test Track",
1318 provider_mappings={
1319 ProviderMapping(
1320 item_id=item_id,
1321 provider_domain=provider,
1322 provider_instance=provider,
1323 )
1324 },
1325 )
1326
1327
1328def _analysis_controller_with_rows(
1329 rows: list[Mapping[str, Any]],
1330) -> AudioAnalysisController:
1331 """Build a stub controller whose DB returns the given analysis rows for any track."""
1332 c, db = _stub_controller()
1333 db.get_rows = AsyncMock(return_value=rows)
1334 music_prov = MagicMock(spec=MusicProvider)
1335 music_prov.is_streaming_provider = True
1336 music_prov.domain = "test-provider"
1337 c.mass.get_provider = MagicMock(return_value=music_prov) # type: ignore[method-assign]
1338 c.mass.get_providers = MagicMock( # type: ignore[method-assign]
1339 return_value=[
1340 _aa_provider_stub(SMART_FADES_ANALYSIS_DOMAIN),
1341 _aa_provider_stub(SONIC_ANALYSIS_DOMAIN),
1342 ]
1343 )
1344 return c
1345
1346
1347@pytest.mark.asyncio
1348async def test_get_track_audio_metadata_skips_corrupt_sqlite_row(
1349 caplog: pytest.LogCaptureFixture,
1350) -> None:
1351 """A corrupt sqlite3.Row is logged and skipped without hiding valid analysis."""
1352 corrupt_data = json_dumps({"spectral_centroid": [100.0, None]})
1353 valid_data = json_dumps(AudioAnalysisData(bpm=128.0).to_dict())
1354 with closing(sqlite3.connect(":memory:")) as db:
1355 db.row_factory = sqlite3.Row
1356 rows = cast(
1357 "list[Mapping[str, Any]]",
1358 db.execute(
1359 """
1360 SELECT 1 AS id, ? AS aa_provider_domain, ? AS analysis_data
1361 UNION ALL
1362 SELECT 2, ?, ?
1363 """,
1364 (
1365 SONIC_ANALYSIS_DOMAIN,
1366 corrupt_data,
1367 SMART_FADES_ANALYSIS_DOMAIN,
1368 valid_data,
1369 ),
1370 ).fetchall(),
1371 )
1372
1373 controller = _analysis_controller_with_rows(rows)
1374 with caplog.at_level("WARNING", logger=audio_analysis_mod.LOGGER.name):
1375 result = await controller.get_track_audio_metadata(_track_with_mapping())
1376
1377 assert result is not None
1378 assert result.bpm == 128.0
1379 warning = next(
1380 record for record in caplog.records if record.name == audio_analysis_mod.LOGGER.name
1381 )
1382 assert "id=1, domain=sonic_analysis" in warning.getMessage()
1383 assert corrupt_data not in warning.getMessage()
1384 assert warning.exc_info is None
1385
1386
1387@pytest.mark.asyncio
1388async def test_get_audio_analysis_deletes_unparsable_rows(
1389 caplog: pytest.LogCaptureFixture,
1390) -> None:
1391 """Corrupt rows are deleted, so their stored version no longer blocks re-analysis."""
1392 rows: list[Mapping[str, Any]] = [
1393 {
1394 "id": 7,
1395 "aa_provider_domain": SMART_FADES_ANALYSIS_DOMAIN,
1396 "analysis_data": json_dumps({"spectral_centroid": [100.0, None]}),
1397 },
1398 _aa_row(SONIC_ANALYSIS_DOMAIN, 8, bpm=101.0),
1399 ]
1400 controller = _analysis_controller_with_rows(rows)
1401 with caplog.at_level("WARNING", logger=audio_analysis_mod.LOGGER.name):
1402 result = await controller.get_audio_analysis("track-1", "test-provider")
1403
1404 assert result is not None
1405 assert result.bpm == 101.0
1406 delete_mock = cast("AsyncMock", controller.mass.music.database.delete)
1407 delete_mock.assert_awaited_once_with(DB_TABLE_AUDIO_ANALYSIS, {"id": 7})
1408 warning = next(r for r in caplog.records if r.name == audio_analysis_mod.LOGGER.name)
1409 assert "in field spectral_centroid" in warning.getMessage()
1410
1411
1412@pytest.mark.asyncio
1413async def test_set_audio_analysis_rejects_non_finite_values() -> None:
1414 """A payload holding non-finite floats is refused before anything reaches the database."""
1415 c, db = _stub_controller()
1416 list_case = AudioAnalysisData(spectral_centroid=[100.0, float("nan"), 200.0])
1417 with pytest.raises(ValueError, match="spectral_centroid"):
1418 await c.set_audio_analysis(
1419 "track-1", "test-provider", SMART_FADES_ANALYSIS_DOMAIN, list_case
1420 )
1421 scalar_case = AudioAnalysisData(bpm=float("inf"))
1422 with pytest.raises(ValueError, match="bpm"):
1423 await c.set_audio_analysis(
1424 "track-1", "test-provider", SMART_FADES_ANALYSIS_DOMAIN, scalar_case
1425 )
1426 db.insert_or_replace.assert_not_called()
1427
1428
1429@pytest.mark.asyncio
1430async def test_get_track_audio_metadata_prefers_smart_fades() -> None:
1431 """bpm/key come from smart_fades even when another AA provider wrote them later."""
1432 c = _analysis_controller_with_rows(
1433 [
1434 _aa_row(SMART_FADES_ANALYSIS_DOMAIN, 1, bpm=128.0, key="F#", mode="minor"),
1435 _aa_row(SONIC_ANALYSIS_DOMAIN, 2, bpm=100.0),
1436 ]
1437 )
1438 result = await c.get_track_audio_metadata(_track_with_mapping())
1439 assert result is not None
1440 assert result.bpm == 128.0
1441 assert result.musical_key == "F# minor"
1442
1443
1444@pytest.mark.asyncio
1445async def test_get_track_audio_metadata_key_without_mode() -> None:
1446 """musical_key falls back to the bare pitch class when no mode was detected."""
1447 c = _analysis_controller_with_rows([_aa_row(SMART_FADES_ANALYSIS_DOMAIN, 1, key="C")])
1448 result = await c.get_track_audio_metadata(_track_with_mapping())
1449 assert result is not None
1450 assert result.bpm is None
1451 assert result.musical_key == "C"
1452
1453
1454@pytest.mark.asyncio
1455async def test_get_track_audio_metadata_none_without_relevant_analysis() -> None:
1456 """No AudioMetadata when stored analysis has neither bpm nor key."""
1457 c = _analysis_controller_with_rows(
1458 [_aa_row(SONIC_ANALYSIS_DOMAIN, 1, loudness_integrated=-7.5)]
1459 )
1460 assert await c.get_track_audio_metadata(_track_with_mapping()) is None
1461
1462
1463@pytest.mark.asyncio
1464async def test_get_wave_form_returns_rms_bins() -> None:
1465 """wave_form returns the stored RMS energy bins as a plain list of floats."""
1466 rms = np.linspace(0.0, 1.0, 1800, dtype=np.float32).tolist()
1467 c = _analysis_controller_with_rows([_aa_row(SMART_FADES_ANALYSIS_DOMAIN, 1, rms_energy=rms)])
1468 result = await c.get_wave_form("track-1", "test-provider")
1469 assert result is not None
1470 assert len(result) == 1800
1471 assert result[0] == pytest.approx(0.0)
1472 assert result[-1] == pytest.approx(1.0)
1473 assert all(isinstance(v, float) for v in result)
1474
1475
1476@pytest.mark.asyncio
1477async def test_get_wave_form_none_without_rms() -> None:
1478 """wave_form returns None when no AA provider stored RMS energy."""
1479 c = _analysis_controller_with_rows([_aa_row(SMART_FADES_ANALYSIS_DOMAIN, 1, bpm=120.0)])
1480 assert await c.get_wave_form("track-1", "test-provider") is None
1481