/
/
1"""Cache controller implementation."""
2
3from __future__ import annotations
4
5import asyncio
6import os
7import time
8from collections.abc import AsyncGenerator
9from contextlib import asynccontextmanager
10from pathlib import Path
11from typing import TYPE_CHECKING, Any
12
13from music_assistant_models.background_task import TaskSchedule
14from music_assistant_models.config_entries import ConfigActionResult, ConfigEntry
15from music_assistant_models.enums import ConfigEntryType
16
17from music_assistant.constants import (
18 DB_TABLE_CACHE,
19 DB_TABLE_SETTINGS,
20 VACUUM_MIN_RECLAIM_RATIO,
21)
22from music_assistant.controllers.cache.constants import (
23 BYPASS_CACHE,
24 CACHE_DATABASE_CLEANUP_TASK_ID,
25 CONF_CLEAR_CACHE,
26 DB_SCHEMA_VERSION,
27 DEFAULT_CACHE_EXPIRATION,
28 LOGGER,
29 MAX_CACHE_DB_SIZE_MB,
30 SWR_FALLBACK_MAX_AGE,
31)
32from music_assistant.controllers.tasks.context import (
33 update_current_task_progress_text,
34)
35from music_assistant.helpers.database import DatabaseConnection
36from music_assistant.helpers.datetime import local_clock_time_to_utc
37from music_assistant.helpers.json import SerializableType, async_json_loads, json_dumps
38from music_assistant.models.core_controller import CoreController
39
40if TYPE_CHECKING:
41 from music_assistant_models.config_entries import CoreConfig
42
43 from music_assistant import MusicAssistant
44
45
46class CacheController(CoreController):
47 """Controller handling caching of data throughout the application."""
48
49 domain: str = "cache"
50
51 def __init__(self, mass: MusicAssistant) -> None:
52 """Initialize core controller."""
53 super().__init__(mass)
54 self.database: DatabaseConnection | None = None
55 self.manifest.name = "Cache controller"
56 self.manifest.description = (
57 "Music Assistant's core controller for caching data throughout the application."
58 )
59 self.manifest.icon = "memory"
60
61 async def get_config_entries(self) -> tuple[ConfigEntry, ...]:
62 """Return all Config Entries for this core module (if any)."""
63 return (
64 ConfigEntry(
65 key=CONF_CLEAR_CACHE,
66 type=ConfigEntryType.ACTION,
67 ),
68 )
69
70 async def handle_config_action(
71 self, action: str
72 ) -> tuple[ConfigEntry, ...] | ConfigActionResult | None:
73 """Handle a one-shot action button press and report its outcome."""
74 if action == CONF_CLEAR_CACHE:
75 await self.clear()
76 return ConfigActionResult(translation_key=f"{CONF_CLEAR_CACHE}.result")
77 return await super().handle_config_action(action)
78
79 async def setup(self, config: CoreConfig) -> None:
80 """Async initialize of cache module."""
81 self.logger.info("Initializing cache controller...")
82 await self._setup_database()
83
84 async def post_setup(self) -> None:
85 """Handle logic after all core controllers have been set up."""
86 self._register_cleanup_task()
87
88 async def close(self) -> None:
89 """Cleanup on exit."""
90 if self.database:
91 await self.database.close()
92
93 async def get_diagnostics(self) -> dict[str, SerializableType]:
94 """Return diagnostics info for this controller to include in diagnostics reports."""
95 return {
96 "db_schema_version": DB_SCHEMA_VERSION,
97 "db_size_mb": round(await self._get_cache_db_size_mb(), 1),
98 "entries": await self.database.get_count(DB_TABLE_CACHE) if self.database else None,
99 }
100
101 async def get(
102 self,
103 key: str,
104 provider: str = "default",
105 category: int = 0,
106 checksum: str | int | None = None,
107 default: Any = None,
108 allow_bypass: bool | None = None,
109 base_class: Any = None,
110 allow_expired_cache: bool = False,
111 ) -> Any:
112 """
113 Get data from cache.
114
115 Returns JSON-deserialized data (dicts, lists, strings, numbers, booleans, None).
116
117 If base_class is provided, the raw data is automatically reconstructed using
118 its from_dict() method. If the cached data is a list of dicts, each item is
119 reconstructed individually.
120
121 :param key: The (unique) lookup key of the cache object.
122 :param provider: Provider id to group cache objects.
123 :param category: Category to group cache objects.
124 :param checksum: If provided, only return data if the stored checksum matches.
125 :param default: Value to return if no cache object is found.
126 :param allow_bypass: Whether to respect the BYPASS_CACHE context variable.
127 :param base_class: If provided, reconstruct data using base_class.from_dict().
128 :param allow_expired_cache: If True, also return entries past their expiration
129 time instead of treating them as cache misses.
130 """
131 data, _, found = await self.get_with_freshness(
132 key,
133 provider=provider,
134 category=category,
135 checksum=checksum,
136 allow_bypass=allow_bypass,
137 base_class=base_class,
138 include_expired=allow_expired_cache,
139 )
140 return data if found else default
141
142 async def get_with_freshness(
143 self,
144 key: str,
145 provider: str = "default",
146 category: int = 0,
147 checksum: str | int | None = None,
148 allow_bypass: bool | None = None,
149 base_class: Any = None,
150 include_expired: bool = False,
151 ) -> tuple[Any, bool, bool]:
152 """
153 Get data from cache together with the freshness and presence of the entry.
154
155 Returns a (data, is_fresh, found) tuple. found is False when there is no usable
156 entry, in which case data is None; is_fresh is False when the entry is expired.
157 Because a stored None value is returned as-is, use the found flag to tell a cache
158 miss from a cached None.
159
160 :param key: The (unique) lookup key of the cache object.
161 :param provider: Provider id to group cache objects.
162 :param category: Category to group cache objects.
163 :param checksum: If provided, only return data if the stored checksum matches.
164 :param allow_bypass: Whether to respect the BYPASS_CACHE context variable.
165 :param base_class: If provided, reconstruct data using base_class.from_dict().
166 :param include_expired: If False (default), an expired entry is reported as not found
167 and is not deserialized; set True to also return expired entries as stale data.
168 """
169 assert self.database is not None
170 assert key, "No key provided"
171 if allow_bypass and BYPASS_CACHE.get():
172 return None, False, False
173 cur_time = int(time.time())
174 if checksum is not None and not isinstance(checksum, str):
175 checksum = str(checksum)
176 if (
177 db_row := await self.database.get_row(
178 DB_TABLE_CACHE, {"category": category, "provider": provider, "key": key}
179 )
180 ) and (not checksum or db_row["checksum"] == checksum):
181 # if allow_bypass is not explicitly set,
182 # determine it based on the 'persistent' flag of the cache entry
183 if allow_bypass is None:
184 allow_bypass = not bool(db_row["persistent"])
185 if allow_bypass and BYPASS_CACHE.get():
186 return None, False, False
187 is_fresh = bool(db_row["expires"] >= cur_time)
188 # skip deserialization for an expired entry the caller will not use
189 if not is_fresh and not include_expired:
190 return None, False, False
191 try:
192 data = await async_json_loads(db_row["data"])
193 except Exception as exc:
194 LOGGER.error(
195 "Error parsing cache data for %s/%s/%s: %s",
196 provider,
197 category,
198 key,
199 str(exc),
200 exc_info=exc if self.logger.isEnabledFor(10) else None,
201 )
202 else:
203 if base_class is not None and data is not None:
204 if isinstance(data, list):
205 return [base_class.from_dict(item) for item in data], is_fresh, True
206 return base_class.from_dict(data), is_fresh, True
207 return data, is_fresh, True
208 return None, False, False
209
210 async def get_expiration(
211 self,
212 key: str,
213 provider: str = "default",
214 category: int = 0,
215 ) -> int | None:
216 """
217 Return the expiration timestamp (epoch seconds) of a cache entry, if any.
218
219 Cheap existence/freshness probe: only the expiration column is read, the
220 stored data is not. Returns None when no entry exists for the given key.
221
222 :param key: The (unique) lookup key of the cache object.
223 :param provider: Provider id to group cache objects.
224 :param category: Category to group cache objects.
225 """
226 assert self.database is not None
227 assert key, "No key provided"
228 rows = await self.database.get_rows_from_query(
229 f"SELECT expires FROM {DB_TABLE_CACHE} "
230 "WHERE category = :category AND provider = :provider AND key = :key",
231 {"category": category, "provider": provider, "key": key},
232 limit=1,
233 )
234 return int(rows[0]["expires"]) if rows else None
235
236 async def set(
237 self,
238 key: str,
239 data: SerializableType,
240 expiration: int = DEFAULT_CACHE_EXPIRATION,
241 provider: str = "default",
242 category: int = 0,
243 checksum: str | None = None,
244 persistent: bool = False,
245 allow_expired_cache: bool = False,
246 ) -> None:
247 """
248 Store data in cache.
249
250 Data must be JSON-serializable (str, int, float, bool, None, list, dict).
251 Do not pass model objects directly — use .to_dict() first.
252 Non-serializable data will raise TypeError.
253
254 :param key: The (unique) lookup key of the cache object.
255 :param data: JSON-serializable data to store.
256 :param expiration: Time in seconds the cache object should be valid.
257 :param provider: Provider id to group cache objects.
258 :param category: Category to group cache objects.
259 :param checksum: Optional checksum to store with the cache object.
260 :param persistent: If True, the entry survives cache clears.
261 :param allow_expired_cache: If True, the entry survives the auto-cleanup task
262 after it expires, so it can still be served as fallback data by the
263 stale-while-revalidate path of `@use_cache`.
264 """
265 assert self.database is not None
266 if not key:
267 return
268 if checksum is not None:
269 checksum = str(checksum)
270 expires = int(time.time() + expiration)
271 # always serialize to JSON to ensure data is serializable
272 # this raises if the data contains non-serializable objects
273 data = await asyncio.to_thread(json_dumps, data)
274 # upsert (update in place on the UNIQUE(category, key, provider) conflict) instead of
275 # INSERT OR REPLACE, which deletes and re-inserts the row and so rewrites every index
276 await self.database.upsert(
277 DB_TABLE_CACHE,
278 {
279 "category": category,
280 "provider": provider,
281 "key": key,
282 "expires": expires,
283 "checksum": checksum,
284 "data": data,
285 "persistent": persistent,
286 "allow_expired_cache": allow_expired_cache,
287 },
288 )
289
290 async def delete(
291 self, key: str | None, category: int | None = None, provider: str | None = None
292 ) -> None:
293 """Delete data from cache."""
294 assert self.database is not None
295 match: dict[str, str | int] = {}
296 if key is not None:
297 match["key"] = key
298 if category is not None:
299 match["category"] = category
300 if provider is not None:
301 match["provider"] = provider
302 await self.database.delete(DB_TABLE_CACHE, match)
303
304 async def clear(
305 self,
306 key_filter: str | None = None,
307 category_filter: int | None = None,
308 provider_filter: str | None = None,
309 include_persistent: bool = False,
310 ) -> None:
311 """Clear all/partial items from cache."""
312 assert self.database is not None
313 self.logger.info("Clearing database...")
314 query_parts: list[str] = []
315 if category_filter is not None:
316 query_parts.append(f"category = {category_filter}")
317 if provider_filter is not None:
318 query_parts.append(f"provider LIKE '%{provider_filter}%'")
319 if key_filter is not None:
320 query_parts.append(f"key LIKE '%{key_filter}%'")
321 if not include_persistent:
322 query_parts.append("persistent = 0")
323 query = "WHERE " + " AND ".join(query_parts) if query_parts else None
324 await self.database.delete(DB_TABLE_CACHE, query=query)
325 self.logger.info("Clearing database DONE")
326
327 async def auto_cleanup(self) -> None:
328 """Run scheduled auto cleanup task."""
329 assert self.database is not None
330 self.logger.debug("Running automatic cleanup...")
331 update_current_task_progress_text("Removing expired cache records")
332 cur_timestamp = int(time.time())
333 # remove expired entries; allow_expired_cache entries are kept as stale-while-revalidate
334 # fallback, but only until they are expired beyond SWR_FALLBACK_MAX_AGE - past that their
335 # key is clearly no longer requested and the row would otherwise live forever
336 swr_cutoff = cur_timestamp - SWR_FALLBACK_MAX_AGE
337 cursor = await self.database.execute(
338 f"DELETE FROM {DB_TABLE_CACHE} WHERE "
339 "(expires < :timestamp AND allow_expired_cache = 0) "
340 "OR (expires < :swr_cutoff AND allow_expired_cache = 1)",
341 {"timestamp": cur_timestamp, "swr_cutoff": swr_cutoff},
342 )
343 await self.database.commit()
344 cleaned_records = cursor.rowcount
345 update_current_task_progress_text(f"Cleaned up {cleaned_records} expired cache record(s)")
346 self.logger.debug("Automatic cleanup finished (cleaned up %s records)", cleaned_records)
347
348 @asynccontextmanager
349 async def handle_refresh(self, bypass: bool) -> AsyncGenerator[None]:
350 """Handle the cache bypass."""
351 try:
352 token = BYPASS_CACHE.set(bypass)
353 yield None
354 finally:
355 BYPASS_CACHE.reset(token)
356
357 async def _check_oversized_cache(self) -> None:
358 """Warn if the cache database exceeds the recommended max size."""
359 db_size_mb = await self._get_cache_db_size_mb()
360 if db_size_mb > MAX_CACHE_DB_SIZE_MB:
361 self.logger.warning(
362 "Cache database size %.2f MB exceeds recommended maximum of %d MB",
363 db_size_mb,
364 MAX_CACHE_DB_SIZE_MB,
365 )
366
367 async def _get_cache_db_size_mb(self) -> float:
368 """Return the on-disk size of the cache database (in MB)."""
369 db_path = os.path.join(self.mass.cache_path, "cache.db")
370 # also include the write ahead log and shared memory db files
371 db_files = [db_path + suffix for suffix in ("", "-wal", "-shm")]
372
373 def _get_db_size() -> float:
374 total = 0
375 for path in db_files:
376 if Path(path).exists():
377 total += Path(path).stat().st_size
378 return total / (1024 * 1024)
379
380 return await asyncio.to_thread(_get_db_size)
381
382 async def _setup_database(self) -> None:
383 """Initialize database."""
384 await self._check_oversized_cache()
385 db_path = os.path.join(self.mass.cache_path, "cache.db")
386 self.database = DatabaseConnection(db_path)
387 await self.database.setup()
388
389 # always create db tables if they don't exist to prevent errors trying to access them later
390 await self.__create_database_tables()
391
392 try:
393 if db_row := await self.database.get_row(DB_TABLE_SETTINGS, {"key": "version"}):
394 prev_version = int(db_row["value"])
395 else:
396 prev_version = 0
397 except KeyError, ValueError:
398 prev_version = 0
399
400 if prev_version not in (0, DB_SCHEMA_VERSION):
401 LOGGER.warning(
402 "Performing database migration from %s to %s",
403 prev_version,
404 DB_SCHEMA_VERSION,
405 )
406 try:
407 await self.__migrate_database(prev_version)
408 except Exception as err:
409 LOGGER.warning("Cache database migration failed: %s, resetting cache", err)
410 await self.database.execute(f"DROP TABLE IF EXISTS {DB_TABLE_CACHE}")
411 await self.__create_database_tables()
412
413 # store current schema version
414 await self.database.insert_or_replace(
415 DB_TABLE_SETTINGS,
416 {"key": "version", "value": str(DB_SCHEMA_VERSION), "type": "str"},
417 )
418 await self.__create_database_indexes()
419
420 # Skip the full rebuild unless a meaningful share of the file can be reclaimed.
421 try:
422 reclaimable_ratio = await self.database.get_reclaimable_ratio()
423 if reclaimable_ratio < VACUUM_MIN_RECLAIM_RATIO:
424 self.logger.debug(
425 "Skipping database compaction (only %.1f%% reclaimable)",
426 reclaimable_ratio * 100,
427 )
428 else:
429 self.logger.debug(
430 "Compacting database (%.1f%% reclaimable)...", reclaimable_ratio * 100
431 )
432 await self.database.vacuum()
433 self.logger.debug("Compacting database done")
434 except Exception as err:
435 self.logger.warning("Database vacuum failed: %s", str(err))
436
437 async def __create_database_tables(self) -> None:
438 """Create database table(s)."""
439 assert self.database is not None
440 await self.database.execute(
441 f"""CREATE TABLE IF NOT EXISTS {DB_TABLE_SETTINGS}(
442 key TEXT PRIMARY KEY,
443 value TEXT,
444 type TEXT
445 );"""
446 )
447 await self.database.execute(
448 f"""CREATE TABLE IF NOT EXISTS {DB_TABLE_CACHE}(
449 [id] INTEGER PRIMARY KEY AUTOINCREMENT,
450 [category] INTEGER NOT NULL DEFAULT 0,
451 [key] TEXT NOT NULL,
452 [provider] TEXT NOT NULL,
453 [expires] INTEGER NOT NULL,
454 [data] TEXT NULL,
455 [checksum] TEXT NULL,
456 [persistent] INTEGER NOT NULL DEFAULT 0,
457 [allow_expired_cache] INTEGER NOT NULL DEFAULT 0,
458 UNIQUE(category, key, provider)
459 )"""
460 )
461
462 await self.database.commit()
463
464 async def __create_database_indexes(self) -> None:
465 """Create database indexes."""
466 assert self.database is not None
467 # The UNIQUE(category, key, provider) constraint already provides an index that serves
468 # every point lookup (get() matches exactly those three columns) and any delete that
469 # includes the category. The only access pattern its column order cannot serve is a
470 # delete that filters by (key, provider) without a category, so that is the single
471 # secondary index kept here.
472 await self.database.execute(
473 f"CREATE INDEX IF NOT EXISTS {DB_TABLE_CACHE}_key_provider_idx "
474 f"ON {DB_TABLE_CACHE}(key,provider);"
475 )
476 await self.database.commit()
477
478 async def __migrate_database(self, prev_version: int) -> None:
479 """Perform a database migration."""
480 assert self.database is not None
481 if prev_version <= 6:
482 # clear spotify cache entries to fix bloated cache from playlist pagination bug
483 await self.database.delete(DB_TABLE_CACHE, query="WHERE provider LIKE '%spotify%'")
484 if prev_version <= 7:
485 await self.database.execute(
486 f"ALTER TABLE {DB_TABLE_CACHE} "
487 "ADD COLUMN allow_expired_cache INTEGER NOT NULL DEFAULT 0"
488 )
489 if prev_version <= 8:
490 # drop the redundant secondary indexes: they either duplicate the
491 # UNIQUE(category, key, provider) autoindex or are a left-prefix of it, so the
492 # autoindex already serves their lookups. The (key, provider) index is (re)created
493 # by __create_database_indexes and intentionally kept.
494 for index_name in (
495 "category_idx",
496 "key_idx",
497 "provider_idx",
498 "category_key_idx",
499 "category_provider_idx",
500 "category_key_provider_idx",
501 ):
502 await self.database.execute(f"DROP INDEX IF EXISTS {DB_TABLE_CACHE}_{index_name}")
503 await self.database.commit()
504
505 def _register_cleanup_task(self) -> None:
506 """Register the recurring cache database cleanup task."""
507 utc_hour, utc_minute = local_clock_time_to_utc(4, 0)
508 desired_schedule = TaskSchedule.daily(hour=utc_hour, minute=utc_minute)
509 self.mass.tasks.register_scheduled_task(
510 task_id=CACHE_DATABASE_CLEANUP_TASK_ID,
511 name="Cache database cleanup",
512 handler=self.auto_cleanup,
513 schedule=desired_schedule,
514 translation_key="cache_database_cleanup",
515 translation_owner=self.translation_owner,
516 metadata={"task_domain": "cache_database_cleanup"},
517 allow_retry=True,
518 )
519