""" Shared one-time SSE ticket store — works across Gunicorn workers. Backend order: 1. Redis (REDIS_URL / TICKET_REDIS_URL) when available 2. SQLite file under TMPDIR (default) with WAL + exclusive lock on consume 3. Process-local dict (last resort; multi-worker unsafe) Tickets are one-shot: mint → consume (atomic) → gone. """ from __future__ import annotations import json import logging import os import secrets import sqlite3 import threading import time from pathlib import Path from typing import Any, Dict, Optional logger = logging.getLogger(__name__) DEFAULT_TTL_SECONDS = int(os.environ.get("STREAM_TICKET_TTL_SECONDS", "90")) _local_lock = threading.Lock() _local_tickets: Dict[str, Dict[str, Any]] = {} _redis_client = None _redis_failed = False _sqlite_path: Optional[Path] = None _sqlite_init_done = False def _redis(): global _redis_client, _redis_failed if _redis_failed: return None if _redis_client is not None: return _redis_client url = os.environ.get("TICKET_REDIS_URL") or os.environ.get("REDIS_URL") if not url: return None try: import redis # type: ignore client = redis.Redis.from_url(url, decode_responses=True, socket_timeout=2) client.ping() _redis_client = client logger.info("[TicketStore] using Redis backend") return client except Exception as e: _redis_failed = True logger.warning("[TicketStore] Redis unavailable (%s) — falling back to SQLite", e) return None def _sqlite_db_path() -> Path: global _sqlite_path if _sqlite_path is not None: return _sqlite_path base = Path(os.environ.get("TICKET_STORE_PATH") or os.environ.get("TMPDIR") or "/tmp") base.mkdir(parents=True, exist_ok=True) _sqlite_path = base / "grantforge_stream_tickets.sqlite3" return _sqlite_path def _sqlite_conn() -> sqlite3.Connection: global _sqlite_init_done path = _sqlite_db_path() conn = sqlite3.connect(str(path), timeout=5, isolation_level=None) conn.row_factory = sqlite3.Row if not _sqlite_init_done: conn.execute("PRAGMA journal_mode=WAL") conn.execute( """ CREATE TABLE IF NOT EXISTS stream_tickets ( ticket TEXT PRIMARY KEY, user_id TEXT NOT NULL, project_id TEXT NOT NULL, exp REAL NOT NULL, payload TEXT, created_at REAL NOT NULL ) """ ) conn.execute("CREATE INDEX IF NOT EXISTS ix_stream_tickets_exp ON stream_tickets(exp)") _sqlite_init_done = True return conn def mint_ticket( user_id: str, project_id: str, *, ttl_seconds: Optional[int] = None, extra: Optional[Dict[str, Any]] = None, ) -> str: ttl = int(ttl_seconds or DEFAULT_TTL_SECONDS) ticket = secrets.token_urlsafe(32) exp = time.time() + ttl payload = {"user_id": user_id, "project_id": project_id, "exp": exp, **(extra or {})} r = _redis() if r is not None: try: r.setex(f"gf:ticket:{ticket}", ttl, json.dumps(payload)) return ticket except Exception as e: logger.warning("[TicketStore] Redis mint failed: %s", e) try: conn = _sqlite_conn() try: conn.execute( "INSERT INTO stream_tickets(ticket, user_id, project_id, exp, payload, created_at) " "VALUES (?,?,?,?,?,?)", (ticket, user_id, project_id, exp, json.dumps(payload), time.time()), ) # opportunistic purge conn.execute("DELETE FROM stream_tickets WHERE exp < ?", (time.time(),)) return ticket finally: conn.close() except Exception as e: logger.warning("[TicketStore] SQLite mint failed: %s — local fallback", e) with _local_lock: _local_tickets[ticket] = payload return ticket def validate_ticket(ticket: str, project_id: str, *, consume: bool = False) -> str: """ Validate ticket and return user_id. consume=False (default for long SSE streams): ticket remains valid until TTL so EventSource auto-reconnects and multi-worker retries work during generation. consume=True: one-shot delete (legacy/high-security). """ if not ticket: raise ValueError("missing_ticket") r = _redis() if r is not None: try: key = f"gf:ticket:{ticket}" raw = r.get(key) if not raw: raise ValueError("invalid_or_expired") data = json.loads(raw) if data.get("exp", 0) < time.time(): try: r.delete(key) except Exception: pass raise ValueError("invalid_or_expired") if str(data.get("project_id") or "") != str(project_id): raise ValueError("project_mismatch") user_id = data.get("user_id") if not user_id: raise ValueError("invalid_or_expired") if consume: try: r.delete(key) except Exception: pass return str(user_id) except ValueError: raise except Exception as e: logger.warning("[TicketStore] Redis validate failed: %s", e) try: conn = _sqlite_conn() try: conn.execute("BEGIN IMMEDIATE") row = conn.execute( "SELECT user_id, project_id, exp FROM stream_tickets WHERE ticket = ?", (ticket,), ).fetchone() if not row: conn.execute("ROLLBACK") raise ValueError("invalid_or_expired") if float(row["exp"]) < time.time(): conn.execute("DELETE FROM stream_tickets WHERE ticket = ?", (ticket,)) conn.execute("COMMIT") raise ValueError("invalid_or_expired") if str(row["project_id"]) != str(project_id): conn.execute("ROLLBACK") raise ValueError("project_mismatch") user_id = str(row["user_id"]) if consume: conn.execute("DELETE FROM stream_tickets WHERE ticket = ?", (ticket,)) conn.execute("COMMIT") return user_id except ValueError: raise except Exception: try: conn.execute("ROLLBACK") except Exception: pass raise finally: conn.close() except ValueError: raise except Exception as e: logger.warning("[TicketStore] SQLite validate failed: %s", e) with _local_lock: data = _local_tickets.get(ticket) if not data: raise ValueError("invalid_or_expired") if data.get("exp", 0) < time.time(): _local_tickets.pop(ticket, None) raise ValueError("invalid_or_expired") if str(data.get("project_id") or "") != str(project_id): raise ValueError("project_mismatch") user_id = data.get("user_id") if not user_id: raise ValueError("invalid_or_expired") if consume: _local_tickets.pop(ticket, None) return str(user_id) def consume_ticket(ticket: str, project_id: str) -> str: """ Atomically validate and consume ticket (one-shot). Returns user_id. Prefer validate_ticket(..., consume=False) for long-lived SSE generator streams. """ return validate_ticket(ticket, project_id, consume=True) def purge_expired() -> int: """Best-effort cleanup; returns approximate number removed.""" removed = 0 now = time.time() r = _redis() # Redis keys expire via TTL — nothing to purge try: conn = _sqlite_conn() try: cur = conn.execute("DELETE FROM stream_tickets WHERE exp < ?", (now,)) removed += cur.rowcount or 0 finally: conn.close() except Exception: pass with _local_lock: expired = [k for k, v in _local_tickets.items() if v.get("exp", 0) < now] for k in expired: _local_tickets.pop(k, None) removed += len(expired) return removed