Files
chisup/app/chis/session_store.py
T

106 lines
3.2 KiB
Python

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