Initial commit: AI Dungeon clone (FastAPI backend + React frontend)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01KFsGHju9szibJJa2YJcdbg
This commit is contained in:
@@ -0,0 +1,611 @@
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import memorybank, models, schemas
|
||||
from ..context import build_context
|
||||
from ..database import get_db
|
||||
from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError
|
||||
from ..scripting import ScriptPipeline
|
||||
from .settings import get_settings
|
||||
|
||||
router = APIRouter(prefix="/api/adventures", tags=["adventures"])
|
||||
|
||||
|
||||
def get_adventure_or_404(adventure_id: int, db: Session) -> models.Adventure:
|
||||
adventure = db.get(models.Adventure, adventure_id)
|
||||
if adventure is None:
|
||||
raise HTTPException(404, "Adventure not found")
|
||||
return adventure
|
||||
|
||||
|
||||
@router.get("", response_model=list[schemas.AdventureListItem])
|
||||
def list_adventures(db: Session = Depends(get_db)):
|
||||
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)
|
||||
.group_by(models.Adventure.id)
|
||||
.order_by(models.Adventure.updated_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
schemas.AdventureListItem(
|
||||
id=adv.id,
|
||||
scenario_id=adv.scenario_id,
|
||||
scenario_title=scenario_title,
|
||||
title=adv.title,
|
||||
updated_at=adv.updated_at,
|
||||
action_count=count,
|
||||
)
|
||||
for adv, count, scenario_title in rows
|
||||
]
|
||||
|
||||
|
||||
PLACEHOLDER_RE = re.compile(r"\$\{([^}]+)\}")
|
||||
|
||||
|
||||
def fill_placeholders(text: str, values: dict[str, str]) -> str:
|
||||
"""Replace ${Name} with the player-provided value; unknown names are left as-is."""
|
||||
if not text or not values:
|
||||
return text
|
||||
return PLACEHOLDER_RE.sub(
|
||||
lambda m: values.get(m.group(1).strip(), m.group(0)), text
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=schemas.AdventureOut, status_code=201)
|
||||
def create_adventure(payload: schemas.AdventureCreate, db: Session = Depends(get_db)):
|
||||
scenario = None
|
||||
if payload.scenario_id is not None:
|
||||
scenario = db.get(models.Scenario, payload.scenario_id)
|
||||
if scenario is None:
|
||||
raise HTTPException(404, "Scenario not found")
|
||||
|
||||
values = payload.placeholders
|
||||
adventure = models.Adventure(
|
||||
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 "",
|
||||
authors_note=fill_placeholders(scenario.authors_note, values) if scenario else "",
|
||||
ai_instructions=fill_placeholders(scenario.ai_instructions, values) if scenario else "",
|
||||
)
|
||||
db.add(adventure)
|
||||
db.flush()
|
||||
|
||||
if scenario:
|
||||
for card in scenario.story_cards:
|
||||
db.add(
|
||||
models.StoryCard(
|
||||
adventure_id=adventure.id,
|
||||
type=card.type,
|
||||
name=card.name,
|
||||
keys=fill_placeholders(card.keys, values),
|
||||
entry=fill_placeholders(card.entry, values),
|
||||
notes=card.notes,
|
||||
)
|
||||
)
|
||||
for position, script in enumerate(scenario.scripts):
|
||||
db.add(
|
||||
models.AdventureScript(
|
||||
adventure_id=adventure.id,
|
||||
position=position,
|
||||
name=script.name,
|
||||
description=script.description,
|
||||
library_js=script.library_js,
|
||||
input_js=script.input_js,
|
||||
context_js=script.context_js,
|
||||
output_js=script.output_js,
|
||||
)
|
||||
)
|
||||
if scenario.prompt.strip():
|
||||
db.add(
|
||||
models.Action(
|
||||
adventure_id=adventure.id,
|
||||
index=0,
|
||||
type="start",
|
||||
text=fill_placeholders(scenario.prompt, values),
|
||||
)
|
||||
)
|
||||
|
||||
db.commit()
|
||||
db.refresh(adventure)
|
||||
return adventure
|
||||
|
||||
|
||||
@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)
|
||||
|
||||
|
||||
@router.patch("/{adventure_id}", response_model=schemas.AdventureOut)
|
||||
def update_adventure(
|
||||
adventure_id: int, payload: schemas.AdventureUpdate, db: Session = Depends(get_db)
|
||||
):
|
||||
adventure = get_adventure_or_404(adventure_id, db)
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(adventure, field, value)
|
||||
db.commit()
|
||||
return 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)
|
||||
db.delete(adventure)
|
||||
db.commit()
|
||||
|
||||
|
||||
# ---------- Turn engine ----------
|
||||
|
||||
# One turn at a time per adventure (in-memory; fine for a single-process local app).
|
||||
# Sync endpoints run in a threadpool, so the check-and-add must be guarded — and
|
||||
# it must happen in the request phase, not when the SSE generator first runs,
|
||||
# or two rapid requests both pass the check and generate concurrently.
|
||||
_active_turns: set[int] = set()
|
||||
_active_turns_guard = threading.Lock()
|
||||
|
||||
|
||||
def acquire_turn_lock(adventure_id: int):
|
||||
"""Atomically claim the adventure's turn slot; with_turn_lock releases it."""
|
||||
with _active_turns_guard:
|
||||
if adventure_id in _active_turns:
|
||||
raise HTTPException(409, "A turn is already generating for this adventure.")
|
||||
_active_turns.add(adventure_id)
|
||||
|
||||
|
||||
async def with_turn_lock(adventure_id: int, gen):
|
||||
"""Wrap an SSE generator so the lock (from acquire_turn_lock) is released."""
|
||||
try:
|
||||
async for event in gen:
|
||||
yield event
|
||||
finally:
|
||||
_active_turns.discard(adventure_id)
|
||||
|
||||
|
||||
def format_player_input(action_type: str, text: str) -> str:
|
||||
"""AI Dungeon input conventions."""
|
||||
text = text.strip()
|
||||
if action_type == "say":
|
||||
text = text.strip('"')
|
||||
if text and text[-1] not in ".!?…":
|
||||
text += "."
|
||||
return f'> You say "{text}"'
|
||||
if action_type == "do":
|
||||
if text.lower().startswith("you "):
|
||||
text = text[4:]
|
||||
if text and text[-1] not in ".!?…":
|
||||
text += "."
|
||||
return f"> You {text}"
|
||||
return text # story: raw text appended
|
||||
|
||||
|
||||
def sse(obj: dict) -> str:
|
||||
return f"data: {json.dumps(obj)}\n\n"
|
||||
|
||||
|
||||
def action_json(action: models.Action) -> dict:
|
||||
return schemas.ActionOut.model_validate(action).model_dump(mode="json")
|
||||
|
||||
|
||||
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):
|
||||
"""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)
|
||||
system_text, story_text, snapshot = build_context(adventure, settings, memories)
|
||||
|
||||
# onModelContext: scripts see (and may rewrite) the whole assembled context.
|
||||
combined = f"{system_text}\n\n{story_text}" if system_text else story_text
|
||||
modified, stop = pipeline.run("context", combined)
|
||||
if stop:
|
||||
yield sse({"type": "stopped", "script": pipeline.report()})
|
||||
return
|
||||
context_changed = modified != combined
|
||||
parts = (
|
||||
PromptParts(system="", story=modified)
|
||||
if context_changed
|
||||
else PromptParts(system=system_text, story=story_text)
|
||||
)
|
||||
snapshot["script"] = pipeline.report() | {
|
||||
"context_changed": context_changed,
|
||||
"context_before": combined if context_changed else None,
|
||||
"context_after": modified if context_changed else None,
|
||||
}
|
||||
|
||||
provider = OpenAICompatibleProvider(
|
||||
settings.endpoint_url, settings.api_key, settings.model, settings.api_mode,
|
||||
settings.reasoning_max_tokens,
|
||||
)
|
||||
chunks: list[str] = []
|
||||
reasoning_chunks: list[str] = []
|
||||
try:
|
||||
async for kind, chunk in provider.generate(
|
||||
parts, temperature=settings.temperature, max_tokens=settings.max_output_tokens
|
||||
):
|
||||
if kind == "reasoning":
|
||||
reasoning_chunks.append(chunk)
|
||||
yield sse({"type": "reasoning", "text": chunk})
|
||||
else:
|
||||
chunks.append(chunk)
|
||||
yield sse({"type": "chunk", "text": chunk})
|
||||
except ProviderError as exc:
|
||||
yield sse({"type": "error", "detail": str(exc)})
|
||||
return
|
||||
|
||||
text = "".join(chunks).strip()
|
||||
if not text:
|
||||
yield sse({"type": "error", "detail": "The AI returned an empty response."})
|
||||
return
|
||||
|
||||
# onOutput
|
||||
text, _ = pipeline.run("output", text)
|
||||
if not text.strip():
|
||||
yield sse({"type": "error", "detail": "A script's output modifier returned empty text."})
|
||||
return
|
||||
snapshot["script"] = snapshot["script"] | pipeline.report()
|
||||
|
||||
ai_action = models.Action(
|
||||
adventure_id=adventure.id,
|
||||
index=next_index(adventure),
|
||||
type="ai",
|
||||
text=text,
|
||||
reasoning="".join(reasoning_chunks).strip() or None,
|
||||
context_snapshot=snapshot,
|
||||
)
|
||||
db.add(ai_action)
|
||||
adventure.updated_at = models.utcnow()
|
||||
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)
|
||||
|
||||
|
||||
async def run_player_turn(
|
||||
adventure: models.Adventure, db: Session, payload: schemas.ActionCreate
|
||||
):
|
||||
pipeline = ScriptPipeline(adventure, db)
|
||||
|
||||
# An empty do/say/story is just a continue.
|
||||
if payload.type != "continue" and payload.text.strip():
|
||||
# onInput sees the formatted text (as in AI Dungeon: "> You ...").
|
||||
formatted = format_player_input(payload.type, payload.text)
|
||||
modified, stop = pipeline.run("input", formatted)
|
||||
if not modified.strip():
|
||||
yield sse({"type": "error", "detail": "A script's input modifier returned empty text.",
|
||||
"script": pipeline.report()})
|
||||
return
|
||||
player_action = models.Action(
|
||||
adventure_id=adventure.id,
|
||||
index=next_index(adventure),
|
||||
type=payload.type,
|
||||
text=modified,
|
||||
)
|
||||
db.add(player_action)
|
||||
db.commit()
|
||||
db.refresh(player_action)
|
||||
# The new action was added via its FK, so the loaded adventure.actions
|
||||
# collection is stale — without this, build_context and next_index for
|
||||
# the AI action would not see the player action just saved.
|
||||
db.expire(adventure, ["actions"])
|
||||
yield sse({"type": "player", "action": action_json(player_action)})
|
||||
if stop:
|
||||
# onInput { stop: true } prevents the AI call.
|
||||
yield sse({"type": "stopped", "script": pipeline.report()})
|
||||
return
|
||||
|
||||
async for event in generate_turn(adventure, db, pipeline):
|
||||
yield event
|
||||
|
||||
|
||||
@router.post("/{adventure_id}/actions")
|
||||
def create_action(
|
||||
adventure_id: int, payload: schemas.ActionCreate, db: Session = Depends(get_db)
|
||||
):
|
||||
adventure = get_adventure_or_404(adventure_id, db)
|
||||
acquire_turn_lock(adventure_id)
|
||||
return StreamingResponse(
|
||||
with_turn_lock(adventure_id, run_player_turn(adventure, db, payload)),
|
||||
media_type="text/event-stream",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/{adventure_id}/retry")
|
||||
def retry_action(adventure_id: int, db: Session = Depends(get_db)):
|
||||
"""Delete the last AI action and regenerate from the same input."""
|
||||
adventure = get_adventure_or_404(adventure_id, db)
|
||||
acquire_turn_lock(adventure_id)
|
||||
try:
|
||||
if adventure.actions and adventure.actions[-1].type == "ai":
|
||||
db.delete(adventure.actions[-1])
|
||||
db.commit()
|
||||
db.refresh(adventure)
|
||||
except BaseException:
|
||||
_active_turns.discard(adventure_id)
|
||||
raise
|
||||
return StreamingResponse(
|
||||
with_turn_lock(adventure_id, generate_turn(adventure, db, ScriptPipeline(adventure, db))),
|
||||
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)):
|
||||
"""Delete the last turn: the trailing AI action plus its player action, if any."""
|
||||
adventure = get_adventure_or_404(adventure_id, db)
|
||||
actions = list(adventure.actions)
|
||||
if not actions or actions[-1].type == "start":
|
||||
raise HTTPException(400, "Nothing to undo")
|
||||
last = actions.pop()
|
||||
db.delete(last)
|
||||
if last.type == "ai" and actions and actions[-1].type in ("do", "say", "story"):
|
||||
db.delete(actions.pop())
|
||||
db.commit()
|
||||
db.refresh(adventure)
|
||||
return adventure.actions
|
||||
|
||||
|
||||
# ---------- Import / Export ----------
|
||||
|
||||
@router.get("/{adventure_id}/export")
|
||||
def export_adventure(adventure_id: int, db: Session = Depends(get_db)):
|
||||
"""Full backup: plot components, story cards, scripts (+state), every action."""
|
||||
adv = get_adventure_or_404(adventure_id, db)
|
||||
return {
|
||||
"format": "ai-dnd-adventure-v1",
|
||||
"title": adv.title,
|
||||
"memory": adv.memory,
|
||||
"authorsNote": adv.authors_note,
|
||||
"aiInstructions": adv.ai_instructions,
|
||||
"storySummary": adv.story_summary,
|
||||
"scriptState": adv.script_state,
|
||||
"autoSummarize": adv.auto_summarize,
|
||||
"memoryBankEnabled": adv.memory_bank_enabled,
|
||||
"memoryCursor": adv.memory_cursor,
|
||||
"summaryCursor": adv.summary_cursor,
|
||||
"memories": [
|
||||
{
|
||||
"text": m.text, "pinned": m.pinned, "forgotten": m.forgotten,
|
||||
"sourceStart": m.source_start, "sourceEnd": m.source_end,
|
||||
"useCount": m.use_count,
|
||||
}
|
||||
for m in adv.memories
|
||||
],
|
||||
"storyCards": [
|
||||
{"type": c.type, "name": c.name, "keys": c.keys, "entry": c.entry, "notes": c.notes}
|
||||
for c in adv.story_cards
|
||||
],
|
||||
"scripts": [
|
||||
{
|
||||
"position": s.position, "enabled": s.enabled,
|
||||
"name": s.name, "description": s.description,
|
||||
"library": s.library_js, "input": s.input_js,
|
||||
"context": s.context_js, "output": s.output_js,
|
||||
}
|
||||
for s in adv.scripts
|
||||
],
|
||||
"actions": [
|
||||
{
|
||||
"index": a.index, "type": a.type, "text": a.text,
|
||||
"reasoning": a.reasoning,
|
||||
"createdAt": a.created_at.isoformat(),
|
||||
}
|
||||
for a in adv.actions
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@router.post("/import", response_model=schemas.AdventureOut, status_code=201)
|
||||
def import_adventure(bundle: dict = Body(...), db: Session = Depends(get_db)):
|
||||
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(
|
||||
title=str(bundle.get("title") or "Imported Adventure"),
|
||||
memory=str(bundle.get("memory") or ""),
|
||||
authors_note=str(bundle.get("authorsNote") or ""),
|
||||
ai_instructions=str(bundle.get("aiInstructions") or ""),
|
||||
story_summary=str(bundle.get("storySummary") or ""),
|
||||
script_state=bundle.get("scriptState") or {},
|
||||
auto_summarize=bool(bundle.get("autoSummarize", False)),
|
||||
memory_bank_enabled=bool(bundle.get("memoryBankEnabled", False)),
|
||||
memory_cursor=int(bundle.get("memoryCursor", 0)),
|
||||
summary_cursor=int(bundle.get("summaryCursor", 0)),
|
||||
)
|
||||
db.add(adventure)
|
||||
db.flush()
|
||||
|
||||
for m in bundle.get("memories") or []:
|
||||
if isinstance(m, dict) and str(m.get("text") or "").strip():
|
||||
db.add(models.Memory(
|
||||
adventure_id=adventure.id,
|
||||
text=str(m["text"]),
|
||||
pinned=bool(m.get("pinned", False)),
|
||||
forgotten=bool(m.get("forgotten", False)),
|
||||
source_start=m.get("sourceStart"),
|
||||
source_end=m.get("sourceEnd"),
|
||||
use_count=int(m.get("useCount", 0)),
|
||||
))
|
||||
|
||||
for card in bundle.get("storyCards") or []:
|
||||
if isinstance(card, dict):
|
||||
db.add(models.StoryCard(
|
||||
adventure_id=adventure.id,
|
||||
type=str(card.get("type") or ""),
|
||||
name=str(card.get("name") or ""),
|
||||
keys=str(card.get("keys") or ""),
|
||||
entry=str(card.get("entry") or ""),
|
||||
notes=str(card.get("notes") or ""),
|
||||
))
|
||||
|
||||
for i, s in enumerate(bundle.get("scripts") or []):
|
||||
if isinstance(s, dict):
|
||||
db.add(models.AdventureScript(
|
||||
adventure_id=adventure.id,
|
||||
position=int(s.get("position", i)),
|
||||
enabled=bool(s.get("enabled", True)),
|
||||
name=str(s.get("name") or "Imported Script"),
|
||||
description=str(s.get("description") or ""),
|
||||
library_js=str(s.get("library") or ""),
|
||||
input_js=str(s.get("input") or ""),
|
||||
context_js=str(s.get("context") or ""),
|
||||
output_js=str(s.get("output") or ""),
|
||||
))
|
||||
|
||||
for i, a in enumerate(bundle.get("actions") or []):
|
||||
if isinstance(a, dict) and str(a.get("text") or ""):
|
||||
db.add(models.Action(
|
||||
adventure_id=adventure.id,
|
||||
index=int(a.get("index", i)),
|
||||
type=str(a.get("type") or "story"),
|
||||
text=str(a["text"]),
|
||||
reasoning=str(a["reasoning"]) if a.get("reasoning") else None,
|
||||
))
|
||||
|
||||
db.commit()
|
||||
db.refresh(adventure)
|
||||
return adventure
|
||||
|
||||
|
||||
# ---------- 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
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/{adventure_id}/scripts/{adv_script_id}", response_model=schemas.AdventureScriptOut
|
||||
)
|
||||
def update_adventure_script(
|
||||
adventure_id: int,
|
||||
adv_script_id: int,
|
||||
payload: schemas.AdventureScriptUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
script = db.get(models.AdventureScript, adv_script_id)
|
||||
if script is None or script.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Script not found")
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(script, field, value)
|
||||
db.commit()
|
||||
return script
|
||||
|
||||
|
||||
# ---------- Insights ----------
|
||||
|
||||
@router.get("/{adventure_id}/context")
|
||||
async def dry_run_context(adventure_id: int, db: Session = Depends(get_db)):
|
||||
"""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)
|
||||
_, _, 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)):
|
||||
action = db.get(models.Action, action_id)
|
||||
if action is None or action.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Action not found")
|
||||
if action.context_snapshot is None:
|
||||
raise HTTPException(404, "No context snapshot for this action")
|
||||
return action.context_snapshot
|
||||
|
||||
|
||||
# ---------- 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
|
||||
|
||||
|
||||
@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)
|
||||
):
|
||||
"""Manually add a memory; it gets embedded by the next post-turn pass."""
|
||||
adventure = get_adventure_or_404(adventure_id, db)
|
||||
if not payload.text.strip():
|
||||
raise HTTPException(400, "Memory text cannot be empty")
|
||||
memory = models.Memory(adventure_id=adventure.id, text=payload.text.strip())
|
||||
db.add(memory)
|
||||
db.commit()
|
||||
db.refresh(memory)
|
||||
return memory
|
||||
|
||||
|
||||
@router.patch("/{adventure_id}/memories/{memory_id}", response_model=schemas.MemoryOut)
|
||||
def update_memory(
|
||||
adventure_id: int,
|
||||
memory_id: int,
|
||||
payload: schemas.MemoryUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
memory = db.get(models.Memory, memory_id)
|
||||
if memory is None or memory.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Memory not found")
|
||||
fields = {k: v for k, v in payload.model_dump(exclude_unset=True).items() if v is not None}
|
||||
if "text" in fields and fields["text"].strip() != memory.text:
|
||||
memory.embedding = None # re-embed on the next post-turn pass
|
||||
for field, value in fields.items():
|
||||
setattr(memory, field, value)
|
||||
db.commit()
|
||||
return 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)):
|
||||
memory = db.get(models.Memory, memory_id)
|
||||
if memory is None or memory.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Memory not found")
|
||||
db.delete(memory)
|
||||
db.commit()
|
||||
|
||||
|
||||
# ---------- 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)
|
||||
return (
|
||||
db.query(models.Action)
|
||||
.filter(models.Action.adventure_id == adventure_id)
|
||||
.order_by(models.Action.index)
|
||||
.all()
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/{adventure_id}/actions/{action_id}", response_model=schemas.ActionOut)
|
||||
def update_action(
|
||||
adventure_id: int,
|
||||
action_id: int,
|
||||
payload: schemas.ActionUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
action = db.get(models.Action, action_id)
|
||||
if action is None or action.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Action not found")
|
||||
action.text = payload.text
|
||||
db.commit()
|
||||
return 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)):
|
||||
action = db.get(models.Action, action_id)
|
||||
if action is None or action.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Action not found")
|
||||
db.delete(action)
|
||||
db.commit()
|
||||
@@ -0,0 +1,11 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from .. import 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)."""
|
||||
return debuglog.recent()
|
||||
@@ -0,0 +1,166 @@
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import 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:
|
||||
scenario = db.get(models.Scenario, scenario_id)
|
||||
if scenario is None:
|
||||
raise HTTPException(404, "Scenario not found")
|
||||
return scenario
|
||||
|
||||
|
||||
@router.get("", response_model=list[schemas.ScenarioListItem])
|
||||
def list_scenarios(db: Session = Depends(get_db)):
|
||||
return (
|
||||
db.query(models.Scenario)
|
||||
.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())
|
||||
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)
|
||||
|
||||
|
||||
@router.patch("/{scenario_id}", response_model=schemas.ScenarioOut)
|
||||
def update_scenario(
|
||||
scenario_id: int, payload: schemas.ScenarioUpdate, db: Session = Depends(get_db)
|
||||
):
|
||||
scenario = get_scenario_or_404(scenario_id, db)
|
||||
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()
|
||||
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)):
|
||||
scenario = get_scenario_or_404(scenario_id, db)
|
||||
db.delete(scenario)
|
||||
db.commit()
|
||||
|
||||
|
||||
# ---------- 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)
|
||||
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,
|
||||
"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",
|
||||
"createdAt", "updatedAt", "id", "publicId", "image", "nsfw", "type", "options"}
|
||||
|
||||
|
||||
@router.post("/import", status_code=201)
|
||||
def import_scenario(bundle: dict = Body(...), db: Session = Depends(get_db)):
|
||||
"""Accepts our export format and AI Dungeon scenario exports best-effort;
|
||||
reports any keys it didn't understand."""
|
||||
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
|
||||
|
||||
scenario = models.Scenario(**fields)
|
||||
if not scenario.title:
|
||||
scenario.title = "Imported Scenario"
|
||||
db.add(scenario)
|
||||
db.flush()
|
||||
|
||||
cards = bundle.get("storyCards") or bundle.get("worldInfo") or []
|
||||
for card in cards:
|
||||
if not isinstance(card, dict):
|
||||
continue
|
||||
db.add(
|
||||
models.StoryCard(
|
||||
scenario_id=scenario.id,
|
||||
type=str(card.get("type") or ""),
|
||||
name=str(card.get("name") or card.get("title") or ""),
|
||||
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(
|
||||
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 ""),
|
||||
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}
|
||||
@@ -0,0 +1,114 @@
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import models, schemas
|
||||
from ..database import get_db
|
||||
from ..scripting import run_hook
|
||||
|
||||
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:
|
||||
script = db.get(models.Script, script_id)
|
||||
if script is None:
|
||||
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()
|
||||
|
||||
|
||||
@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())
|
||||
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)
|
||||
|
||||
|
||||
@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)
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(script, field, value)
|
||||
db.commit()
|
||||
return script
|
||||
|
||||
|
||||
@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))
|
||||
db.commit()
|
||||
|
||||
|
||||
@router.post("/{script_id}/test")
|
||||
def test_script(
|
||||
script_id: int, payload: schemas.ScriptTestRequest, db: Session = Depends(get_db)
|
||||
):
|
||||
"""Dry-run one hook against sample text — no AI call, no persistence."""
|
||||
script = get_script_or_404(script_id, db)
|
||||
result = run_hook(
|
||||
script.library_js,
|
||||
getattr(script, HOOK_FIELDS[payload.hook]),
|
||||
payload.text,
|
||||
payload.state,
|
||||
history=[],
|
||||
story_cards=[],
|
||||
info={"actionCount": 0, "characterNames": [], "memoryLength": 0, "maxChars": 0},
|
||||
)
|
||||
return {
|
||||
"text": result.text,
|
||||
"stop": result.stop,
|
||||
"state": result.state,
|
||||
"storyCards": result.story_cards,
|
||||
"logs": result.logs,
|
||||
"error": result.error,
|
||||
}
|
||||
|
||||
|
||||
# ---------- Import / Export ----------
|
||||
|
||||
@router.get("/{script_id}/export")
|
||||
def export_script(script_id: int, db: Session = Depends(get_db)):
|
||||
"""JSON bundle matching how AI Dungeon scripts circulate."""
|
||||
script = get_script_or_404(script_id, db)
|
||||
return {
|
||||
"name": script.name,
|
||||
"description": script.description,
|
||||
"library": script.library_js,
|
||||
"input": script.input_js,
|
||||
"context": script.context_js,
|
||||
"output": script.output_js,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/import", response_model=schemas.ScriptOut, status_code=201)
|
||||
def import_script(bundle: dict = Body(...), db: Session = Depends(get_db)):
|
||||
"""Accepts our export bundle; tolerates *_js key names too."""
|
||||
def pick(*keys: str) -> str:
|
||||
for key in keys:
|
||||
value = bundle.get(key)
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return ""
|
||||
|
||||
script = models.Script(
|
||||
name=pick("name") or "Imported Script",
|
||||
description=pick("description"),
|
||||
library_js=pick("library", "library_js", "sharedLibrary"),
|
||||
input_js=pick("input", "input_js", "onInput"),
|
||||
context_js=pick("context", "context_js", "onModelContext"),
|
||||
output_js=pick("output", "output_js", "onOutput"),
|
||||
)
|
||||
db.add(script)
|
||||
db.commit()
|
||||
return script
|
||||
@@ -0,0 +1,57 @@
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import models, schemas
|
||||
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)
|
||||
if settings is None:
|
||||
settings = models.Settings(id=1)
|
||||
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)
|
||||
|
||||
|
||||
@router.put("", response_model=schemas.SettingsOut)
|
||||
def update_settings(payload: schemas.SettingsUpdate, db: Session = Depends(get_db)):
|
||||
settings = get_settings(db)
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(settings, field, value)
|
||||
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"
|
||||
headers = {}
|
||||
if settings.api_key:
|
||||
headers["Authorization"] = f"Bearer {settings.api_key}"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
resp = await client.get(url, headers=headers)
|
||||
except httpx.HTTPError as exc:
|
||||
return {"ok": False, "detail": f"Connection failed: {exc}"}
|
||||
|
||||
if resp.status_code != 200:
|
||||
return {"ok": False, "detail": f"HTTP {resp.status_code}: {resp.text[:300]}"}
|
||||
|
||||
models_available: list[str] = []
|
||||
try:
|
||||
data = resp.json()
|
||||
models_available = [m.get("id", "?") for m in data.get("data", [])]
|
||||
except ValueError:
|
||||
pass
|
||||
return {"ok": True, "models": models_available}
|
||||
@@ -0,0 +1,57 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import models, schemas
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/story-cards", tags=["story-cards"])
|
||||
|
||||
|
||||
@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),
|
||||
):
|
||||
query = db.query(models.StoryCard)
|
||||
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()
|
||||
|
||||
|
||||
@router.post("", response_model=schemas.StoryCardOut, status_code=201)
|
||||
def create_story_card(payload: schemas.StoryCardCreate, db: Session = Depends(get_db)):
|
||||
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:
|
||||
raise HTTPException(404, "Owner not found")
|
||||
card = models.StoryCard(**payload.model_dump())
|
||||
db.add(card)
|
||||
db.commit()
|
||||
return card
|
||||
|
||||
|
||||
@router.patch("/{card_id}", response_model=schemas.StoryCardOut)
|
||||
def update_story_card(
|
||||
card_id: int, payload: schemas.StoryCardUpdate, db: Session = Depends(get_db)
|
||||
):
|
||||
card = db.get(models.StoryCard, card_id)
|
||||
if card is None:
|
||||
raise HTTPException(404, "Story card not found")
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(card, field, value)
|
||||
db.commit()
|
||||
return 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")
|
||||
db.delete(card)
|
||||
db.commit()
|
||||
Reference in New Issue
Block a user