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:
parththakkar106
2026-07-07 12:30:06 +05:30
co-authored by Claude Opus 4.8
parent de4db373f2
commit 4772171b6c
17 changed files with 612 additions and 139 deletions
+6 -1
View File
@@ -33,4 +33,9 @@ VOLUME /data
EXPOSE 8000 EXPOSE 8000
WORKDIR /app/backend 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
View File
@@ -9,6 +9,22 @@
# Docker compose sets this to /data/data.db (a named volume). # Docker compose sets this to /data/data.db (a named volume).
AIDND_DB_PATH= 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 # Phase 8 — optional accounts & multi-user (all optional; defaults keep the
# app in frictionless single-user "local mode") # app in frictionless single-user "local mode")
@@ -19,12 +35,17 @@ AIDND_DB_PATH=
AIDND_MULTI_USER= AIDND_MULTI_USER=
# Secret for signing session cookies and encrypting stored API keys at rest. # 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 # If unset in local mode, one is auto-generated into `secret.key` next to the
# (fine for local/docker-volume runs). Set it explicitly on hosted deploys so # database (fine for local/docker-volume runs). REQUIRED when
# sessions survive redeploys when the disk is ephemeral or replaced. # 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= 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= AIDND_COOKIE_SECURE=
# --- Shared demo key (BYOK fallback; only active when AIDND_MULTI_USER=1) --- # --- 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 # 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. # 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
View File
@@ -16,8 +16,6 @@ whitelist and a per-day turn cap.
""" """
import os import os
import time
from collections import defaultdict, deque
from dataclasses import dataclass from dataclasses import dataclass
from datetime import timezone from datetime import timezone
@@ -35,7 +33,15 @@ def _env_flag(name: str) -> bool:
MULTI_USER = _env_flag("AIDND_MULTI_USER") MULTI_USER = _env_flag("AIDND_MULTI_USER")
SESSION_COOKIE = "aidnd_session" 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 COOKIE_MAX_AGE = 60 * 60 * 24 * 365
# ---------- Shared demo key (BYOK fallback) ---------- # ---------- 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.") raise HTTPException(401, "No session. Call GET /api/auth/me first.")
_touch(user, db) _touch(user, db)
return user 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)
+34 -3
View File
@@ -5,21 +5,50 @@ from sqlalchemy import create_engine, event
from sqlalchemy.orm import DeclarativeBase, sessionmaker from sqlalchemy.orm import DeclarativeBase, sessionmaker
# AIDND_DB_PATH lets deployments (Docker volume, hosted disk) relocate the # 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") _env_db_path = os.environ.get("AIDND_DB_PATH")
DB_PATH = ( DB_PATH = (
Path(_env_db_path).resolve() Path(_env_db_path).resolve()
if _env_db_path if _env_db_path
else Path(__file__).resolve().parent.parent / "data.db" else Path(__file__).resolve().parent.parent / "data.db"
) )
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
# 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()
)
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( engine = create_engine(
f"sqlite:///{DB_PATH}", f"sqlite:///{DB_PATH}",
connect_args={"check_same_thread": False}, connect_args={"check_same_thread": False},
) )
@event.listens_for(engine, "connect") @event.listens_for(engine, "connect")
def _enable_sqlite_foreign_keys(dbapi_connection, _record): def _enable_sqlite_foreign_keys(dbapi_connection, _record):
# SQLite ships with foreign keys OFF per connection; without this every # SQLite ships with foreign keys OFF per connection; without this every
@@ -27,6 +56,8 @@ def _enable_sqlite_foreign_keys(dbapi_connection, _record):
cursor = dbapi_connection.cursor() cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA foreign_keys=ON") cursor.execute("PRAGMA foreign_keys=ON")
cursor.close() cursor.close()
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
+221
View File
@@ -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
View File
@@ -1,3 +1,4 @@
import os
from pathlib import Path from pathlib import Path
from fastapi import FastAPI from fastapi import FastAPI
@@ -5,21 +6,80 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles from fastapi.staticfiles import StaticFiles
from starlette.exceptions import HTTPException as StarletteHTTPException from starlette.exceptions import HTTPException as StarletteHTTPException
from .auth import MULTI_USER
from .database import engine from .database import engine
from .limits import BodySizeLimitMiddleware
from .migrations import bootstrap from .migrations import bootstrap
from .routers import adventures, auth, debug, scenarios, scripts, settings, story_cards from .routers import adventures, auth, debug, scenarios, scripts, settings, story_cards
bootstrap(engine) 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( app.add_middleware(
CORSMiddleware, CORSMiddleware,
allow_origins=["http://localhost:5173", "http://127.0.0.1:5173"], allow_origins=CORS_ORIGINS,
allow_methods=["*"], allow_methods=["*"],
allow_headers=["*"], 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(auth.router)
app.include_router(scenarios.router) app.include_router(scenarios.router)
app.include_router(adventures.router) app.include_router(adventures.router)
+39 -6
View File
@@ -1,14 +1,20 @@
"""Lightweight versioned schema migrations over SQLite's PRAGMA user_version. """Lightweight versioned schema migrations.
How it works: How it works:
- A fresh database is created by `Base.metadata.create_all()` (always current) - A fresh database is created by `Base.metadata.create_all()` (always current)
and stamped with LATEST_VERSION. and stamped with LATEST_VERSION.
- An existing database runs every migration whose version is greater than its - 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 To change the schema: update models.py (keeps fresh DBs current) AND append a
(version, sql) pair here (upgrades existing DBs). Keep migrations idempotent (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 from sqlalchemy import inspect, text
@@ -70,19 +76,46 @@ MIGRATIONS: list[tuple[int, str]] = [
LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1) 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: def bootstrap(engine: Engine) -> None:
fresh = not inspect(engine).get_table_names() fresh = not inspect(engine).get_table_names()
Base.metadata.create_all(bind=engine) Base.metadata.create_all(bind=engine)
with engine.begin() as conn: with engine.begin() as conn:
if fresh: if fresh:
conn.execute(text(f"PRAGMA user_version = {LATEST_VERSION}")) _set_version(conn, LATEST_VERSION)
return return
current = conn.execute(text("PRAGMA user_version")).scalar() or 1 current = _get_version(conn)
for version, sql in MIGRATIONS: for version, sql in MIGRATIONS:
if version > current: if version > current:
conn.execute(text(sql)) conn.execute(text(sql))
current = version current = version
conn.execute(text(f"PRAGMA user_version = {current}")) _set_version(conn, current)
_encrypt_plaintext_api_keys(conn) _encrypt_plaintext_api_keys(conn)
+35 -8
View File
@@ -2,12 +2,12 @@ import json
import re import re
import threading import threading
from fastapi import APIRouter, Body, Depends, HTTPException from fastapi import APIRouter, Body, Depends, HTTPException, Request
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from sqlalchemy import func from sqlalchemy import func
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import auth, memorybank, models, schemas from .. import auth, limits, memorybank, models, schemas
from ..context import build_context from ..context import build_context
from ..database import get_db from ..database import get_db
from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError
@@ -70,6 +70,7 @@ def create_adventure(
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
): ):
limits.check_row_cap("adventures", db, user)
scenario = None scenario = None
if payload.scenario_id is not None: if payload.scenario_id is not None:
scenario = db.get(models.Scenario, payload.scenario_id) 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" 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: def action_json(action: models.Action) -> dict:
return schemas.ActionOut.model_validate(action).model_dump(mode="json") return schemas.ActionOut.model_validate(action).model_dump(mode="json")
@@ -363,24 +370,32 @@ async def run_player_turn(
def create_action( def create_action(
adventure_id: int, adventure_id: int,
payload: schemas.ActionCreate, payload: schemas.ActionCreate,
request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
): ):
adventure = get_adventure_or_404(adventure_id, db, user) 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) check_demo_cap(db, user)
acquire_turn_lock(adventure_id) acquire_turn_lock(adventure_id)
return StreamingResponse( return StreamingResponse(
with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)), with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)),
media_type="text/event-stream", media_type="text/event-stream",
headers=SSE_HEADERS,
) )
@router.post("/{adventure_id}/retry") @router.post("/{adventure_id}/retry")
def retry_action( 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.""" """Delete the last AI action and regenerate from the same input."""
adventure = get_adventure_or_404(adventure_id, db, user) adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user)
check_demo_cap(db, user) check_demo_cap(db, user)
acquire_turn_lock(adventure_id) acquire_turn_lock(adventure_id)
try: try:
@@ -397,6 +412,7 @@ def retry_action(
generate_turn(adventure, db, ScriptPipeline(adventure, db), user), generate_turn(adventure, db, ScriptPipeline(adventure, db), user),
), ),
media_type="text/event-stream", 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) @router.post("/import", response_model=schemas.AdventureOut, status_code=201)
def import_adventure( def import_adventure(
request: Request,
bundle: dict = Body(...), bundle: dict = Body(...),
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
): ):
if bundle.get("format") != "ai-dnd-adventure-v1": if bundle.get("format") != "ai-dnd-adventure-v1":
raise HTTPException(400, "Not an adventure export file (expected 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( adventure = models.Adventure(
user_id=user.id, 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 ""), memory=str(bundle.get("memory") or ""),
authors_note=str(bundle.get("authorsNote") or ""), authors_note=str(bundle.get("authorsNote") or ""),
ai_instructions=str(bundle.get("aiInstructions") or ""), ai_instructions=str(bundle.get("aiInstructions") or ""),
@@ -511,8 +537,8 @@ def import_adventure(
if isinstance(card, dict): if isinstance(card, dict):
db.add(models.StoryCard( db.add(models.StoryCard(
adventure_id=adventure.id, adventure_id=adventure.id,
type=str(card.get("type") or ""), type=str(card.get("type") or "")[:schemas.CARD_TYPE_MAX],
name=str(card.get("name") or ""), name=str(card.get("name") or "")[:schemas.NAME_MAX],
keys=str(card.get("keys") or ""), keys=str(card.get("keys") or ""),
entry=str(card.get("entry") or ""), entry=str(card.get("entry") or ""),
notes=str(card.get("notes") or ""), notes=str(card.get("notes") or ""),
@@ -524,7 +550,7 @@ def import_adventure(
adventure_id=adventure.id, adventure_id=adventure.id,
position=int(s.get("position", i)), position=int(s.get("position", i)),
enabled=bool(s.get("enabled", True)), 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 ""), description=str(s.get("description") or ""),
library_js=str(s.get("library") or ""), library_js=str(s.get("library") or ""),
input_js=str(s.get("input") or ""), input_js=str(s.get("input") or ""),
@@ -537,7 +563,7 @@ def import_adventure(
db.add(models.Action( db.add(models.Action(
adventure_id=adventure.id, adventure_id=adventure.id,
index=int(a.get("index", i)), 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"]), text=str(a["text"]),
reasoning=str(a["reasoning"]) if a.get("reasoning") else None, 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.""" """Manually add a memory; it gets embedded by the next post-turn pass."""
adventure = get_adventure_or_404(adventure_id, db, user) adventure = get_adventure_or_404(adventure_id, db, user)
limits.check_row_cap("memories", db, user, adventure=adventure)
if not payload.text.strip(): if not payload.text.strip():
raise HTTPException(400, "Memory text cannot be empty") raise HTTPException(400, "Memory text cannot be empty")
memory = models.Memory(adventure_id=adventure.id, text=payload.text.strip()) memory = models.Memory(adventure_id=adventure.id, text=payload.text.strip())
+5 -3
View File
@@ -3,7 +3,7 @@ import re
from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi import APIRouter, Depends, HTTPException, Request, Response
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import auth, models, schemas, security from .. import auth, limits, models, schemas, security
from ..database import get_db from ..database import get_db
from .settings import get_settings from .settings import get_settings
@@ -53,6 +53,8 @@ def me(request: Request, response: Response, db: Session = Depends(get_db)):
else: else:
user = auth.resolve_session_user(request, db) user = auth.resolve_session_user(request, db)
if user is None: 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) user = models.User(is_guest=True)
db.add(user) db.add(user)
db.commit() db.commit()
@@ -71,7 +73,7 @@ def register(
scenario, script and setting they created as a guest is kept.""" scenario, script and setting they created as a guest is kept."""
if not auth.MULTI_USER: if not auth.MULTI_USER:
raise HTTPException(400, "Accounts are disabled in local mode.") raise HTTPException(400, "Accounts are disabled in local mode.")
auth.rate_limit_auth(request) limits.rate_limit("auth", request)
email = payload.email.strip().lower() email = payload.email.strip().lower()
if not EMAIL_RE.match(email): if not EMAIL_RE.match(email):
raise HTTPException(422, "Enter a valid email address.") 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).""" session is simply abandoned (its data stays under the guest user)."""
if not auth.MULTI_USER: if not auth.MULTI_USER:
raise HTTPException(400, "Accounts are disabled in local mode.") raise HTTPException(400, "Accounts are disabled in local mode.")
auth.rate_limit_auth(request) limits.rate_limit("auth", request)
email = payload.email.strip().lower() email = payload.email.strip().lower()
user = db.query(models.User).filter(models.User.email == email).first() user = db.query(models.User).filter(models.User.email == email).first()
if ( if (
+14 -5
View File
@@ -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 import or_
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import auth, models, schemas from .. import auth, limits, models, schemas
from ..database import get_db from ..database import get_db
router = APIRouter(prefix="/api/scenarios", tags=["scenarios"]) router = APIRouter(prefix="/api/scenarios", tags=["scenarios"])
@@ -39,6 +39,7 @@ def create_scenario(
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), 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) scenario = models.Scenario(**payload.model_dump(), user_id=user.id)
db.add(scenario) db.add(scenario)
db.commit() db.commit()
@@ -141,12 +142,15 @@ _IGNORED_KEYS = {"format", "storyCards", "worldInfo", "worldInformation", "scrip
@router.post("/import", status_code=201) @router.post("/import", status_code=201)
def import_scenario( def import_scenario(
request: Request,
bundle: dict = Body(...), bundle: dict = Body(...),
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), user: models.User = Depends(auth.get_current_user),
): ):
"""Accepts our export format and AI Dungeon scenario exports best-effort; """Accepts our export format and AI Dungeon scenario exports best-effort;
reports any keys it didn't understand.""" reports any keys it didn't understand."""
limits.rate_limit("import", request, user)
limits.check_row_cap("scenarios", db, user)
fields: dict = {} fields: dict = {}
unmapped: list[str] = [] unmapped: list[str] = []
for key, value in bundle.items(): for key, value in bundle.items():
@@ -164,6 +168,10 @@ def import_scenario(
scenario = models.Scenario(**fields, user_id=user.id) scenario = models.Scenario(**fields, user_id=user.id)
if not scenario.title: if not scenario.title:
scenario.title = "Imported Scenario" 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.add(scenario)
db.flush() db.flush()
@@ -174,14 +182,15 @@ def import_scenario(
or bundle.get("worldInformation") or bundle.get("worldInformation")
or [] or []
) )
limits.check_bundle_lists(story_cards=cards)
for card in cards: for card in cards:
if not isinstance(card, dict): if not isinstance(card, dict):
continue continue
db.add( db.add(
models.StoryCard( models.StoryCard(
scenario_id=scenario.id, scenario_id=scenario.id,
type=str(card.get("type") or ""), type=str(card.get("type") or "")[:schemas.CARD_TYPE_MAX],
name=str(card.get("name") or card.get("title") or ""), name=str(card.get("name") or card.get("title") or "")[:schemas.NAME_MAX],
keys=str(card.get("keys") or ""), keys=str(card.get("keys") or ""),
# AI Dungeon world info uses "value"; story cards use "entry". # AI Dungeon world info uses "value"; story cards use "entry".
entry=str(card.get("entry") or card.get("value") or ""), entry=str(card.get("entry") or card.get("value") or ""),
@@ -194,7 +203,7 @@ def import_scenario(
continue continue
script = models.Script( script = models.Script(
user_id=user.id, 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 ""), description=str(item.get("description") or ""),
library_js=str(item.get("library") or item.get("sharedLibrary") or ""), library_js=str(item.get("library") or item.get("sharedLibrary") or ""),
input_js=str(item.get("input") or item.get("onInput") or ""), input_js=str(item.get("input") or item.get("onInput") or ""),
+10 -3
View File
@@ -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 sqlalchemy.orm import Session
from .. import auth, models, schemas from .. import auth, limits, models, schemas
from ..database import get_db from ..database import get_db
from ..scripting import run_hook from ..scripting import run_hook
@@ -36,6 +36,7 @@ def create_script(
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), 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) script = models.Script(**payload.model_dump(), user_id=user.id)
db.add(script) db.add(script)
db.commit() db.commit()
@@ -79,11 +80,13 @@ def delete_script(
def test_script( def test_script(
script_id: int, script_id: int,
payload: schemas.ScriptTestRequest, payload: schemas.ScriptTestRequest,
request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), user: models.User = Depends(auth.get_current_user),
): ):
"""Dry-run one hook against sample text — no AI call, no persistence.""" """Dry-run one hook against sample text — no AI call, no persistence."""
script = get_script_or_404(script_id, db, user) script = get_script_or_404(script_id, db, user)
limits.rate_limit("script-test", request, user)
result = run_hook( result = run_hook(
script.library_js, script.library_js,
getattr(script, HOOK_FIELDS[payload.hook]), getattr(script, HOOK_FIELDS[payload.hook]),
@@ -125,11 +128,14 @@ def export_script(
@router.post("/import", response_model=schemas.ScriptOut, status_code=201) @router.post("/import", response_model=schemas.ScriptOut, status_code=201)
def import_script( def import_script(
request: Request,
bundle: dict = Body(...), bundle: dict = Body(...),
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), user: models.User = Depends(auth.get_current_user),
): ):
"""Accepts our export bundle; tolerates *_js key names too.""" """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: def pick(*keys: str) -> str:
for key in keys: for key in keys:
value = bundle.get(key) value = bundle.get(key)
@@ -139,7 +145,8 @@ def import_script(
script = models.Script( script = models.Script(
user_id=user.id, 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"), description=pick("description"),
library_js=pick("library", "library_js", "sharedLibrary"), library_js=pick("library", "library_js", "sharedLibrary"),
input_js=pick("input", "input_js", "onInput"), input_js=pick("input", "input_js", "onInput"),
+4 -2
View File
@@ -1,8 +1,8 @@
import httpx import httpx
from fastapi import APIRouter, Depends from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import auth, models, schemas, security from .. import auth, limits, models, schemas, security
from ..database import get_db from ..database import get_db
router = APIRouter(prefix="/api/settings", tags=["settings"]) router = APIRouter(prefix="/api/settings", tags=["settings"])
@@ -64,12 +64,14 @@ def update_settings(
@router.post("/test") @router.post("/test")
async def test_connection( async def test_connection(
request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user), user: models.User = Depends(auth.get_current_user),
): ):
"""Hit the endpoint's /models listing as a cheap connectivity check. """Hit the endpoint's /models listing as a cheap connectivity check.
Tests whatever the turn engine would actually use — including the shared Tests whatever the turn engine would actually use — including the shared
demo endpoint when the user has no key of their own.""" demo endpoint when the user has no key of their own."""
limits.rate_limit("connection-test", request, user)
settings = get_settings(db, user) settings = get_settings(db, user)
cfg = auth.resolve_provider_config(settings) cfg = auth.resolve_provider_config(settings)
url = cfg.endpoint_url.rstrip("/") + "/models" url = cfg.endpoint_url.rstrip("/") + "/models"
+5 -1
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import auth, models, schemas from .. import auth, limits, models, schemas
from ..database import get_db from ..database import get_db
router = APIRouter(prefix="/api/story-cards", tags=["story-cards"]) router = APIRouter(prefix="/api/story-cards", tags=["story-cards"])
@@ -50,6 +50,10 @@ def create_story_card(
owner = db.get(owner_model, owner_id) owner = db.get(owner_model, owner_id)
if owner is None or owner.user_id != user.id: if owner is None or owner.user_id != user.id:
raise HTTPException(404, "Owner not found") 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()) card = models.StoryCard(**payload.model_dump())
db.add(card) db.add(card)
db.commit() db.commit()
+91 -68
View File
@@ -1,7 +1,26 @@
from datetime import datetime 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): class ORMModel(BaseModel):
@@ -11,11 +30,11 @@ class ORMModel(BaseModel):
# ---------- Story cards ---------- # ---------- Story cards ----------
class StoryCardBase(BaseModel): class StoryCardBase(BaseModel):
type: str = "" type: CardType = ""
name: str = "" name: Name = ""
keys: str = "" keys: Prose = ""
entry: str = "" entry: Prose = ""
notes: str = "" notes: Prose = ""
class StoryCardCreate(StoryCardBase): class StoryCardCreate(StoryCardBase):
@@ -24,11 +43,11 @@ class StoryCardCreate(StoryCardBase):
class StoryCardUpdate(BaseModel): class StoryCardUpdate(BaseModel):
type: str | None = None type: CardType | None = None
name: str | None = None name: Name | None = None
keys: str | None = None keys: Prose | None = None
entry: str | None = None entry: Prose | None = None
notes: str | None = None notes: Prose | None = None
class StoryCardOut(ORMModel, StoryCardBase): class StoryCardOut(ORMModel, StoryCardBase):
@@ -40,13 +59,13 @@ class StoryCardOut(ORMModel, StoryCardBase):
# ---------- Scenarios ---------- # ---------- Scenarios ----------
class ScenarioBase(BaseModel): class ScenarioBase(BaseModel):
title: str = "Untitled Scenario" title: Name = "Untitled Scenario"
description: str = "" description: Prose = ""
prompt: str = "" prompt: Prose = ""
memory: str = "" memory: Prose = ""
authors_note: str = "" authors_note: Prose = ""
ai_instructions: str = "" ai_instructions: Prose = ""
tags: str = "" tags: Tags = ""
class ScenarioCreate(ScenarioBase): class ScenarioCreate(ScenarioBase):
@@ -54,13 +73,13 @@ class ScenarioCreate(ScenarioBase):
class ScenarioUpdate(BaseModel): class ScenarioUpdate(BaseModel):
title: str | None = None title: Name | None = None
description: str | None = None description: Prose | None = None
prompt: str | None = None prompt: Prose | None = None
memory: str | None = None memory: Prose | None = None
authors_note: str | None = None authors_note: Prose | None = None
ai_instructions: str | None = None ai_instructions: Prose | None = None
tags: str | None = None tags: Tags | None = None
script_ids: list[int] | None = None script_ids: list[int] | None = None
@@ -86,17 +105,17 @@ class ScenarioListItem(ORMModel):
class AdventureCreate(BaseModel): class AdventureCreate(BaseModel):
scenario_id: int | None = None scenario_id: int | None = None
title: str | None = None title: Name | None = None
# ${Placeholder} values collected from the player at start (AI Dungeon behavior). # ${Placeholder} values collected from the player at start (AI Dungeon behavior).
placeholders: dict[str, str] = {} placeholders: dict[str, str] = {}
class AdventureUpdate(BaseModel): class AdventureUpdate(BaseModel):
title: str | None = None title: Name | None = None
memory: str | None = None memory: Prose | None = None
authors_note: str | None = None authors_note: Prose | None = None
ai_instructions: str | None = None ai_instructions: Prose | None = None
story_summary: str | None = None story_summary: Prose | None = None
auto_summarize: bool | None = None auto_summarize: bool | None = None
memory_bank_enabled: bool | None = None memory_bank_enabled: bool | None = None
@@ -112,12 +131,12 @@ class ActionOut(ORMModel):
class ActionUpdate(BaseModel): class ActionUpdate(BaseModel):
text: str text: ActionText
class ActionCreate(BaseModel): class ActionCreate(BaseModel):
type: Literal["do", "say", "story", "continue"] type: Literal["do", "say", "story", "continue"]
text: str = "" text: ActionText = ""
class AdventureOut(ORMModel): class AdventureOut(ORMModel):
@@ -153,11 +172,11 @@ class MemoryOut(ORMModel):
class MemoryCreate(BaseModel): class MemoryCreate(BaseModel):
text: str text: Annotated[str, Field(max_length=MEMORY_TEXT_MAX)]
class MemoryUpdate(BaseModel): class MemoryUpdate(BaseModel):
text: str | None = None text: Annotated[str, Field(max_length=MEMORY_TEXT_MAX)] | None = None
pinned: bool | None = None pinned: bool | None = None
forgotten: bool | None = None forgotten: bool | None = None
@@ -174,12 +193,12 @@ class AdventureListItem(ORMModel):
# ---------- Scripts ---------- # ---------- Scripts ----------
class ScriptBase(BaseModel): class ScriptBase(BaseModel):
name: str = "Untitled Script" name: Name = "Untitled Script"
description: str = "" description: Prose = ""
library_js: str = "" library_js: ScriptSource = ""
input_js: str = "" input_js: ScriptSource = ""
context_js: str = "" context_js: ScriptSource = ""
output_js: str = "" output_js: ScriptSource = ""
class ScriptCreate(ScriptBase): class ScriptCreate(ScriptBase):
@@ -187,12 +206,12 @@ class ScriptCreate(ScriptBase):
class ScriptUpdate(BaseModel): class ScriptUpdate(BaseModel):
name: str | None = None name: Name | None = None
description: str | None = None description: Prose | None = None
library_js: str | None = None library_js: ScriptSource | None = None
input_js: str | None = None input_js: ScriptSource | None = None
context_js: str | None = None context_js: ScriptSource | None = None
output_js: str | None = None output_js: ScriptSource | None = None
class ScriptOut(ORMModel, ScriptBase): class ScriptOut(ORMModel, ScriptBase):
@@ -203,7 +222,7 @@ class ScriptOut(ORMModel, ScriptBase):
class ScriptTestRequest(BaseModel): class ScriptTestRequest(BaseModel):
hook: Literal["input", "context", "output"] hook: Literal["input", "context", "output"]
text: str = "" text: Prose = ""
state: dict = {} state: dict = {}
@@ -222,17 +241,19 @@ class AdventureScriptOut(ORMModel):
class AdventureScriptUpdate(BaseModel): class AdventureScriptUpdate(BaseModel):
enabled: bool | None = None enabled: bool | None = None
library_js: str | None = None library_js: ScriptSource | None = None
input_js: str | None = None input_js: ScriptSource | None = None
context_js: str | None = None context_js: ScriptSource | None = None
output_js: str | None = None output_js: ScriptSource | None = None
# ---------- Auth (Phase 8) ---------- # ---------- Auth (Phase 8) ----------
class AuthCredentials(BaseModel): class AuthCredentials(BaseModel):
email: str email: Annotated[str, Field(max_length=320)] # VARCHAR(320)
password: str # 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 ---------- # ---------- Settings ----------
@@ -259,17 +280,19 @@ ScenarioOut.model_rebuild()
class SettingsUpdate(BaseModel): class SettingsUpdate(BaseModel):
endpoint_url: str | None = None endpoint_url: Annotated[str, Field(max_length=500)] | None = None # VARCHAR(500)
api_key: str | None = None # Encryption expands the stored value ~4/3 into the same VARCHAR(500):
model: str | None = None # 256 plaintext chars is the largest safe input ("enc:" + Fernet + base64).
api_mode: str | None = None api_key: Annotated[str, Field(max_length=256)] | None = None
temperature: float | None = None model: Name | None = None
max_output_tokens: int | None = None api_mode: Annotated[str, Field(max_length=20)] | None = None
reasoning_max_tokens: int | None = None temperature: Annotated[float, Field(ge=0, le=5)] | None = None
context_token_budget: int | None = None max_output_tokens: Annotated[int, Field(ge=1, le=100_000)] | None = None
narrator_prompt: str | 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 stream: bool | None = None
summary_model: str | None = None summary_model: Name | None = None
embedding_model: str | None = None embedding_model: Name | None = None
memory_bank_capacity: int | None = None memory_bank_capacity: Annotated[int, Field(ge=1, le=1000)] | None = None
memory_top_k: int | None = None memory_top_k: Annotated[int, Field(ge=1, le=50)] | None = None
+14 -2
View File
@@ -7,7 +7,9 @@ Everything keys off one server-side secret:
The secret comes from AIDND_SECRET_KEY, or is auto-generated once into 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 `secret.key` next to the database so local installs and Docker volumes work
with zero configuration (losing the file logs everyone out and orphans 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 Passwords use hashlib.scrypt (stdlib, OpenSSL-backed) so we don't need a
separate hashing dependency. separate hashing dependency.
@@ -27,9 +29,19 @@ _SECRET_FILE = DB_PATH.parent / "secret.key"
def _load_secret() -> bytes: def _load_secret() -> bytes:
env = os.environ.get("AIDND_SECRET_KEY") env = os.environ.get("AIDND_SECRET_KEY", "").strip()
if env: if env:
return env.encode() 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(): if _SECRET_FILE.exists():
return _SECRET_FILE.read_bytes().strip() return _SECRET_FILE.read_bytes().strip()
secret = secrets.token_urlsafe(48).encode() secret = secrets.token_urlsafe(48).encode()
+2
View File
@@ -5,3 +5,5 @@ pydantic>=2.7
httpx>=0.27 httpx>=0.27
tiktoken>=0.7 tiktoken>=0.7
quickjs>=1.19 quickjs>=1.19
cryptography>=42
psycopg[binary]>=3.2
+23
View File
@@ -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 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 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. 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.