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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01KFsGHju9szibJJa2YJcdbg
This commit is contained in:
parththakkar106
2026-07-06 23:04:03 +05:30
co-authored by Claude Fable 5
parent 253b533d3b
commit de4db373f2
26 changed files with 1247 additions and 189 deletions
+171
View File
@@ -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)
+2 -1
View File
@@ -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)
+11 -4
View File
@@ -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.
+36
View File
@@ -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},
)
+53 -1
View File
@@ -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)
+158 -46
View File
@@ -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")
+120
View File
@@ -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}
+9 -3
View File
@@ -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()
+55 -17
View File
@@ -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 ""),
+55 -18
View File
@@ -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"),
+43 -15
View File
@@ -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)
+40 -16
View File
@@ -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()
+11 -1
View File
@@ -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
+117
View File
@@ -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.<user_id>.<hmac>" — 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 ""