Phase 9: production hardening
Config via env, abuse/resource limits, and production serving so the app is safe to expose publicly: - Fail-fast on missing SECRET_KEY when MULTI_USER=true - quickjs per-execution time/memory limits (while(true) can't hang server) - Per-user/per-IP rate limiting on turn/script/auth endpoints - Request body size limit + per-user row caps - Security headers (CSP, X-Frame-Options, nosniff, referrer-policy) incl. SSE - Debug router 403 and /docs disabled in multi-user mode - DATABASE_URL support (defaults to Neon Postgres) alongside SQLite - Documented all env vars in backend/.env.example Verified locally via uvicorn (MULTI_USER=1, SQLite); see plan/09-phase-hardening.md. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017e6tQuojBLYPetUfmhit4X
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
de4db373f2
commit
4772171b6c
+6
-1
@@ -33,4 +33,9 @@ VOLUME /data
|
||||
|
||||
EXPOSE 8000
|
||||
WORKDIR /app/backend
|
||||
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||
# --proxy-headers: behind a reverse proxy (any hosted deploy), trust
|
||||
# X-Forwarded-For so per-IP rate limits key on the client, not the proxy.
|
||||
# Single worker on purpose: the turn lock, rate limiter, and debug log are
|
||||
# in-process state.
|
||||
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", \
|
||||
"--proxy-headers", "--forwarded-allow-ips", "*"]
|
||||
|
||||
+29
-5
@@ -9,6 +9,22 @@
|
||||
# Docker compose sets this to /data/data.db (a named volume).
|
||||
AIDND_DB_PATH=
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 9 — production hardening
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Switch from SQLite to a server database (hosted deploys use Neon Postgres).
|
||||
# Any SQLAlchemy URL; postgres:// and postgresql:// schemes are rewritten to
|
||||
# the psycopg3 driver automatically. The platform-conventional DATABASE_URL
|
||||
# is honored too (AIDND_DATABASE_URL wins if both are set). Unset = SQLite.
|
||||
AIDND_DATABASE_URL=
|
||||
|
||||
# Comma-separated list of allowed CORS origins. Only needed when the frontend
|
||||
# is served from a different origin than the API; the production build is
|
||||
# served same-origin by FastAPI, so hosted deploys can leave this unset.
|
||||
# Default: http://localhost:5173,http://127.0.0.1:5173 (the Vite dev server).
|
||||
AIDND_CORS_ORIGINS=
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 8 — optional accounts & multi-user (all optional; defaults keep the
|
||||
# app in frictionless single-user "local mode")
|
||||
@@ -19,12 +35,17 @@ AIDND_DB_PATH=
|
||||
AIDND_MULTI_USER=
|
||||
|
||||
# Secret for signing session cookies and encrypting stored API keys at rest.
|
||||
# If unset, one is auto-generated into `secret.key` next to the database
|
||||
# (fine for local/docker-volume runs). Set it explicitly on hosted deploys so
|
||||
# sessions survive redeploys when the disk is ephemeral or replaced.
|
||||
# If unset in local mode, one is auto-generated into `secret.key` next to the
|
||||
# database (fine for local/docker-volume runs). REQUIRED when
|
||||
# AIDND_MULTI_USER is on — the app refuses to start without it, because a
|
||||
# regenerated secret on an ephemeral hosted filesystem would log out every
|
||||
# user on each deploy. Generate one:
|
||||
# python -c "import secrets; print(secrets.token_urlsafe(48))"
|
||||
AIDND_SECRET_KEY=
|
||||
|
||||
# "1" marks session cookies Secure (HTTPS-only). Turn on in production.
|
||||
# Session cookie Secure flag (HTTPS-only). Defaults to on when
|
||||
# AIDND_MULTI_USER is on, off otherwise — set 0/1 only to override (e.g. 0
|
||||
# when testing multi-user mode over plain http on a LAN address).
|
||||
AIDND_COOKIE_SECURE=
|
||||
|
||||
# --- Shared demo key (BYOK fallback; only active when AIDND_MULTI_USER=1) ---
|
||||
@@ -41,4 +62,7 @@ AIDND_DEMO_TURNS_PER_DAY=
|
||||
|
||||
# The AI endpoint/API key/model are NOT env vars — they are configured at
|
||||
# runtime in the app's Settings page and stored (encrypted) in the database.
|
||||
# Phase 9 will add: rate limiting and CORS_ORIGINS.
|
||||
#
|
||||
# Rate limits, request size limits, and per-user row caps are hardcoded with
|
||||
# generous values (see backend/app/limits.py) and active only in multi-user
|
||||
# mode — local installs are never throttled.
|
||||
|
||||
+9
-21
@@ -16,8 +16,6 @@ whitelist and a per-day turn cap.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import defaultdict, deque
|
||||
from dataclasses import dataclass
|
||||
from datetime import timezone
|
||||
|
||||
@@ -35,7 +33,15 @@ def _env_flag(name: str) -> bool:
|
||||
MULTI_USER = _env_flag("AIDND_MULTI_USER")
|
||||
|
||||
SESSION_COOKIE = "aidnd_session"
|
||||
COOKIE_SECURE = _env_flag("AIDND_COOKIE_SECURE") # enable behind HTTPS in prod
|
||||
# Secure cookies default on in multi-user (hosted = HTTPS; browsers also
|
||||
# accept Secure on http://localhost). AIDND_COOKIE_SECURE=0/1 overrides —
|
||||
# e.g. 0 when testing multi-user over plain http on a LAN address.
|
||||
_cookie_secure_env = os.environ.get("AIDND_COOKIE_SECURE", "").strip().lower()
|
||||
COOKIE_SECURE = (
|
||||
_cookie_secure_env in ("1", "true", "yes", "on")
|
||||
if _cookie_secure_env
|
||||
else MULTI_USER
|
||||
)
|
||||
COOKIE_MAX_AGE = 60 * 60 * 24 * 365
|
||||
|
||||
# ---------- Shared demo key (BYOK fallback) ----------
|
||||
@@ -151,21 +157,3 @@ def get_current_user(request: Request, db: Session = Depends(get_db)) -> models.
|
||||
raise HTTPException(401, "No session. Call GET /api/auth/me first.")
|
||||
_touch(user, db)
|
||||
return user
|
||||
|
||||
|
||||
# ---------- Brute-force limiter for register/login ----------
|
||||
|
||||
_ATTEMPT_LIMIT = 10
|
||||
_ATTEMPT_WINDOW = 300 # seconds
|
||||
_attempts: dict[str, deque] = defaultdict(deque)
|
||||
|
||||
|
||||
def rate_limit_auth(request: Request) -> None:
|
||||
ip = request.client.host if request.client else "unknown"
|
||||
now = time.time()
|
||||
window = _attempts[ip]
|
||||
while window and window[0] < now - _ATTEMPT_WINDOW:
|
||||
window.popleft()
|
||||
if len(window) >= _ATTEMPT_LIMIT:
|
||||
raise HTTPException(429, "Too many attempts. Try again in a few minutes.")
|
||||
window.append(now)
|
||||
|
||||
+43
-12
@@ -5,28 +5,59 @@ from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.orm import DeclarativeBase, sessionmaker
|
||||
|
||||
# AIDND_DB_PATH lets deployments (Docker volume, hosted disk) relocate the
|
||||
# database; default stays backend/data.db for local runs.
|
||||
# SQLite database; default stays backend/data.db for local runs. The parent
|
||||
# directory also hosts the auto-generated secret.key (see security.py), so
|
||||
# DB_PATH stays defined even when Postgres is in use.
|
||||
_env_db_path = os.environ.get("AIDND_DB_PATH")
|
||||
DB_PATH = (
|
||||
Path(_env_db_path).resolve()
|
||||
if _env_db_path
|
||||
else Path(__file__).resolve().parent.parent / "data.db"
|
||||
)
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
engine = create_engine(
|
||||
f"sqlite:///{DB_PATH}",
|
||||
connect_args={"check_same_thread": False},
|
||||
# AIDND_DATABASE_URL (or the platform-conventional DATABASE_URL) switches the
|
||||
# app to a server database — any SQLAlchemy URL works, but Postgres is what
|
||||
# hosted deploys use (Phase 9 decision: Neon). Unset = SQLite, as always.
|
||||
DATABASE_URL = (
|
||||
os.environ.get("AIDND_DATABASE_URL", "").strip()
|
||||
or os.environ.get("DATABASE_URL", "").strip()
|
||||
)
|
||||
|
||||
|
||||
@event.listens_for(engine, "connect")
|
||||
def _enable_sqlite_foreign_keys(dbapi_connection, _record):
|
||||
# SQLite ships with foreign keys OFF per connection; without this every
|
||||
# ondelete=CASCADE/SET NULL in models.py is silently ignored.
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
def _normalize_url(url: str) -> str:
|
||||
"""Map the postgres:// / postgresql:// schemes hosts hand out to the
|
||||
psycopg3 driver installed in requirements.txt."""
|
||||
for prefix in ("postgres://", "postgresql://"):
|
||||
if url.startswith(prefix):
|
||||
return "postgresql+psycopg://" + url[len(prefix):]
|
||||
return url
|
||||
|
||||
|
||||
if DATABASE_URL:
|
||||
engine = create_engine(
|
||||
_normalize_url(DATABASE_URL),
|
||||
# Serverless Postgres (Neon) suspends idle databases; pre-ping
|
||||
# replaces silently-dead pooled connections instead of erroring.
|
||||
pool_pre_ping=True,
|
||||
# Store/read naive UTC like SQLite does, regardless of server default.
|
||||
connect_args={"options": "-c timezone=UTC"},
|
||||
)
|
||||
else:
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
engine = create_engine(
|
||||
f"sqlite:///{DB_PATH}",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
@event.listens_for(engine, "connect")
|
||||
def _enable_sqlite_foreign_keys(dbapi_connection, _record):
|
||||
# SQLite ships with foreign keys OFF per connection; without this every
|
||||
# ondelete=CASCADE/SET NULL in models.py is silently ignored.
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
|
||||
|
||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
"""Phase 9 — abuse guards for hosted (multi-user) deployments.
|
||||
|
||||
Rate limits and row caps are no-ops in local mode: a single local player
|
||||
should never be throttled by their own app. Values are hardcoded on purpose —
|
||||
generous enough that a legitimate player never notices, tight enough that a
|
||||
hostile visitor can't burn the demo key, peg the CPU, or bloat the database.
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
from collections import defaultdict, deque
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import auth, models
|
||||
|
||||
# ---------- Rate limiting ----------
|
||||
# Fixed windows per (scope, caller). In-memory: fine for the single-process
|
||||
# deployment this app targets (and the worst case after a restart is a brief
|
||||
# extra allowance).
|
||||
|
||||
# scope -> (max requests, window seconds)
|
||||
RATE_LIMITS: dict[str, tuple[int, int]] = {
|
||||
"turn": (10, 60), # AI turn generation (demo key also has a daily cap)
|
||||
"script-test": (30, 60), # sandboxed, but each run costs up to 2s CPU
|
||||
"connection-test": (10, 60), # outbound HTTP to a user-supplied URL
|
||||
"import": (30, 60), # large writes
|
||||
"auth": (10, 300), # register/login attempts, per IP
|
||||
"guest": (30, 300), # new guest users, per IP (each is a DB row)
|
||||
}
|
||||
|
||||
_windows: dict[tuple[str, str], deque] = defaultdict(deque)
|
||||
_windows_guard = threading.Lock()
|
||||
|
||||
|
||||
def _client_ip(request: Request) -> str:
|
||||
return request.client.host if request.client else "unknown"
|
||||
|
||||
|
||||
def rate_limit(scope: str, request: Request, user: models.User | None = None) -> None:
|
||||
"""429 when the caller exceeds the scope's window. Keyed per user when one
|
||||
is known (accounts survive IP changes), per IP otherwise."""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
limit, window_seconds = RATE_LIMITS[scope]
|
||||
key = (scope, f"u{user.id}" if user else f"ip{_client_ip(request)}")
|
||||
now = time.time()
|
||||
with _windows_guard:
|
||||
window = _windows[key]
|
||||
while window and window[0] < now - window_seconds:
|
||||
window.popleft()
|
||||
if len(window) >= limit:
|
||||
raise HTTPException(
|
||||
429, "You're doing that too fast — wait a minute and try again."
|
||||
)
|
||||
window.append(now)
|
||||
if len(_windows) > 10_000:
|
||||
_prune(now)
|
||||
|
||||
|
||||
def _prune(now: float) -> None:
|
||||
"""Drop callers whose whole window has expired (call with guard held) so
|
||||
the per-IP dict can't grow without bound."""
|
||||
longest = max(seconds for _, seconds in RATE_LIMITS.values())
|
||||
stale = [key for key, window in _windows.items()
|
||||
if not window or window[-1] < now - longest]
|
||||
for key in stale:
|
||||
del _windows[key]
|
||||
|
||||
|
||||
# ---------- Per-user row caps ----------
|
||||
|
||||
MAX_ADVENTURES_PER_USER = 100
|
||||
MAX_SCENARIOS_PER_USER = 200
|
||||
MAX_SCRIPTS_PER_USER = 200
|
||||
MAX_STORY_CARDS_PER_OWNER = 200 # per scenario or adventure
|
||||
MAX_MEMORIES_PER_ADVENTURE = 1000
|
||||
MAX_ACTIONS_PER_ADVENTURE = 5000
|
||||
|
||||
|
||||
def check_row_cap(
|
||||
kind: str,
|
||||
db: Session,
|
||||
user: models.User,
|
||||
*,
|
||||
adventure: models.Adventure | None = None,
|
||||
scenario_id: int | None = None,
|
||||
adventure_id: int | None = None,
|
||||
) -> None:
|
||||
"""409 with a friendly message when creating one more row of `kind` would
|
||||
exceed its cap. Ownership of the passed scenario/adventure has already
|
||||
been checked by the caller."""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
if kind == "adventures":
|
||||
count = _count(db, models.Adventure, models.Adventure.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_ADVENTURES_PER_USER, "adventures",
|
||||
"delete one you no longer play to make room",
|
||||
)
|
||||
elif kind == "scenarios":
|
||||
count = _count(db, models.Scenario, models.Scenario.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_SCENARIOS_PER_USER, "scenarios", "delete one to make room"
|
||||
)
|
||||
elif kind == "scripts":
|
||||
count = _count(db, models.Script, models.Script.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_SCRIPTS_PER_USER, "scripts", "delete one to make room"
|
||||
)
|
||||
elif kind == "story_cards":
|
||||
owner_filter = (
|
||||
models.StoryCard.scenario_id == scenario_id
|
||||
if scenario_id is not None
|
||||
else models.StoryCard.adventure_id == adventure_id
|
||||
)
|
||||
count = _count(db, models.StoryCard, owner_filter)
|
||||
cap, subject, hint = (
|
||||
MAX_STORY_CARDS_PER_OWNER, "story cards here", "delete one to make room"
|
||||
)
|
||||
elif kind == "memories":
|
||||
count = _count(db, models.Memory, models.Memory.adventure_id == adventure.id)
|
||||
cap, subject, hint = (
|
||||
MAX_MEMORIES_PER_ADVENTURE, "memories in this adventure",
|
||||
"delete some to make room",
|
||||
)
|
||||
elif kind == "actions":
|
||||
count = _count(db, models.Action, models.Action.adventure_id == adventure.id)
|
||||
cap, subject, hint = (
|
||||
MAX_ACTIONS_PER_ADVENTURE, "actions in this adventure",
|
||||
"export it and continue in a new adventure",
|
||||
)
|
||||
else: # pragma: no cover — programming error, not user input
|
||||
raise ValueError(f"Unknown row cap kind: {kind}")
|
||||
if count >= cap:
|
||||
raise HTTPException(409, f"You've reached the limit of {cap} {subject} — {hint}.")
|
||||
|
||||
|
||||
def _count(db: Session, model, condition) -> int:
|
||||
return db.query(func.count(model.id)).filter(condition).scalar() or 0
|
||||
|
||||
|
||||
_BUNDLE_LIST_CAPS = {
|
||||
"story_cards": MAX_STORY_CARDS_PER_OWNER,
|
||||
"memories": MAX_MEMORIES_PER_ADVENTURE,
|
||||
"actions": MAX_ACTIONS_PER_ADVENTURE,
|
||||
}
|
||||
|
||||
|
||||
def check_bundle_lists(**lists) -> None:
|
||||
"""409 when an import bundle's lists exceed the same caps live creation
|
||||
enforces (kwargs: story_cards=, memories=, actions=)."""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
for name, value in lists.items():
|
||||
cap = _BUNDLE_LIST_CAPS[name]
|
||||
if isinstance(value, list) and len(value) > cap:
|
||||
noun = name.replace("_", " ")
|
||||
raise HTTPException(
|
||||
409, f"This file contains {len(value)} {noun} — the limit is {cap}."
|
||||
)
|
||||
|
||||
|
||||
# ---------- Request body size ----------
|
||||
# Generous enough for the biggest legitimate payload (an adventure export with
|
||||
# thousands of actions), applied in every mode — no honest request comes close.
|
||||
|
||||
MAX_BODY_BYTES = 2 * 1024 * 1024
|
||||
MAX_IMPORT_BODY_BYTES = 20 * 1024 * 1024
|
||||
|
||||
|
||||
class BodySizeLimitMiddleware:
|
||||
"""Rejects oversized request bodies by declared Content-Length. Pure ASGI
|
||||
(not BaseHTTPMiddleware) so SSE responses stream through untouched.
|
||||
Chunked uploads without a length are refused — every real client of this
|
||||
API (browser fetch, curl with a file) sends Content-Length."""
|
||||
|
||||
def __init__(self, app):
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] == "http" and scope.get("method") in ("POST", "PUT", "PATCH"):
|
||||
headers = {k.decode("latin-1").lower(): v.decode("latin-1")
|
||||
for k, v in scope.get("headers", [])}
|
||||
limit = (
|
||||
MAX_IMPORT_BODY_BYTES
|
||||
if scope.get("path", "").endswith("/import")
|
||||
else MAX_BODY_BYTES
|
||||
)
|
||||
length = headers.get("content-length")
|
||||
problem = None
|
||||
if length is None:
|
||||
if "chunked" in headers.get("transfer-encoding", "").lower():
|
||||
problem = (411, "Content-Length is required.")
|
||||
else:
|
||||
try:
|
||||
if int(length) > limit:
|
||||
problem = (
|
||||
413,
|
||||
f"Request too large (limit {limit // (1024 * 1024)} MB).",
|
||||
)
|
||||
except ValueError:
|
||||
problem = (400, "Invalid Content-Length.")
|
||||
if problem:
|
||||
await _send_json_error(send, *problem)
|
||||
return
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
|
||||
async def _send_json_error(send, status: int, detail: str) -> None:
|
||||
body = json.dumps({"detail": detail}).encode()
|
||||
await send({
|
||||
"type": "http.response.start",
|
||||
"status": status,
|
||||
"headers": [(b"content-type", b"application/json"),
|
||||
(b"content-length", str(len(body)).encode())],
|
||||
})
|
||||
await send({"type": "http.response.body", "body": body})
|
||||
+62
-2
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import FastAPI
|
||||
@@ -5,21 +6,80 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
|
||||
from .auth import MULTI_USER
|
||||
from .database import engine
|
||||
from .limits import BodySizeLimitMiddleware
|
||||
from .migrations import bootstrap
|
||||
from .routers import adventures, auth, debug, scenarios, scripts, settings, story_cards
|
||||
|
||||
bootstrap(engine)
|
||||
|
||||
app = FastAPI(title="AI D&D")
|
||||
# Production serves the SPA same-origin, so CORS only matters for the Vite dev
|
||||
# server; AIDND_CORS_ORIGINS overrides for any other cross-origin setup.
|
||||
CORS_ORIGINS = [
|
||||
o.strip()
|
||||
for o in os.environ.get("AIDND_CORS_ORIGINS", "").split(",")
|
||||
if o.strip()
|
||||
] or ["http://localhost:5173", "http://127.0.0.1:5173"]
|
||||
|
||||
# The interactive API docs stay local-only: in multi-user mode they just hand
|
||||
# strangers a map of the API surface.
|
||||
app = FastAPI(
|
||||
title="AI D&D",
|
||||
docs_url=None if MULTI_USER else "/docs",
|
||||
redoc_url=None,
|
||||
openapi_url=None if MULTI_USER else "/openapi.json",
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["http://localhost:5173", "http://127.0.0.1:5173"],
|
||||
allow_origins=CORS_ORIGINS,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.add_middleware(BodySizeLimitMiddleware)
|
||||
|
||||
|
||||
class SecurityHeadersMiddleware:
|
||||
"""Standard hardening headers on every response. Pure ASGI (wraps `send`)
|
||||
so SSE streams pass through unbuffered. The CSP allows exactly what the
|
||||
SPA uses: same-origin everything, inline styles (React), Google Fonts."""
|
||||
|
||||
_HEADERS = [
|
||||
(b"x-content-type-options", b"nosniff"),
|
||||
(b"referrer-policy", b"same-origin"),
|
||||
(b"x-frame-options", b"DENY"),
|
||||
(
|
||||
b"content-security-policy",
|
||||
b"default-src 'self'; "
|
||||
b"script-src 'self'; "
|
||||
b"style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; "
|
||||
b"font-src https://fonts.gstatic.com; "
|
||||
b"img-src 'self' data:; "
|
||||
b"connect-src 'self'; "
|
||||
b"frame-ancestors 'none'",
|
||||
),
|
||||
]
|
||||
|
||||
def __init__(self, app):
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] != "http":
|
||||
return await self.app(scope, receive, send)
|
||||
|
||||
async def send_with_headers(message):
|
||||
if message["type"] == "http.response.start":
|
||||
message.setdefault("headers", [])
|
||||
message["headers"] = list(message["headers"]) + self._HEADERS
|
||||
await send(message)
|
||||
|
||||
await self.app(scope, receive, send_with_headers)
|
||||
|
||||
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
|
||||
app.include_router(auth.router)
|
||||
app.include_router(scenarios.router)
|
||||
app.include_router(adventures.router)
|
||||
|
||||
@@ -1,14 +1,20 @@
|
||||
"""Lightweight versioned schema migrations over SQLite's PRAGMA user_version.
|
||||
"""Lightweight versioned schema migrations.
|
||||
|
||||
How it works:
|
||||
- A fresh database is created by `Base.metadata.create_all()` (always current)
|
||||
and stamped with LATEST_VERSION.
|
||||
- An existing database runs every migration whose version is greater than its
|
||||
stored user_version, in order, then is stamped.
|
||||
stored version, in order, then is stamped.
|
||||
|
||||
The version lives in SQLite's PRAGMA user_version, or a one-row
|
||||
`schema_version` table on Postgres (no PRAGMA there).
|
||||
|
||||
To change the schema: update models.py (keeps fresh DBs current) AND append a
|
||||
(version, sql) pair here (upgrades existing DBs). Keep migrations idempotent
|
||||
where cheap (IF NOT EXISTS etc.).
|
||||
where cheap (IF NOT EXISTS etc.). Migrations up to 23 predate Postgres support
|
||||
and use SQLite-only syntax — that's fine because every Postgres database
|
||||
starts fresh (created by create_all, stamped LATEST, never replays them), but
|
||||
migrations added from Phase 9 on must run on both dialects.
|
||||
"""
|
||||
|
||||
from sqlalchemy import inspect, text
|
||||
@@ -70,19 +76,46 @@ MIGRATIONS: list[tuple[int, str]] = [
|
||||
LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1)
|
||||
|
||||
|
||||
def _get_version(conn) -> int:
|
||||
if conn.dialect.name == "sqlite":
|
||||
return conn.execute(text("PRAGMA user_version")).scalar() or 1
|
||||
conn.execute(text(
|
||||
"CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL)"
|
||||
))
|
||||
version = conn.execute(text("SELECT version FROM schema_version")).scalar()
|
||||
# A non-fresh database with no stamp can only have been created by an
|
||||
# earlier create_all of this same codebase — i.e. already at LATEST.
|
||||
return version if version is not None else LATEST_VERSION
|
||||
|
||||
|
||||
def _set_version(conn, version: int) -> None:
|
||||
if conn.dialect.name == "sqlite":
|
||||
conn.execute(text(f"PRAGMA user_version = {version}"))
|
||||
return
|
||||
conn.execute(text(
|
||||
"CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL)"
|
||||
))
|
||||
if conn.execute(text("SELECT version FROM schema_version")).scalar() is None:
|
||||
conn.execute(
|
||||
text("INSERT INTO schema_version (version) VALUES (:v)"), {"v": version}
|
||||
)
|
||||
else:
|
||||
conn.execute(text("UPDATE schema_version SET version = :v"), {"v": version})
|
||||
|
||||
|
||||
def bootstrap(engine: Engine) -> None:
|
||||
fresh = not inspect(engine).get_table_names()
|
||||
Base.metadata.create_all(bind=engine)
|
||||
with engine.begin() as conn:
|
||||
if fresh:
|
||||
conn.execute(text(f"PRAGMA user_version = {LATEST_VERSION}"))
|
||||
_set_version(conn, LATEST_VERSION)
|
||||
return
|
||||
current = conn.execute(text("PRAGMA user_version")).scalar() or 1
|
||||
current = _get_version(conn)
|
||||
for version, sql in MIGRATIONS:
|
||||
if version > current:
|
||||
conn.execute(text(sql))
|
||||
current = version
|
||||
conn.execute(text(f"PRAGMA user_version = {current}"))
|
||||
_set_version(conn, current)
|
||||
_encrypt_plaintext_api_keys(conn)
|
||||
|
||||
|
||||
|
||||
@@ -2,12 +2,12 @@ import json
|
||||
import re
|
||||
import threading
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, memorybank, models, schemas
|
||||
from .. import auth, limits, memorybank, models, schemas
|
||||
from ..context import build_context
|
||||
from ..database import get_db
|
||||
from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError
|
||||
@@ -70,6 +70,7 @@ def create_adventure(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = CurrentUser,
|
||||
):
|
||||
limits.check_row_cap("adventures", db, user)
|
||||
scenario = None
|
||||
if payload.scenario_id is not None:
|
||||
scenario = db.get(models.Scenario, payload.scenario_id)
|
||||
@@ -207,6 +208,12 @@ def sse(obj: dict) -> str:
|
||||
return f"data: {json.dumps(obj)}\n\n"
|
||||
|
||||
|
||||
# no-cache defeats any intermediary caching; X-Accel-Buffering makes
|
||||
# nginx-style reverse proxies (hosted deploys) flush each event immediately
|
||||
# instead of buffering the stream.
|
||||
SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
|
||||
|
||||
|
||||
def action_json(action: models.Action) -> dict:
|
||||
return schemas.ActionOut.model_validate(action).model_dump(mode="json")
|
||||
|
||||
@@ -363,24 +370,32 @@ async def run_player_turn(
|
||||
def create_action(
|
||||
adventure_id: int,
|
||||
payload: schemas.ActionCreate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = CurrentUser,
|
||||
):
|
||||
adventure = get_adventure_or_404(adventure_id, db, user)
|
||||
limits.rate_limit("turn", request, user)
|
||||
limits.check_row_cap("actions", db, user, adventure=adventure)
|
||||
check_demo_cap(db, user)
|
||||
acquire_turn_lock(adventure_id)
|
||||
return StreamingResponse(
|
||||
with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)),
|
||||
media_type="text/event-stream",
|
||||
headers=SSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/{adventure_id}/retry")
|
||||
def retry_action(
|
||||
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
|
||||
adventure_id: int,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = CurrentUser,
|
||||
):
|
||||
"""Delete the last AI action and regenerate from the same input."""
|
||||
adventure = get_adventure_or_404(adventure_id, db, user)
|
||||
limits.rate_limit("turn", request, user)
|
||||
check_demo_cap(db, user)
|
||||
acquire_turn_lock(adventure_id)
|
||||
try:
|
||||
@@ -397,6 +412,7 @@ def retry_action(
|
||||
generate_turn(adventure, db, ScriptPipeline(adventure, db), user),
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers=SSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@@ -472,16 +488,26 @@ def export_adventure(
|
||||
|
||||
@router.post("/import", response_model=schemas.AdventureOut, status_code=201)
|
||||
def import_adventure(
|
||||
request: Request,
|
||||
bundle: dict = Body(...),
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = CurrentUser,
|
||||
):
|
||||
if bundle.get("format") != "ai-dnd-adventure-v1":
|
||||
raise HTTPException(400, "Not an adventure export file (expected format ai-dnd-adventure-v1).")
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("adventures", db, user)
|
||||
limits.check_bundle_lists(
|
||||
story_cards=bundle.get("storyCards"),
|
||||
memories=bundle.get("memories"),
|
||||
actions=bundle.get("actions"),
|
||||
)
|
||||
|
||||
# Raw-dict import bypasses the schemas — clamp strings headed for VARCHAR
|
||||
# columns (Postgres enforces the widths; see schemas.py).
|
||||
adventure = models.Adventure(
|
||||
user_id=user.id,
|
||||
title=str(bundle.get("title") or "Imported Adventure"),
|
||||
title=str(bundle.get("title") or "Imported Adventure")[:schemas.NAME_MAX],
|
||||
memory=str(bundle.get("memory") or ""),
|
||||
authors_note=str(bundle.get("authorsNote") or ""),
|
||||
ai_instructions=str(bundle.get("aiInstructions") or ""),
|
||||
@@ -511,8 +537,8 @@ def import_adventure(
|
||||
if isinstance(card, dict):
|
||||
db.add(models.StoryCard(
|
||||
adventure_id=adventure.id,
|
||||
type=str(card.get("type") or ""),
|
||||
name=str(card.get("name") or ""),
|
||||
type=str(card.get("type") or "")[:schemas.CARD_TYPE_MAX],
|
||||
name=str(card.get("name") or "")[:schemas.NAME_MAX],
|
||||
keys=str(card.get("keys") or ""),
|
||||
entry=str(card.get("entry") or ""),
|
||||
notes=str(card.get("notes") or ""),
|
||||
@@ -524,7 +550,7 @@ def import_adventure(
|
||||
adventure_id=adventure.id,
|
||||
position=int(s.get("position", i)),
|
||||
enabled=bool(s.get("enabled", True)),
|
||||
name=str(s.get("name") or "Imported Script"),
|
||||
name=str(s.get("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=str(s.get("description") or ""),
|
||||
library_js=str(s.get("library") or ""),
|
||||
input_js=str(s.get("input") or ""),
|
||||
@@ -537,7 +563,7 @@ def import_adventure(
|
||||
db.add(models.Action(
|
||||
adventure_id=adventure.id,
|
||||
index=int(a.get("index", i)),
|
||||
type=str(a.get("type") or "story"),
|
||||
type=str(a.get("type") or "story")[:20], # VARCHAR(20)
|
||||
text=str(a["text"]),
|
||||
reasoning=str(a["reasoning"]) if a.get("reasoning") else None,
|
||||
))
|
||||
@@ -631,6 +657,7 @@ def create_memory(
|
||||
):
|
||||
"""Manually add a memory; it gets embedded by the next post-turn pass."""
|
||||
adventure = get_adventure_or_404(adventure_id, db, user)
|
||||
limits.check_row_cap("memories", db, user, adventure=adventure)
|
||||
if not payload.text.strip():
|
||||
raise HTTPException(400, "Memory text cannot be empty")
|
||||
memory = models.Memory(adventure_id=adventure.id, text=payload.text.strip())
|
||||
|
||||
@@ -3,7 +3,7 @@ import re
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, models, schemas, security
|
||||
from .. import auth, limits, models, schemas, security
|
||||
from ..database import get_db
|
||||
from .settings import get_settings
|
||||
|
||||
@@ -53,6 +53,8 @@ def me(request: Request, response: Response, db: Session = Depends(get_db)):
|
||||
else:
|
||||
user = auth.resolve_session_user(request, db)
|
||||
if user is None:
|
||||
# Each new guest is a database row — cap how fast one IP can mint them.
|
||||
limits.rate_limit("guest", request)
|
||||
user = models.User(is_guest=True)
|
||||
db.add(user)
|
||||
db.commit()
|
||||
@@ -71,7 +73,7 @@ def register(
|
||||
scenario, script and setting they created as a guest is kept."""
|
||||
if not auth.MULTI_USER:
|
||||
raise HTTPException(400, "Accounts are disabled in local mode.")
|
||||
auth.rate_limit_auth(request)
|
||||
limits.rate_limit("auth", request)
|
||||
email = payload.email.strip().lower()
|
||||
if not EMAIL_RE.match(email):
|
||||
raise HTTPException(422, "Enter a valid email address.")
|
||||
@@ -99,7 +101,7 @@ def login(
|
||||
session is simply abandoned (its data stays under the guest user)."""
|
||||
if not auth.MULTI_USER:
|
||||
raise HTTPException(400, "Accounts are disabled in local mode.")
|
||||
auth.rate_limit_auth(request)
|
||||
limits.rate_limit("auth", request)
|
||||
email = payload.email.strip().lower()
|
||||
user = db.query(models.User).filter(models.User.email == email).first()
|
||||
if (
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, models, schemas
|
||||
from .. import auth, limits, models, schemas
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/scenarios", tags=["scenarios"])
|
||||
@@ -39,6 +39,7 @@ def create_scenario(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
limits.check_row_cap("scenarios", db, user)
|
||||
scenario = models.Scenario(**payload.model_dump(), user_id=user.id)
|
||||
db.add(scenario)
|
||||
db.commit()
|
||||
@@ -141,12 +142,15 @@ _IGNORED_KEYS = {"format", "storyCards", "worldInfo", "worldInformation", "scrip
|
||||
|
||||
@router.post("/import", status_code=201)
|
||||
def import_scenario(
|
||||
request: Request,
|
||||
bundle: dict = Body(...),
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Accepts our export format and AI Dungeon scenario exports best-effort;
|
||||
reports any keys it didn't understand."""
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("scenarios", db, user)
|
||||
fields: dict = {}
|
||||
unmapped: list[str] = []
|
||||
for key, value in bundle.items():
|
||||
@@ -164,6 +168,10 @@ def import_scenario(
|
||||
scenario = models.Scenario(**fields, user_id=user.id)
|
||||
if not scenario.title:
|
||||
scenario.title = "Imported Scenario"
|
||||
# Raw-dict import bypasses the schemas — clamp to VARCHAR widths
|
||||
# (Postgres enforces them; see schemas.py).
|
||||
scenario.title = scenario.title[:schemas.NAME_MAX]
|
||||
scenario.tags = scenario.tags[:schemas.TAGS_MAX]
|
||||
db.add(scenario)
|
||||
db.flush()
|
||||
|
||||
@@ -174,14 +182,15 @@ def import_scenario(
|
||||
or bundle.get("worldInformation")
|
||||
or []
|
||||
)
|
||||
limits.check_bundle_lists(story_cards=cards)
|
||||
for card in cards:
|
||||
if not isinstance(card, dict):
|
||||
continue
|
||||
db.add(
|
||||
models.StoryCard(
|
||||
scenario_id=scenario.id,
|
||||
type=str(card.get("type") or ""),
|
||||
name=str(card.get("name") or card.get("title") or ""),
|
||||
type=str(card.get("type") or "")[:schemas.CARD_TYPE_MAX],
|
||||
name=str(card.get("name") or card.get("title") or "")[:schemas.NAME_MAX],
|
||||
keys=str(card.get("keys") or ""),
|
||||
# AI Dungeon world info uses "value"; story cards use "entry".
|
||||
entry=str(card.get("entry") or card.get("value") or ""),
|
||||
@@ -194,7 +203,7 @@ def import_scenario(
|
||||
continue
|
||||
script = models.Script(
|
||||
user_id=user.id,
|
||||
name=str(item.get("name") or "Imported Script"),
|
||||
name=str(item.get("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=str(item.get("description") or ""),
|
||||
library_js=str(item.get("library") or item.get("sharedLibrary") or ""),
|
||||
input_js=str(item.get("input") or item.get("onInput") or ""),
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, models, schemas
|
||||
from .. import auth, limits, models, schemas
|
||||
from ..database import get_db
|
||||
from ..scripting import run_hook
|
||||
|
||||
@@ -36,6 +36,7 @@ def create_script(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
limits.check_row_cap("scripts", db, user)
|
||||
script = models.Script(**payload.model_dump(), user_id=user.id)
|
||||
db.add(script)
|
||||
db.commit()
|
||||
@@ -79,11 +80,13 @@ def delete_script(
|
||||
def test_script(
|
||||
script_id: int,
|
||||
payload: schemas.ScriptTestRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Dry-run one hook against sample text — no AI call, no persistence."""
|
||||
script = get_script_or_404(script_id, db, user)
|
||||
limits.rate_limit("script-test", request, user)
|
||||
result = run_hook(
|
||||
script.library_js,
|
||||
getattr(script, HOOK_FIELDS[payload.hook]),
|
||||
@@ -125,11 +128,14 @@ def export_script(
|
||||
|
||||
@router.post("/import", response_model=schemas.ScriptOut, status_code=201)
|
||||
def import_script(
|
||||
request: Request,
|
||||
bundle: dict = Body(...),
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Accepts our export bundle; tolerates *_js key names too."""
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("scripts", db, user)
|
||||
def pick(*keys: str) -> str:
|
||||
for key in keys:
|
||||
value = bundle.get(key)
|
||||
@@ -139,7 +145,8 @@ def import_script(
|
||||
|
||||
script = models.Script(
|
||||
user_id=user.id,
|
||||
name=pick("name") or "Imported Script",
|
||||
# Raw-dict import bypasses the schemas — clamp to the VARCHAR width.
|
||||
name=(pick("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=pick("description"),
|
||||
library_js=pick("library", "library_js", "sharedLibrary"),
|
||||
input_js=pick("input", "input_js", "onInput"),
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, models, schemas, security
|
||||
from .. import auth, limits, models, schemas, security
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/settings", tags=["settings"])
|
||||
@@ -64,12 +64,14 @@ def update_settings(
|
||||
|
||||
@router.post("/test")
|
||||
async def test_connection(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Hit the endpoint's /models listing as a cheap connectivity check.
|
||||
Tests whatever the turn engine would actually use — including the shared
|
||||
demo endpoint when the user has no key of their own."""
|
||||
limits.rate_limit("connection-test", request, user)
|
||||
settings = get_settings(db, user)
|
||||
cfg = auth.resolve_provider_config(settings)
|
||||
url = cfg.endpoint_url.rstrip("/") + "/models"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, models, schemas
|
||||
from .. import auth, limits, models, schemas
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/story-cards", tags=["story-cards"])
|
||||
@@ -50,6 +50,10 @@ def create_story_card(
|
||||
owner = db.get(owner_model, owner_id)
|
||||
if owner is None or owner.user_id != user.id:
|
||||
raise HTTPException(404, "Owner not found")
|
||||
limits.check_row_cap(
|
||||
"story_cards", db, user,
|
||||
scenario_id=payload.scenario_id, adventure_id=payload.adventure_id,
|
||||
)
|
||||
card = models.StoryCard(**payload.model_dump())
|
||||
db.add(card)
|
||||
db.commit()
|
||||
|
||||
+91
-68
@@ -1,7 +1,26 @@
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# Length caps (Phase 9). The VARCHAR ones are correctness, not just abuse
|
||||
# limits: Postgres enforces column lengths (SQLite never did), so anything
|
||||
# longer must be a 422 here rather than a 500 at INSERT. Text-column caps are
|
||||
# generous abuse ceilings a legitimate player won't hit.
|
||||
NAME_MAX = 200 # titles/names — VARCHAR(200)
|
||||
TAGS_MAX = 500 # VARCHAR(500)
|
||||
CARD_TYPE_MAX = 100 # VARCHAR(100)
|
||||
PROSE_MAX = 50_000 # memory, author's note, prompts, entries, notes...
|
||||
SCRIPT_MAX = 200_000 # one JS source
|
||||
ACTION_MAX = 20_000 # one player action
|
||||
MEMORY_TEXT_MAX = 5_000
|
||||
|
||||
Name = Annotated[str, Field(max_length=NAME_MAX)]
|
||||
Tags = Annotated[str, Field(max_length=TAGS_MAX)]
|
||||
CardType = Annotated[str, Field(max_length=CARD_TYPE_MAX)]
|
||||
Prose = Annotated[str, Field(max_length=PROSE_MAX)]
|
||||
ScriptSource = Annotated[str, Field(max_length=SCRIPT_MAX)]
|
||||
ActionText = Annotated[str, Field(max_length=ACTION_MAX)]
|
||||
|
||||
|
||||
class ORMModel(BaseModel):
|
||||
@@ -11,11 +30,11 @@ class ORMModel(BaseModel):
|
||||
# ---------- Story cards ----------
|
||||
|
||||
class StoryCardBase(BaseModel):
|
||||
type: str = ""
|
||||
name: str = ""
|
||||
keys: str = ""
|
||||
entry: str = ""
|
||||
notes: str = ""
|
||||
type: CardType = ""
|
||||
name: Name = ""
|
||||
keys: Prose = ""
|
||||
entry: Prose = ""
|
||||
notes: Prose = ""
|
||||
|
||||
|
||||
class StoryCardCreate(StoryCardBase):
|
||||
@@ -24,11 +43,11 @@ class StoryCardCreate(StoryCardBase):
|
||||
|
||||
|
||||
class StoryCardUpdate(BaseModel):
|
||||
type: str | None = None
|
||||
name: str | None = None
|
||||
keys: str | None = None
|
||||
entry: str | None = None
|
||||
notes: str | None = None
|
||||
type: CardType | None = None
|
||||
name: Name | None = None
|
||||
keys: Prose | None = None
|
||||
entry: Prose | None = None
|
||||
notes: Prose | None = None
|
||||
|
||||
|
||||
class StoryCardOut(ORMModel, StoryCardBase):
|
||||
@@ -40,13 +59,13 @@ class StoryCardOut(ORMModel, StoryCardBase):
|
||||
# ---------- Scenarios ----------
|
||||
|
||||
class ScenarioBase(BaseModel):
|
||||
title: str = "Untitled Scenario"
|
||||
description: str = ""
|
||||
prompt: str = ""
|
||||
memory: str = ""
|
||||
authors_note: str = ""
|
||||
ai_instructions: str = ""
|
||||
tags: str = ""
|
||||
title: Name = "Untitled Scenario"
|
||||
description: Prose = ""
|
||||
prompt: Prose = ""
|
||||
memory: Prose = ""
|
||||
authors_note: Prose = ""
|
||||
ai_instructions: Prose = ""
|
||||
tags: Tags = ""
|
||||
|
||||
|
||||
class ScenarioCreate(ScenarioBase):
|
||||
@@ -54,13 +73,13 @@ class ScenarioCreate(ScenarioBase):
|
||||
|
||||
|
||||
class ScenarioUpdate(BaseModel):
|
||||
title: str | None = None
|
||||
description: str | None = None
|
||||
prompt: str | None = None
|
||||
memory: str | None = None
|
||||
authors_note: str | None = None
|
||||
ai_instructions: str | None = None
|
||||
tags: str | None = None
|
||||
title: Name | None = None
|
||||
description: Prose | None = None
|
||||
prompt: Prose | None = None
|
||||
memory: Prose | None = None
|
||||
authors_note: Prose | None = None
|
||||
ai_instructions: Prose | None = None
|
||||
tags: Tags | None = None
|
||||
script_ids: list[int] | None = None
|
||||
|
||||
|
||||
@@ -86,17 +105,17 @@ class ScenarioListItem(ORMModel):
|
||||
|
||||
class AdventureCreate(BaseModel):
|
||||
scenario_id: int | None = None
|
||||
title: str | None = None
|
||||
title: Name | None = None
|
||||
# ${Placeholder} values collected from the player at start (AI Dungeon behavior).
|
||||
placeholders: dict[str, str] = {}
|
||||
|
||||
|
||||
class AdventureUpdate(BaseModel):
|
||||
title: str | None = None
|
||||
memory: str | None = None
|
||||
authors_note: str | None = None
|
||||
ai_instructions: str | None = None
|
||||
story_summary: str | None = None
|
||||
title: Name | None = None
|
||||
memory: Prose | None = None
|
||||
authors_note: Prose | None = None
|
||||
ai_instructions: Prose | None = None
|
||||
story_summary: Prose | None = None
|
||||
auto_summarize: bool | None = None
|
||||
memory_bank_enabled: bool | None = None
|
||||
|
||||
@@ -112,12 +131,12 @@ class ActionOut(ORMModel):
|
||||
|
||||
|
||||
class ActionUpdate(BaseModel):
|
||||
text: str
|
||||
text: ActionText
|
||||
|
||||
|
||||
class ActionCreate(BaseModel):
|
||||
type: Literal["do", "say", "story", "continue"]
|
||||
text: str = ""
|
||||
text: ActionText = ""
|
||||
|
||||
|
||||
class AdventureOut(ORMModel):
|
||||
@@ -153,11 +172,11 @@ class MemoryOut(ORMModel):
|
||||
|
||||
|
||||
class MemoryCreate(BaseModel):
|
||||
text: str
|
||||
text: Annotated[str, Field(max_length=MEMORY_TEXT_MAX)]
|
||||
|
||||
|
||||
class MemoryUpdate(BaseModel):
|
||||
text: str | None = None
|
||||
text: Annotated[str, Field(max_length=MEMORY_TEXT_MAX)] | None = None
|
||||
pinned: bool | None = None
|
||||
forgotten: bool | None = None
|
||||
|
||||
@@ -174,12 +193,12 @@ class AdventureListItem(ORMModel):
|
||||
# ---------- Scripts ----------
|
||||
|
||||
class ScriptBase(BaseModel):
|
||||
name: str = "Untitled Script"
|
||||
description: str = ""
|
||||
library_js: str = ""
|
||||
input_js: str = ""
|
||||
context_js: str = ""
|
||||
output_js: str = ""
|
||||
name: Name = "Untitled Script"
|
||||
description: Prose = ""
|
||||
library_js: ScriptSource = ""
|
||||
input_js: ScriptSource = ""
|
||||
context_js: ScriptSource = ""
|
||||
output_js: ScriptSource = ""
|
||||
|
||||
|
||||
class ScriptCreate(ScriptBase):
|
||||
@@ -187,12 +206,12 @@ class ScriptCreate(ScriptBase):
|
||||
|
||||
|
||||
class ScriptUpdate(BaseModel):
|
||||
name: str | None = None
|
||||
description: str | None = None
|
||||
library_js: str | None = None
|
||||
input_js: str | None = None
|
||||
context_js: str | None = None
|
||||
output_js: str | None = None
|
||||
name: Name | None = None
|
||||
description: Prose | None = None
|
||||
library_js: ScriptSource | None = None
|
||||
input_js: ScriptSource | None = None
|
||||
context_js: ScriptSource | None = None
|
||||
output_js: ScriptSource | None = None
|
||||
|
||||
|
||||
class ScriptOut(ORMModel, ScriptBase):
|
||||
@@ -203,7 +222,7 @@ class ScriptOut(ORMModel, ScriptBase):
|
||||
|
||||
class ScriptTestRequest(BaseModel):
|
||||
hook: Literal["input", "context", "output"]
|
||||
text: str = ""
|
||||
text: Prose = ""
|
||||
state: dict = {}
|
||||
|
||||
|
||||
@@ -222,17 +241,19 @@ class AdventureScriptOut(ORMModel):
|
||||
|
||||
class AdventureScriptUpdate(BaseModel):
|
||||
enabled: bool | None = None
|
||||
library_js: str | None = None
|
||||
input_js: str | None = None
|
||||
context_js: str | None = None
|
||||
output_js: str | None = None
|
||||
library_js: ScriptSource | None = None
|
||||
input_js: ScriptSource | None = None
|
||||
context_js: ScriptSource | None = None
|
||||
output_js: ScriptSource | None = None
|
||||
|
||||
|
||||
# ---------- Auth (Phase 8) ----------
|
||||
|
||||
class AuthCredentials(BaseModel):
|
||||
email: str
|
||||
password: str
|
||||
email: Annotated[str, Field(max_length=320)] # VARCHAR(320)
|
||||
# Upper bound keeps scrypt cost flat — hashing megabyte "passwords" is CPU
|
||||
# an attacker would otherwise get for free.
|
||||
password: Annotated[str, Field(max_length=128)]
|
||||
|
||||
|
||||
# ---------- Settings ----------
|
||||
@@ -259,17 +280,19 @@ ScenarioOut.model_rebuild()
|
||||
|
||||
|
||||
class SettingsUpdate(BaseModel):
|
||||
endpoint_url: str | None = None
|
||||
api_key: str | None = None
|
||||
model: str | None = None
|
||||
api_mode: str | None = None
|
||||
temperature: float | None = None
|
||||
max_output_tokens: int | None = None
|
||||
reasoning_max_tokens: int | None = None
|
||||
context_token_budget: int | None = None
|
||||
narrator_prompt: str | None = None
|
||||
endpoint_url: Annotated[str, Field(max_length=500)] | None = None # VARCHAR(500)
|
||||
# Encryption expands the stored value ~4/3 into the same VARCHAR(500):
|
||||
# 256 plaintext chars is the largest safe input ("enc:" + Fernet + base64).
|
||||
api_key: Annotated[str, Field(max_length=256)] | None = None
|
||||
model: Name | None = None
|
||||
api_mode: Annotated[str, Field(max_length=20)] | None = None
|
||||
temperature: Annotated[float, Field(ge=0, le=5)] | None = None
|
||||
max_output_tokens: Annotated[int, Field(ge=1, le=100_000)] | None = None
|
||||
reasoning_max_tokens: Annotated[int, Field(ge=0, le=100_000)] | None = None
|
||||
context_token_budget: Annotated[int, Field(ge=256, le=200_000)] | None = None
|
||||
narrator_prompt: Prose | None = None
|
||||
stream: bool | None = None
|
||||
summary_model: str | None = None
|
||||
embedding_model: str | None = None
|
||||
memory_bank_capacity: int | None = None
|
||||
memory_top_k: int | None = None
|
||||
summary_model: Name | None = None
|
||||
embedding_model: Name | None = None
|
||||
memory_bank_capacity: Annotated[int, Field(ge=1, le=1000)] | None = None
|
||||
memory_top_k: Annotated[int, Field(ge=1, le=50)] | None = None
|
||||
|
||||
+14
-2
@@ -7,7 +7,9 @@ Everything keys off one server-side secret:
|
||||
The secret comes from AIDND_SECRET_KEY, or is auto-generated once into
|
||||
`secret.key` next to the database so local installs and Docker volumes work
|
||||
with zero configuration (losing the file logs everyone out and orphans
|
||||
stored API keys — users just re-enter them).
|
||||
stored API keys — users just re-enter them). Multi-user deploys must set the
|
||||
env var: hosted filesystems are ephemeral, and a secret.key regenerated on
|
||||
every deploy would silently log out all users each time.
|
||||
|
||||
Passwords use hashlib.scrypt (stdlib, OpenSSL-backed) so we don't need a
|
||||
separate hashing dependency.
|
||||
@@ -27,9 +29,19 @@ _SECRET_FILE = DB_PATH.parent / "secret.key"
|
||||
|
||||
|
||||
def _load_secret() -> bytes:
|
||||
env = os.environ.get("AIDND_SECRET_KEY")
|
||||
env = os.environ.get("AIDND_SECRET_KEY", "").strip()
|
||||
if env:
|
||||
return env.encode()
|
||||
# Same flag parse as auth.MULTI_USER (auth imports this module, so it
|
||||
# can't be imported from there).
|
||||
if os.environ.get("AIDND_MULTI_USER", "").strip().lower() in ("1", "true", "yes", "on"):
|
||||
raise RuntimeError(
|
||||
"AIDND_SECRET_KEY must be set when AIDND_MULTI_USER is on: an "
|
||||
"auto-generated secret.key on an ephemeral hosted filesystem would "
|
||||
"rotate on every deploy, logging out every user and orphaning "
|
||||
"their stored API keys. Generate one with: "
|
||||
"python -c \"import secrets; print(secrets.token_urlsafe(48))\""
|
||||
)
|
||||
if _SECRET_FILE.exists():
|
||||
return _SECRET_FILE.read_bytes().strip()
|
||||
secret = secrets.token_urlsafe(48).encode()
|
||||
|
||||
@@ -5,3 +5,5 @@ pydantic>=2.7
|
||||
httpx>=0.27
|
||||
tiktoken>=0.7
|
||||
quickjs>=1.19
|
||||
cryptography>=42
|
||||
psycopg[binary]>=3.2
|
||||
|
||||
@@ -53,3 +53,26 @@ Render tier: free tier sleeps after idle + has no persistent disk).
|
||||
Running the production Docker image locally with `MULTI_USER=true`: a hostile user cannot hang
|
||||
the server with a `while(true)` script, cannot see another user's data or the debug log, gets
|
||||
rate-limited instead of burning the demo key, and the app streams turns normally the whole time.
|
||||
|
||||
### Verified 2026-07-07 (uvicorn, `MULTI_USER=1`, fresh SQLite DB, curl)
|
||||
|
||||
- **Fail-fast secret:** `import app.main` with `MULTI_USER=1` and no `AIDND_SECRET_KEY` raises
|
||||
the RuntimeError as designed (won't boot).
|
||||
- **`while(true)` script:** `POST /api/scripts/{id}/test` on an `input_js` infinite loop returns
|
||||
`InternalError: interrupted` (engine time limit) — server stays responsive afterward.
|
||||
- **Cross-user isolation:** guest B sees `[]` for scripts, gets 404 on guest A's script id;
|
||||
guest A keeps its own row. No leakage.
|
||||
- **Debug log:** `GET /api/debug/requests` → 403 in multi-user mode.
|
||||
- **Rate limiting:** 12 rapid `POST /api/auth/register` → 429 after the 10th (auth scope, 10/300s).
|
||||
- **Body size:** 3 MB body to `POST /api/scenarios` → 413 (limit 2 MB) via BodySizeLimitMiddleware.
|
||||
- **Security headers:** CSP, `x-frame-options: DENY`, `x-content-type-options: nosniff`,
|
||||
`referrer-policy: same-origin` on every response — including the SSE stream.
|
||||
- **Docs disabled:** Swagger UI and OpenAPI schema not served (`/docs`, `/openapi.json` fall
|
||||
through to the SPA `index.html`; no `swagger-ui`, no API schema exposed).
|
||||
- **SSE streaming:** `POST /api/adventures/{id}/actions` streams `text/event-stream` with
|
||||
`x-accel-buffering: no`, chunked, incremental events — the pure-ASGI middlewares don't buffer.
|
||||
(No LLM key configured here, so it streams the "No model configured" error event; a *live*
|
||||
provider turn through this path was verified end-to-end in Phase 8.)
|
||||
|
||||
All Phase 9 exit criteria met. Not yet exercised: Postgres (`DATABASE_URL`) path and the
|
||||
Docker production image specifically — both are Phase 10 deploy steps.
|
||||
|
||||
Reference in New Issue
Block a user