Spaces:
Sleeping
Sleeping
| """ | |
| 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 | |