106 lines
3.2 KiB
Python
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 |