You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

139 lines
4.0 KiB
Python

"""
Redis-backed response caching utilities.
Design goals:
- Async Redis client with graceful degradation when Redis is unavailable
- Stable cache key builder
- JSON storage with ETag computation
- Simple helpers for endpoints to get/set cached responses
"""
from __future__ import annotations
import logging
from typing import Any, Dict, Optional, Tuple
import hashlib
import asyncio
import orjson
logger = logging.getLogger(__name__)
try:
from redis import asyncio as aioredis
except Exception: # pragma: no cover - library import safety
aioredis = None # type: ignore
from app.core.config import settings
_redis_client: Optional["aioredis.Redis"] = None
_redis_lock = asyncio.Lock()
def _serialize(obj: Any) -> bytes:
"""Serialize Python object to JSON bytes using orjson."""
return orjson.dumps(obj)
def _deserialize(data: Optional[bytes]) -> Optional[Dict[str, Any]]:
if not data:
return None
try:
return orjson.loads(data)
except Exception as e:
logger.debug("Cache deserialization failed: %s", e)
return None
def compute_etag(payload_bytes: bytes) -> str:
"""Compute strong ETag for given payload bytes."""
return hashlib.sha256(payload_bytes).hexdigest()
async def get_redis() -> Optional["aioredis.Redis"]:
"""Get a shared async Redis client. Returns None if Redis unavailable."""
global _redis_client
if aioredis is None:
return None
if _redis_client is not None:
return _redis_client
async with _redis_lock:
if _redis_client is not None:
return _redis_client
try:
_redis_client = aioredis.from_url(settings.REDIS_URL, encoding="utf-8", decode_responses=False)
# Light-touch ping to verify connectivity (do not raise)
try:
await _redis_client.ping()
except Exception as e:
logger.debug("Redis ping failed: %s", e)
pass
return _redis_client
except Exception as e:
logger.debug("Redis connection failed: %s", e)
return None
def build_cache_key(namespace: str, *parts: Any) -> str:
"""Build a stable cache key using namespace and parts.
Each part is converted to string and stripped. Empty parts are skipped.
"""
key_parts = [namespace]
for p in parts:
if p is None:
continue
s = str(p).strip()
if not s:
continue
key_parts.append(s)
return ":".join(key_parts)
async def get_cached_response(key: str) -> Optional[Tuple[Dict[str, Any], str]]:
"""Get cached response body and its ETag. Returns None if missing or on error.
The cached value is stored as JSON with shape: {"etag": str, "body": {...}}.
"""
client = await get_redis()
if client is None:
return None
try:
raw = await client.get(key)
data = _deserialize(raw)
if not data or "body" not in data:
return None
etag = data.get("etag")
# If etag missing, compute it from body
if not etag:
etag = compute_etag(_serialize(data["body"]))
return data["body"], etag
except Exception as e:
logger.debug("Cache get failed for key: %s", e)
return None
async def set_cached_response(key: str, body: Dict[str, Any], ttl_seconds: Optional[int] = None) -> str:
"""Cache response body with ETag. Returns the computed ETag.
If Redis is unavailable, this function is a no-op and returns the ETag anyway.
"""
payload_bytes = _serialize(body)
etag = compute_etag(payload_bytes)
record = {"etag": etag, "body": body}
client = await get_redis()
if client is None:
return etag
try:
if ttl_seconds is None:
ttl_seconds = max(60, int(getattr(settings, "CACHE_TTL", 3600)))
await client.set(key, _serialize(record), ex=ttl_seconds)
except Exception as e:
logger.debug("Cache set failed: %s", e)
return etag