import json from dataclasses import dataclass from datetime import datetime, timezone from typing import Callable, Dict, Optional @dataclass class ChisSession: base_url: str uid: str role_id: Optional[str] manage_unit: Optional[str] cookies: Dict[str, str] login_at: datetime expires_at: datetime last_validated_at: Optional[datetime] = None class RedisChisSessionStore: key_prefix = "chis:session:" def __init__(self, redis_client, now: Optional[Callable[[], datetime]] = None): self.redis = redis_client self.now = now or _utc_now @classmethod def from_url(cls, redis_url, now: Optional[Callable[[], datetime]] = None): import redis return cls(redis.Redis.from_url(redis_url), now=now) def save(self, account_ref, session: ChisSession): ttl_seconds = self._ttl_seconds(session) if ttl_seconds <= 0: self.delete(account_ref) return self.redis.setex(self._key(account_ref), ttl_seconds, self.dumps(session)) def get(self, account_ref): raw_value = self.redis.get(self._key(account_ref)) if not raw_value: return None session = self.loads(raw_value) if session.expires_at <= self.now(): self.delete(account_ref) return None return session def delete(self, account_ref): self.redis.delete(self._key(account_ref)) def dumps(self, session: ChisSession): return json.dumps( { "base_url": session.base_url, "uid": session.uid, "role_id": session.role_id, "manage_unit": session.manage_unit, "cookies": session.cookies, "login_at": _datetime_to_text(session.login_at), "expires_at": _datetime_to_text(session.expires_at), "last_validated_at": _datetime_to_text(session.last_validated_at), }, ensure_ascii=False, sort_keys=True, ) def loads(self, raw_value): if isinstance(raw_value, bytes): raw_value = raw_value.decode("utf-8") payload = json.loads(raw_value) return ChisSession( base_url=payload["base_url"], uid=payload["uid"], role_id=payload.get("role_id"), manage_unit=payload.get("manage_unit"), cookies=payload.get("cookies", {}), login_at=_datetime_from_text(payload["login_at"]), expires_at=_datetime_from_text(payload["expires_at"]), last_validated_at=_datetime_from_text(payload.get("last_validated_at")), ) def _ttl_seconds(self, session): return int((session.expires_at - self.now()).total_seconds()) def _key(self, account_ref): return f"{self.key_prefix}{account_ref}" def _utc_now(): return datetime.now(timezone.utc) def _datetime_to_text(value): if value is None: return None return value.isoformat() def _datetime_from_text(value): if value is None: return None parsed = datetime.fromisoformat(value) if parsed.tzinfo is None: return parsed.replace(tzinfo=timezone.utc) return parsed