Remove three kinds of duplication in the backend

Stage 2, items 1, 4, and 5 of `plan/17-refactor.md`.

**One path resolver in `worldstate`.** `apply_delta` and `apply_override` routed
`flags.<name>`, `milestones.<id>`, `world.<stat>`, `player.<stat>`, and
`npc.<id>.<stat>` with parallel code, about 100 lines each. `_resolve` now says
what a path points at and returns either a target or the rejection to report.
Each function keeps its own write rule, because the rules genuinely differ: an
override sets a number rather than adding to it, ignores `cooldown`,
`max_delta_per_turn`, and the rule that a counter only counts up, and can un-set
a milestone.

A differential check ran both implementations over 3960 payloads: twenty paths,
fourteen values, three starting states, plus every three-path combination. The
results are identical except that 674 rejections from `apply_override` now carry
a `fix` string. `apply_delta` already worded those, and the world-state editor
renders them, so an override that names an unknown flag now explains itself the
way a delta does.

**`sse`, `SSE_HEADERS`, and `turn_error` move to `app/sse.py`.** Two routers
stream, and `chat.py` had to import from `routers.adventures` to reach them.

**`get_adventure_or_404` becomes the `current_adventure` dependency.** All 32
handlers repeated the call as their first statement. The ownership check now
reads in the signature and runs before the body. FastAPI caches a dependency for
one request, so the handler's `db` is the session the adventure came from.

The generated OpenAPI document is byte-identical except on `rename_branch`,
where `branch_id` is now listed before `adventure_id`, because that handler no
longer names `adventure_id` itself. Parameter order in the document is
cosmetic.

Six tests in `test_state_revert.py` call `undo_turn` and `retry_action`
directly rather than over HTTP. They pass the adventure they already hold
instead of an id.

549 tests pass.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014Dix4oGV3njgWRdu7P9t6r
This commit is contained in:
parththakkar106
2026-08-29 02:04:02 +05:30
co-authored by Claude Opus 5
parent 0623dd8b78
commit 2c57b1ceab
16 changed files with 258 additions and 261 deletions
+1 -3
View File
@@ -42,17 +42,15 @@ from ... import limits # noqa: F401 `adventures.limits` is patched by tests.
from .crud import SNIPPET_MAX, _snippet from .crud import SNIPPET_MAX, _snippet
from .paging import ACTION_PAGE from .paging import ACTION_PAGE
from .takes import retry_action, undo_turn from .takes import retry_action, undo_turn
from .turns import SSE_HEADERS, sse, world_delta_of from .turns import world_delta_of
__all__ = [ __all__ = [
"ACTION_PAGE", "ACTION_PAGE",
"SNIPPET_MAX", "SNIPPET_MAX",
"SSE_HEADERS",
"_snippet", "_snippet",
"limits", "limits",
"retry_action", "retry_action",
"router", "router",
"sse",
"undo_turn", "undo_turn",
"world_delta_of", "world_delta_of",
] ]
+4 -8
View File
@@ -11,18 +11,17 @@ from sqlalchemy.orm import Session
from ... import models, schemas, tree from ... import models, schemas, tree
from ...database import get_db from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .nodes import delete_turn from .nodes import delete_turn
from .paging import ACTION_PAGE, action_window, annotate_takes from .paging import ACTION_PAGE, action_window, annotate_takes
@router.get("/{adventure_id}/actions", response_model=schemas.ActionPage) @router.get("/{adventure_id}/actions", response_model=schemas.ActionPage)
def list_actions( def list_actions(
adventure_id: int,
before_id: int | None = None, before_id: int | None = None,
limit: int = ACTION_PAGE, limit: int = ACTION_PAGE,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns a page of the story, working backwards from the newest action. """Returns a page of the story, working backwards from the newest action.
@@ -30,7 +29,6 @@ def list_actions(
asks for what comes before it. Omit `before_id` for the newest window. See asks for what comes before it. Omit `before_id` for the newest window. See
`action_window` for why this anchors on a row rather than an offset. `action_window` for why this anchors on a row rather than an offset.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
limit = max(1, min(limit, ACTION_PAGE * 4)) limit = max(1, min(limit, ACTION_PAGE * 4))
actions, total, has_more = action_window( actions, total, has_more = action_window(
db, adventure, before_id=before_id, limit=limit db, adventure, before_id=before_id, limit=limit
@@ -51,9 +49,8 @@ def update_action(
action_id: int, action_id: int,
payload: schemas.ActionUpdate, payload: schemas.ActionUpdate,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
@@ -70,9 +67,8 @@ def delete_action(
adventure_id: int, adventure_id: int,
action_id: int, action_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
+6 -10
View File
@@ -18,14 +18,15 @@ from ...context import lineage
from ...database import get_db from ...database import get_db
from . import turns from . import turns
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .nodes import db_tip from .nodes import db_tip
from .paging import current_window from .paging import current_window
@router.get("/{adventure_id}/branches", response_model=list[schemas.BranchOut]) @router.get("/{adventure_id}/branches", response_model=list[schemas.BranchOut])
def list_branches( def list_branches(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns every branch of the adventure and where each one leaves its parent. """Returns every branch of the adventure and where each one leaves its parent.
@@ -35,7 +36,6 @@ def list_branches(
`actions`, never one query per branch, so a view of a hundred forks does not `actions`, never one query per branch, so a view of a hundred forks does not
cost a hundred round trips. cost a hundred round trips.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
branches = ( branches = (
db.query(models.Branch) db.query(models.Branch)
.filter(models.Branch.adventure_id == adventure.id) .filter(models.Branch.adventure_id == adventure.id)
@@ -95,11 +95,10 @@ def get_branch_or_404(
"/{adventure_id}/branches/{branch_id}", response_model=schemas.BranchOut "/{adventure_id}/branches/{branch_id}", response_model=schemas.BranchOut
) )
def rename_branch( def rename_branch(
adventure_id: int,
branch_id: int, branch_id: int,
payload: schemas.BranchRename, payload: schemas.BranchRename,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Names a branch, or clears the name to leave it unnamed. """Names a branch, or clears the name to leave it unnamed.
@@ -107,7 +106,6 @@ def rename_branch(
name anyone chose, and storing one gives the client an empty label to draw name anyone chose, and storing one gives the client an empty label to draw
instead of the fork depth. instead of the fork depth.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
branch = get_branch_or_404(adventure, branch_id, db) branch = get_branch_or_404(adventure, branch_id, db)
name = (payload.name or "").strip() name = (payload.name or "").strip()
branch.name = name or None branch.name = name or None
@@ -145,7 +143,7 @@ def delete_branch(
adventure_id: int, adventure_id: int,
branch_id: int, branch_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Deletes a branch and everything forked from it. """Deletes a branch and everything forked from it.
@@ -163,7 +161,6 @@ def delete_branch(
cascade on `branches.parent_branch_id`, so the delete is a single statement cascade on `branches.parent_branch_id`, so the delete is a single statement
however deep the subtree is. however deep the subtree is.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
branch = get_branch_or_404(adventure, branch_id, db) branch = get_branch_or_404(adventure, branch_id, db)
if branch.parent_branch_id is None: if branch.parent_branch_id is None:
raise HTTPException( raise HTTPException(
@@ -238,7 +235,7 @@ def switch_branch(
adventure_id: int, adventure_id: int,
branch_id: int, branch_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Reads and plays a different branch of the story. """Reads and plays a different branch of the story.
@@ -249,7 +246,6 @@ def switch_branch(
another branch's numbers, including the world-state cooldown clock inside the another branch's numbers, including the world-state cooldown clock inside the
snapshot. snapshot.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
branch = db.get(models.Branch, branch_id) branch = db.get(models.Branch, branch_id)
if branch is None or branch.adventure_id != adventure.id: if branch is None or branch.adventure_id != adventure.id:
raise HTTPException(404, "Branch not found") raise HTTPException(404, "Branch not found")
+3 -3
View File
@@ -10,19 +10,19 @@ from sqlalchemy.orm import Session
from ... import analytics, bundle, limits, models, schemas from ... import analytics, bundle, limits, models, schemas
from ...database import get_db from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
@router.get("/{adventure_id}/export") @router.get("/{adventure_id}/export")
def export_adventure( def export_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
adv: models.Adventure = Depends(current_adventure),
): ):
"""Returns a full backup: plot components, story cards, scripts, state, and tree. """Returns a full backup: plot components, story cards, scripts, state, and tree.
`app/bundle.py` owns the format, in both of its versions. A backup outlives `app/bundle.py` owns the format, in both of its versions. A backup outlives
the schema, so no call site decides anything about its shape. the schema, so no call site decides anything about its shape.
""" """
adv = get_adventure_or_404(adventure_id, db, user)
return bundle.export(db, adv) return bundle.export(db, adv)
+10 -15
View File
@@ -12,7 +12,7 @@ from sqlalchemy.orm.attributes import set_committed_value
from ... import analytics, attempts, images, limits, memorybank, models, schemas, tree, worldstate from ... import analytics, attempts, images, limits, memorybank, models, schemas, tree, worldstate
from ...database import get_db from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .paging import action_window, annotate_takes from .paging import action_window, annotate_takes
from .scenario_text import fill_placeholders, scenario_card_specs from .scenario_text import fill_placeholders, scenario_card_specs
@@ -229,7 +229,8 @@ def create_adventure(
@router.get("/{adventure_id}", response_model=schemas.AdventureOut) @router.get("/{adventure_id}", response_model=schemas.AdventureOut)
def get_adventure( def get_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns the adventure and the newest window of its story. """Returns the adventure and the newest window of its story.
@@ -238,7 +239,6 @@ def get_adventure(
exist above. `GET /{id}/actions` serves the older pages as the reader exist above. `GET /{id}/actions` serves the older pages as the reader
scrolls up. scrolls up.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
actions, total, _ = action_window(db, adventure) actions, total, _ = action_window(db, adventure)
# Annotate before handing over the window. This path serializes through the # Annotate before handing over the window. This path serializes through the
# relationship rather than building `ActionOut` itself, so the pager numbers # relationship rather than building `ActionOut` itself, so the pager numbers
@@ -259,7 +259,7 @@ def get_adventure(
@router.get("/{adventure_id}/script-state") @router.get("/{adventure_id}/script-state")
def get_script_state( def get_script_state(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns the scripting `state` object. """Returns the scripting `state` object.
@@ -267,21 +267,19 @@ def get_script_state(
`state.x`, persisted after each hook. It stays `{}` until a script sets a `state.x`, persisted after each hook. It stays `{}` until a script sets a
variable. variable.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
state = adventure.script_state if isinstance(adventure.script_state, dict) else {} state = adventure.script_state if isinstance(adventure.script_state, dict) else {}
return {"state": state} return {"state": state}
@router.get("/{adventure_id}/world-state") @router.get("/{adventure_id}/world-state")
def get_world_state( def get_world_state(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns the live RPG world state and the scenario's `stat_schema`. """Returns the live RPG world state and the scenario's `stat_schema`.
The play view uses both to render the character sheet and the milestones. The play view uses both to render the character sheet and the milestones.
`schema` is null when the adventure has no RPG layer. `schema` is null when the adventure has no RPG layer.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
schema = adventure.scenario.stat_schema if adventure.scenario else None schema = adventure.scenario.stat_schema if adventure.scenario else None
state = adventure.world_state if isinstance(adventure.world_state, dict) else {} state = adventure.world_state if isinstance(adventure.world_state, dict) else {}
return { return {
@@ -292,10 +290,9 @@ def get_world_state(
@router.put("/{adventure_id}/world-state") @router.put("/{adventure_id}/world-state")
def override_world_state( def override_world_state(
adventure_id: int,
overrides: dict = Body(...), overrides: dict = Body(...),
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Edits the live RPG values directly, as a manual correction rather than a turn. """Edits the live RPG values directly, as a manual correction rather than a turn.
@@ -303,7 +300,6 @@ def override_world_state(
`milestones.y` to their new absolute values. The endpoint rejects unknown `milestones.y` to their new absolute values. The endpoint rejects unknown
paths and wrong types one at a time, and applies the rest. paths and wrong types one at a time, and applies the rest.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
schema = adventure.scenario.stat_schema if adventure.scenario else None schema = adventure.scenario.stat_schema if adventure.scenario else None
if not worldstate.has_schema(schema): if not worldstate.has_schema(schema):
raise HTTPException(400, "This adventure has no RPG world-state layer") raise HTTPException(400, "This adventure has no RPG world-state layer")
@@ -316,12 +312,10 @@ def override_world_state(
@router.patch("/{adventure_id}", response_model=schemas.AdventureOut) @router.patch("/{adventure_id}", response_model=schemas.AdventureOut)
def update_adventure( def update_adventure(
adventure_id: int,
payload: schemas.AdventureUpdate, payload: schemas.AdventureUpdate,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
for field, value in payload.model_dump(exclude_unset=True).items(): for field, value in payload.model_dump(exclude_unset=True).items():
setattr(adventure, field, value) setattr(adventure, field, value)
db.commit() db.commit()
@@ -330,9 +324,10 @@ def update_adventure(
@router.delete("/{adventure_id}", status_code=204) @router.delete("/{adventure_id}", status_code=204)
def delete_adventure( def delete_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
db.delete(adventure) db.delete(adventure)
db.commit() db.commit()
# No later request reads this adventure's vectors, so drop them now. The # No later request reads this adventure's vectors, so drop them now. The
+17
View File
@@ -9,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from ... import auth, models from ... import auth, models
from ...database import get_db
router = APIRouter(prefix="/api/adventures", tags=["adventures"]) router = APIRouter(prefix="/api/adventures", tags=["adventures"])
@@ -23,3 +24,19 @@ def get_adventure_or_404(
if adventure is None or adventure.user_id != user.id: if adventure is None or adventure.user_id != user.id:
raise HTTPException(404, "Adventure not found") raise HTTPException(404, "Adventure not found")
return adventure return adventure
def current_adventure(
adventure_id: int,
db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user),
) -> models.Adventure:
"""Resolves the `{adventure_id}` in the path to the caller's own adventure.
Declare this in a handler's signature. The ownership check then reads in the
signature, where you look for it, and it runs before the handler body.
FastAPI caches a dependency for the length of one request, so the handler's
own `db` is the same session this one loaded the adventure from.
"""
return get_adventure_or_404(adventure_id, db, user)
+5 -5
View File
@@ -12,15 +12,16 @@ from ...context import build_context
from ...database import get_db from ...database import get_db
from ..settings import get_settings from ..settings import get_settings
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
@router.get("/{adventure_id}/context") @router.get("/{adventure_id}/context")
async def dry_run_context( async def dry_run_context(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns what the app would send to the AI if the player continued now.""" """Returns what the app would send to the AI if the player continued now."""
adventure = get_adventure_or_404(adventure_id, db, user)
settings = get_settings(db, user) settings = get_settings(db, user)
if auth.resolve_provider_config(settings).using_demo: if auth.resolve_provider_config(settings).using_demo:
memories = ( memories = (
@@ -39,9 +40,8 @@ def action_context(
adventure_id: int, adventure_id: int,
action_id: int, action_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
+7 -9
View File
@@ -11,7 +11,7 @@ from ... import limits, memorybank, models, schemas, tree
from ...context import lineage from ...context import lineage
from ...database import get_db from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
# The columns `schemas.MemoryOut` renders. `embedded` is a real column and # The columns `schemas.MemoryOut` renders. `embedded` is a real column and
@@ -32,9 +32,10 @@ MEMORY_LIST_COLUMNS = (
@router.get("/{adventure_id}/memories", response_model=list[schemas.MemoryOut]) @router.get("/{adventure_id}/memories", response_model=list[schemas.MemoryOut])
def list_memories( def list_memories(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
# Name the columns in a query rather than walk `adventure.memories`. # Name the columns in a query rather than walk `adventure.memories`.
# Retrieval used to walk the relationship, which is why a turn cost # Retrieval used to walk the relationship, which is why a turn cost
# megabytes: a relationship load returns whole entities, so it reads # megabytes: a relationship load returns whole entities, so it reads
@@ -63,13 +64,12 @@ def list_memories(
@router.post("/{adventure_id}/memories", response_model=schemas.MemoryOut, status_code=201) @router.post("/{adventure_id}/memories", response_model=schemas.MemoryOut, status_code=201)
def create_memory( def create_memory(
adventure_id: int,
payload: schemas.MemoryCreate, payload: schemas.MemoryCreate,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Adds a memory manually. The next post-turn pass embeds it.""" """Adds a memory manually. The next post-turn pass embeds it."""
adventure = get_adventure_or_404(adventure_id, db, user)
limits.check_row_cap("memories", db, user, adventure=adventure) limits.check_row_cap("memories", db, user, adventure=adventure)
if not payload.text.strip(): if not payload.text.strip():
raise HTTPException(400, "Memory text cannot be empty") raise HTTPException(400, "Memory text cannot be empty")
@@ -88,9 +88,8 @@ def update_memory(
memory_id: int, memory_id: int,
payload: schemas.MemoryUpdate, payload: schemas.MemoryUpdate,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
get_adventure_or_404(adventure_id, db, user)
memory = db.get(models.Memory, memory_id) memory = db.get(models.Memory, memory_id)
if memory is None or memory.adventure_id != adventure_id: if memory is None or memory.adventure_id != adventure_id:
raise HTTPException(404, "Memory not found") raise HTTPException(404, "Memory not found")
@@ -108,9 +107,8 @@ def delete_memory(
adventure_id: int, adventure_id: int,
memory_id: int, memory_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
get_adventure_or_404(adventure_id, db, user)
memory = db.get(models.Memory, memory_id) memory = db.get(models.Memory, memory_id)
if memory is None or memory.adventure_id != adventure_id: if memory is None or memory.adventure_id != adventure_id:
raise HTTPException(404, "Memory not found") raise HTTPException(404, "Memory not found")
+5 -4
View File
@@ -14,7 +14,7 @@ from ... import models, schemas, worldstate
from ...database import get_db from ...database import get_db
from . import turns from . import turns
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .scenario_text import ( from .scenario_text import (
CARD_FIELDS, SCENARIO_TEXT_FIELDS, fill_placeholders, scenario_card_specs, CARD_FIELDS, SCENARIO_TEXT_FIELDS, fill_placeholders, scenario_card_specs,
scenario_placeholder_names, scenario_placeholder_names,
@@ -113,10 +113,11 @@ def plan_refresh(
@router.get("/{adventure_id}/refresh", response_model=schemas.RefreshPlan) @router.get("/{adventure_id}/refresh", response_model=schemas.RefreshPlan)
def preview_refresh( def preview_refresh(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns what "Update from scenario" would change, for the confirm dialog.""" """Returns what "Update from scenario" would change, for the confirm dialog."""
adventure = get_adventure_or_404(adventure_id, db, user)
scenario = resolve_source_scenario(adventure, db, user) scenario = resolve_source_scenario(adventure, db, user)
if scenario is None: if scenario is None:
raise HTTPException(404, "No scenario to update from") raise HTTPException(404, "No scenario to update from")
@@ -137,6 +138,7 @@ def refresh_from_scenario(
payload: schemas.AdventureRefresh = Body(default=schemas.AdventureRefresh()), payload: schemas.AdventureRefresh = Body(default=schemas.AdventureRefresh()),
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Copies the scenario's current plot text, story cards, and stat schema over """Copies the scenario's current plot text, story cards, and stat schema over
this adventure's copy. this adventure's copy.
@@ -150,7 +152,6 @@ def refresh_from_scenario(
cards; and, through `worldstate.reconcile`, the live value of every stat the cards; and, through `worldstate.reconcile`, the live value of every stat the
schema still defines. schema still defines.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
scenario = resolve_source_scenario(adventure, db, user) scenario = resolve_source_scenario(adventure, db, user)
if scenario is None: if scenario is None:
raise HTTPException(404, "No scenario to update from") raise HTTPException(404, "No scenario to update from")
+6 -6
View File
@@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
from ... import models, schemas from ... import models, schemas
from ...database import get_db from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
# Fields that are copied from a library Script into its adventure-script # Fields that are copied from a library Script into its adventure-script
@@ -59,9 +59,10 @@ def _mark_out_of_date(
@router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut]) @router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut])
def list_adventure_scripts( def list_adventure_scripts(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
return [_mark_out_of_date(s, db, user) for s in adventure.scripts] return [_mark_out_of_date(s, db, user) for s in adventure.scripts]
@@ -74,12 +75,12 @@ def sync_adventure_script(
adv_script_id: int, adv_script_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Overwrites this copy's code with the latest from its library script. """Overwrites this copy's code with the latest from its library script.
`enabled`, `position`, and the adventure's shared `script_state` are kept. `enabled`, `position`, and the adventure's shared `script_state` are kept.
""" """
get_adventure_or_404(adventure_id, db, user)
script = db.get(models.AdventureScript, adv_script_id) script = db.get(models.AdventureScript, adv_script_id)
if script is None or script.adventure_id != adventure_id: if script is None or script.adventure_id != adventure_id:
raise HTTPException(404, "Script not found") raise HTTPException(404, "Script not found")
@@ -104,9 +105,8 @@ def update_adventure_script(
adv_script_id: int, adv_script_id: int,
payload: schemas.AdventureScriptUpdate, payload: schemas.AdventureScriptUpdate,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
get_adventure_or_404(adventure_id, db, user)
script = db.get(models.AdventureScript, adv_script_id) script = db.get(models.AdventureScript, adv_script_id)
if script is None or script.adventure_id != adventure_id: if script is None or script.adventure_id != adventure_id:
raise HTTPException(404, "Script not found") raise HTTPException(404, "Script not found")
+10 -12
View File
@@ -16,12 +16,12 @@ from ...context import cursors
from ...context import lineage from ...context import lineage
from ...database import get_db from ...database import get_db
from ...scripting import ScriptPipeline from ...scripting import ScriptPipeline
from ...sse import SSE_HEADERS
from . import turns from . import turns
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .nodes import delete_turn, last_action, stand_on from .nodes import delete_turn, last_action, stand_on
from .paging import action_window, annotate_takes, current_window from .paging import action_window, annotate_takes, current_window
from .turns import SSE_HEADERS
@router.post("/{adventure_id}/retry") @router.post("/{adventure_id}/retry")
@@ -30,6 +30,7 @@ def retry_action(
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Regenerates the last AI action and keeps the discarded attempt. """Regenerates the last AI action and keeps the discarded attempt.
@@ -38,7 +39,6 @@ def retry_action(
attempt is stored as a sibling at the same coordinate. No text the AI wrote attempt is stored as a sibling at the same coordinate. No text the AI wrote
is rewritten or deleted. is rewritten or deleted.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user) limits.rate_limit("turn", request, user)
turns.check_demo_cap(db, user) turns.check_demo_cap(db, user)
turns.acquire_turn_lock(adventure_id) turns.acquire_turn_lock(adventure_id)
@@ -79,7 +79,7 @@ def list_variants(
adventure_id: int, adventure_id: int,
action_id: int, action_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Returns every attempt made for one AI turn. """Returns every attempt made for one AI turn.
@@ -90,7 +90,6 @@ def list_variants(
Switching changes which row the story tells, and a client that holds an id it Switching changes which row the story tells, and a client that holds an id it
received a moment ago still has to be able to ask about the same turn. received a moment ago still has to be able to ask about the same turn.
""" """
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
@@ -119,7 +118,7 @@ def select_variant(
action_id: int, action_id: int,
payload: schemas.VariantSelect, payload: schemas.VariantSelect,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Makes an earlier attempt live again and restores the state it produced. """Makes an earlier attempt live again and restores the state it produced.
@@ -130,7 +129,6 @@ def select_variant(
would leave the story contradicting itself. The attempts of earlier turns would leave the story contradicting itself. The attempts of earlier turns
stay readable through `list_variants`. stay readable through `list_variants`.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
@@ -172,7 +170,7 @@ def fork_from_attempt(
adventure_id: int, adventure_id: int,
action_id: int, action_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, adventure: models.Adventure = Depends(current_adventure),
): ):
"""Continues the story from this attempt, forking a branch if one is needed. """Continues the story from this attempt, forking a branch if one is needed.
@@ -184,7 +182,6 @@ def fork_from_attempt(
* The story has moved past its turn, so the endpoint forks. The attempt gets * The story has moved past its turn, so the endpoint forks. The attempt gets
a branch of its own, and the line it leaves keeps every turn it has. a branch of its own, and the line it leaves keeps every turn it has.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id) action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id: if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found") raise HTTPException(404, "Action not found")
@@ -236,6 +233,7 @@ def add_take(
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Plays a turn again, whoever wrote it. """Plays a turn again, whoever wrote it.
@@ -256,7 +254,6 @@ def add_take(
turn, so the new attempt is written at the same depth under the same parent, turn, so the new attempt is written at the same depth under the same parent,
and the line it leaves is unchanged. No node below is copied. and the line it leaves is unchanged. No node below is copied.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user) limits.rate_limit("turn", request, user)
limits.check_row_cap("actions", db, user, adventure=adventure) limits.check_row_cap("actions", db, user, adventure=adventure)
turns.check_demo_cap(db, user) turns.check_demo_cap(db, user)
@@ -317,7 +314,9 @@ def add_take(
@router.post("/{adventure_id}/undo", response_model=schemas.ActionPage) @router.post("/{adventure_id}/undo", response_model=schemas.ActionPage)
def undo_turn( def undo_turn(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
): ):
"""Deletes the last turn: the trailing AI action and its player action, if any. """Deletes the last turn: the trailing AI action and its player action, if any.
@@ -325,7 +324,6 @@ def undo_turn(
ran, and it prunes any memory that summarized the removed actions. The turn ran, and it prunes any memory that summarized the removed actions. The turn
lock prevents an undo while a turn is still generating. lock prevents an undo while a turn is still generating.
""" """
adventure = get_adventure_or_404(adventure_id, db, user)
turns.acquire_turn_lock(adventure_id) turns.acquire_turn_lock(adventure_id)
try: try:
# Only the last turn is removed, so fetch the two actions it can # Only the last turn is removed, so fetch the two actions it can
+3 -24
View File
@@ -6,7 +6,6 @@ lock guards one set only while one module owns it. And a test that replaces
`OpenAICompatibleProvider`, `generate_turn`, or `check_demo_cap` patches this `OpenAICompatibleProvider`, `generate_turn`, or `check_demo_cap` patches this
module, which every caller reads through. module, which every caller reads through.
""" """
import json
import threading import threading
from fastapi import Depends, HTTPException, Request from fastapi import Depends, HTTPException, Request
@@ -20,9 +19,10 @@ from ...context import build_context, cursors
from ...database import get_db from ...database import get_db
from ...providers import OpenAICompatibleProvider, PromptParts, ProviderError from ...providers import OpenAICompatibleProvider, PromptParts, ProviderError
from ...scripting import ScriptPipeline from ...scripting import ScriptPipeline
from ...sse import SSE_HEADERS, sse, turn_error
from ..settings import get_settings from ..settings import get_settings
from .deps import CurrentUser, get_adventure_or_404, router from .deps import CurrentUser, current_adventure, router
from .nodes import _move_to_after, next_depth, next_index from .nodes import _move_to_after, next_depth, next_index
from .paging import annotate_takes from .paging import annotate_takes
@@ -98,27 +98,6 @@ def format_player_input(action_type: str, text: str) -> str:
return text # The "story" type is appended as raw text. return text # The "story" type is appended as raw text.
def sse(obj: dict) -> str:
return f"data: {json.dumps(obj)}\n\n"
def turn_error(detail: str, **extra) -> str:
"""Returns an SSE error for a turn that could not be produced, and counts it.
A failed turn is still an HTTP 200 response, so the middleware's status-code
tally cannot see it. This metric exists so that a demo whose model refuses
every request does not report as healthy.
"""
analytics.record(analytics.M_EVENT, analytics.EV_TURN_ERROR)
return sse({"type": "error", "detail": detail, **extra})
# `no-cache` stops an intermediary from caching the stream. `X-Accel-Buffering`
# makes nginx-style reverse proxies, which hosted deploys use, flush each event
# immediately rather than buffer it.
SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
def action_json(action: models.Action, db: Session | None = None) -> dict: def action_json(action: models.Action, db: Session | None = None) -> dict:
"""Serializes one action for the wire. """Serializes one action for the wire.
@@ -425,8 +404,8 @@ def create_action(
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
user: models.User = CurrentUser, user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
): ):
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user) limits.rate_limit("turn", request, user)
limits.check_row_cap("actions", db, user, adventure=adventure) limits.check_row_cap("actions", db, user, adventure=adventure)
check_demo_cap(db, user) check_demo_cap(db, user)
+1 -1
View File
@@ -19,7 +19,7 @@ from sqlalchemy.orm import Session
from .. import auth, limits, models, schemas from .. import auth, limits, models, schemas
from ..database import get_db from ..database import get_db
from ..providers import OpenAICompatibleProvider, ProviderError from ..providers import OpenAICompatibleProvider, ProviderError
from .adventures import SSE_HEADERS, sse from ..sse import SSE_HEADERS, sse
from .settings import get_settings, list_endpoint_models from .settings import get_settings, list_endpoint_models
router = APIRouter(prefix="/api/chat", tags=["chat"]) router = APIRouter(prefix="/api/chat", tags=["chat"])
+30
View File
@@ -0,0 +1,30 @@
"""Server-sent events: the wire format, the headers, and the error frame.
Two routers stream: the turn engine in `routers/adventures/turns.py` and the
chat scratchpad in `routers/chat.py`. Both send JSON objects as SSE `data:`
frames, so the format lives here rather than in either one.
"""
import json
from . import analytics
# `no-cache` stops an intermediary from caching the stream. `X-Accel-Buffering`
# makes nginx-style reverse proxies, which hosted deploys use, flush each event
# immediately rather than buffer it.
SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
def sse(obj: dict) -> str:
"""Returns one SSE frame carrying `obj` as JSON."""
return f"data: {json.dumps(obj)}\n\n"
def turn_error(detail: str, **extra) -> str:
"""Returns an SSE error for a turn that could not be produced, and counts it.
A failed turn is still an HTTP 200 response, so the middleware's status-code
tally cannot see it. This metric exists so that a demo whose model refuses
every request does not report as healthy.
"""
analytics.record(analytics.M_EVENT, analytics.EV_TURN_ERROR)
return sse({"type": "error", "detail": detail, **extra})
+140 -151
View File
@@ -5,10 +5,105 @@ place. A proposed change that breaks a limit is recorded as refused rather than
dropped, because the player and the model both need to see that it did not land. dropped, because the player and the model both need to see that it did not land.
""" """
import copy import copy
from typing import NamedTuple
from .schema import STAT_SECTIONS, _initials, instantiate, npc_name from .schema import STAT_SECTIONS, _initials, instantiate, npc_name
class _Target(NamedTuple):
"""Where one path writes, and what the schema says about the value there."""
kind: str # "flag", "milestone", or "stat"
section: str # "flags", "milestones", "player", "world", or "npc"
key: str # the key inside the container
stat_def: dict | None # the stat's definition, or None for a flag or milestone
npc_id: str | None # the character, for `npc.<id>.<stat>`, else None
npc_stats: dict | None # that character's stat defs, used to build a new block
def container(self, ws: dict) -> dict:
"""Returns the dict this path writes into, and creates it if it is missing.
Call this only when you are about to write. Creating the container is a
side effect, and a rejected path must not leave an empty section behind.
"""
if self.npc_id is None:
return ws.setdefault(self.section, {})
return ws.setdefault("npc", {}).setdefault(self.npc_id, _initials(self.npc_stats))
def _resolve(path: str, stat_schema: dict) -> tuple[_Target | None, dict | None]:
"""Routes one path to the value it names, without deciding what happens to it.
Returns `(target, None)` when the schema defines the path, and
`(None, rejection)` when it does not. The rejection is the entry that goes
into a report's `rejected` list, with its reason and its fix already worded.
`apply_delta` and `apply_override` route paths identically and differ only in
what they write, so the routing lives here and each function keeps its own
write rule.
"""
parts = path.split(".")
# `flags.<name>` is a boolean the scenario declares.
if parts[0] == "flags" and len(parts) == 2:
flag_defs = stat_schema.get("flags") or {}
if parts[1] not in flag_defs:
return None, {
"path": path, "reason": "unknown flag",
"fix": f"There is no flag `{parts[1]}`. {_names_phrase('flag', flag_defs)}",
}
return _Target("flag", "flags", parts[1], None, None, None), None
# `milestones.<id>` records that a milestone is reached.
if parts[0] == "milestones" and len(parts) == 2:
milestones = stat_schema.get("milestones") or {}
if parts[1] not in milestones:
return None, {
"path": path, "reason": "unknown milestone",
"fix": f"There is no milestone `{parts[1]}`. "
f"{_names_phrase('milestone', milestones)}",
}
return _Target("milestone", "milestones", parts[1], None, None, None), None
# `world.<stat>` and `player.<stat>`.
if parts[0] in STAT_SECTIONS and len(parts) == 2:
stat_defs = stat_schema.get(parts[0]) or {}
stat_def = stat_defs.get(parts[1])
if not isinstance(stat_def, dict):
return None, {
"path": path, "reason": "unknown stat",
"fix": f"`{parts[0]}` has no stat `{parts[1]}`. "
f"{_names_phrase('stat', stat_defs)}",
}
return _Target("stat", parts[0], parts[1], stat_def, None, None), None
# `npc.<id>.<stat>`. Each character carries its own stat definitions.
if parts[0] == "npc" and len(parts) == 3:
npcs = stat_schema.get("npcs") or {}
ndef = npcs.get(parts[1])
if not isinstance(ndef, dict):
return None, {
"path": path, "reason": "unknown npc",
"fix": f"There is no character `{parts[1]}`. "
f"{_names_phrase('character', npcs)}",
}
stat_defs = ndef.get("stats") or {}
stat_def = stat_defs.get(parts[2])
if not isinstance(stat_def, dict):
return None, {
"path": path, "reason": "unknown npc stat",
"fix": f"`{npc_name(ndef, parts[1])}` has no stat `{parts[2]}`. "
f"{_names_phrase('stat', stat_defs)}",
}
return _Target("stat", "npc", parts[2], stat_def, parts[1], stat_defs), None
return None, {
"path": path, "reason": "unknown path",
"fix": f"`{path}` is not a tracked value. Use player.<stat>, "
f"world.<stat>, npc.<id>.<stat>, flags.<name> or milestones.<id>.",
}
def _coerce_number(value): def _coerce_number(value):
if isinstance(value, bool): # `bool` is an `int` subclass, so reject it here. if isinstance(value, bool): # `bool` is an `int` subclass, so reject it here.
return None return None
@@ -177,10 +272,6 @@ def apply_override(world_state: dict, stat_schema: dict, overrides: dict) -> tup
if not isinstance(overrides, dict): if not isinstance(overrides, dict):
return ws, report return ws, report
milestones = stat_schema.get("milestones") or {}
flag_defs = stat_schema.get("flags") or {}
npcs = stat_schema.get("npcs") or {}
def set_stat(container: dict, key: str, stat_def: dict, value, path: str) -> None: def set_stat(container: dict, key: str, stat_def: dict, value, path: str) -> None:
if stat_def.get("type") == "text": if stat_def.get("type") == "text":
if not isinstance(value, str): if not isinstance(value, str):
@@ -212,80 +303,42 @@ def apply_override(world_state: dict, stat_schema: dict, overrides: dict) -> tup
for raw_path, value in overrides.items(): for raw_path, value in overrides.items():
path = str(raw_path) path = str(raw_path)
parts = path.split(".") target, rejection = _resolve(path, stat_schema)
if rejection is not None:
report["rejected"].append(rejection)
continue
if parts[0] == "flags" and len(parts) == 2: if target.kind == "flag":
fid = parts[1]
if fid not in flag_defs:
report["rejected"].append({"path": path, "reason": "unknown flag"})
continue
if not isinstance(value, bool): if not isinstance(value, bool):
report["rejected"].append({"path": path, "reason": "not a boolean"}) report["rejected"].append({
"path": path, "reason": "not a boolean",
"fix": f"`{path}` takes true or false.",
})
continue continue
flags = ws.setdefault("flags", {}) flags = target.container(ws)
old = bool(flags.get(fid, False)) old = bool(flags.get(target.key, False))
flags[fid] = value flags[target.key] = value
report["applied"].append({"path": path, "old": old, "new": value}) report["applied"].append({"path": path, "old": old, "new": value})
continue
if parts[0] == "milestones" and len(parts) == 2: elif target.kind == "milestone":
mid = parts[1] # An override toggles a milestone in both directions, so unlike a
if mid not in milestones: # delta it takes false as well as true.
report["rejected"].append({"path": path, "reason": "unknown milestone"})
continue
if not isinstance(value, bool): if not isinstance(value, bool):
report["rejected"].append({"path": path, "reason": "not a boolean"}) report["rejected"].append({
"path": path, "reason": "not a boolean",
"fix": f"`{path}` takes true or false.",
})
continue continue
reached = ws.setdefault("milestones", {}) reached = target.container(ws)
old = bool(reached.get(mid, {}).get("reached")) old = bool(reached.get(target.key, {}).get("reached"))
if value: if value:
reached[mid] = {"reached": True} reached[target.key] = {"reached": True}
else: else:
reached.pop(mid, None) reached.pop(target.key, None)
report["applied"].append({"path": path, "old": old, "new": value}) report["applied"].append({"path": path, "old": old, "new": value})
continue
if parts[0] in STAT_SECTIONS and len(parts) == 2: else:
stat_def = (stat_schema.get(parts[0]) or {}).get(parts[1]) set_stat(target.container(ws), target.key, target.stat_def, value, path)
if not isinstance(stat_def, dict):
report["rejected"].append({
"path": path, "reason": "unknown stat",
"fix": f"`{parts[0]}` has no stat `{parts[1]}`. "
f"{_names_phrase('stat', stat_schema.get(parts[0]) or {})}",
})
continue
container = ws.setdefault(parts[0], {})
set_stat(container, parts[1], stat_def, value, path)
continue
if parts[0] == "npc" and len(parts) == 3:
ndef = npcs.get(parts[1])
if not isinstance(ndef, dict):
report["rejected"].append({
"path": path, "reason": "unknown npc",
"fix": f"There is no character `{parts[1]}`. "
f"{_names_phrase('character', npcs)}",
})
continue
stat_defs = ndef.get("stats") or {}
stat_def = stat_defs.get(parts[2])
if not isinstance(stat_def, dict):
report["rejected"].append({
"path": path, "reason": "unknown npc stat",
"fix": f"`{npc_name(ndef, parts[1])}` has no stat `{parts[2]}`. "
f"{_names_phrase('stat', stat_defs)}",
})
continue
npc_state = ws.setdefault("npc", {})
container = npc_state.setdefault(parts[1], _initials(stat_defs))
set_stat(container, parts[2], stat_def, value, path)
continue
report["rejected"].append({
"path": path, "reason": "unknown path",
"fix": f"`{path}` is not a tracked value. Use player.<stat>, "
f"world.<stat>, npc.<id>.<stat>, flags.<name> or milestones.<id>.",
})
return ws, report return ws, report
@@ -307,46 +360,29 @@ def apply_delta(world_state: dict, stat_schema: dict, delta: dict,
if not isinstance(delta, dict): if not isinstance(delta, dict):
return ws, report return ws, report
milestones = stat_schema.get("milestones") or {}
flag_defs = stat_schema.get("flags") or {}
npcs = stat_schema.get("npcs") or {}
for raw_path, change in delta.items(): for raw_path, change in delta.items():
path = str(raw_path) path = str(raw_path)
parts = path.split(".") target, rejection = _resolve(path, stat_schema)
if rejection is not None:
# `flags.<name>` is a two-way boolean, and either value is accepted. report["rejected"].append(rejection)
if parts[0] == "flags" and len(parts) == 2:
fid = parts[1]
if fid not in flag_defs:
report["rejected"].append({
"path": path, "reason": "unknown flag",
"fix": f"There is no flag `{fid}`. {_names_phrase('flag', flag_defs)}",
})
continue continue
if target.kind == "flag":
# A flag goes both ways, so either value is accepted.
if not isinstance(change, bool): if not isinstance(change, bool):
report["rejected"].append({ report["rejected"].append({
"path": path, "reason": "not a boolean", "path": path, "reason": "not a boolean",
"fix": f"`{path}` takes true or false.", "fix": f"`{path}` takes true or false.",
}) })
continue continue
flags = ws.setdefault("flags", {}) flags = target.container(ws)
old = bool(flags.get(fid, False)) old = bool(flags.get(target.key, False))
if change != old: if change != old:
flags[fid] = change flags[target.key] = change
report["applied"].append({"path": path, "old": old, "new": change}) report["applied"].append({"path": path, "old": old, "new": change})
continue
# `milestones.<id>` is a sticky boolean, and only `true` is accepted. elif target.kind == "milestone":
if parts[0] == "milestones" and len(parts) == 2: # A milestone is sticky, so only true is accepted.
mid = parts[1]
if mid not in milestones:
report["rejected"].append({
"path": path, "reason": "unknown milestone",
"fix": f"There is no milestone `{mid}`. "
f"{_names_phrase('milestone', milestones)}",
})
continue
if change is not True: if change is not True:
report["rejected"].append({ report["rejected"].append({
"path": path, "reason": "not true", "path": path, "reason": "not true",
@@ -354,65 +390,18 @@ def apply_delta(world_state: dict, stat_schema: dict, delta: dict,
f"reached once and never taken back.", f"reached once and never taken back.",
}) })
continue continue
reached = ws.setdefault("milestones", {}) reached = target.container(ws)
if reached.get(mid, {}).get("reached"): if reached.get(target.key, {}).get("reached"):
continue # Already reached, so do nothing. continue # Already reached, so do nothing.
reached[mid] = {"reached": True, "at": action_index} reached[target.key] = {"reached": True, "at": action_index}
report["applied"].append({"path": path, "old": False, "new": True}) report["applied"].append({"path": path, "old": False, "new": True})
continue
# world.<stat> / player.<stat> elif target.stat_def.get("type") == "text":
if parts[0] in STAT_SECTIONS and len(parts) == 2: _apply_text_stat(target.container(ws), target.key, target.stat_def, change,
stat_def = (stat_schema.get(parts[0]) or {}).get(parts[1]) path, action_index, meta, report)
if not isinstance(stat_def, dict):
report["rejected"].append({
"path": path, "reason": "unknown stat",
"fix": f"`{parts[0]}` has no stat `{parts[1]}`. "
f"{_names_phrase('stat', stat_schema.get(parts[0]) or {})}",
})
continue
container = ws.setdefault(parts[0], {})
if stat_def.get("type") == "text":
_apply_text_stat(container, parts[1], stat_def, change, path,
action_index, meta, report)
else: else:
_apply_stat(container, parts[1], stat_def, change, path, _apply_stat(target.container(ws), target.key, target.stat_def, change,
action_index, meta, report) path, action_index, meta, report)
continue
# `npc.<npcId>.<stat>`. Each NPC has its own stat definitions.
if parts[0] == "npc" and len(parts) == 3:
ndef = npcs.get(parts[1])
if not isinstance(ndef, dict):
report["rejected"].append({
"path": path, "reason": "unknown npc",
"fix": f"There is no character `{parts[1]}`. "
f"{_names_phrase('character', npcs)}",
})
continue
stat_defs = ndef.get("stats") or {}
stat_def = stat_defs.get(parts[2])
if not isinstance(stat_def, dict):
report["rejected"].append({
"path": path, "reason": "unknown npc stat",
"fix": f"`{npc_name(ndef, parts[1])}` has no stat `{parts[2]}`. "
f"{_names_phrase('stat', stat_defs)}",
})
continue
npc_state = ws.setdefault("npc", {})
container = npc_state.setdefault(parts[1], _initials(stat_defs))
if stat_def.get("type") == "text":
_apply_text_stat(container, parts[2], stat_def, change, path,
action_index, meta, report)
else:
_apply_stat(container, parts[2], stat_def, change, path,
action_index, meta, report)
continue
report["rejected"].append({
"path": path, "reason": "unknown path",
"fix": f"`{path}` is not a tracked value. Use player.<stat>, "
f"world.<stat>, npc.<id>.<stat>, flags.<name> or milestones.<id>.",
})
return ws, report return ws, report
+7 -7
View File
@@ -81,7 +81,7 @@ def test_undo_reverts_state_to_before_the_turn(db):
_add(db, adv, 2, "ai", state_after={"gold": 10}) _add(db, adv, 2, "ai", state_after={"gold": 10})
db.commit() db.commit()
adventures.undo_turn(adv.id, db=db, user=user) adventures.undo_turn(adv.id, db=db, adventure=adv)
assert adv.script_state == {"gold": 0} assert adv.script_state == {"gold": 0}
assert [a.type for a in adv.actions] == ["start"] assert [a.type for a in adv.actions] == ["start"]
@@ -95,7 +95,7 @@ def test_undo_of_bare_continue_uses_the_node_in_front(db):
_add(db, adv, 1, "ai", state_after={"gold": 5}) _add(db, adv, 1, "ai", state_after={"gold": 5})
db.commit() db.commit()
adventures.undo_turn(adv.id, db=db, user=user) adventures.undo_turn(adv.id, db=db, adventure=adv)
assert adv.script_state == {"gold": 0} assert adv.script_state == {"gold": 0}
assert [a.type for a in adv.actions] == ["start"] assert [a.type for a in adv.actions] == ["start"]
@@ -110,7 +110,7 @@ def test_undo_leaves_state_untouched_when_snapshot_missing(db):
_add(db, adv, 2, "ai") _add(db, adv, 2, "ai")
_forget_snapshots(db, adv) _forget_snapshots(db, adv)
adventures.undo_turn(adv.id, db=db, user=user) adventures.undo_turn(adv.id, db=db, adventure=adv)
assert adv.script_state == {"gold": 10} assert adv.script_state == {"gold": 10}
@@ -120,7 +120,7 @@ def test_undo_raises_when_nothing_to_undo(db):
_add(db, adv, 0, "start") _add(db, adv, 0, "start")
db.commit() db.commit()
with pytest.raises(HTTPException) as exc: with pytest.raises(HTTPException) as exc:
adventures.undo_turn(adv.id, db=db, user=user) adventures.undo_turn(adv.id, db=db, adventure=adv)
assert exc.value.status_code == 400 assert exc.value.status_code == 400
@@ -133,7 +133,7 @@ def test_undo_blocked_by_active_turn_lock(db):
adventures.turns.acquire_turn_lock(adv.id) # a turn is "generating" adventures.turns.acquire_turn_lock(adv.id) # a turn is "generating"
try: try:
with pytest.raises(HTTPException) as exc: with pytest.raises(HTTPException) as exc:
adventures.undo_turn(adv.id, db=db, user=user) adventures.undo_turn(adv.id, db=db, adventure=adv)
assert exc.value.status_code == 409 assert exc.value.status_code == 409
# The failed undo must not have released someone else's lock. # The failed undo must not have released someone else's lock.
assert adv.id in adventures.turns._active_turns assert adv.id in adventures.turns._active_turns
@@ -151,7 +151,7 @@ def test_undo_prunes_memory_covering_removed_actions(db):
db.add_all([covering, keep]) db.add_all([covering, keep])
db.commit() db.commit()
adventures.undo_turn(adv.id, db=db, user=user) # removes indexes 2 & 3 adventures.undo_turn(adv.id, db=db, adventure=adv) # removes indexes 2 & 3
texts = {m.text for m in adv.memories} texts = {m.text for m in adv.memories}
assert texts == {"k"} assert texts == {"k"}
@@ -228,7 +228,7 @@ def test_retry_restores_the_state_the_turn_started_from(db, monkeypatch):
yield # make it an async generator yield # make it an async generator
monkeypatch.setattr(adventures.turns, "generate_turn", _noop) monkeypatch.setattr(adventures.turns, "generate_turn", _noop)
adventures.retry_action(adv.id, request=None, db=db, user=user) adventures.retry_action(adv.id, request=None, db=db, user=user, adventure=adv)
assert adv.script_state == {"gold": 10} assert adv.script_state == {"gold": 10}
# Nothing is written until a replacement actually arrives: the attempt on # Nothing is written until a replacement actually arrives: the attempt on