import base64 import json import sqlite3 import time from dataclasses import dataclass from functools import cached_property from typing import Any from database.connection import DatabasePool def _decode_jwt_payload(token: str) -> dict[str, Any]: try: _, claims, _ = token.split(".") claims = claims + "=" * (4 - len(claims) % 4) if len(claims) % 4 else claims return json.loads(base64.urlsafe_b64decode(claims)) # type: ignore[no-any-return] except Exception: return {} @dataclass class Session: access_jwt: str refresh_jwt: str handle: str did: str pds: str email: str | None = None email_confirmed: bool = False email_auth_factor: bool = False active: bool = True status: str | None = None @cached_property def access_payload(self) -> dict[str, Any]: return _decode_jwt_payload(self.access_jwt) @cached_property def refresh_payload(self) -> dict[str, Any]: return _decode_jwt_payload(self.refresh_jwt) def is_access_token_expired(self, buffer_seconds: int = 60) -> bool: exp = self.access_payload.get("exp", 0) return bool(time.time() >= (exp - buffer_seconds)) def is_refresh_token_expired(self, buffer_seconds: int = 60) -> bool: exp = self.refresh_payload.get("exp", 0) return bool(time.time() >= (exp - buffer_seconds)) @classmethod def from_row(cls, row: sqlite3.Row) -> "Session": return cls( access_jwt=row["access_jwt"], refresh_jwt=row["refresh_jwt"], handle=row["handle"], did=row["did"], pds=row["pds"], email=row["email"], email_confirmed=bool(row["email_confirmed"]), email_auth_factor=bool(row["email_auth_factor"]), active=bool(row["active"]), status=row["status"], ) @classmethod def from_dict(cls, data: dict[str, Any], pds: str) -> "Session": return cls( access_jwt=data["accessJwt"], refresh_jwt=data["refreshJwt"], handle=data["handle"], did=data["did"], pds=pds, email=data.get("email"), email_confirmed=data.get("emailConfirmed", False), email_auth_factor=data.get("emailAuthFactor", False), active=data.get("active", True), status=data.get("status"), ) @dataclass class IdentityInfo: did: str handle: str pds: str signing_key: str @classmethod def from_row(cls, row: sqlite3.Row) -> "IdentityInfo": return cls( did=row["did"], handle=row["handle"], pds=row["pds"], signing_key=row["signing_key"], ) @classmethod def from_dict(cls, data: dict[str, Any]) -> "IdentityInfo": return cls( did=data["did"], handle=data["handle"], pds=data["pds"], signing_key=data["signing_key"], ) class AtprotoStore: def __init__( self, db: sqlite3.Connection, identity_ttl: int = 12 * 60 * 60, ) -> None: self.db = db self.db.row_factory = sqlite3.Row self.identity_ttl = identity_ttl def get_session(self, did: str) -> Session | None: row = self.db.execute( "SELECT * FROM atproto_sessions WHERE did = ?", (did,) ).fetchone() return Session.from_row(row) if row else None def set_session(self, session: Session) -> None: now = time.time() self.db.execute( """ INSERT OR REPLACE INTO atproto_sessions (did, pds, handle, access_jwt, refresh_jwt, email, email_confirmed, email_auth_factor, active, status, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) """, ( session.did, session.pds, session.handle, session.access_jwt, session.refresh_jwt, session.email, session.email_confirmed, session.email_auth_factor, session.active, session.status, now, ), ) self.db.commit() def get_session_by_pds(self, pds: str, identifier: str) -> Session | None: row = self.db.execute( """ SELECT * FROM atproto_sessions WHERE pds = ? AND (did = ? OR handle = ?) """, (pds, identifier, identifier), ).fetchone() return Session.from_row(row) if row else None def list_sessions_by_pds(self, pds: str) -> list[Session]: rows = self.db.execute( "SELECT * FROM atproto_sessions WHERE pds = ?", (pds,) ).fetchall() return [Session.from_row(row) for row in rows] def remove_session(self, did: str) -> None: self.db.execute("DELETE FROM atproto_sessions WHERE did = ?", (did,)) self.db.commit() def get_identity(self, identifier: str) -> IdentityInfo | None: row = self.db.execute( "SELECT * FROM atproto_identities WHERE identifier = ? AND created_at + ? > ?", (identifier, self.identity_ttl, time.time()), ).fetchone() return IdentityInfo.from_row(row) if row else None def set_identity(self, identifier: str, identity: IdentityInfo) -> None: now = time.time() for key in (identifier, identity.did, identity.handle): self.db.execute( """ INSERT OR REPLACE INTO atproto_identities (identifier, did, handle, pds, signing_key, created_at) VALUES (?, ?, ?, ?, ?, ?) """, ( key, identity.did, identity.handle, identity.pds, identity.signing_key, now, ), ) self.db.commit() def remove_identity(self, identifier: str) -> None: self.db.execute( "DELETE FROM atproto_identities WHERE identifier = ?", (identifier,) ) self.db.commit() def cleanup_expired(self) -> None: cutoff = time.time() - self.identity_ttl self.db.execute( "DELETE FROM atproto_identities WHERE created_at + ? < ?", (self.identity_ttl, cutoff), ) self.db.commit() def flush_all(self) -> tuple[int, int]: sessions = self.db.execute("SELECT COUNT(*) FROM atproto_sessions").fetchone()[ 0 ] identities = self.db.execute( "SELECT COUNT(*) FROM atproto_identities" ).fetchone()[0] self.db.execute("DELETE FROM atproto_sessions") self.db.execute("DELETE FROM atproto_identities") self.db.commit() return sessions, identities _store: AtprotoStore | None = None def get_store(db: DatabasePool) -> AtprotoStore: global _store if _store is None: _store = AtprotoStore(db.get_conn()) return _store def flush_caches() -> tuple[int, int]: if _store is not None: return _store.flush_all() return 0, 0