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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user