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:
co-authored by
Claude Fable 5
parent
253b533d3b
commit
de4db373f2
+32
-5
@@ -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.
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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}
|
||||
@@ -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()
|
||||
|
||||
@@ -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 ""),
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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 ""
|
||||
+14
-4
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user