Newer
Older
gn-synapse / app / auth / stores.py
"""Redis-хранилища state/PKCE для OAuth-потока (TTL = срок state в gnexus-gauth).

Redis вместо in-memory: переживает рестарт api-контейнера и
работает при нескольких воркерах uvicorn/gunicorn.
"""

import json
from datetime import datetime

from redis import Redis

from app.config import get_settings

STATE_PREFIX = "synapse:auth:state:"
PKCE_PREFIX = "synapse:auth:pkce:"

_redis: Redis | None = None


def get_redis() -> Redis:
    global _redis
    if _redis is None:
        _redis = Redis.from_url(get_settings().redis_url, decode_responses=True)
    return _redis


def _ttl(expires_at: datetime) -> int:
    return max(1, int(expires_at.timestamp() - datetime.now(expires_at.tzinfo).timestamp()))


class RedisStateStore:
    """StateStoreInterface: state -> context (return_to, scopes)."""

    def put(self, state: str, expires_at: datetime, context: dict | None = None) -> None:
        get_redis().setex(STATE_PREFIX + state, _ttl(expires_at), json.dumps(context or {}))

    def has(self, state: str) -> bool:
        return bool(get_redis().exists(STATE_PREFIX + state))

    def get_context(self, state: str) -> dict:
        raw = get_redis().get(STATE_PREFIX + state)
        return json.loads(raw) if raw else {}

    def forget(self, state: str) -> None:
        get_redis().delete(STATE_PREFIX + state)


class RedisPkceStore:
    """PkceStoreInterface: state -> PKCE verifier."""

    def put(self, state: str, verifier: str, expires_at: datetime) -> None:
        get_redis().setex(PKCE_PREFIX + state, _ttl(expires_at), verifier)

    def get(self, state: str) -> str | None:
        return get_redis().get(PKCE_PREFIX + state)

    def forget(self, state: str) -> None:
        get_redis().delete(PKCE_PREFIX + state)


# --- Учёт выданных токенов: чтобы webhook-выход мог их отозвать ---

TOKENS_PREFIX = "synapse:auth:tokens:"  # :user_id -> hash {login_id: {access, refresh}}
TOKENS_TTL = 60 * 60 * 24 * 30  # TTL = жизни refresh-токена (30 дней на стороне auth)


def store_login(user_id: str, login_id: str, access_token: str, refresh_token: str | None) -> None:
    """Запоминаем выдачy токена при SSO-входе: на logout-webhook по user_id отдадим revoke."""
    key = TOKENS_PREFIX + str(user_id)
    entry = {"access": access_token, "refresh": refresh_token or ""}
    get_redis().hset(key, login_id, json.dumps(entry))
    get_redis().expire(key, TOKENS_TTL)


def revoke_logins_for_user(user_id: str, client) -> int:
    """Все сохранённые входы пользователя: revoke access+refresh на gnexus-auth, очистка.

    Возвращает, сколько входов отзывает (0 — этот пользователь к нам не входил).
    """
    key = TOKENS_PREFIX + str(user_id)
    redis = get_redis()
    stored = redis.hgetall(key)
    if not stored:
        return 0
    for login_id, raw in stored.items():
        entry = json.loads(raw)
        for token, hint in ((entry.get("access"), "access_token"), (entry.get("refresh"), "refresh_token")):
            if token:
                try:
                    client.revoke_token(token, hint)
                except Exception:  # noqa: BLE001 — токен может быть уже отозван на стороне auth
                    pass
        redis.hdel(key, login_id)
    return len(stored)