* Rewrite comments in Google developer documentation style Rewrite the comments and docstrings across the backend core modules so they read plainly. The previous prose was accurate but dense and figurative, which made it slow to skim. Applies the Google developer documentation style guide: short sentences, active voice, present tense, American spelling, and no metaphors, idioms, or rhetorical asides. Replaces em-dash chains with separate sentences.
272 lines
9.7 KiB
Python
272 lines
9.7 KiB
Python
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
|
from fastapi.responses import Response
|
|
from sqlalchemy import or_
|
|
from sqlalchemy.orm import Session
|
|
|
|
from .. import analytics, auth, images, limits, models, schemas
|
|
from ..database import get_db
|
|
|
|
router = APIRouter(prefix="/api/scenarios", tags=["scenarios"])
|
|
|
|
|
|
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 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),
|
|
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),
|
|
user: models.User = Depends(auth.get_current_user),
|
|
):
|
|
limits.check_row_cap("scenarios", db, user)
|
|
scenario = models.Scenario(**payload.model_dump(), user_id=user.id)
|
|
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),
|
|
user: models.User = Depends(auth.get_current_user),
|
|
):
|
|
scenario = get_scenario_or_404(scenario_id, db, user)
|
|
# A funnel step, recorded for shared scenarios only. Opening one is the
|
|
# first sign that a visitor is interested, and someone editing their own
|
|
# scenario is already past this point. Their titles are theirs rather than a
|
|
# statistic.
|
|
if scenario.is_public:
|
|
analytics.record_event(analytics.EV_SCENARIO_OPEN, user)
|
|
return scenario
|
|
|
|
|
|
@router.get("/{scenario_id}/image")
|
|
def get_scenario_image(
|
|
scenario_id: int,
|
|
db: Session = Depends(get_db),
|
|
user: models.User = Depends(auth.get_current_user),
|
|
):
|
|
"""Serve an uploaded cover image as real bytes.
|
|
|
|
Lists point here instead of inlining the data URI. The response is marked
|
|
immutable and the URL carries a `?v=<updated_at>` stamp, so browsers cache
|
|
it indefinitely but pick up a new picture the moment the author saves one.
|
|
"""
|
|
scenario = get_scenario_or_404(scenario_id, db, user)
|
|
decoded = images.decode(scenario.image)
|
|
if decoded is None:
|
|
raise HTTPException(404, "This scenario has no uploaded image")
|
|
data, content_type = decoded
|
|
return Response(
|
|
content=data,
|
|
media_type=content_type,
|
|
headers={"Cache-Control": "private, max-age=31536000, immutable"},
|
|
)
|
|
|
|
|
|
@router.patch("/{scenario_id}", response_model=schemas.ScenarioOut)
|
|
def update_scenario(
|
|
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, 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), 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))
|
|
db.commit()
|
|
return scenario
|
|
|
|
|
|
@router.delete("/{scenario_id}", status_code=204)
|
|
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()
|
|
|
|
|
|
# ---------- Import / Export ----------
|
|
|
|
@router.get("/{scenario_id}/export")
|
|
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,
|
|
"description": s.description,
|
|
"prompt": s.prompt,
|
|
"memory": s.memory,
|
|
"authorsNote": s.authors_note,
|
|
"aiInstructions": s.ai_instructions,
|
|
"tags": s.tags,
|
|
"image": s.image,
|
|
"icon": s.icon,
|
|
"statSchema": s.stat_schema,
|
|
"storyCards": [
|
|
{"type": c.type, "name": c.name, "keys": c.keys, "entry": c.entry, "notes": c.notes}
|
|
for c in s.story_cards
|
|
],
|
|
"scripts": [
|
|
{
|
|
"name": sc.name, "description": sc.description, "library": sc.library_js,
|
|
"input": sc.input_js, "context": sc.context_js, "output": sc.output_js,
|
|
}
|
|
for sc in s.scripts
|
|
],
|
|
}
|
|
|
|
|
|
# Key aliases seen in AI Dungeon scenario exports, mapped best-effort.
|
|
_SCENARIO_KEYS = {
|
|
"title": "title",
|
|
"description": "description",
|
|
"prompt": "prompt",
|
|
"memory": "memory",
|
|
"authorsNote": "authors_note",
|
|
"authors_note": "authors_note",
|
|
"authorsNoteText": "authors_note",
|
|
"aiInstructions": "ai_instructions",
|
|
"ai_instructions": "ai_instructions",
|
|
"instructions": "ai_instructions",
|
|
}
|
|
_IGNORED_KEYS = {"format", "storyCards", "worldInfo", "worldInformation", "scripts", "tags",
|
|
"statSchema", "stat_schema", "image", "icon",
|
|
"createdAt", "updatedAt", "id", "publicId", "nsfw", "type", "options"}
|
|
|
|
|
|
@router.post("/import", status_code=201)
|
|
def import_scenario(
|
|
request: Request,
|
|
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."""
|
|
limits.rate_limit("import", request, user)
|
|
limits.check_row_cap("scenarios", db, user)
|
|
fields: dict = {}
|
|
unmapped: list[str] = []
|
|
for key, value in bundle.items():
|
|
if key in _SCENARIO_KEYS and isinstance(value, str):
|
|
fields[_SCENARIO_KEYS[key]] = value
|
|
elif key not in _IGNORED_KEYS:
|
|
unmapped.append(key)
|
|
|
|
tags = bundle.get("tags")
|
|
if isinstance(tags, list):
|
|
fields["tags"] = ", ".join(str(t) for t in tags)
|
|
elif isinstance(tags, str):
|
|
fields["tags"] = tags
|
|
|
|
schema = bundle.get("statSchema") or bundle.get("stat_schema")
|
|
if isinstance(schema, dict):
|
|
fields["stat_schema"] = schema
|
|
|
|
# An AI Dungeon bundle also carries an `image`, so this reads it. The value
|
|
# is untrusted input, so it goes through `sanitize()` rather than a direct
|
|
# assignment.
|
|
image = images.sanitize(bundle.get("image"), schemas.IMAGE_MAX)
|
|
if image:
|
|
fields["image"] = image
|
|
icon = bundle.get("icon")
|
|
if isinstance(icon, str) and icon:
|
|
fields["icon"] = icon[:schemas.ICON_MAX]
|
|
|
|
scenario = models.Scenario(**fields, user_id=user.id)
|
|
if not scenario.title:
|
|
scenario.title = "Imported Scenario"
|
|
# A raw-dict import bypasses the schemas, so truncate to the VARCHAR widths.
|
|
# Postgres enforces them. See `schemas.py`. Column defaults have not been
|
|
# applied yet, because that happens at flush, so a bundle with no `tags` key
|
|
# leaves the attribute None, which is why the code says `or ""`.
|
|
scenario.title = scenario.title[:schemas.NAME_MAX]
|
|
scenario.tags = (scenario.tags or "")[:schemas.TAGS_MAX]
|
|
db.add(scenario)
|
|
db.flush()
|
|
|
|
# AI Dungeon exports have used all three names for the same list.
|
|
cards = (
|
|
bundle.get("storyCards")
|
|
or bundle.get("worldInfo")
|
|
or bundle.get("worldInformation")
|
|
or []
|
|
)
|
|
limits.check_bundle_lists(story_cards=cards)
|
|
for card in cards:
|
|
if not isinstance(card, dict):
|
|
continue
|
|
db.add(
|
|
models.StoryCard(
|
|
scenario_id=scenario.id,
|
|
type=str(card.get("type") or "")[:schemas.CARD_TYPE_MAX],
|
|
name=str(card.get("name") or card.get("title") or "")[:schemas.NAME_MAX],
|
|
keys=str(card.get("keys") or ""),
|
|
# AI Dungeon world info uses "value"; story cards use "entry".
|
|
entry=str(card.get("entry") or card.get("value") or ""),
|
|
notes=str(card.get("notes") or card.get("description") or ""),
|
|
)
|
|
)
|
|
|
|
for item in bundle.get("scripts") or []:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
script = models.Script(
|
|
user_id=user.id,
|
|
name=str(item.get("name") or "Imported Script")[:schemas.NAME_MAX],
|
|
description=str(item.get("description") or ""),
|
|
library_js=str(item.get("library") or item.get("sharedLibrary") or ""),
|
|
input_js=str(item.get("input") or item.get("onInput") or ""),
|
|
context_js=str(item.get("context") or item.get("onModelContext") or ""),
|
|
output_js=str(item.get("output") or item.get("onOutput") or ""),
|
|
)
|
|
db.add(script)
|
|
db.flush()
|
|
scenario.scripts.append(script)
|
|
|
|
db.commit()
|
|
out = schemas.ScenarioOut.model_validate(scenario).model_dump(mode="json")
|
|
return {"scenario": out, "unmapped_keys": unmapped}
|