From de4db373f22249dfebb0b85b05fda4f55f99da8b Mon Sep 17 00:00:00 2001 From: parththakkar106 Date: Mon, 6 Jul 2026 23:04:03 +0530 Subject: [PATCH] Phase 8: optional accounts, per-user data, shared demo key Guest-first multi-user mode behind AIDND_MULTI_USER (local installs unchanged): signed-cookie guest sessions bootstrapped by /api/auth/me, register upgrades the guest in place, login/logout, per-IP rate limits. Every router scoped by user_id; Settings become per-user with the API key Fernet-encrypted at rest and write-only through the API. Users without a key get a server-funded demo key (OpenRouter free models, 20 turns/day, memory bank disabled on demo turns). Public read-only demo scenarios (seed_demo.py); debug log restricted to local mode. Frontend: auth modal + guest nudge, 401 re-establish/retry, demo banner and key management in Settings. Migrations 13-23 adopt existing data under a local user and encrypt stored keys. Verified: migration on a copy of real data.db, two-session isolation + register/login via curl and Chrome, demo cap 429, live OpenRouter turn through the encrypted-key path, vite build + oxlint. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01KFsGHju9szibJJa2YJcdbg --- .gitignore | 3 + README.md | 12 +- backend/.env.example | 37 ++++- backend/app/auth.py | 171 +++++++++++++++++++++ backend/app/main.py | 3 +- backend/app/memorybank.py | 15 +- backend/app/migrations.py | 36 +++++ backend/app/models.py | 54 ++++++- backend/app/routers/adventures.py | 204 ++++++++++++++++++++------ backend/app/routers/auth.py | 120 +++++++++++++++ backend/app/routers/debug.py | 12 +- backend/app/routers/scenarios.py | 72 ++++++--- backend/app/routers/scripts.py | 73 ++++++--- backend/app/routers/settings.py | 58 ++++++-- backend/app/routers/story_cards.py | 56 +++++-- backend/app/schemas.py | 12 +- backend/app/security.py | 117 +++++++++++++++ backend/seed_demo.py | 18 ++- frontend/src/App.jsx | 47 +++++- frontend/src/api.js | 27 +++- frontend/src/components.jsx | 62 ++++++++ frontend/src/index.css | 50 +++++++ frontend/src/pages/ScenarioEditor.jsx | 17 ++- frontend/src/pages/Scenarios.jsx | 1 + frontend/src/pages/Settings.jsx | 61 +++++++- plan/08-phase-accounts.md | 98 +++++++------ 26 files changed, 1247 insertions(+), 189 deletions(-) create mode 100644 backend/app/auth.py create mode 100644 backend/app/routers/auth.py create mode 100644 backend/app/security.py diff --git a/.gitignore b/.gitignore index f945aea..8712e3d 100644 --- a/.gitignore +++ b/.gitignore @@ -16,3 +16,6 @@ frontend/dist/ # Misc .claude/ + +# Phase 8: auto-generated session/encryption secret (lives next to the DB) +secret.key diff --git a/README.md b/README.md index 8464d64..a3edcee 100644 --- a/README.md +++ b/README.md @@ -33,6 +33,12 @@ models make the whole experience $0. (`backend/app/memorybank.py`). - **Import/export** — AI Dungeon-compatible formats for scripts and scenarios; JSON for everything. +- **Optional accounts for hosted deployments** — by default the app is single-user with zero + auth friction; set `AIDND_MULTI_USER=1` and visitors play instantly as guests (signed + session cookie), can register (email + password) at any point to keep their data, and each + user gets isolated data plus their own encrypted-at-rest API key. A server-funded **shared + demo key** with a daily turn cap lets people try it without bringing a key + (`backend/app/auth.py`). ## Quick start @@ -95,8 +101,10 @@ player input ``` frontend/ React + Vite SPA ──HTTP/SSE──► backend/ FastAPI - ├─ routers/ scenarios, adventures, story cards, scripts, settings, debug - ├─ models.py SQLAlchemy: Scenario, Adventure, Action, StoryCard, Script, Settings, Memory + ├─ routers/ auth, scenarios, adventures, story cards, scripts, settings, debug + ├─ models.py SQLAlchemy: User, Scenario, Adventure, Action, StoryCard, Script, Settings, Memory + ├─ auth.py guest/registered users, sessions, shared demo key + ├─ security.py password hashing, cookie signing, API-key encryption ├─ context/ prompt assembly under a token budget ├─ scripting/ quickjs sandbox + AI Dungeon API surface ├─ memorybank.py auto-summarization + embedding retrieval diff --git a/backend/.env.example b/backend/.env.example index 6edfe2c..2ecdc96 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -9,9 +9,36 @@ # Docker compose sets this to /data/data.db (a named volume). AIDND_DB_PATH= -# --- Coming in later phases (documented here as they land) --- -# Phase 8/9 will add: SECRET_KEY, MULTI_USER, DEMO_API_KEY, DEMO_ENDPOINT_URL, -# DEMO_MODEL_WHITELIST, DEMO_TURNS_PER_DAY, CORS_ORIGINS. -# +# --------------------------------------------------------------------------- +# Phase 8 — optional accounts & multi-user (all optional; defaults keep the +# app in frictionless single-user "local mode") +# --------------------------------------------------------------------------- + +# "1"/"true" turns on multi-user mode: guest sessions via signed cookies, +# register/login UI, per-user data. Leave unset for local installs. +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. +AIDND_SECRET_KEY= + +# "1" marks session cookies Secure (HTTPS-only). Turn on in production. +AIDND_COOKIE_SECURE= + +# --- Shared demo key (BYOK fallback; only active when AIDND_MULTI_USER=1) --- +# Users with no API key of their own get this server-funded endpoint with a +# model whitelist and a per-day turn cap. Unset = no demo, users must bring +# their own key. Memory bank/auto-summarization are disabled on demo turns. +AIDND_DEMO_API_KEY= +# Default endpoint if unset: https://openrouter.ai/api/v1 +AIDND_DEMO_ENDPOINT_URL= +# Comma-separated model whitelist. Default: google/gemma-4-26b-a4b-it:free +AIDND_DEMO_MODELS= +# Successful AI turns per user per day on the demo key. Default: 20 +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 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. diff --git a/backend/app/auth.py b/backend/app/auth.py new file mode 100644 index 0000000..47063a3 --- /dev/null +++ b/backend/app/auth.py @@ -0,0 +1,171 @@ +"""Phase 8 — user resolution, sessions, and the shared demo key. + +Two modes, chosen by the AIDND_MULTI_USER env var: + +- Local mode (default): every request resolves to one auto-created "local + user". No cookies, no login UI — a clone/docker-compose behaves exactly + like the pre-Phase-8 single-user app. +- Multi-user mode (hosted): requests carry a signed session cookie. GET + /api/auth/me creates a guest user on first visit; registering upgrades the + guest in place so their data survives. Requests without a valid session get + 401 and the frontend re-establishes via /me. + +The shared demo key (BYOK fallback) is also configured here: users whose +settings have no API key are routed to a server-funded endpoint with a model +whitelist and a per-day turn cap. +""" + +import os +import time +from collections import defaultdict, deque +from dataclasses import dataclass +from datetime import timezone + +from fastapi import Depends, HTTPException, Request +from sqlalchemy.orm import Session + +from . import models, security +from .database import get_db + + +def _env_flag(name: str) -> bool: + return os.environ.get(name, "").strip().lower() in ("1", "true", "yes", "on") + + +MULTI_USER = _env_flag("AIDND_MULTI_USER") + +SESSION_COOKIE = "aidnd_session" +COOKIE_SECURE = _env_flag("AIDND_COOKIE_SECURE") # enable behind HTTPS in prod +COOKIE_MAX_AGE = 60 * 60 * 24 * 365 + +# ---------- Shared demo key (BYOK fallback) ---------- + +DEMO_API_KEY = os.environ.get("AIDND_DEMO_API_KEY", "").strip() +DEMO_ENDPOINT_URL = ( + os.environ.get("AIDND_DEMO_ENDPOINT_URL", "").strip() + or "https://openrouter.ai/api/v1" +) +DEMO_MODELS = [ + m.strip() + for m in os.environ.get("AIDND_DEMO_MODELS", "").split(",") + if m.strip() +] or ["google/gemma-4-26b-a4b-it:free"] +DEMO_TURNS_PER_DAY = int(os.environ.get("AIDND_DEMO_TURNS_PER_DAY", "20") or 20) + +DEMO_CAP_MESSAGE = ( + f"You've used all {DEMO_TURNS_PER_DAY} free demo turns for today. " + "Add your own API key in Settings to keep playing (it resets tomorrow)." +) + + +def demo_enabled() -> bool: + # The demo key is a hosted-deployment feature; local installs talk to + # whatever endpoint Settings points at, even with no API key (Ollama). + return MULTI_USER and bool(DEMO_API_KEY) + + +@dataclass +class ProviderConfig: + """What the turn engine should actually connect with, after the + BYOK-vs-demo decision.""" + + endpoint_url: str + api_key: str + model: str + using_demo: bool + + +def resolve_provider_config(settings: models.Settings) -> ProviderConfig: + key = settings.api_key_plain + if key or not demo_enabled(): + return ProviderConfig(settings.endpoint_url, key, settings.model, False) + model = settings.model if settings.model in DEMO_MODELS else DEMO_MODELS[0] + return ProviderConfig(DEMO_ENDPOINT_URL, DEMO_API_KEY, model, True) + + +def _today() -> str: + return models.utcnow().date().isoformat() + + +def demo_turns_left(user: models.User) -> int: + used = user.demo_turns_used if user.demo_turns_date == _today() else 0 + return max(0, DEMO_TURNS_PER_DAY - used) + + +def count_demo_turn(user: models.User) -> None: + """Record one demo turn; the caller's commit persists it.""" + today = _today() + if user.demo_turns_date != today: + user.demo_turns_date = today + user.demo_turns_used = 0 + user.demo_turns_used += 1 + + +# ---------- User resolution ---------- + +def local_user(db: Session) -> models.User: + """The single implicit user in local mode (owns pre-Phase-8 data via + migration; created lazily on a fresh database).""" + user = ( + db.query(models.User) + .filter(models.User.email.is_(None), models.User.is_guest.is_(False)) + .order_by(models.User.id) + .first() + ) + if user is None: + user = models.User(is_guest=False) + db.add(user) + db.commit() + return user + + +def _touch(user: models.User, db: Session) -> None: + now = models.utcnow() + last = user.last_seen_at + if last is not None and last.tzinfo is None: + # SQLite hands DateTime columns back naive; they were stored as UTC. + last = last.replace(tzinfo=timezone.utc) + if last is None or (now - last).total_seconds() > 3600: + user.last_seen_at = now + db.commit() + + +def resolve_session_user(request: Request, db: Session) -> models.User | None: + token = request.cookies.get(SESSION_COOKIE) + if not token: + return None + user_id = security.verify_session(token) + if user_id is None: + return None + return db.get(models.User, user_id) + + +def get_current_user(request: Request, db: Session = Depends(get_db)) -> models.User: + """Dependency used by every router. 401 in multi-user mode means the + frontend must (re)establish a session via GET /api/auth/me.""" + if not MULTI_USER: + user = local_user(db) + else: + user = resolve_session_user(request, db) + if user is None: + 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) diff --git a/backend/app/main.py b/backend/app/main.py index ffd4969..72eab8c 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -7,7 +7,7 @@ from starlette.exceptions import HTTPException as StarletteHTTPException from .database import engine from .migrations import bootstrap -from .routers import adventures, debug, scenarios, scripts, settings, story_cards +from .routers import adventures, auth, debug, scenarios, scripts, settings, story_cards bootstrap(engine) @@ -20,6 +20,7 @@ app.add_middleware( allow_headers=["*"], ) +app.include_router(auth.router) app.include_router(scenarios.router) app.include_router(adventures.router) app.include_router(story_cards.router) diff --git a/backend/app/memorybank.py b/backend/app/memorybank.py index fd49767..c1f9979 100644 --- a/backend/app/memorybank.py +++ b/backend/app/memorybank.py @@ -57,7 +57,7 @@ _tasks: set[asyncio.Task] = set() def summary_provider(settings: models.Settings) -> OpenAICompatibleProvider: return OpenAICompatibleProvider( settings.endpoint_url, - settings.api_key, + settings.api_key_plain, settings.summary_model or settings.model, settings.api_mode, settings.reasoning_max_tokens, @@ -66,7 +66,7 @@ def summary_provider(settings: models.Settings) -> OpenAICompatibleProvider: def embedding_provider(settings: models.Settings) -> OpenAICompatibleProvider: return OpenAICompatibleProvider( - settings.endpoint_url, settings.api_key, settings.embedding_model + settings.endpoint_url, settings.api_key_plain, settings.embedding_model ) @@ -165,8 +165,15 @@ async def run_post_turn(adventure_id: int) -> None: db = SessionLocal() try: adventure = db.get(models.Adventure, adventure_id) - settings = db.get(models.Settings, 1) - if adventure is None or settings is None: + if adventure is None: + return + # Settings are per-user (Phase 8): use the adventure owner's row. + settings = ( + db.query(models.Settings) + .filter(models.Settings.user_id == adventure.user_id) + .first() + ) + if settings is None: return # Undo/retry can shrink the action list below a stored cursor, which # would stall summarization until the story grew past it again. diff --git a/backend/app/migrations.py b/backend/app/migrations.py index 4e6f9ca..a7ee7b2 100644 --- a/backend/app/migrations.py +++ b/backend/app/migrations.py @@ -46,6 +46,25 @@ MIGRATIONS: list[tuple[int, str]] = [ # Reasoning-model support: separate thinking budget + stored reasoning text. (11, "ALTER TABLE settings ADD COLUMN reasoning_max_tokens INTEGER NOT NULL DEFAULT 0"), (12, "ALTER TABLE actions ADD COLUMN reasoning TEXT"), + # Phase 8: optional accounts. The `users` table itself comes from + # create_all; these adopt all pre-existing rows under a "local user" + # (id=1) so a single-user install keeps working unchanged. + (13, """ + INSERT INTO users (id, email, password_hash, is_guest, created_at, + demo_turns_used, demo_turns_date) + SELECT 1, NULL, NULL, 0, CURRENT_TIMESTAMP, 0, '' + WHERE NOT EXISTS (SELECT 1 FROM users) + """), + (14, "ALTER TABLE scenarios ADD COLUMN user_id INTEGER REFERENCES users(id) ON DELETE CASCADE"), + (15, "UPDATE scenarios SET user_id = 1"), + (16, "ALTER TABLE scenarios ADD COLUMN is_public BOOLEAN NOT NULL DEFAULT 0"), + (17, "ALTER TABLE scripts ADD COLUMN user_id INTEGER REFERENCES users(id) ON DELETE CASCADE"), + (18, "UPDATE scripts SET user_id = 1"), + (19, "ALTER TABLE adventures ADD COLUMN user_id INTEGER REFERENCES users(id) ON DELETE CASCADE"), + (20, "UPDATE adventures SET user_id = 1"), + (21, "ALTER TABLE settings ADD COLUMN user_id INTEGER REFERENCES users(id) ON DELETE CASCADE"), + (22, "UPDATE settings SET user_id = 1"), + (23, "CREATE UNIQUE INDEX IF NOT EXISTS ix_settings_user_id ON settings (user_id)"), ] LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1) @@ -64,3 +83,20 @@ def bootstrap(engine: Engine) -> None: conn.execute(text(sql)) current = version conn.execute(text(f"PRAGMA user_version = {current}")) + _encrypt_plaintext_api_keys(conn) + + +def _encrypt_plaintext_api_keys(conn) -> None: + """Phase 8 data migration (can't be plain SQL): API keys saved before + encryption-at-rest existed are stored bare; wrap them in Fernet. Runs on + every start but matches nothing once all rows carry the enc: prefix.""" + from . import security # deferred: security derives its key from DB_PATH setup + + rows = conn.execute(text( + "SELECT id, api_key FROM settings WHERE api_key != '' AND api_key NOT LIKE 'enc:%'" + )).all() + for row_id, plain in rows: + conn.execute( + text("UPDATE settings SET api_key = :key WHERE id = :id"), + {"key": security.encrypt_secret(plain), "id": row_id}, + ) diff --git a/backend/app/models.py b/backend/app/models.py index ef20ef6..ab6d7bf 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -12,6 +12,31 @@ def utcnow() -> datetime: return datetime.now(timezone.utc) +class User(Base): + """Phase 8 — optional accounts. + + Three kinds of rows share this table: + - the "local user" (email NULL, is_guest False): auto-created in + single-user/local mode; owns everything a pre-Phase-8 DB had; + - guests (email NULL, is_guest True): created on first visit in + multi-user mode, identified only by their session cookie; + - registered users (email set): a guest upgraded in place, so their + data survives registration with no re-parenting. + """ + + __tablename__ = "users" + + id: Mapped[int] = mapped_column(primary_key=True) + email: Mapped[str | None] = mapped_column(String(320), unique=True, nullable=True) + password_hash: Mapped[str | None] = mapped_column(String(300), nullable=True) + is_guest: Mapped[bool] = mapped_column(Boolean, default=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow) + last_seen_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + # Shared demo key usage (resets when the UTC date changes). + demo_turns_used: Mapped[int] = mapped_column(Integer, default=0) + demo_turns_date: Mapped[str] = mapped_column(String(10), default="") + + scenario_scripts = Table( "scenario_scripts", Base.metadata, @@ -24,6 +49,11 @@ class Scenario(Base): __tablename__ = "scenarios" id: Mapped[int] = mapped_column(primary_key=True) + # NULL owner + is_public = seeded demo content, readable by everyone. + user_id: Mapped[int | None] = mapped_column( + ForeignKey("users.id", ondelete="CASCADE"), nullable=True + ) + is_public: Mapped[bool] = mapped_column(Boolean, default=False) title: Mapped[str] = mapped_column(String(200), default="Untitled Scenario") description: Mapped[str] = mapped_column(Text, default="") prompt: Mapped[str] = mapped_column(Text, default="") @@ -46,6 +76,9 @@ class Adventure(Base): __tablename__ = "adventures" id: Mapped[int] = mapped_column(primary_key=True) + user_id: Mapped[int | None] = mapped_column( + ForeignKey("users.id", ondelete="CASCADE"), nullable=True + ) scenario_id: Mapped[int | None] = mapped_column( ForeignKey("scenarios.id", ondelete="SET NULL"), nullable=True ) @@ -157,6 +190,9 @@ class Script(Base): __tablename__ = "scripts" id: Mapped[int] = mapped_column(primary_key=True) + user_id: Mapped[int | None] = mapped_column( + ForeignKey("users.id", ondelete="CASCADE"), nullable=True + ) name: Mapped[str] = mapped_column(String(200), default="Untitled Script") description: Mapped[str] = mapped_column(Text, default="") library_js: Mapped[str] = mapped_column(Text, default="") @@ -191,8 +227,14 @@ class AdventureScript(Base): class Settings(Base): __tablename__ = "settings" - id: Mapped[int] = mapped_column(primary_key=True) # single row, id=1 + id: Mapped[int] = mapped_column(primary_key=True) + # Phase 8: one row per user (pre-Phase-8 DBs had a single id=1 row, which + # the migration assigns to the local user). + user_id: Mapped[int | None] = mapped_column( + ForeignKey("users.id", ondelete="CASCADE"), nullable=True, unique=True + ) endpoint_url: Mapped[str] = mapped_column(String(500), default="http://localhost:11434/v1") + # Fernet-encrypted at rest ("enc:..." — see security.py); use api_key_plain. api_key: Mapped[str] = mapped_column(String(500), default="") model: Mapped[str] = mapped_column(String(200), default="") api_mode: Mapped[str] = mapped_column(String(20), default="chat") # chat|completion @@ -219,3 +261,13 @@ class Settings(Base): embedding_model: Mapped[str] = mapped_column(String(200), default="") # "" = bank disabled memory_bank_capacity: Mapped[int] = mapped_column(Integer, default=200) memory_top_k: Mapped[int] = mapped_column(Integer, default=5) + + @property + def has_api_key(self) -> bool: + return bool(self.api_key) + + @property + def api_key_plain(self) -> str: + from . import security # local import: models is imported before security + + return security.decrypt_secret(self.api_key) diff --git a/backend/app/routers/adventures.py b/backend/app/routers/adventures.py index f2c6c3c..2446123 100644 --- a/backend/app/routers/adventures.py +++ b/backend/app/routers/adventures.py @@ -7,7 +7,7 @@ from fastapi.responses import StreamingResponse from sqlalchemy import func from sqlalchemy.orm import Session -from .. import memorybank, models, schemas +from .. import auth, memorybank, models, schemas from ..context import build_context from ..database import get_db from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError @@ -16,20 +16,25 @@ from .settings import get_settings router = APIRouter(prefix="/api/adventures", tags=["adventures"]) +CurrentUser = Depends(auth.get_current_user) -def get_adventure_or_404(adventure_id: int, db: Session) -> models.Adventure: + +def get_adventure_or_404( + adventure_id: int, db: Session, user: models.User +) -> models.Adventure: adventure = db.get(models.Adventure, adventure_id) - if adventure is None: + if adventure is None or adventure.user_id != user.id: raise HTTPException(404, "Adventure not found") return adventure @router.get("", response_model=list[schemas.AdventureListItem]) -def list_adventures(db: Session = Depends(get_db)): +def list_adventures(db: Session = Depends(get_db), user: models.User = CurrentUser): rows = ( db.query(models.Adventure, func.count(models.Action.id), models.Scenario.title) .outerjoin(models.Action) .outerjoin(models.Scenario, models.Adventure.scenario_id == models.Scenario.id) + .filter(models.Adventure.user_id == user.id) .group_by(models.Adventure.id) .order_by(models.Adventure.updated_at.desc()) .all() @@ -60,15 +65,21 @@ def fill_placeholders(text: str, values: dict[str, str]) -> str: @router.post("", response_model=schemas.AdventureOut, status_code=201) -def create_adventure(payload: schemas.AdventureCreate, db: Session = Depends(get_db)): +def create_adventure( + payload: schemas.AdventureCreate, + db: Session = Depends(get_db), + user: models.User = CurrentUser, +): scenario = None if payload.scenario_id is not None: scenario = db.get(models.Scenario, payload.scenario_id) - if scenario is None: + # Playable = your own scenario or a shared demo one. + if scenario is None or (scenario.user_id != user.id and not scenario.is_public): raise HTTPException(404, "Scenario not found") values = payload.placeholders adventure = models.Adventure( + user_id=user.id, scenario_id=scenario.id if scenario else None, title=payload.title or (scenario.title if scenario else "Untitled Adventure"), memory=fill_placeholders(scenario.memory, values) if scenario else "", @@ -119,15 +130,20 @@ def create_adventure(payload: schemas.AdventureCreate, db: Session = Depends(get @router.get("/{adventure_id}", response_model=schemas.AdventureOut) -def get_adventure(adventure_id: int, db: Session = Depends(get_db)): - return get_adventure_or_404(adventure_id, db) +def get_adventure( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): + return get_adventure_or_404(adventure_id, db, user) @router.patch("/{adventure_id}", response_model=schemas.AdventureOut) def update_adventure( - adventure_id: int, payload: schemas.AdventureUpdate, db: Session = Depends(get_db) + adventure_id: int, + payload: schemas.AdventureUpdate, + db: Session = Depends(get_db), + user: models.User = CurrentUser, ): - adventure = get_adventure_or_404(adventure_id, db) + adventure = get_adventure_or_404(adventure_id, db, user) for field, value in payload.model_dump(exclude_unset=True).items(): setattr(adventure, field, value) db.commit() @@ -135,8 +151,10 @@ def update_adventure( @router.delete("/{adventure_id}", status_code=204) -def delete_adventure(adventure_id: int, db: Session = Depends(get_db)): - adventure = get_adventure_or_404(adventure_id, db) +def delete_adventure( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): + adventure = get_adventure_or_404(adventure_id, db, user) db.delete(adventure) db.commit() @@ -197,11 +215,26 @@ def next_index(adventure: models.Adventure) -> int: return max((a.index for a in adventure.actions), default=-1) + 1 -async def generate_turn(adventure: models.Adventure, db: Session, pipeline: ScriptPipeline): +async def generate_turn( + adventure: models.Adventure, + db: Session, + pipeline: ScriptPipeline, + user: models.User, +): """SSE generator: streams the AI continuation through the context/output script hooks, then stores the result.""" - settings = get_settings(db) - memories = await memorybank.retrieve_memories(adventure, settings, update_stats=True) + settings = get_settings(db, user) + cfg = auth.resolve_provider_config(settings) + if cfg.using_demo: + # No embedding/summarization calls on the server-funded key: memory + # retrieval is skipped (with a visible note when the bank is on). + memories = ( + {"used": [], "error": "Memory bank is unavailable on the shared demo key — add your own API key in Settings."} + if adventure.memory_bank_enabled + else None + ) + else: + memories = await memorybank.retrieve_memories(adventure, settings, update_stats=True) system_text, story_text, snapshot = build_context(adventure, settings, memories) # onModelContext: scripts see (and may rewrite) the whole assembled context. @@ -223,7 +256,7 @@ async def generate_turn(adventure: models.Adventure, db: Session, pipeline: Scri } provider = OpenAICompatibleProvider( - settings.endpoint_url, settings.api_key, settings.model, settings.api_mode, + cfg.endpoint_url, cfg.api_key, cfg.model, settings.api_mode, settings.reasoning_max_tokens, ) chunks: list[str] = [] @@ -264,15 +297,33 @@ async def generate_turn(adventure: models.Adventure, db: Session, pipeline: Scri ) db.add(ai_action) adventure.updated_at = models.utcnow() + if cfg.using_demo: + # Successful demo turns count against the daily cap (checked up front + # in the endpoint); failed provider calls above don't reach here. + auth.count_demo_turn(user) db.commit() db.refresh(ai_action) yield sse({"type": "done", "action": action_json(ai_action), "script": pipeline.report()}) - # Phase 6: fire-and-forget summarization/embedding (opens its own DB session). - memorybank.schedule_post_turn(adventure) + # Phase 6: fire-and-forget summarization/embedding (opens its own DB + # session). Skipped on the demo key — background AI calls would be + # unmetered spend on the server-funded key. + if not cfg.using_demo: + memorybank.schedule_post_turn(adventure) + + +def check_demo_cap(db: Session, user: models.User) -> None: + """409/429-style guard before a turn starts, so a capped player's input + isn't stored and then left without a reply.""" + settings = get_settings(db, user) + if auth.resolve_provider_config(settings).using_demo and auth.demo_turns_left(user) <= 0: + raise HTTPException(429, auth.DEMO_CAP_MESSAGE) async def run_player_turn( - adventure: models.Adventure, db: Session, payload: schemas.ActionCreate + adventure: models.Adventure, + db: Session, + payload: schemas.ActionCreate, + user: models.User, ): pipeline = ScriptPipeline(adventure, db) @@ -304,26 +355,33 @@ async def run_player_turn( yield sse({"type": "stopped", "script": pipeline.report()}) return - async for event in generate_turn(adventure, db, pipeline): + async for event in generate_turn(adventure, db, pipeline, user): yield event @router.post("/{adventure_id}/actions") def create_action( - adventure_id: int, payload: schemas.ActionCreate, db: Session = Depends(get_db) + adventure_id: int, + payload: schemas.ActionCreate, + db: Session = Depends(get_db), + user: models.User = CurrentUser, ): - adventure = get_adventure_or_404(adventure_id, db) + adventure = get_adventure_or_404(adventure_id, db, user) + check_demo_cap(db, user) acquire_turn_lock(adventure_id) return StreamingResponse( - with_turn_lock(adventure_id, run_player_turn(adventure, db, payload)), + with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)), media_type="text/event-stream", ) @router.post("/{adventure_id}/retry") -def retry_action(adventure_id: int, db: Session = Depends(get_db)): +def retry_action( + adventure_id: int, 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) + adventure = get_adventure_or_404(adventure_id, db, user) + check_demo_cap(db, user) acquire_turn_lock(adventure_id) try: if adventure.actions and adventure.actions[-1].type == "ai": @@ -334,15 +392,20 @@ def retry_action(adventure_id: int, db: Session = Depends(get_db)): _active_turns.discard(adventure_id) raise return StreamingResponse( - with_turn_lock(adventure_id, generate_turn(adventure, db, ScriptPipeline(adventure, db))), + with_turn_lock( + adventure_id, + generate_turn(adventure, db, ScriptPipeline(adventure, db), user), + ), media_type="text/event-stream", ) @router.post("/{adventure_id}/undo", response_model=list[schemas.ActionOut]) -def undo_turn(adventure_id: int, db: Session = Depends(get_db)): +def undo_turn( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): """Delete the last turn: the trailing AI action plus its player action, if any.""" - adventure = get_adventure_or_404(adventure_id, db) + adventure = get_adventure_or_404(adventure_id, db, user) actions = list(adventure.actions) if not actions or actions[-1].type == "start": raise HTTPException(400, "Nothing to undo") @@ -358,9 +421,11 @@ def undo_turn(adventure_id: int, db: Session = Depends(get_db)): # ---------- Import / Export ---------- @router.get("/{adventure_id}/export") -def export_adventure(adventure_id: int, db: Session = Depends(get_db)): +def export_adventure( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): """Full backup: plot components, story cards, scripts (+state), every action.""" - adv = get_adventure_or_404(adventure_id, db) + adv = get_adventure_or_404(adventure_id, db, user) return { "format": "ai-dnd-adventure-v1", "title": adv.title, @@ -406,11 +471,16 @@ def export_adventure(adventure_id: int, db: Session = Depends(get_db)): @router.post("/import", response_model=schemas.AdventureOut, status_code=201) -def import_adventure(bundle: dict = Body(...), db: Session = Depends(get_db)): +def import_adventure( + 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).") adventure = models.Adventure( + user_id=user.id, title=str(bundle.get("title") or "Imported Adventure"), memory=str(bundle.get("memory") or ""), authors_note=str(bundle.get("authorsNote") or ""), @@ -480,8 +550,10 @@ def import_adventure(bundle: dict = Body(...), db: Session = Depends(get_db)): # ---------- Adventure scripts ---------- @router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut]) -def list_adventure_scripts(adventure_id: int, db: Session = Depends(get_db)): - return get_adventure_or_404(adventure_id, db).scripts +def list_adventure_scripts( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): + return get_adventure_or_404(adventure_id, db, user).scripts @router.patch( @@ -492,7 +564,9 @@ def update_adventure_script( adv_script_id: int, payload: schemas.AdventureScriptUpdate, db: Session = Depends(get_db), + user: models.User = CurrentUser, ): + get_adventure_or_404(adventure_id, db, user) script = db.get(models.AdventureScript, adv_script_id) if script is None or script.adventure_id != adventure_id: raise HTTPException(404, "Script not found") @@ -505,17 +579,32 @@ def update_adventure_script( # ---------- Insights ---------- @router.get("/{adventure_id}/context") -async def dry_run_context(adventure_id: int, db: Session = Depends(get_db)): +async def dry_run_context( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): """What would be sent to the AI if the player continued right now.""" - adventure = get_adventure_or_404(adventure_id, db) - settings = get_settings(db) - memories = await memorybank.retrieve_memories(adventure, settings, update_stats=False) + adventure = get_adventure_or_404(adventure_id, db, user) + settings = get_settings(db, user) + if auth.resolve_provider_config(settings).using_demo: + memories = ( + {"used": [], "error": "Memory bank is unavailable on the shared demo key."} + if adventure.memory_bank_enabled + else None + ) + else: + memories = await memorybank.retrieve_memories(adventure, settings, update_stats=False) _, _, report = build_context(adventure, settings, memories) return report @router.get("/{adventure_id}/actions/{action_id}/context") -def action_context(adventure_id: int, action_id: int, db: Session = Depends(get_db)): +def action_context( + adventure_id: int, + action_id: int, + db: Session = Depends(get_db), + user: models.User = CurrentUser, +): + get_adventure_or_404(adventure_id, db, user) action = db.get(models.Action, action_id) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -527,16 +616,21 @@ def action_context(adventure_id: int, action_id: int, db: Session = Depends(get_ # ---------- Memory bank (Phase 6) ---------- @router.get("/{adventure_id}/memories", response_model=list[schemas.MemoryOut]) -def list_memories(adventure_id: int, db: Session = Depends(get_db)): - return get_adventure_or_404(adventure_id, db).memories +def list_memories( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): + return get_adventure_or_404(adventure_id, db, user).memories @router.post("/{adventure_id}/memories", response_model=schemas.MemoryOut, status_code=201) def create_memory( - adventure_id: int, payload: schemas.MemoryCreate, db: Session = Depends(get_db) + adventure_id: int, + payload: schemas.MemoryCreate, + db: Session = Depends(get_db), + user: models.User = CurrentUser, ): """Manually add a memory; it gets embedded by the next post-turn pass.""" - adventure = get_adventure_or_404(adventure_id, db) + adventure = get_adventure_or_404(adventure_id, db, user) if not payload.text.strip(): raise HTTPException(400, "Memory text cannot be empty") memory = models.Memory(adventure_id=adventure.id, text=payload.text.strip()) @@ -552,7 +646,9 @@ def update_memory( memory_id: int, payload: schemas.MemoryUpdate, db: Session = Depends(get_db), + user: models.User = CurrentUser, ): + get_adventure_or_404(adventure_id, db, user) memory = db.get(models.Memory, memory_id) if memory is None or memory.adventure_id != adventure_id: raise HTTPException(404, "Memory not found") @@ -566,7 +662,13 @@ def update_memory( @router.delete("/{adventure_id}/memories/{memory_id}", status_code=204) -def delete_memory(adventure_id: int, memory_id: int, db: Session = Depends(get_db)): +def delete_memory( + adventure_id: int, + memory_id: int, + db: Session = Depends(get_db), + user: models.User = CurrentUser, +): + get_adventure_or_404(adventure_id, db, user) memory = db.get(models.Memory, memory_id) if memory is None or memory.adventure_id != adventure_id: raise HTTPException(404, "Memory not found") @@ -577,8 +679,10 @@ def delete_memory(adventure_id: int, memory_id: int, db: Session = Depends(get_d # ---------- Actions (CRUD) ---------- @router.get("/{adventure_id}/actions", response_model=list[schemas.ActionOut]) -def list_actions(adventure_id: int, db: Session = Depends(get_db)): - get_adventure_or_404(adventure_id, db) +def list_actions( + adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser +): + get_adventure_or_404(adventure_id, db, user) return ( db.query(models.Action) .filter(models.Action.adventure_id == adventure_id) @@ -593,7 +697,9 @@ def update_action( action_id: int, payload: schemas.ActionUpdate, db: Session = Depends(get_db), + user: models.User = CurrentUser, ): + get_adventure_or_404(adventure_id, db, user) action = db.get(models.Action, action_id) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -603,7 +709,13 @@ def update_action( @router.delete("/{adventure_id}/actions/{action_id}", status_code=204) -def delete_action(adventure_id: int, action_id: int, db: Session = Depends(get_db)): +def delete_action( + adventure_id: int, + action_id: int, + db: Session = Depends(get_db), + user: models.User = CurrentUser, +): + get_adventure_or_404(adventure_id, db, user) action = db.get(models.Action, action_id) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") diff --git a/backend/app/routers/auth.py b/backend/app/routers/auth.py new file mode 100644 index 0000000..4141425 --- /dev/null +++ b/backend/app/routers/auth.py @@ -0,0 +1,120 @@ +import re + +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from sqlalchemy.orm import Session + +from .. import auth, models, schemas, security +from ..database import get_db +from .settings import get_settings + +router = APIRouter(prefix="/api/auth", tags=["auth"]) + +EMAIL_RE = re.compile(r"^[^@\s]+@[^@\s]+\.[^@\s]+$") + + +def _set_session_cookie(response: Response, user_id: int) -> None: + response.set_cookie( + auth.SESSION_COOKIE, + security.sign_session(user_id), + max_age=auth.COOKIE_MAX_AGE, + httponly=True, + samesite="lax", + secure=auth.COOKIE_SECURE, + path="/", + ) + + +def me_payload(user: models.User, db: Session) -> dict: + settings = get_settings(db, user) + cfg = auth.resolve_provider_config(settings) + return { + "multi_user": auth.MULTI_USER, + "id": user.id, + "email": user.email, + "is_guest": user.is_guest, + "demo": { + "enabled": auth.demo_enabled(), + "using_demo": cfg.using_demo, + "model": cfg.model if cfg.using_demo else None, + "turns_per_day": auth.DEMO_TURNS_PER_DAY, + "turns_left": auth.demo_turns_left(user) if auth.demo_enabled() else None, + "models": auth.DEMO_MODELS if auth.demo_enabled() else [], + }, + } + + +@router.get("/me") +def me(request: Request, response: Response, db: Session = Depends(get_db)): + """Who am I? In multi-user mode this also bootstraps the session: with no + (or an invalid) cookie it creates a guest user and sets one — the + frontend calls this on load and after any 401.""" + if not auth.MULTI_USER: + user = auth.local_user(db) + else: + user = auth.resolve_session_user(request, db) + if user is None: + user = models.User(is_guest=True) + db.add(user) + db.commit() + _set_session_cookie(response, user.id) + return me_payload(user, db) + + +@router.post("/register") +def register( + payload: schemas.AuthCredentials, + request: Request, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + """Upgrade the current guest in place — same user_id, so every adventure, + 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) + email = payload.email.strip().lower() + if not EMAIL_RE.match(email): + raise HTTPException(422, "Enter a valid email address.") + if len(payload.password) < 8: + raise HTTPException(422, "Password must be at least 8 characters.") + if not user.is_guest: + raise HTTPException(400, "This session is already registered.") + if db.query(models.User).filter(models.User.email == email).first(): + raise HTTPException(409, "An account with this email already exists — log in instead.") + user.email = email + user.password_hash = security.hash_password(payload.password) + user.is_guest = False + db.commit() + return me_payload(user, db) + + +@router.post("/login") +def login( + payload: schemas.AuthCredentials, + request: Request, + response: Response, + db: Session = Depends(get_db), +): + """Point this browser's session at an existing account. Any current guest + 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) + email = payload.email.strip().lower() + user = db.query(models.User).filter(models.User.email == email).first() + if ( + user is None + or not user.password_hash + or not security.verify_password(payload.password, user.password_hash) + ): + raise HTTPException(401, "Incorrect email or password.") + _set_session_cookie(response, user.id) + return me_payload(user, db) + + +@router.post("/logout") +def logout(response: Response): + if not auth.MULTI_USER: + raise HTTPException(400, "Accounts are disabled in local mode.") + response.delete_cookie(auth.SESSION_COOKIE, path="/") + return {"ok": True} diff --git a/backend/app/routers/debug.py b/backend/app/routers/debug.py index e5b697b..1d53667 100644 --- a/backend/app/routers/debug.py +++ b/backend/app/routers/debug.py @@ -1,11 +1,17 @@ -from fastapi import APIRouter +from fastapi import APIRouter, HTTPException -from .. import debuglog +from .. import auth, debuglog router = APIRouter(prefix="/api/debug", tags=["debug"]) @router.get("/requests") def recent_requests(): - """Most-recent-first log of provider requests/responses (no API keys).""" + """Most-recent-first log of provider requests/responses (no API keys). + + The log is a single process-wide ring buffer with no per-user + attribution, so in multi-user (hosted) mode it would leak other players' + prompts — disabled there, available on local installs.""" + if auth.MULTI_USER: + raise HTTPException(403, "The debug log is only available on local installs.") return debuglog.recent() diff --git a/backend/app/routers/scenarios.py b/backend/app/routers/scenarios.py index 8ee4a1b..61a016f 100644 --- a/backend/app/routers/scenarios.py +++ b/backend/app/routers/scenarios.py @@ -1,52 +1,77 @@ from fastapi import APIRouter, Body, Depends, HTTPException +from sqlalchemy import or_ from sqlalchemy.orm import Session -from .. import models, schemas +from .. import auth, models, schemas from ..database import get_db router = APIRouter(prefix="/api/scenarios", tags=["scenarios"]) -def get_scenario_or_404(scenario_id: int, db: Session) -> models.Scenario: +def get_scenario_or_404( + scenario_id: int, db: Session, user: models.User, *, edit: bool = False +) -> models.Scenario: + """Visible = owned or public; editable = owned only.""" scenario = db.get(models.Scenario, scenario_id) - if scenario is None: + if scenario is None or (scenario.user_id != user.id and not scenario.is_public): raise HTTPException(404, "Scenario not found") + if edit and scenario.user_id != user.id: + raise HTTPException(403, "This is a shared demo scenario — it can't be edited. Start an adventure from it, or duplicate it.") return scenario @router.get("", response_model=list[schemas.ScenarioListItem]) -def list_scenarios(db: Session = Depends(get_db)): +def list_scenarios( + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): return ( db.query(models.Scenario) + .filter(or_(models.Scenario.user_id == user.id, models.Scenario.is_public)) .order_by(models.Scenario.updated_at.desc()) .all() ) @router.post("", response_model=schemas.ScenarioOut, status_code=201) -def create_scenario(payload: schemas.ScenarioCreate, db: Session = Depends(get_db)): - scenario = models.Scenario(**payload.model_dump()) +def create_scenario( + payload: schemas.ScenarioCreate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + scenario = models.Scenario(**payload.model_dump(), user_id=user.id) db.add(scenario) db.commit() return scenario @router.get("/{scenario_id}", response_model=schemas.ScenarioOut) -def get_scenario(scenario_id: int, db: Session = Depends(get_db)): - return get_scenario_or_404(scenario_id, db) +def get_scenario( + scenario_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + return get_scenario_or_404(scenario_id, db, user) @router.patch("/{scenario_id}", response_model=schemas.ScenarioOut) def update_scenario( - scenario_id: int, payload: schemas.ScenarioUpdate, db: Session = Depends(get_db) + scenario_id: int, + payload: schemas.ScenarioUpdate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), ): - scenario = get_scenario_or_404(scenario_id, db) + scenario = get_scenario_or_404(scenario_id, db, user, edit=True) data = payload.model_dump(exclude_unset=True) script_ids = data.pop("script_ids", None) for field, value in data.items(): setattr(scenario, field, value) if script_ids is not None: - scripts = db.query(models.Script).filter(models.Script.id.in_(script_ids)).all() + scripts = ( + db.query(models.Script) + .filter(models.Script.id.in_(script_ids), models.Script.user_id == user.id) + .all() + ) if len(scripts) != len(set(script_ids)): raise HTTPException(404, "One or more scripts not found") scenario.scripts = sorted(scripts, key=lambda s: script_ids.index(s.id)) @@ -55,8 +80,12 @@ def update_scenario( @router.delete("/{scenario_id}", status_code=204) -def delete_scenario(scenario_id: int, db: Session = Depends(get_db)): - scenario = get_scenario_or_404(scenario_id, db) +def delete_scenario( + scenario_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + scenario = get_scenario_or_404(scenario_id, db, user, edit=True) db.delete(scenario) db.commit() @@ -64,8 +93,12 @@ def delete_scenario(scenario_id: int, db: Session = Depends(get_db)): # ---------- Import / Export ---------- @router.get("/{scenario_id}/export") -def export_scenario(scenario_id: int, db: Session = Depends(get_db)): - s = get_scenario_or_404(scenario_id, db) +def export_scenario( + scenario_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + s = get_scenario_or_404(scenario_id, db, user) return { "format": "ai-dnd-scenario-v1", "title": s.title, @@ -107,7 +140,11 @@ _IGNORED_KEYS = {"format", "storyCards", "worldInfo", "worldInformation", "scrip @router.post("/import", status_code=201) -def import_scenario(bundle: dict = Body(...), db: Session = Depends(get_db)): +def import_scenario( + 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.""" fields: dict = {} @@ -124,7 +161,7 @@ def import_scenario(bundle: dict = Body(...), db: Session = Depends(get_db)): elif isinstance(tags, str): fields["tags"] = tags - scenario = models.Scenario(**fields) + scenario = models.Scenario(**fields, user_id=user.id) if not scenario.title: scenario.title = "Imported Scenario" db.add(scenario) @@ -156,6 +193,7 @@ def import_scenario(bundle: dict = Body(...), db: Session = Depends(get_db)): if not isinstance(item, dict): continue script = models.Script( + user_id=user.id, name=str(item.get("name") or "Imported Script"), description=str(item.get("description") or ""), library_js=str(item.get("library") or item.get("sharedLibrary") or ""), diff --git a/backend/app/routers/scripts.py b/backend/app/routers/scripts.py index 67f389c..3e92453 100644 --- a/backend/app/routers/scripts.py +++ b/backend/app/routers/scripts.py @@ -1,7 +1,7 @@ from fastapi import APIRouter, Body, Depends, HTTPException from sqlalchemy.orm import Session -from .. import models, schemas +from .. import auth, models, schemas from ..database import get_db from ..scripting import run_hook @@ -10,34 +10,55 @@ router = APIRouter(prefix="/api/scripts", tags=["scripts"]) HOOK_FIELDS = {"input": "input_js", "context": "context_js", "output": "output_js"} -def get_script_or_404(script_id: int, db: Session) -> models.Script: +def get_script_or_404(script_id: int, db: Session, user: models.User) -> models.Script: script = db.get(models.Script, script_id) - if script is None: + if script is None or script.user_id != user.id: raise HTTPException(404, "Script not found") return script @router.get("", response_model=list[schemas.ScriptOut]) -def list_scripts(db: Session = Depends(get_db)): - return db.query(models.Script).order_by(models.Script.updated_at.desc()).all() +def list_scripts( + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + return ( + db.query(models.Script) + .filter(models.Script.user_id == user.id) + .order_by(models.Script.updated_at.desc()) + .all() + ) @router.post("", response_model=schemas.ScriptOut, status_code=201) -def create_script(payload: schemas.ScriptCreate, db: Session = Depends(get_db)): - script = models.Script(**payload.model_dump()) +def create_script( + payload: schemas.ScriptCreate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + script = models.Script(**payload.model_dump(), user_id=user.id) db.add(script) db.commit() return script @router.get("/{script_id}", response_model=schemas.ScriptOut) -def get_script(script_id: int, db: Session = Depends(get_db)): - return get_script_or_404(script_id, db) +def get_script( + script_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + return get_script_or_404(script_id, db, user) @router.patch("/{script_id}", response_model=schemas.ScriptOut) -def update_script(script_id: int, payload: schemas.ScriptUpdate, db: Session = Depends(get_db)): - script = get_script_or_404(script_id, db) +def update_script( + script_id: int, + payload: schemas.ScriptUpdate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + script = get_script_or_404(script_id, db, user) for field, value in payload.model_dump(exclude_unset=True).items(): setattr(script, field, value) db.commit() @@ -45,17 +66,24 @@ def update_script(script_id: int, payload: schemas.ScriptUpdate, db: Session = D @router.delete("/{script_id}", status_code=204) -def delete_script(script_id: int, db: Session = Depends(get_db)): - db.delete(get_script_or_404(script_id, db)) +def delete_script( + script_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + db.delete(get_script_or_404(script_id, db, user)) db.commit() @router.post("/{script_id}/test") def test_script( - script_id: int, payload: schemas.ScriptTestRequest, db: Session = Depends(get_db) + script_id: int, + payload: schemas.ScriptTestRequest, + 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) + script = get_script_or_404(script_id, db, user) result = run_hook( script.library_js, getattr(script, HOOK_FIELDS[payload.hook]), @@ -78,9 +106,13 @@ def test_script( # ---------- Import / Export ---------- @router.get("/{script_id}/export") -def export_script(script_id: int, db: Session = Depends(get_db)): +def export_script( + script_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): """JSON bundle matching how AI Dungeon scripts circulate.""" - script = get_script_or_404(script_id, db) + script = get_script_or_404(script_id, db, user) return { "name": script.name, "description": script.description, @@ -92,7 +124,11 @@ def export_script(script_id: int, db: Session = Depends(get_db)): @router.post("/import", response_model=schemas.ScriptOut, status_code=201) -def import_script(bundle: dict = Body(...), db: Session = Depends(get_db)): +def import_script( + 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.""" def pick(*keys: str) -> str: for key in keys: @@ -102,6 +138,7 @@ def import_script(bundle: dict = Body(...), db: Session = Depends(get_db)): return "" script = models.Script( + user_id=user.id, name=pick("name") or "Imported Script", description=pick("description"), library_js=pick("library", "library_js", "sharedLibrary"), diff --git a/backend/app/routers/settings.py b/backend/app/routers/settings.py index a25d173..fc41954 100644 --- a/backend/app/routers/settings.py +++ b/backend/app/routers/settings.py @@ -2,30 +2,44 @@ import httpx from fastapi import APIRouter, Depends from sqlalchemy.orm import Session -from .. import models, schemas +from .. import auth, models, schemas, security from ..database import get_db router = APIRouter(prefix="/api/settings", tags=["settings"]) -def get_settings(db: Session) -> models.Settings: - settings = db.get(models.Settings, 1) +def get_settings(db: Session, user: models.User) -> models.Settings: + """Per-user settings row, created on first access (Phase 8: settings — + endpoint, key, models, memory config — are per user, not global).""" + settings = ( + db.query(models.Settings).filter(models.Settings.user_id == user.id).first() + ) if settings is None: - settings = models.Settings(id=1) + settings = models.Settings(user_id=user.id) db.add(settings) db.commit() return settings @router.get("", response_model=schemas.SettingsOut) -def read_settings(db: Session = Depends(get_db)): - return get_settings(db) +def read_settings( + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + return get_settings(db, user) @router.put("", response_model=schemas.SettingsOut) -def update_settings(payload: schemas.SettingsUpdate, db: Session = Depends(get_db)): - settings = get_settings(db) +def update_settings( + payload: schemas.SettingsUpdate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + settings = get_settings(db, user) fields = payload.model_dump(exclude_unset=True) + # Write-only API key: absent = unchanged, "" = cleared, else encrypted. + if "api_key" in fields: + fields["api_key"] = security.encrypt_secret(fields["api_key"].strip()) embedding_model_changed = ( "embedding_model" in fields and fields["embedding_model"] != settings.embedding_model @@ -35,19 +49,33 @@ def update_settings(payload: schemas.SettingsUpdate, db: Session = Depends(get_d if embedding_model_changed: # Vectors from the old model have a different dimensionality/space; # clear them so the post-turn task re-embeds with the new model. - db.query(models.Memory).update({"embedding": None}) + # (This user's adventures only — settings are per-user now.) + owned = ( + db.query(models.Adventure.id) + .filter(models.Adventure.user_id == user.id) + .scalar_subquery() + ) + db.query(models.Memory).filter(models.Memory.adventure_id.in_(owned)).update( + {"embedding": None}, synchronize_session=False + ) db.commit() return settings @router.post("/test") -async def test_connection(db: Session = Depends(get_db)): - """Hit the endpoint's /models listing as a cheap connectivity check.""" - settings = get_settings(db) - url = settings.endpoint_url.rstrip("/") + "/models" +async def test_connection( + 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.""" + settings = get_settings(db, user) + cfg = auth.resolve_provider_config(settings) + url = cfg.endpoint_url.rstrip("/") + "/models" headers = {} - if settings.api_key: - headers["Authorization"] = f"Bearer {settings.api_key}" + if cfg.api_key: + headers["Authorization"] = f"Bearer {cfg.api_key}" try: async with httpx.AsyncClient(timeout=10) as client: resp = await client.get(url, headers=headers) diff --git a/backend/app/routers/story_cards.py b/backend/app/routers/story_cards.py index 906266e..74a4e32 100644 --- a/backend/app/routers/story_cards.py +++ b/backend/app/routers/story_cards.py @@ -1,33 +1,54 @@ from fastapi import APIRouter, Depends, HTTPException from sqlalchemy.orm import Session -from .. import models, schemas +from .. import auth, models, schemas from ..database import get_db router = APIRouter(prefix="/api/story-cards", tags=["story-cards"]) +def _card_editable_or_404(card: models.StoryCard | None, user: models.User) -> models.StoryCard: + """Cards inherit their scope from the owning scenario/adventure. Public + (demo) scenarios are visible to everyone but editable by no one.""" + if card is not None: + owner = card.scenario if card.scenario_id is not None else card.adventure + if owner is not None and owner.user_id == user.id: + return card + raise HTTPException(404, "Story card not found") + + @router.get("", response_model=list[schemas.StoryCardOut]) def list_story_cards( scenario_id: int | None = None, adventure_id: int | None = None, db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), ): - query = db.query(models.StoryCard) + if (scenario_id is None) == (adventure_id is None): + raise HTTPException(422, "Provide exactly one of scenario_id or adventure_id") if scenario_id is not None: - query = query.filter(models.StoryCard.scenario_id == scenario_id) - if adventure_id is not None: - query = query.filter(models.StoryCard.adventure_id == adventure_id) - return query.order_by(models.StoryCard.id).all() + scenario = db.get(models.Scenario, scenario_id) + if scenario is None or (scenario.user_id != user.id and not scenario.is_public): + raise HTTPException(404, "Owner not found") + return sorted(scenario.story_cards, key=lambda c: c.id) + adventure = db.get(models.Adventure, adventure_id) + if adventure is None or adventure.user_id != user.id: + raise HTTPException(404, "Owner not found") + return sorted(adventure.story_cards, key=lambda c: c.id) @router.post("", response_model=schemas.StoryCardOut, status_code=201) -def create_story_card(payload: schemas.StoryCardCreate, db: Session = Depends(get_db)): +def create_story_card( + payload: schemas.StoryCardCreate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): if (payload.scenario_id is None) == (payload.adventure_id is None): raise HTTPException(422, "Provide exactly one of scenario_id or adventure_id") owner_model = models.Scenario if payload.scenario_id else models.Adventure owner_id = payload.scenario_id or payload.adventure_id - if db.get(owner_model, owner_id) is None: + owner = db.get(owner_model, owner_id) + if owner is None or owner.user_id != user.id: raise HTTPException(404, "Owner not found") card = models.StoryCard(**payload.model_dump()) db.add(card) @@ -37,11 +58,12 @@ def create_story_card(payload: schemas.StoryCardCreate, db: Session = Depends(ge @router.patch("/{card_id}", response_model=schemas.StoryCardOut) def update_story_card( - card_id: int, payload: schemas.StoryCardUpdate, db: Session = Depends(get_db) + card_id: int, + payload: schemas.StoryCardUpdate, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), ): - card = db.get(models.StoryCard, card_id) - if card is None: - raise HTTPException(404, "Story card not found") + card = _card_editable_or_404(db.get(models.StoryCard, card_id), user) for field, value in payload.model_dump(exclude_unset=True).items(): setattr(card, field, value) db.commit() @@ -49,9 +71,11 @@ def update_story_card( @router.delete("/{card_id}", status_code=204) -def delete_story_card(card_id: int, db: Session = Depends(get_db)): - card = db.get(models.StoryCard, card_id) - if card is None: - raise HTTPException(404, "Story card not found") +def delete_story_card( + card_id: int, + db: Session = Depends(get_db), + user: models.User = Depends(auth.get_current_user), +): + card = _card_editable_or_404(db.get(models.StoryCard, card_id), user) db.delete(card) db.commit() diff --git a/backend/app/schemas.py b/backend/app/schemas.py index 66f0ef6..f8aed8c 100644 --- a/backend/app/schemas.py +++ b/backend/app/schemas.py @@ -66,6 +66,7 @@ class ScenarioUpdate(BaseModel): class ScenarioOut(ORMModel, ScenarioBase): id: int + is_public: bool = False # shared demo content — read-only for everyone created_at: datetime updated_at: datetime story_cards: list[StoryCardOut] = [] @@ -77,6 +78,7 @@ class ScenarioListItem(ORMModel): title: str description: str tags: str + is_public: bool = False updated_at: datetime @@ -226,11 +228,19 @@ class AdventureScriptUpdate(BaseModel): output_js: str | None = None +# ---------- Auth (Phase 8) ---------- + +class AuthCredentials(BaseModel): + email: str + password: str + + # ---------- Settings ---------- class SettingsOut(ORMModel): endpoint_url: str - api_key: str + # The key itself is never echoed back (encrypted at rest, write-only). + has_api_key: bool model: str api_mode: str temperature: float diff --git a/backend/app/security.py b/backend/app/security.py new file mode 100644 index 0000000..d50e1c6 --- /dev/null +++ b/backend/app/security.py @@ -0,0 +1,117 @@ +"""Phase 8 — secrets and crypto primitives for optional accounts. + +Everything keys off one server-side secret: + - session cookies are HMAC-signed with it, + - stored LLM API keys are Fernet-encrypted with a key derived from it. + +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). + +Passwords use hashlib.scrypt (stdlib, OpenSSL-backed) so we don't need a +separate hashing dependency. +""" + +import base64 +import hashlib +import hmac +import os +import secrets + +from cryptography.fernet import Fernet, InvalidToken + +from .database import DB_PATH + +_SECRET_FILE = DB_PATH.parent / "secret.key" + + +def _load_secret() -> bytes: + env = os.environ.get("AIDND_SECRET_KEY") + if env: + return env.encode() + if _SECRET_FILE.exists(): + return _SECRET_FILE.read_bytes().strip() + secret = secrets.token_urlsafe(48).encode() + _SECRET_FILE.write_bytes(secret) + return secret + + +SECRET_KEY = _load_secret() +_fernet = Fernet(base64.urlsafe_b64encode(hashlib.sha256(SECRET_KEY).digest())) + + +# ---------- Password hashing (scrypt) ---------- + +_SCRYPT_N, _SCRYPT_R, _SCRYPT_P = 2**14, 8, 1 + + +def hash_password(password: str) -> str: + salt = secrets.token_bytes(16) + key = hashlib.scrypt( + password.encode(), salt=salt, n=_SCRYPT_N, r=_SCRYPT_R, p=_SCRYPT_P + ) + return f"scrypt${_SCRYPT_N}${_SCRYPT_R}${_SCRYPT_P}${salt.hex()}${key.hex()}" + + +def verify_password(password: str, stored: str) -> bool: + try: + scheme, n, r, p, salt_hex, key_hex = stored.split("$") + if scheme != "scrypt": + return False + key = hashlib.scrypt( + password.encode(), salt=bytes.fromhex(salt_hex), + n=int(n), r=int(r), p=int(p), + ) + return hmac.compare_digest(key, bytes.fromhex(key_hex)) + except (ValueError, AttributeError): + return False + + +# ---------- Session tokens ---------- +# "v1.." — no expiry (long-lived guest sessions are the point). + +def sign_session(user_id: int) -> str: + payload = f"v1.{user_id}" + sig = hmac.new(SECRET_KEY, payload.encode(), hashlib.sha256).hexdigest() + return f"{payload}.{sig}" + + +def verify_session(token: str) -> int | None: + try: + version, user_id, sig = token.split(".") + if version != "v1": + return None + payload = f"{version}.{user_id}" + expected = hmac.new(SECRET_KEY, payload.encode(), hashlib.sha256).hexdigest() + if not hmac.compare_digest(sig, expected): + return None + return int(user_id) + except (ValueError, AttributeError): + return None + + +# ---------- API-key encryption at rest ---------- +# Stored values carry an "enc:" prefix so plaintext keys from pre-Phase-8 +# databases can be recognized and migrated. + +ENC_PREFIX = "enc:" + + +def encrypt_secret(plain: str) -> str: + if not plain: + return "" + return ENC_PREFIX + _fernet.encrypt(plain.encode()).decode() + + +def decrypt_secret(stored: str) -> str: + """Returns the plaintext key. Tolerates legacy plaintext values (returned + as-is) and undecryptable tokens (secret rotated → treated as unset).""" + if not stored: + return "" + if not stored.startswith(ENC_PREFIX): + return stored + try: + return _fernet.decrypt(stored[len(ENC_PREFIX):].encode()).decode() + except (InvalidToken, ValueError): + return "" diff --git a/backend/seed_demo.py b/backend/seed_demo.py index 0639da4..a18b3c0 100644 --- a/backend/seed_demo.py +++ b/backend/seed_demo.py @@ -2,9 +2,14 @@ Run from the backend folder: .venv\\Scripts\\python.exe seed_demo.py Safe to rerun: it deletes any previous rows titled "[Demo] ..." first. + +Phase 8: the scenario is seeded as PUBLIC (user_id NULL + is_public), so in +multi-user mode every guest sees it as read-only starter content. The sample +adventure and script-library copies belong to the local user (only relevant +on single-user installs). """ -from app import models, migrations +from app import auth, models, migrations from app.database import SessionLocal, engine # create_all + user_version stamp; plain create_all would leave a fresh DB at @@ -205,6 +210,8 @@ STORY_CARDS = [ db = SessionLocal() try: + owner = auth.local_user(db) + # Remove earlier demo rows so reruns stay clean. for adv in db.query(models.Adventure).filter(models.Adventure.title.like(f"{DEMO_PREFIX}%")): db.delete(adv) @@ -214,12 +221,14 @@ try: db.delete(s) db.commit() - # Script library + # Scripts attached to the public scenario are unowned (user_id NULL) so + # they ship with it everywhere; they're copied into each adventure at + # creation, so they never need to appear in anyone's script library. scripts = [models.Script(**s) for s in SCRIPTS] db.add_all(scripts) - # Scenario with cards and scripts attached - scenario = models.Scenario(**SCENARIO) + # Scenario with cards and scripts attached — public starter content. + scenario = models.Scenario(**SCENARIO, is_public=True) scenario.scripts = scripts db.add(scenario) db.flush() @@ -228,6 +237,7 @@ try: # Adventure created from the scenario, mirroring POST /api/adventures adventure = models.Adventure( + user_id=owner.id, scenario_id=scenario.id, title=scenario.title, memory=scenario.memory, diff --git a/frontend/src/App.jsx b/frontend/src/App.jsx index 01af37f..3b571ff 100644 --- a/frontend/src/App.jsx +++ b/frontend/src/App.jsx @@ -1,6 +1,32 @@ +import { useEffect, useState } from 'react' import { NavLink, Outlet } from 'react-router-dom' +import { api } from './api' +import { AuthModal } from './components' export default function App() { + // null until /auth/me resolves; in local mode multi_user=false hides all auth UI. + const [me, setMe] = useState(null) + const [authMode, setAuthMode] = useState(null) // 'register' | 'login' | null + + useEffect(() => { + api.getMe().then(setMe).catch(() => {}) + }, []) + + const onAuthed = (newMe, mode) => { + setAuthMode(null) + if (mode === 'login') { + // Different user now — reload so every page refetches its scoped data. + window.location.reload() + } else { + setMe(newMe) // register upgrades the same user in place; data unchanged + } + } + + const logout = async () => { + try { await api.logout() } catch { /* already logged out */ } + window.location.reload() + } + return ( <> - + + {authMode && ( + setAuthMode(null)} onAuthed={onAuthed} /> + )} ) } diff --git a/frontend/src/api.js b/frontend/src/api.js index 5b37a8b..34089e1 100644 --- a/frontend/src/api.js +++ b/frontend/src/api.js @@ -1,8 +1,19 @@ -async function request(path, options = {}) { +// Multi-user mode: a 401 means our session cookie is missing/stale. Hitting +// /api/auth/me creates a fresh guest session, after which the original call +// is retried once. +async function ensureSession() { + await fetch('/api/auth/me') +} + +async function request(path, options = {}, isRetry = false) { const resp = await fetch(`/api${path}`, { headers: { 'Content-Type': 'application/json' }, ...options, }) + if (resp.status === 401 && !isRetry && path !== '/auth/me') { + await ensureSession() + return request(path, options, true) + } if (!resp.ok) { let detail = resp.statusText try { @@ -16,13 +27,17 @@ async function request(path, options = {}) { } // POSTs to an SSE endpoint and dispatches events: {type: 'player'|'chunk'|'done'|'error', ...} -async function streamSSE(path, payload, onEvent, signal) { +async function streamSSE(path, payload, onEvent, signal, isRetry = false) { const resp = await fetch(`/api${path}`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify(payload), signal, }) + if (resp.status === 401 && !isRetry) { + await ensureSession() + return streamSSE(path, payload, onEvent, signal, true) + } if (!resp.ok) { let detail = resp.statusText try { detail = (await resp.json()).detail || detail } catch { /* non-JSON */ } @@ -46,6 +61,14 @@ async function streamSSE(path, payload, onEvent, signal) { } export const api = { + // Auth (Phase 8 — no-ops in local mode beyond getMe) + getMe: () => request('/auth/me'), + register: (email, password) => + request('/auth/register', { method: 'POST', body: JSON.stringify({ email, password }) }), + login: (email, password) => + request('/auth/login', { method: 'POST', body: JSON.stringify({ email, password }) }), + logout: () => request('/auth/logout', { method: 'POST' }), + // Scenarios listScenarios: () => request('/scenarios'), getScenario: (id) => request(`/scenarios/${id}`), diff --git a/frontend/src/components.jsx b/frontend/src/components.jsx index ef9a376..1bcf15a 100644 --- a/frontend/src/components.jsx +++ b/frontend/src/components.jsx @@ -1,4 +1,5 @@ import { useState } from 'react' +import { api } from './api' export function downloadJSON(obj, filename) { const blob = new Blob([JSON.stringify(obj, null, 2)], { type: 'application/json' }) @@ -75,6 +76,67 @@ export function PlaceholderModal({ title, names, onSubmit, onCancel }) { ) } +// Phase 8: register/login for the hosted multi-user mode. `onAuthed(me)` gets +// the fresh /auth/me payload after success. +export function AuthModal({ mode: initialMode, onClose, onAuthed }) { + const [mode, setMode] = useState(initialMode || 'register') + const [email, setEmail] = useState('') + const [password, setPassword] = useState('') + const [error, setError] = useState('') + const [busy, setBusy] = useState(false) + const registering = mode === 'register' + + const submit = async (e) => { + e.preventDefault() + setBusy(true) + setError('') + try { + const me = registering + ? await api.register(email, password) + : await api.login(email, password) + onAuthed(me, mode) + } catch (err) { + setError(err.message) + setBusy(false) + } + } + + return ( +
+
e.stopPropagation()} onSubmit={submit}> +

{registering ? 'Create an account' : 'Log in'}

+

+ {registering + ? 'Everything you’ve played as a guest stays with your new account, and you can pick it up from any device.' + : 'Welcome back — log in to reach your adventures.'} +

+ + + {error &&
{error}
} +
+ +
+ + +
+
+
+
+ ) +} + export function Field({ label, value, onChange, textarea, rows, placeholder }) { return (