Files
video-gen/video-gen-api/app/services/system_config_cache.py
T
2026-07-24 09:18:05 +08:00

91 lines
2.6 KiB
Python

from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Iterable
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.system_config import SystemConfig
from app.utils.redis import get_redis
_CACHE_VERSION_KEY = "system_config_cache_version"
_DEFAULT_TTL_SECONDS = 60
@dataclass
class _CacheState:
values: dict[str, str | None]
expires_at: float
version: str | None
_cache = _CacheState(values={}, expires_at=0.0, version=None)
async def _get_remote_version() -> str | None:
redis = get_redis()
if not redis:
return None
try:
value = await redis.get(_CACHE_VERSION_KEY)
return str(value or "0")
except Exception:
return None
async def invalidate_system_config_cache(keys: Iterable[str] | None = None) -> None:
"""Invalidate local cache and notify other workers through a Redis version bump when available."""
key_set = set(keys or [])
if key_set:
for key in key_set:
_cache.values.pop(key, None)
else:
_cache.values.clear()
_cache.expires_at = 0.0
redis = get_redis()
if redis:
try:
await redis.incr(_CACHE_VERSION_KEY)
except Exception:
pass
async def get_system_config_values(
db: AsyncSession,
keys: Iterable[str],
*,
ttl_seconds: int = _DEFAULT_TTL_SECONDS,
) -> dict[str, str | None]:
key_list = [str(key) for key in keys if str(key)]
if not key_list:
return {}
now = time.monotonic()
remote_version = await _get_remote_version()
if remote_version is not None and remote_version != _cache.version:
_cache.values.clear()
_cache.expires_at = 0.0
_cache.version = remote_version
missing = [key for key in key_list if key not in _cache.values]
if now >= _cache.expires_at:
missing = key_list
if missing:
result = await db.execute(select(SystemConfig).where(SystemConfig.key.in_(missing)))
rows = {row.key: row.value for row in result.scalars().all()}
for key in missing:
_cache.values[key] = rows.get(key)
_cache.expires_at = now + max(1, int(ttl_seconds or _DEFAULT_TTL_SECONDS))
if remote_version is not None:
_cache.version = remote_version
return {key: _cache.values.get(key) for key in key_list}
async def get_system_config_value(db: AsyncSession, key: str, *, ttl_seconds: int = _DEFAULT_TTL_SECONDS) -> str | None:
return (await get_system_config_values(db, [key], ttl_seconds=ttl_seconds)).get(key)