# retoor <retoor@molodetz.nl>
import logging
import threading
import time
from dataclasses import dataclass
from typing import Generic, TypeVar
logger = logging.getLogger(__name__)
T = TypeVar("T")
@dataclass
class CacheEntry(Generic[T]):
value: T
expires_at: float
class TTLCache(Generic[T]):
def __init__(self, name: str, ttl_seconds: float) -> None:
self._name = name
self._ttl_seconds = ttl_seconds
self._entries: dict[str, CacheEntry[T]] = {}
self._lock = threading.Lock()
def get(self, key: str) -> T | None:
with self._lock:
entry = self._entries.get(key)
if entry is None:
logger.debug("cache %s miss key=%s", self._name, key)
return None
if time.monotonic() >= entry.expires_at:
del self._entries[key]
logger.debug("cache %s expired key=%s", self._name, key)
return None
logger.debug("cache %s hit key=%s", self._name, key)
return entry.value
def set(self, key: str, value: T) -> None:
with self._lock:
self._entries[key] = CacheEntry(value=value, expires_at=time.monotonic() + self._ttl_seconds)
logger.debug("cache %s set key=%s ttl=%.0fs", self._name, key, self._ttl_seconds)
def clear(self) -> None:
with self._lock:
count = len(self._entries)
self._entries.clear()
logger.debug("cache %s cleared %d entries", self._name, count)