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