from collections.abc import Awaitable, Callable from typing import TypeVar from pydantic import BaseModel from redis.asyncio import Redis ModelT = TypeVar("ModelT", bound=BaseModel) NAMESPACE = "{{ project_slug }}" class RedisCache: def __init__( self, redis: Redis, *, namespace: str = NAMESPACE, ttl: int = 300 ) -> None: self._redis = redis self._namespace = namespace self._ttl = ttl def _key(self, key: str) -> str: return f"{self._namespace}:{key}" async def get_model(self, key: str, model: type[ModelT]) -> ModelT | None: raw = await self._redis.get(self._key(key)) if raw is None: return None return model.model_validate_json(raw) async def set_model( self, key: str, value: BaseModel, *, ttl: int | None = None ) -> None: await self._redis.set( self._key(key), value.model_dump_json(), ex=ttl or self._ttl ) async def get_or_set_model( self, key: str, model: type[ModelT], loader: Callable[[], Awaitable[ModelT]], *, ttl: int | None = None, ) -> ModelT: cached = await self.get_model(key, model) if cached is not None: return cached value = await loader() await self.set_model(key, value, ttl=ttl) return value async def invalidate(self, *keys: str) -> None: if keys: await self._redis.delete(*(self._key(key) for key in keys))