/
/
/
1"""Tests for the AudioAnalysisProvider base class lifecycle."""
2
3from __future__ import annotations
4
5import asyncio
6import sqlite3
7import threading
8import time
9from concurrent.futures import ThreadPoolExecutor
10from datetime import UTC, datetime
11from typing import TYPE_CHECKING, cast
12from unittest.mock import AsyncMock, MagicMock, patch
13
14import pytest
15
16from music_assistant.models.audio_analysis import AudioAnalysisData, AudioAnalysisError
17from music_assistant.models.audio_analysis_provider import (
18 AudioAnalysisProvider,
19 InstrumentedSemaphore,
20)
21from tests.common import collect_loop_errors
22
23if TYPE_CHECKING:
24 from music_assistant_models.media_items import AudioFormat
25 from music_assistant_models.streamdetails import StreamDetails
26
27
28class _StubProvider(AudioAnalysisProvider):
29 """Minimal concrete provider for base-class tests."""
30
31 async def _start_analysis(
32 self, session_id: str, streamdetails: StreamDetails, audio_format: AudioFormat
33 ) -> bool:
34 return True
35
36 async def process_pcm_chunk(self, session_id: str, pcm_chunk: bytes) -> None:
37 return None
38
39 async def _finalize(self, session_id: str) -> AudioAnalysisData | None:
40 return None
41
42
43def _make_provider() -> _StubProvider:
44 mass = MagicMock()
45 mass.streams.audio_analysis.get_audio_analysis_version = AsyncMock(return_value=None)
46 mass.streams.audio_analysis.set_audio_analysis = AsyncMock()
47 mass.streams.audio_analysis.record_analysis_failure = AsyncMock()
48 mass.streams.audio_analysis.clear_analysis_failure = AsyncMock()
49 manifest = MagicMock()
50 manifest.domain = "test_stub_provider"
51 config = MagicMock()
52 config.get_value = MagicMock(return_value="GLOBAL")
53 return _StubProvider(mass, manifest, config, supported_features=set())
54
55
56@pytest.mark.asyncio
57async def test_post_analysis_default_is_noop() -> None:
58 """Default post_analysis must be a no-op that returns None."""
59 provider = _make_provider()
60 streamdetails = MagicMock()
61 analysis = AudioAnalysisData()
62 await provider.post_analysis(streamdetails, analysis)
63
64
65@pytest.mark.asyncio
66async def test_finalize_calls_post_analysis_when_finalize_returns_analysis() -> None:
67 """When _finalize returns analysis, finalize must call post_analysis with it."""
68 provider = _make_provider()
69 streamdetails = MagicMock()
70 audio_format = MagicMock()
71 analysis = AudioAnalysisData(loudness_integrated=-14.0)
72
73 provider._finalize = AsyncMock(return_value=analysis) # type: ignore[method-assign]
74 provider.post_analysis = AsyncMock(return_value=None) # type: ignore[method-assign]
75
76 await provider.start_analysis("session-1", streamdetails, audio_format)
77 await provider.finalize("session-1")
78
79 provider.post_analysis.assert_awaited_once_with(streamdetails, analysis)
80 assert "session-1" not in provider._sessions
81
82
83@pytest.mark.asyncio
84async def test_finalize_skips_post_analysis_when_finalize_returns_none() -> None:
85 """When _finalize returns None, post_analysis must NOT be called."""
86 provider = _make_provider()
87 streamdetails = MagicMock()
88 audio_format = MagicMock()
89
90 provider._finalize = AsyncMock(return_value=None) # type: ignore[method-assign]
91 provider.post_analysis = AsyncMock(return_value=None) # type: ignore[method-assign]
92
93 await provider.start_analysis("session-2", streamdetails, audio_format)
94 await provider.finalize("session-2")
95
96 provider.post_analysis.assert_not_awaited()
97 assert "session-2" not in provider._sessions
98
99
100@pytest.mark.asyncio
101async def test_finalize_swallows_finalize_exception_and_skips_post_analysis() -> None:
102 """If _finalize raises, post_analysis must not be called and the exception must not propagate."""
103 provider = _make_provider()
104 streamdetails = MagicMock()
105 audio_format = MagicMock()
106
107 provider._finalize = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
108 provider.post_analysis = AsyncMock(return_value=None) # type: ignore[method-assign]
109
110 await provider.start_analysis("session-3", streamdetails, audio_format)
111 await provider.finalize("session-3")
112
113 provider.post_analysis.assert_not_awaited()
114 assert "session-3" not in provider._sessions
115
116
117@pytest.mark.asyncio
118async def test_finalize_swallows_post_analysis_exception() -> None:
119 """post_analysis raising must be caught; the analysis row stays valid."""
120 provider = _make_provider()
121 streamdetails = MagicMock()
122 audio_format = MagicMock()
123 analysis = AudioAnalysisData()
124
125 provider._finalize = AsyncMock(return_value=analysis) # type: ignore[method-assign]
126 provider.post_analysis = AsyncMock(side_effect=RuntimeError("tag write failed")) # type: ignore[method-assign]
127
128 await provider.start_analysis("session-4", streamdetails, audio_format)
129 # Must not raise
130 await provider.finalize("session-4")
131
132 provider.post_analysis.assert_awaited_once()
133 assert "session-4" not in provider._sessions
134
135
136@pytest.mark.asyncio
137async def test_start_analysis_skips_tracks_over_max_duration() -> None:
138 """A provider with max_analysis_duration set rejects longer tracks before any DB/work."""
139 provider = _make_provider()
140 provider.max_analysis_duration = 1800.0
141 provider._start_analysis = AsyncMock(return_value=True) # type: ignore[method-assign]
142 streamdetails = MagicMock()
143 streamdetails.duration = 7200 # 2 hours
144
145 accepted = await provider.start_analysis("s", streamdetails, MagicMock())
146
147 assert accepted is False
148 provider._start_analysis.assert_not_awaited()
149 # The duration gate comes first, so the version lookup is skipped too.
150 version_lookup = provider.mass.streams.audio_analysis.get_audio_analysis_version
151 version_lookup.assert_not_awaited() # type: ignore[attr-defined]
152
153
154@pytest.mark.asyncio
155async def test_start_analysis_allows_tracks_under_max_duration() -> None:
156 """Tracks within the cap proceed to the provider's own _start_analysis."""
157 provider = _make_provider()
158 provider.max_analysis_duration = 1800.0
159 provider._start_analysis = AsyncMock(return_value=True) # type: ignore[method-assign]
160 streamdetails = MagicMock()
161 streamdetails.duration = 240
162
163 assert await provider.start_analysis("s", streamdetails, MagicMock()) is True
164 provider._start_analysis.assert_awaited_once()
165
166
167@pytest.mark.asyncio
168async def test_start_analysis_no_duration_cap_by_default() -> None:
169 """With max_analysis_duration unset, even a very long track is accepted."""
170 provider = _make_provider()
171 provider._start_analysis = AsyncMock(return_value=True) # type: ignore[method-assign]
172 streamdetails = MagicMock()
173 streamdetails.duration = 99999
174
175 assert await provider.start_analysis("s", streamdetails, MagicMock()) is True
176
177
178@pytest.mark.asyncio
179async def test_run_offloaded_acquires_semaphore_when_present() -> None:
180 """_run_offloaded holds the controller's analysis semaphore while the work runs."""
181 provider = _make_provider()
182 semaphore = InstrumentedSemaphore(1)
183 provider.mass.streams.audio_analysis.analysis_semaphore = semaphore
184
185 def _work() -> bool:
186 # Runs in the worker thread; the cap must already be held here.
187 return semaphore.locked()
188
189 held_during = await provider._run_offloaded(_work)
190 assert held_during is True
191 assert not semaphore.locked()
192
193
194@pytest.mark.asyncio
195async def test_run_offloaded_without_cap_runs_plainly() -> None:
196 """With no real semaphore configured, _run_offloaded still runs the work and forwards args."""
197 provider = _make_provider()
198 # The MagicMock attribute is not an asyncio.Semaphore, so this falls back to a plain thread.
199 result = await provider._run_offloaded(lambda value: value * 2, 21)
200 assert result == 42
201
202
203@pytest.mark.asyncio
204async def test_run_offloaded_timed_returns_result_and_execution_seconds() -> None:
205 """_run_offloaded_timed forwards args and reports the callable's own execution time."""
206 provider = _make_provider()
207
208 def _work(value: int) -> int:
209 time.sleep(0.05)
210 return value * 2
211
212 result, seconds = await provider._run_offloaded_timed(_work, 21)
213 assert result == 42
214 assert seconds >= 0.05
215
216
217@pytest.mark.asyncio
218async def test_run_offloaded_holds_permit_until_thread_finishes_on_cancel() -> None:
219 """Cancelling the awaiter must not free the permit while the worker thread is still running."""
220 provider = _make_provider()
221 semaphore = InstrumentedSemaphore(1)
222 provider.mass.streams.audio_analysis.analysis_semaphore = semaphore
223
224 started = threading.Event()
225 may_finish = threading.Event()
226
227 def _blocking() -> str:
228 started.set()
229 may_finish.wait(timeout=5)
230 return "done"
231
232 task = asyncio.create_task(provider._run_offloaded(_blocking))
233 assert await asyncio.to_thread(started.wait, 5) # thread is running, permit acquired
234 assert semaphore.locked()
235
236 task.cancel()
237 with pytest.raises(asyncio.CancelledError):
238 await task
239 # Awaiter cancelled, but the thread is still running â the permit must NOT be freed yet.
240 assert semaphore.locked()
241
242 may_finish.set() # let the thread complete; the done-callback then releases the permit
243 for _ in range(100):
244 if not semaphore.locked():
245 break
246 await asyncio.sleep(0.02)
247 assert not semaphore.locked()
248
249
250@pytest.mark.asyncio
251async def test_run_offloaded_worker_failure_after_cancel_logs_no_loop_error() -> None:
252 """A worker failing after the awaiter was cancelled must not be reported to the loop."""
253 provider = _make_provider()
254 semaphore = InstrumentedSemaphore(1)
255 provider.mass.streams.audio_analysis.analysis_semaphore = semaphore
256
257 started = threading.Event()
258 may_finish = threading.Event()
259
260 def _failing() -> str:
261 started.set()
262 may_finish.wait(timeout=5)
263 raise RuntimeError("boom")
264
265 with collect_loop_errors() as reported:
266 task = asyncio.create_task(provider._run_offloaded(_failing))
267 assert await asyncio.to_thread(started.wait, 5) # thread is running, permit acquired
268 assert semaphore.locked()
269
270 task.cancel()
271 with pytest.raises(asyncio.CancelledError):
272 await task
273 # release the worker only once the cancellation is fully processed, so the failure
274 # reliably lands after the awaiter has already given up -- releasing earlier would
275 # let the failure race the cancellation and pass even on the buggy shield-based code
276 may_finish.set()
277
278 # the done-callback releases the permit once the worker (and its failure) is
279 # observed, so waiting for that also gives any (buggy) loop report time to land
280 for _ in range(100):
281 if not semaphore.locked():
282 break
283 await asyncio.sleep(0.02)
284 assert not semaphore.locked() # permit still released despite the failure
285
286 assert reported == []
287
288
289@pytest.mark.asyncio
290async def test_run_offloaded_serializes_to_one_while_streaming() -> None:
291 """While a player is streaming, offloads run one at a time even if the semaphore allows more."""
292 provider = _make_provider()
293 ctrl = provider.mass.streams.audio_analysis
294 ctrl.analysis_semaphore = InstrumentedSemaphore(4) # plenty of permits
295 ctrl.analysis_solo_lock = asyncio.Lock()
296 ctrl.analysis_executor = None # fall back to asyncio.to_thread
297 ctrl.playback_active = MagicMock(return_value=True) # type: ignore[method-assign]
298
299 first_started = threading.Event()
300 first_may_finish = threading.Event()
301 second_started = threading.Event()
302
303 def _first() -> str:
304 first_started.set()
305 first_may_finish.wait(timeout=5)
306 return "first"
307
308 def _second() -> str:
309 second_started.set()
310 return "second"
311
312 t1 = asyncio.create_task(provider._run_offloaded(_first))
313 assert await asyncio.to_thread(first_started.wait, 5)
314
315 # The second offload must block on the solo lock while the first holds it.
316 t2 = asyncio.create_task(provider._run_offloaded(_second))
317 await asyncio.sleep(0.1)
318 assert not second_started.is_set()
319
320 first_may_finish.set()
321 assert await t1 == "first"
322 assert await t2 == "second"
323 assert second_started.is_set()
324
325
326@pytest.mark.asyncio
327async def test_run_offloaded_runs_concurrently_when_idle() -> None:
328 """With no player streaming, the solo lock is not taken and offloads overlap up to the cap."""
329 provider = _make_provider()
330 ctrl = provider.mass.streams.audio_analysis
331 ctrl.analysis_semaphore = InstrumentedSemaphore(4)
332 ctrl.analysis_solo_lock = asyncio.Lock()
333 ctrl.analysis_executor = None
334 ctrl.playback_active = MagicMock(return_value=False) # type: ignore[method-assign]
335
336 first_started = threading.Event()
337 first_may_finish = threading.Event()
338 second_started = threading.Event()
339
340 def _first() -> str:
341 first_started.set()
342 first_may_finish.wait(timeout=5)
343 return "first"
344
345 def _second() -> str:
346 second_started.set()
347 return "second"
348
349 t1 = asyncio.create_task(provider._run_offloaded(_first))
350 assert await asyncio.to_thread(first_started.wait, 5)
351 t2 = asyncio.create_task(provider._run_offloaded(_second))
352 # Idle: the second offload runs without waiting for the first to finish.
353 assert await asyncio.to_thread(second_started.wait, 5)
354
355 first_may_finish.set()
356 await asyncio.gather(t1, t2)
357
358
359@pytest.mark.asyncio
360async def test_run_offloaded_uses_dedicated_executor_when_present() -> None:
361 """Offloads run on the controller's dedicated (niced) analysis pool when configured."""
362 provider = _make_provider()
363 ctrl = provider.mass.streams.audio_analysis
364 ctrl.analysis_semaphore = InstrumentedSemaphore(1)
365 ctrl.analysis_solo_lock = None # not an asyncio.Lock -> solo step skipped
366 executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="analysis")
367 ctrl.analysis_executor = executor
368 try:
369 thread_name = await provider._run_offloaded(lambda: threading.current_thread().name)
370 finally:
371 executor.shutdown(wait=False)
372 assert thread_name.startswith("analysis")
373
374
375@pytest.mark.asyncio
376async def test_run_offloaded_releases_permit_if_scheduling_fails() -> None:
377 """A failure to schedule the worker must release the permit, not leak it."""
378 provider = _make_provider()
379 semaphore = InstrumentedSemaphore(1)
380 provider.mass.streams.audio_analysis.analysis_semaphore = semaphore
381
382 with (
383 patch(
384 "music_assistant.models.audio_analysis_provider.asyncio.to_thread",
385 side_effect=RuntimeError("cannot schedule"),
386 ),
387 pytest.raises(RuntimeError, match="cannot schedule"),
388 ):
389 await provider._run_offloaded(lambda: "x")
390
391 assert not semaphore.locked() # permit released despite the failure
392
393
394def test_audio_analysis_error_carries_reason_and_retry() -> None:
395 """AudioAnalysisError exposes reason and retry_at; retry_at defaults to None."""
396 err = AudioAnalysisError("bad file")
397 assert err.reason == "bad file"
398 assert err.retry_at is None
399
400 when = datetime(2030, 1, 1, tzinfo=UTC)
401 err2 = AudioAnalysisError("offline", retry_at=when)
402 assert err2.retry_at == when
403 assert str(err2) == "offline"
404
405
406def test_audio_analysis_error_rejects_naive_retry_at() -> None:
407 """A naive (tz-unaware) retry_at is rejected to avoid silent epoch skew."""
408 with pytest.raises(ValueError, match="timezone-aware"):
409 AudioAnalysisError("x", retry_at=datetime(2030, 1, 1)) # noqa: DTZ001
410
411
412@pytest.mark.asyncio
413async def test_finalize_records_classified_failure_on_audio_analysis_error() -> None:
414 """A raised AudioAnalysisError in _finalize records reason + retry_at and skips persist."""
415 provider = _make_provider()
416 streamdetails = MagicMock()
417 streamdetails.item_id = "track-1"
418 streamdetails.provider = "test_prov"
419 streamdetails.media_type = "track"
420 when = datetime(2030, 1, 1, tzinfo=UTC)
421
422 provider._finalize = AsyncMock( # type: ignore[method-assign]
423 side_effect=AudioAnalysisError("no usable audio frames extracted", retry_at=when)
424 )
425
426 await provider.start_analysis("s1", streamdetails, MagicMock())
427 await provider.finalize("s1")
428
429 rec = cast("AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure)
430 rec.assert_awaited_once()
431 kwargs = rec.call_args.kwargs
432 assert kwargs["reason"] == "no usable audio frames extracted"
433 assert kwargs["retry_at"] == when
434 assert kwargs["aa_provider_domain"] == "test_stub_provider"
435 cast("AsyncMock", provider.mass.streams.audio_analysis.set_audio_analysis).assert_not_awaited()
436 assert "s1" not in provider._sessions
437
438
439@pytest.mark.asyncio
440async def test_record_failure_passes_provider_analysis_version() -> None:
441 """Failures are recorded at the provider's analysis_version so a version bump unblocks them."""
442 provider = _make_provider()
443 provider.analysis_version = 7
444 streamdetails = MagicMock()
445 streamdetails.item_id = "track-3"
446 streamdetails.provider = "test_prov"
447 streamdetails.media_type = "track"
448
449 provider._finalize = AsyncMock(side_effect=AudioAnalysisError("boom")) # type: ignore[method-assign]
450
451 await provider.start_analysis("s5", streamdetails, MagicMock())
452 await provider.finalize("s5")
453
454 rec = cast("AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure)
455 rec.assert_awaited_once()
456 assert rec.call_args.kwargs["analysis_version"] == 7
457
458
459@pytest.mark.asyncio
460async def test_finalize_records_never_retry_on_generic_exception() -> None:
461 """A generic exception in _finalize records str(err) with retry_at None."""
462 provider = _make_provider()
463 streamdetails = MagicMock()
464 streamdetails.item_id = "track-2"
465 streamdetails.provider = "test_prov"
466 streamdetails.media_type = "track"
467
468 provider._finalize = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
469
470 await provider.start_analysis("s2", streamdetails, MagicMock())
471 await provider.finalize("s2")
472
473 rec = cast("AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure)
474 rec.assert_awaited_once()
475 assert rec.call_args.kwargs["reason"] == "boom"
476 assert rec.call_args.kwargs["retry_at"] is None
477 assert "s2" not in provider._sessions
478
479
480@pytest.mark.asyncio
481async def test_finalize_no_record_on_none_return() -> None:
482 """A plain None return from _finalize records nothing (deliberate skip)."""
483 provider = _make_provider()
484 streamdetails = MagicMock()
485 provider._finalize = AsyncMock(return_value=None) # type: ignore[method-assign]
486
487 await provider.start_analysis("s3", streamdetails, MagicMock())
488 await provider.finalize("s3")
489
490 cast(
491 "AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure
492 ).assert_not_awaited()
493
494
495@pytest.mark.asyncio
496async def test_start_analysis_records_on_audio_analysis_error() -> None:
497 """A raised AudioAnalysisError in _start_analysis records the failure and rejects."""
498 provider = _make_provider()
499 streamdetails = MagicMock()
500 streamdetails.item_id = "track-4"
501 streamdetails.provider = "test_prov"
502 streamdetails.media_type = "track"
503
504 provider._start_analysis = AsyncMock( # type: ignore[method-assign]
505 side_effect=AudioAnalysisError("unsupported codec")
506 )
507
508 accepted = await provider.start_analysis("s4", streamdetails, MagicMock())
509
510 assert accepted is False
511 rec = cast("AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure)
512 rec.assert_awaited_once()
513 assert rec.call_args.kwargs["reason"] == "unsupported codec"
514 assert "s4" not in provider._sessions
515
516
517@pytest.mark.asyncio
518async def test_finalize_swallows_recorder_error() -> None:
519 """If record_analysis_failure raises, finalize must not propagate and must still clean up."""
520 provider = _make_provider()
521 provider.logger = MagicMock()
522 streamdetails = MagicMock()
523 streamdetails.item_id = "track-x"
524 streamdetails.provider = "test_prov"
525 streamdetails.media_type = "track"
526 provider._finalize = AsyncMock( # type: ignore[method-assign]
527 side_effect=AudioAnalysisError("boom")
528 )
529 aa = cast("MagicMock", provider.mass.streams.audio_analysis)
530 aa.record_analysis_failure = AsyncMock(side_effect=sqlite3.OperationalError("db down"))
531
532 await provider.start_analysis("sx", streamdetails, MagicMock())
533 # Must not raise despite the recorder failing.
534 await provider.finalize("sx")
535
536 provider.logger.warning.assert_called()
537 assert "sx" not in provider._sessions
538
539
540@pytest.mark.asyncio
541async def test_ensure_models_loaded_loads_once() -> None:
542 """Concurrent ensure_models_loaded calls load the heavy models a single time."""
543 provider = _make_provider()
544 provider.has_unloadable_models = True
545 calls = 0
546
547 async def _load() -> None:
548 nonlocal calls
549 calls += 1
550 await asyncio.sleep(0.01)
551
552 provider._load_models = _load # type: ignore[method-assign]
553
554 results = await asyncio.gather(*(provider.ensure_models_loaded() for _ in range(5)))
555
556 assert all(results)
557 assert calls == 1
558 assert provider._models_loaded is True
559
560
561@pytest.mark.asyncio
562async def test_unload_idle_models_frees_then_reloads() -> None:
563 """unload_idle_models frees the models; the next ensure reloads them."""
564 provider = _make_provider()
565 provider.has_unloadable_models = True
566 load_calls = 0
567 free_calls = 0
568
569 async def _load() -> None:
570 nonlocal load_calls
571 load_calls += 1
572
573 def _free() -> None:
574 nonlocal free_calls
575 free_calls += 1
576
577 provider._load_models = _load # type: ignore[method-assign]
578 provider._free_models = _free # type: ignore[method-assign]
579
580 await provider.ensure_models_loaded()
581 await provider.unload_idle_models()
582 assert free_calls == 1
583 assert provider._models_loaded is False
584
585 await provider.ensure_models_loaded()
586 assert load_calls == 2 # reloaded on demand
587 assert provider._models_loaded is True
588
589
590@pytest.mark.asyncio
591async def test_start_analysis_rejects_when_model_load_fails() -> None:
592 """A provider whose model load fails declines the session instead of starting it."""
593 provider = _make_provider()
594 provider.has_unloadable_models = True
595 provider._start_analysis = AsyncMock(return_value=True) # type: ignore[method-assign]
596
597 async def _load() -> None:
598 raise RuntimeError("model load boom")
599
600 provider._load_models = _load # type: ignore[method-assign]
601 streamdetails = MagicMock()
602 streamdetails.duration = 120
603
604 assert await provider.start_analysis("s", streamdetails, MagicMock()) is False
605 provider._start_analysis.assert_not_awaited()
606 assert provider._models_loaded is False
607
608
609@pytest.mark.asyncio
610async def test_unload_cancels_in_flight_finalize_before_freeing_models() -> None:
611 """unload() cancels a running finalize and frees the models only once it has unwound."""
612 provider = _make_provider()
613 provider.has_unloadable_models = True
614 started = asyncio.Event()
615 running = False
616 cancelled = False
617 freed_while_running = False
618
619 async def _finalize(session_id: str) -> AudioAnalysisData | None: # noqa: ARG001
620 nonlocal running, cancelled
621 running = True
622 started.set()
623 try:
624 await asyncio.sleep(60)
625 except asyncio.CancelledError:
626 cancelled = True
627 raise
628 finally:
629 running = False
630 return None
631
632 def _free() -> None:
633 nonlocal freed_while_running
634 if running:
635 freed_while_running = True
636
637 provider._load_models = AsyncMock() # type: ignore[method-assign]
638 provider._free_models = _free # type: ignore[method-assign]
639 provider._finalize = _finalize # type: ignore[method-assign]
640
641 await provider.ensure_models_loaded()
642 task = asyncio.create_task(provider.finalize("s-inflight"))
643 await started.wait()
644 assert provider._finalize_tasks == {task}
645
646 await asyncio.wait_for(provider.unload(), timeout=5)
647
648 assert cancelled is True
649 assert freed_while_running is False
650 assert task.cancelled()
651 assert not provider._finalize_tasks
652 assert provider._models_loaded is False
653 recorder = cast("AsyncMock", provider.mass.streams.audio_analysis.record_analysis_failure)
654 recorder.assert_not_awaited()
655
656
657@pytest.mark.asyncio
658async def test_unload_frees_models_when_no_finalize_in_flight() -> None:
659 """unload() with nothing in flight frees the models as before."""
660 provider = _make_provider()
661 provider.has_unloadable_models = True
662 free_calls = 0
663
664 def _free() -> None:
665 nonlocal free_calls
666 free_calls += 1
667
668 provider._load_models = AsyncMock() # type: ignore[method-assign]
669 provider._free_models = _free # type: ignore[method-assign]
670
671 await provider.ensure_models_loaded()
672 assert provider._models_loaded is True
673
674 await asyncio.wait_for(provider.unload(), timeout=5)
675
676 assert free_calls == 1
677 assert provider._models_loaded is False
678
679
680@pytest.mark.asyncio
681async def test_unload_from_within_finalize_does_not_cancel_itself() -> None:
682 """A finalize that reaches unload() must not cancel the task it is running on."""
683 provider = _make_provider()
684 provider.has_unloadable_models = True
685
686 async def _finalize(session_id: str) -> AudioAnalysisData | None: # noqa: ARG001
687 await provider.unload()
688 return None
689
690 provider._load_models = AsyncMock() # type: ignore[method-assign]
691 provider._free_models = MagicMock() # type: ignore[method-assign]
692 provider._finalize = _finalize # type: ignore[method-assign]
693
694 await provider.ensure_models_loaded()
695 await asyncio.wait_for(provider.finalize("s-self"), timeout=5)
696
697 assert provider._models_loaded is False
698 assert not provider._finalize_tasks
699
700
701@pytest.mark.asyncio
702async def test_finalize_is_skipped_while_unloading() -> None:
703 """A finalize issued while unloading runs no inference, registers nothing, clears its session."""
704 provider = _make_provider()
705 provider._finalize = AsyncMock(return_value=None) # type: ignore[method-assign]
706 provider._sessions["s-gate"] = MagicMock()
707 provider.unloading = True
708
709 await provider.finalize("s-gate")
710
711 provider._finalize.assert_not_awaited()
712 assert not provider._finalize_tasks
713 assert "s-gate" not in provider._sessions
714
715
716@pytest.mark.asyncio
717async def test_finalize_started_during_unload_does_not_run_inference() -> None:
718 """A finalize registered while unload() unwinds must not infer against models being freed."""
719 provider = _make_provider()
720 provider.has_unloadable_models = True
721 finalized: list[str] = []
722 late_task: asyncio.Task[None] | None = None
723 started = asyncio.Event()
724
725 async def _finalize(session_id: str) -> AudioAnalysisData | None:
726 nonlocal late_task
727 finalized.append(session_id)
728 if session_id != "s-first":
729 return None
730 started.set()
731 try:
732 await asyncio.sleep(60)
733 except asyncio.CancelledError:
734 # mirrors the controller resolving the provider and finalizing another session
735 # while unload() is suspended awaiting this cancellation
736 late_task = asyncio.create_task(provider.finalize("s-late"))
737 raise
738 return None
739
740 provider._load_models = AsyncMock() # type: ignore[method-assign]
741 provider._free_models = MagicMock() # type: ignore[method-assign]
742 provider._finalize = _finalize # type: ignore[method-assign]
743
744 await provider.ensure_models_loaded()
745 first = asyncio.create_task(provider.finalize("s-first"))
746 await started.wait()
747
748 # mass.unload_provider sets this before it awaits provider.unload()
749 provider.unloading = True
750 await asyncio.wait_for(provider.unload(), timeout=5)
751
752 assert late_task is not None
753 await asyncio.wait_for(late_task, timeout=5)
754 assert finalized == ["s-first"]
755 assert first.cancelled()
756 assert not provider._finalize_tasks
757
758
759@pytest.mark.asyncio
760async def test_start_analysis_declines_while_unloading() -> None:
761 """A session arriving from a stale provider snapshot must not start on a dying provider."""
762 provider = _make_provider()
763 provider._start_analysis = AsyncMock(return_value=True) # type: ignore[method-assign]
764 provider.unloading = True
765
766 accepted = await provider.start_analysis("s", MagicMock(), MagicMock())
767
768 assert accepted is False
769 provider._start_analysis.assert_not_awaited()
770 # The gate comes first, so nothing downstream of it runs either.
771 version_lookup = provider.mass.streams.audio_analysis.get_audio_analysis_version
772 version_lookup.assert_not_awaited() # type: ignore[attr-defined]
773 assert "s" not in provider._sessions
774
775
776@pytest.mark.asyncio
777async def test_ensure_models_loaded_declines_while_unloading() -> None:
778 """Models must never be reloaded onto a provider that is on its way out."""
779 provider = _make_provider()
780 provider._load_models = AsyncMock() # type: ignore[method-assign]
781 provider.unloading = True
782
783 assert await provider.ensure_models_loaded() is False
784 provider._load_models.assert_not_awaited()
785 assert provider._models_loaded is False
786
787
788@pytest.mark.asyncio
789async def test_ensure_models_loaded_declines_when_queued_behind_unload() -> None:
790 """A loader already waiting on the models lock must decline once unload() has run."""
791 provider = _make_provider()
792 provider.has_unloadable_models = True
793 loaded = False
794
795 async def _load() -> None:
796 nonlocal loaded
797 loaded = True
798
799 provider._load_models = _load # type: ignore[method-assign]
800 provider._free_models = MagicMock() # type: ignore[method-assign]
801
802 # Hold the lock so the loader is queued behind the unload that follows.
803 await provider._models_lock.acquire()
804 loader = asyncio.create_task(provider.ensure_models_loaded())
805 await asyncio.sleep(0)
806 # mass.unload_provider sets this before it awaits provider.unload()
807 provider.unloading = True
808 provider._models_lock.release()
809 await provider.unload()
810
811 assert await loader is False
812 assert loaded is False
813 assert provider._models_loaded is False
814