/
/
1"""Tests for cache controller."""
2
3import os
4import time
5from collections.abc import Callable
6from dataclasses import dataclass
7from pathlib import Path
8from typing import Any
9from unittest.mock import AsyncMock, patch
10
11import aiofiles
12import pytest
13
14from music_assistant.constants import DB_TABLE_CACHE, DB_TABLE_SETTINGS, VACUUM_MIN_RECLAIM_RATIO
15from music_assistant.controllers.cache import MAX_CACHE_DB_SIZE_MB, CacheController
16from music_assistant.controllers.cache.constants import DB_SCHEMA_VERSION, SWR_FALLBACK_MAX_AGE
17from music_assistant.helpers.database import DatabaseConnection
18from music_assistant.mass import MusicAssistant
19
20
21@dataclass
22class _FakeModel:
23 """Simple model for testing base_class reconstruction."""
24
25 name: str = ""
26 value: int = 0
27
28 @classmethod
29 def from_dict(cls, data: dict[str, Any]) -> _FakeModel:
30 """Reconstruct from dict."""
31 return cls(name=data.get("name", ""), value=data.get("value", 0))
32
33
34async def _create_db_files(cache_path: str) -> list[str]:
35 """
36 Create small cache.db, cache.db-wal, and cache.db-shm files.
37
38 :param cache_path: Path to the cache directory.
39 """
40 db_path = os.path.join(cache_path, "cache.db")
41 paths = [db_path + suffix for suffix in ("", "-wal", "-shm")]
42 for path in paths:
43 async with aiofiles.open(path, "wb") as f:
44 await f.write(b"\0")
45 return paths
46
47
48# --- Core get/set behavior ---
49
50
51async def test_set_and_get_string(cache_controller: CacheController) -> None:
52 """Test storing and retrieving a string value."""
53 await cache_controller.set("test_key", "hello", provider="test")
54 result = await cache_controller.get("test_key", provider="test")
55 assert result == "hello"
56
57
58async def test_get_expiration(cache_controller: CacheController) -> None:
59 """get_expiration returns the stored expiry epoch, also for expired rows."""
60 before = int(time.time())
61 await cache_controller.set("exp_key", {"a": 1}, provider="test", expiration=500)
62 expires = await cache_controller.get_expiration("exp_key", provider="test")
63 assert expires is not None
64 assert abs(expires - (before + 500)) <= 5
65 # a missing key yields None
66 assert await cache_controller.get_expiration("missing_key", provider="test") is None
67 # an expired-but-present row still reports its (past) expiration
68 await cache_controller.set("expired_key", "x", provider="test", expiration=-100)
69 expired = await cache_controller.get_expiration("expired_key", provider="test")
70 assert expired is not None
71 assert expired < int(time.time())
72
73
74async def test_set_and_get_int(cache_controller: CacheController) -> None:
75 """Test storing and retrieving an integer value."""
76 await cache_controller.set("num", 42, provider="test")
77 result = await cache_controller.get("num", provider="test")
78 assert result == 42
79
80
81async def test_set_and_get_float(cache_controller: CacheController) -> None:
82 """Test storing and retrieving a float value."""
83 await cache_controller.set("pi", 3.14, provider="test")
84 result = await cache_controller.get("pi", provider="test")
85 assert result == pytest.approx(3.14)
86
87
88async def test_set_and_get_bool(cache_controller: CacheController) -> None:
89 """Test storing and retrieving a boolean value."""
90 await cache_controller.set("flag", True, provider="test")
91 result = await cache_controller.get("flag", provider="test")
92 assert result is True
93
94
95async def test_set_and_get_none(cache_controller: CacheController) -> None:
96 """Test storing and retrieving None."""
97 await cache_controller.set("empty", None, provider="test")
98 result = await cache_controller.get("empty", provider="test", default="MISSING")
99 assert result is None
100
101
102async def test_set_and_get_dict(cache_controller: CacheController) -> None:
103 """Test storing and retrieving a dict value."""
104 data = {"name": "test", "count": 5, "nested": {"a": 1}}
105 await cache_controller.set("dict_key", data, provider="test")
106 result = await cache_controller.get("dict_key", provider="test")
107 assert result == data
108
109
110async def test_set_and_get_list(cache_controller: CacheController) -> None:
111 """Test storing and retrieving a list value."""
112 data = [1, "two", 3.0, None, True]
113 await cache_controller.set("list_key", data, provider="test")
114 result = await cache_controller.get("list_key", provider="test")
115 assert result == data
116
117
118async def test_set_and_get_nested_structure(cache_controller: CacheController) -> None:
119 """Test storing and retrieving a deeply nested structure."""
120 data = {"items": [{"id": 1, "tags": ["a", "b"]}, {"id": 2, "tags": []}]}
121 await cache_controller.set("nested", data, provider="test")
122 result = await cache_controller.get("nested", provider="test")
123 assert result == data
124
125
126# --- JSON serialization guarantees ---
127
128
129async def test_data_always_deserialized_from_json(cache_controller: CacheController) -> None:
130 """Test that data is always returned as JSON-deserialized (no Python objects)."""
131 await cache_controller.set("list_data", [1, 2, 3], provider="test")
132 result = await cache_controller.get("list_data", provider="test")
133 assert isinstance(result, list)
134 assert result == [1, 2, 3]
135
136
137async def test_base_class_single_dict(cache_controller: CacheController) -> None:
138 """Test that base_class reconstructs a single dict into a model."""
139 await cache_controller.set("model", {"name": "test", "value": 42}, provider="test")
140 result = await cache_controller.get("model", provider="test", base_class=_FakeModel)
141 assert isinstance(result, _FakeModel)
142 assert result.name == "test"
143 assert result.value == 42
144
145
146async def test_base_class_list_of_dicts(cache_controller: CacheController) -> None:
147 """Test that base_class reconstructs each item in a list of dicts."""
148 await cache_controller.set("models", [{"name": "a"}, {"name": "b"}], provider="test")
149 result = await cache_controller.get("models", provider="test", base_class=_FakeModel)
150 assert isinstance(result, list)
151 assert len(result) == 2
152 assert all(isinstance(item, _FakeModel) for item in result)
153 assert result[0].name == "a"
154 assert result[1].name == "b"
155
156
157async def test_base_class_not_applied_to_none(cache_controller: CacheController) -> None:
158 """Test that base_class is not applied when cache returns default."""
159 result = await cache_controller.get("nonexistent", provider="test", base_class=_FakeModel)
160 assert result is None
161
162
163async def test_non_serializable_raises(cache_controller: CacheController) -> None:
164 """Test that non-serializable data raises on set."""
165 with pytest.raises(TypeError):
166 await cache_controller.set("bad", object(), provider="test") # type: ignore[arg-type]
167
168
169# --- Expiration ---
170
171
172async def test_expired_cache_returns_default(cache_controller: CacheController) -> None:
173 """Test that expired cache entries return the default value."""
174 await cache_controller.set("expiring", "data", provider="test", expiration=-1)
175 result = await cache_controller.get("expiring", provider="test", default="gone")
176 assert result == "gone"
177
178
179# --- Checksum validation ---
180
181
182async def test_checksum_match(cache_controller: CacheController) -> None:
183 """Test that data is returned when checksum matches."""
184 await cache_controller.set("ck", "val", provider="test", checksum="abc")
185 result = await cache_controller.get("ck", provider="test", checksum="abc")
186 assert result == "val"
187
188
189async def test_checksum_mismatch(cache_controller: CacheController) -> None:
190 """Test that default is returned when checksum doesn't match."""
191 await cache_controller.set("ck2", "val", provider="test", checksum="abc")
192 result = await cache_controller.get("ck2", provider="test", checksum="xyz", default="nope")
193 assert result == "nope"
194
195
196async def test_checksum_as_int(cache_controller: CacheController) -> None:
197 """Test that integer checksums are converted to strings."""
198 await cache_controller.set("ck3", "val", provider="test", checksum="123")
199 result = await cache_controller.get("ck3", provider="test", checksum=123)
200 assert result == "val"
201
202
203# --- Category and provider isolation ---
204
205
206async def test_different_providers_isolated(cache_controller: CacheController) -> None:
207 """Test that the same key in different providers returns different data."""
208 await cache_controller.set("key", "from_a", provider="prov_a")
209 await cache_controller.set("key", "from_b", provider="prov_b")
210 assert await cache_controller.get("key", provider="prov_a") == "from_a"
211 assert await cache_controller.get("key", provider="prov_b") == "from_b"
212
213
214async def test_different_categories_isolated(cache_controller: CacheController) -> None:
215 """Test that the same key in different categories returns different data."""
216 await cache_controller.set("key", "cat_1", provider="test", category=1)
217 await cache_controller.set("key", "cat_2", provider="test", category=2)
218 assert await cache_controller.get("key", provider="test", category=1) == "cat_1"
219 assert await cache_controller.get("key", provider="test", category=2) == "cat_2"
220
221
222# --- Delete and clear ---
223
224
225async def test_delete_specific_key(cache_controller: CacheController) -> None:
226 """Test deleting a specific cache entry."""
227 await cache_controller.set("del_me", "data", provider="test")
228 await cache_controller.delete("del_me", provider="test")
229 result = await cache_controller.get("del_me", provider="test", default="gone")
230 assert result == "gone"
231
232
233async def test_clear_removes_entries(cache_controller: CacheController) -> None:
234 """Test that clear removes all non-persistent entries."""
235 await cache_controller.set("a", "1", provider="test")
236 await cache_controller.set("b", "2", provider="test")
237 await cache_controller.clear()
238 assert await cache_controller.get("a", provider="test") is None
239 assert await cache_controller.get("b", provider="test") is None
240
241
242async def test_clear_preserves_persistent(cache_controller: CacheController) -> None:
243 """Test that clear preserves persistent entries."""
244 await cache_controller.set("persist", "keep", provider="test", persistent=True)
245 await cache_controller.set("temp", "drop", provider="test")
246 await cache_controller.clear()
247 assert await cache_controller.get("persist", provider="test") == "keep"
248 assert await cache_controller.get("temp", provider="test") is None
249
250
251async def test_clear_with_provider_filter(cache_controller: CacheController) -> None:
252 """Test that clear with provider filter only removes matching entries."""
253 await cache_controller.set("k", "v1", provider="spotify")
254 await cache_controller.set("k", "v2", provider="tidal")
255 await cache_controller.clear(provider_filter="spotify")
256 assert await cache_controller.get("k", provider="spotify") is None
257 assert await cache_controller.get("k", provider="tidal") == "v2"
258
259
260# --- Bypass ---
261
262
263async def test_bypass_cache(cache_controller: CacheController) -> None:
264 """Test that the bypass context manager skips cache reads."""
265 await cache_controller.set("bypass_key", "data", provider="test")
266 async with cache_controller.handle_refresh(bypass=True):
267 result = await cache_controller.get("bypass_key", provider="test", default="bypassed")
268 assert result == "bypassed"
269 # outside bypass, cache should still work
270 result = await cache_controller.get("bypass_key", provider="test")
271 assert result == "data"
272
273
274# --- Overwrite ---
275
276
277async def test_overwrite_existing_key(cache_controller: CacheController) -> None:
278 """Test that setting the same key overwrites the previous value."""
279 await cache_controller.set("ow", "old", provider="test")
280 await cache_controller.set("ow", "new", provider="test")
281 assert await cache_controller.get("ow", provider="test") == "new"
282
283
284# --- Empty key handling ---
285
286
287async def test_get_with_empty_key_raises(cache_controller: CacheController) -> None:
288 """Test that getting with empty key raises."""
289 with pytest.raises(AssertionError):
290 await cache_controller.get("", provider="test")
291
292
293async def test_set_with_empty_key_is_noop(cache_controller: CacheController) -> None:
294 """Test that setting with empty key is silently ignored."""
295 await cache_controller.set("", "data", provider="test")
296 # should not raise, just be a no-op
297
298
299# --- Oversized cache detection ---
300
301
302async def test_cache_warns_when_exceeding_limit(
303 mass_minimal: MusicAssistant,
304 caplog: pytest.LogCaptureFixture,
305) -> None:
306 """Test that a warning is logged (and files kept) when the db exceeds the limit."""
307 cache = mass_minimal.cache
308 db_files = await _create_db_files(mass_minimal.cache_path)
309
310 with patch("asyncio.to_thread", new_callable=AsyncMock) as mock_to_thread:
311
312 async def _side_effect(func: Callable[..., Any], *args: Any) -> Any:
313 if getattr(func, "__name__", "") == "_get_db_size":
314 return float(MAX_CACHE_DB_SIZE_MB + 100)
315 return func(*args)
316
317 mock_to_thread.side_effect = _side_effect
318 await cache._check_oversized_cache()
319
320 assert "exceeds recommended maximum" in caplog.text
321 for path in db_files:
322 assert Path(path).exists()
323
324
325async def test_cache_does_not_warn_when_under_limit(
326 mass_minimal: MusicAssistant,
327 caplog: pytest.LogCaptureFixture,
328) -> None:
329 """Test that no warning is logged when the db is under the limit."""
330 cache = mass_minimal.cache
331 db_files = await _create_db_files(mass_minimal.cache_path)
332
333 with patch("asyncio.to_thread", new_callable=AsyncMock) as mock_to_thread:
334
335 async def _side_effect(func: Callable[..., Any], *args: Any) -> Any:
336 if getattr(func, "__name__", "") == "_get_db_size":
337 return 1.0
338 return func(*args)
339
340 mock_to_thread.side_effect = _side_effect
341 await cache._check_oversized_cache()
342
343 assert "exceeds recommended maximum" not in caplog.text
344 for path in db_files:
345 assert Path(path).exists()
346
347
348async def test_all_three_db_files_included_in_size(
349 mass_minimal: MusicAssistant,
350 caplog: pytest.LogCaptureFixture,
351) -> None:
352 """Test that cache.db, cache.db-wal, and cache.db-shm are all summed for size check."""
353 cache = mass_minimal.cache
354 db_path = os.path.join(mass_minimal.cache_path, "cache.db")
355
356 for suffix in ("", "-wal", "-shm"):
357 async with aiofiles.open(db_path + suffix, "wb") as f:
358 await f.write(b"\0" * 100)
359
360 size_threshold_mb = 0.0002
361 with patch(
362 "music_assistant.controllers.cache.controller.MAX_CACHE_DB_SIZE_MB", size_threshold_mb
363 ):
364 await cache._check_oversized_cache()
365
366 assert "exceeds recommended maximum" in caplog.text
367 for suffix in ("", "-wal", "-shm"):
368 assert Path(db_path + suffix).exists()
369
370
371# --- allow_expired_cache flag ---
372
373
374async def test_get_with_allow_expired_cache_returns_expired_data(
375 cache_controller: CacheController,
376) -> None:
377 """Test that get(allow_expired_cache=True) returns data past its expiration."""
378 await cache_controller.set("stale", "old_data", provider="test", expiration=-1)
379 assert await cache_controller.get("stale", provider="test", default="gone") == "gone"
380 assert (
381 await cache_controller.get(
382 "stale", provider="test", default="gone", allow_expired_cache=True
383 )
384 == "old_data"
385 )
386
387
388async def test_get_with_allow_expired_cache_still_returns_default_when_missing(
389 cache_controller: CacheController,
390) -> None:
391 """Test that allow_expired_cache=True still returns default when nothing is cached."""
392 result = await cache_controller.get(
393 "nonexistent", provider="test", default="gone", allow_expired_cache=True
394 )
395 assert result == "gone"
396
397
398async def test_auto_cleanup_removes_expired_entries(cache_controller: CacheController) -> None:
399 """Test that auto_cleanup removes expired entries by default."""
400 await cache_controller.set("evict", "data", provider="test", expiration=-1)
401 await cache_controller.auto_cleanup()
402 result = await cache_controller.get(
403 "evict", provider="test", default="gone", allow_expired_cache=True
404 )
405 assert result == "gone"
406
407
408async def test_auto_cleanup_keeps_allow_expired_cache_entries(
409 cache_controller: CacheController,
410) -> None:
411 """Test that auto_cleanup keeps expired entries with allow_expired_cache=True."""
412 await cache_controller.set(
413 "keep", "data", provider="test", expiration=-1, allow_expired_cache=True
414 )
415 await cache_controller.auto_cleanup()
416 result = await cache_controller.get("keep", provider="test", allow_expired_cache=True)
417 assert result == "data"
418
419
420async def test_auto_cleanup_keeps_fresh_entries(cache_controller: CacheController) -> None:
421 """Test that auto_cleanup keeps fresh entries regardless of the flag."""
422 await cache_controller.set("alive", "data", provider="test", expiration=3600)
423 await cache_controller.auto_cleanup()
424 assert await cache_controller.get("alive", provider="test") == "data"
425
426
427async def test_auto_cleanup_scans_all_records(cache_controller: CacheController) -> None:
428 """
429 Test that auto_cleanup removes expired entries beyond the row-fetch page size.
430
431 Regression test: cleanup previously fetched rows via the default 500-row limit,
432 so large caches kept most of their expired entries forever.
433 """
434 expired_count = 1200
435 for i in range(expired_count):
436 await cache_controller.set(f"expired_{i}", "data", provider="test", expiration=-1)
437 await cache_controller.set("fresh", "data", provider="test", expiration=3600)
438
439 await cache_controller.auto_cleanup()
440
441 assert cache_controller.database is not None
442 assert await cache_controller.database.get_count(DB_TABLE_CACHE) == 1
443 assert await cache_controller.get("fresh", provider="test") == "data"
444
445
446# --- Startup vacuum ---
447
448
449async def test_setup_skips_vacuum_when_little_reclaimable(
450 mass_minimal: MusicAssistant,
451) -> None:
452 """Test that the startup vacuum is skipped when little space can be reclaimed."""
453 cache = mass_minimal.cache
454 with (
455 patch.object(
456 DatabaseConnection,
457 "get_reclaimable_ratio",
458 AsyncMock(return_value=VACUUM_MIN_RECLAIM_RATIO / 2),
459 ),
460 patch.object(DatabaseConnection, "vacuum", AsyncMock()) as mock_vacuum,
461 ):
462 await cache._setup_database()
463 mock_vacuum.assert_not_called()
464
465
466async def test_setup_runs_vacuum_when_reclaimable(
467 mass_minimal: MusicAssistant,
468) -> None:
469 """Test that the startup vacuum runs when enough space can be reclaimed."""
470 cache = mass_minimal.cache
471 with (
472 patch.object(
473 DatabaseConnection,
474 "get_reclaimable_ratio",
475 AsyncMock(return_value=VACUUM_MIN_RECLAIM_RATIO + 0.1),
476 ),
477 patch.object(DatabaseConnection, "vacuum", AsyncMock()) as mock_vacuum,
478 ):
479 await cache._setup_database()
480 mock_vacuum.assert_awaited_once_with()
481
482
483# --- upsert (in-place write) ---
484
485
486async def test_set_upserts_in_place(cache_controller: CacheController) -> None:
487 """Test that overwriting a key updates the row in place instead of replacing it."""
488 assert cache_controller.database is not None
489 await cache_controller.set("k", "v1", provider="test")
490 row = await cache_controller.database.get_row(
491 DB_TABLE_CACHE, {"category": 0, "provider": "test", "key": "k"}
492 )
493 assert row is not None
494 row_id = row["id"]
495
496 await cache_controller.set("k", "v2", provider="test")
497 assert await cache_controller.get("k", provider="test") == "v2"
498 # a single row that kept its id â an INSERT OR REPLACE would delete it and assign a new id
499 assert await cache_controller.database.get_count(DB_TABLE_CACHE) == 1
500 row = await cache_controller.database.get_row(
501 DB_TABLE_CACHE, {"category": 0, "provider": "test", "key": "k"}
502 )
503 assert row is not None
504 assert row["id"] == row_id
505
506
507# --- secondary indexes ---
508
509
510async def _index_names(cache_controller: CacheController) -> set[str]:
511 """Return the names of all indexes on the cache table."""
512 assert cache_controller.database is not None
513 rows = await cache_controller.database.get_rows_from_query(
514 "SELECT name FROM sqlite_master WHERE type = 'index' AND tbl_name = :table",
515 {"table": DB_TABLE_CACHE},
516 )
517 return {str(row["name"]) for row in rows}
518
519
520async def test_only_key_provider_index_is_created(cache_controller: CacheController) -> None:
521 """Test that only the (key, provider) secondary index is created besides the autoindex."""
522 names = await _index_names(cache_controller)
523 assert f"{DB_TABLE_CACHE}_key_provider_idx" in names
524 # the UNIQUE(category, key, provider) constraint provides an autoindex
525 assert any(name.startswith("sqlite_autoindex") for name in names)
526 # the redundant indexes are not (re)created
527 for removed in (
528 "category_idx",
529 "key_idx",
530 "provider_idx",
531 "category_key_idx",
532 "category_provider_idx",
533 "category_key_provider_idx",
534 ):
535 assert f"{DB_TABLE_CACHE}_{removed}" not in names
536
537
538async def test_migration_drops_redundant_indexes(mass_minimal: MusicAssistant) -> None:
539 """Test that opening a pre-v9 database migrates cleanly and drops the redundant indexes."""
540 db_path = os.path.join(mass_minimal.cache_path, "cache.db")
541 redundant = (
542 ("category_idx", "category"),
543 ("key_idx", "key"),
544 ("provider_idx", "provider"),
545 ("category_key_idx", "category,key"),
546 ("category_provider_idx", "category,provider"),
547 ("category_key_provider_idx", "category,key,provider"),
548 ("key_provider_idx", "key,provider"),
549 )
550 # build a v8 database with the old (full) index set and one row
551 old_db = DatabaseConnection(db_path)
552 await old_db.setup()
553 await old_db.execute(
554 f"CREATE TABLE {DB_TABLE_SETTINGS}(key TEXT PRIMARY KEY, value TEXT, type TEXT)"
555 )
556 await old_db.execute(
557 f"""CREATE TABLE {DB_TABLE_CACHE}(
558 [id] INTEGER PRIMARY KEY AUTOINCREMENT,
559 [category] INTEGER NOT NULL DEFAULT 0,
560 [key] TEXT NOT NULL,
561 [provider] TEXT NOT NULL,
562 [expires] INTEGER NOT NULL,
563 [data] TEXT NULL,
564 [checksum] TEXT NULL,
565 [persistent] INTEGER NOT NULL DEFAULT 0,
566 [allow_expired_cache] INTEGER NOT NULL DEFAULT 0,
567 UNIQUE(category, key, provider)
568 )"""
569 )
570 for suffix, columns in redundant:
571 await old_db.execute(
572 f"CREATE INDEX {DB_TABLE_CACHE}_{suffix} ON {DB_TABLE_CACHE}({columns})"
573 )
574 await old_db.execute(
575 f"INSERT INTO {DB_TABLE_SETTINGS}(key, value, type) VALUES ('version', '8', 'str')"
576 )
577 await old_db.execute(
578 f"INSERT INTO {DB_TABLE_CACHE}(category, key, provider, expires, data) "
579 "VALUES (0, 'kept', 'test', 9999999999, '\"payload\"')"
580 )
581 await old_db.commit()
582 await old_db.close()
583
584 # opening the controller runs the migration
585 await mass_minimal.cache._setup_database()
586 cache = mass_minimal.cache
587 assert cache.database is not None
588
589 version_row = await cache.database.get_row(DB_TABLE_SETTINGS, {"key": "version"})
590 assert version_row is not None
591 assert version_row["value"] == str(DB_SCHEMA_VERSION)
592
593 names = await _index_names(cache)
594 assert f"{DB_TABLE_CACHE}_key_provider_idx" in names
595 for suffix, _ in redundant[:-1]: # every index except key_provider is dropped
596 assert f"{DB_TABLE_CACHE}_{suffix}" not in names
597
598 # existing data survived the migration
599 assert await cache.get("kept", provider="test") == "payload"
600
601
602# --- stale-while-revalidate cleanup ---
603
604
605async def test_auto_cleanup_removes_stale_swr_rows(cache_controller: CacheController) -> None:
606 """Test that auto_cleanup removes SWR fallback rows expired beyond the grace window."""
607 # expired but within the grace window -> kept as fallback
608 await cache_controller.set(
609 "recent", "data", provider="test", expiration=-1, allow_expired_cache=True
610 )
611 # expired well beyond the grace window -> removed
612 await cache_controller.set(
613 "ancient",
614 "data",
615 provider="test",
616 expiration=-(SWR_FALLBACK_MAX_AGE + 86400),
617 allow_expired_cache=True,
618 )
619 await cache_controller.auto_cleanup()
620
621 assert await cache_controller.get("recent", provider="test", allow_expired_cache=True) == "data"
622 assert (
623 await cache_controller.get(
624 "ancient", provider="test", default="gone", allow_expired_cache=True
625 )
626 == "gone"
627 )
628