grantforge-api / backend /core /ticket_store.py
GrantForge Bot
Deploy sha-565ad85979610064f6d1c18ab3b6404357d61073 — source build (no GHCR)
ce8f04a
Raw
History Blame Contribute Delete
8.37 kB
"""
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