From 2c57b1ceab01686acec4707994f7d1fc6ec1cd88 Mon Sep 17 00:00:00 2001 From: parththakkar106 Date: Sat, 29 Aug 2026 02:04:02 +0530 Subject: [PATCH] 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.`, `milestones.`, `world.`, `player.`, and `npc..` 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 Claude-Session: https://claude.ai/code/session_014Dix4oGV3njgWRdu7P9t6r --- backend/app/routers/adventures/__init__.py | 4 +- backend/app/routers/adventures/actions.py | 12 +- backend/app/routers/adventures/branches.py | 16 +- backend/app/routers/adventures/bundle_io.py | 6 +- backend/app/routers/adventures/crud.py | 25 +- backend/app/routers/adventures/deps.py | 17 ++ backend/app/routers/adventures/insights.py | 10 +- backend/app/routers/adventures/memories.py | 16 +- backend/app/routers/adventures/refresh.py | 9 +- backend/app/routers/adventures/scripts.py | 12 +- backend/app/routers/adventures/takes.py | 22 +- backend/app/routers/adventures/turns.py | 27 +- backend/app/routers/chat.py | 2 +- backend/app/sse.py | 30 ++ backend/app/worldstate/apply.py | 297 ++++++++++---------- backend/tests/test_state_revert.py | 14 +- 16 files changed, 258 insertions(+), 261 deletions(-) create mode 100644 backend/app/sse.py diff --git a/backend/app/routers/adventures/__init__.py b/backend/app/routers/adventures/__init__.py index afbf4de..0996c11 100644 --- a/backend/app/routers/adventures/__init__.py +++ b/backend/app/routers/adventures/__init__.py @@ -42,17 +42,15 @@ from ... import limits # noqa: F401 `adventures.limits` is patched by tests. from .crud import SNIPPET_MAX, _snippet from .paging import ACTION_PAGE from .takes import retry_action, undo_turn -from .turns import SSE_HEADERS, sse, world_delta_of +from .turns import world_delta_of __all__ = [ "ACTION_PAGE", "SNIPPET_MAX", - "SSE_HEADERS", "_snippet", "limits", "retry_action", "router", - "sse", "undo_turn", "world_delta_of", ] diff --git a/backend/app/routers/adventures/actions.py b/backend/app/routers/adventures/actions.py index 553634b..d14f3f2 100644 --- a/backend/app/routers/adventures/actions.py +++ b/backend/app/routers/adventures/actions.py @@ -11,18 +11,17 @@ from sqlalchemy.orm import Session from ... import models, schemas, tree 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 .paging import ACTION_PAGE, action_window, annotate_takes @router.get("/{adventure_id}/actions", response_model=schemas.ActionPage) def list_actions( - adventure_id: int, before_id: int | None = None, limit: int = ACTION_PAGE, 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. @@ -30,7 +29,6 @@ def list_actions( 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. """ - adventure = get_adventure_or_404(adventure_id, db, user) limit = max(1, min(limit, ACTION_PAGE * 4)) actions, total, has_more = action_window( db, adventure, before_id=before_id, limit=limit @@ -51,9 +49,8 @@ def update_action( action_id: int, payload: schemas.ActionUpdate, 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) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -70,9 +67,8 @@ def delete_action( adventure_id: int, action_id: int, 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) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") diff --git a/backend/app/routers/adventures/branches.py b/backend/app/routers/adventures/branches.py index 5453531..086715c 100644 --- a/backend/app/routers/adventures/branches.py +++ b/backend/app/routers/adventures/branches.py @@ -18,14 +18,15 @@ from ...context import lineage from ...database import get_db 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 .paging import current_window @router.get("/{adventure_id}/branches", response_model=list[schemas.BranchOut]) 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. @@ -35,7 +36,6 @@ def list_branches( `actions`, never one query per branch, so a view of a hundred forks does not cost a hundred round trips. """ - adventure = get_adventure_or_404(adventure_id, db, user) branches = ( db.query(models.Branch) .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 ) def rename_branch( - adventure_id: int, branch_id: int, payload: schemas.BranchRename, 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. @@ -107,7 +106,6 @@ def rename_branch( name anyone chose, and storing one gives the client an empty label to draw instead of the fork depth. """ - adventure = get_adventure_or_404(adventure_id, db, user) branch = get_branch_or_404(adventure, branch_id, db) name = (payload.name or "").strip() branch.name = name or None @@ -145,7 +143,7 @@ def delete_branch( adventure_id: int, branch_id: int, db: Session = Depends(get_db), - user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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 however deep the subtree is. """ - adventure = get_adventure_or_404(adventure_id, db, user) branch = get_branch_or_404(adventure, branch_id, db) if branch.parent_branch_id is None: raise HTTPException( @@ -238,7 +235,7 @@ def switch_branch( adventure_id: int, branch_id: int, db: Session = Depends(get_db), - user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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 snapshot. """ - adventure = get_adventure_or_404(adventure_id, db, user) branch = db.get(models.Branch, branch_id) if branch is None or branch.adventure_id != adventure.id: raise HTTPException(404, "Branch not found") diff --git a/backend/app/routers/adventures/bundle_io.py b/backend/app/routers/adventures/bundle_io.py index a0ce3ec..3c1fd50 100644 --- a/backend/app/routers/adventures/bundle_io.py +++ b/backend/app/routers/adventures/bundle_io.py @@ -10,19 +10,19 @@ from sqlalchemy.orm import Session from ... import analytics, bundle, limits, models, schemas 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") 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. `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. """ - adv = get_adventure_or_404(adventure_id, db, user) return bundle.export(db, adv) diff --git a/backend/app/routers/adventures/crud.py b/backend/app/routers/adventures/crud.py index 24c216d..4c26b67 100644 --- a/backend/app/routers/adventures/crud.py +++ b/backend/app/routers/adventures/crud.py @@ -12,7 +12,7 @@ from sqlalchemy.orm.attributes import set_committed_value from ... import analytics, attempts, images, limits, memorybank, models, schemas, tree, worldstate 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 .scenario_text import fill_placeholders, scenario_card_specs @@ -229,7 +229,8 @@ def create_adventure( @router.get("/{adventure_id}", response_model=schemas.AdventureOut) 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. @@ -238,7 +239,6 @@ def get_adventure( exist above. `GET /{id}/actions` serves the older pages as the reader scrolls up. """ - adventure = get_adventure_or_404(adventure_id, db, user) actions, total, _ = action_window(db, adventure) # Annotate before handing over the window. This path serializes through the # relationship rather than building `ActionOut` itself, so the pager numbers @@ -259,7 +259,7 @@ def get_adventure( @router.get("/{adventure_id}/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. @@ -267,21 +267,19 @@ def get_script_state( `state.x`, persisted after each hook. It stays `{}` until a script sets a variable. """ - adventure = get_adventure_or_404(adventure_id, db, user) state = adventure.script_state if isinstance(adventure.script_state, dict) else {} return {"state": state} @router.get("/{adventure_id}/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`. The play view uses both to render the character sheet and the milestones. `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 state = adventure.world_state if isinstance(adventure.world_state, dict) else {} return { @@ -292,10 +290,9 @@ def get_world_state( @router.put("/{adventure_id}/world-state") def override_world_state( - adventure_id: int, overrides: dict = Body(...), 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. @@ -303,7 +300,6 @@ def override_world_state( `milestones.y` to their new absolute values. The endpoint rejects unknown 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 if not worldstate.has_schema(schema): 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) def update_adventure( - adventure_id: int, payload: schemas.AdventureUpdate, 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(): setattr(adventure, field, value) db.commit() @@ -330,9 +324,10 @@ def update_adventure( @router.delete("/{adventure_id}", status_code=204) 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.commit() # No later request reads this adventure's vectors, so drop them now. The diff --git a/backend/app/routers/adventures/deps.py b/backend/app/routers/adventures/deps.py index 02faeba..d5e3137 100644 --- a/backend/app/routers/adventures/deps.py +++ b/backend/app/routers/adventures/deps.py @@ -9,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException from sqlalchemy.orm import Session from ... import auth, models +from ...database import get_db 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: raise HTTPException(404, "Adventure not found") 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) diff --git a/backend/app/routers/adventures/insights.py b/backend/app/routers/adventures/insights.py index 5526bd1..254f676 100644 --- a/backend/app/routers/adventures/insights.py +++ b/backend/app/routers/adventures/insights.py @@ -12,15 +12,16 @@ from ...context import build_context from ...database import get_db 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") 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.""" - adventure = get_adventure_or_404(adventure_id, db, user) settings = get_settings(db, user) if auth.resolve_provider_config(settings).using_demo: memories = ( @@ -39,9 +40,8 @@ def action_context( adventure_id: int, action_id: int, 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) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") diff --git a/backend/app/routers/adventures/memories.py b/backend/app/routers/adventures/memories.py index e6acf9b..119b8d3 100644 --- a/backend/app/routers/adventures/memories.py +++ b/backend/app/routers/adventures/memories.py @@ -11,7 +11,7 @@ from ... import limits, memorybank, models, schemas, tree from ...context import lineage 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 @@ -32,9 +32,10 @@ MEMORY_LIST_COLUMNS = ( @router.get("/{adventure_id}/memories", response_model=list[schemas.MemoryOut]) 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`. # Retrieval used to walk the relationship, which is why a turn cost # 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) def create_memory( - adventure_id: int, payload: schemas.MemoryCreate, db: Session = Depends(get_db), user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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) if not payload.text.strip(): raise HTTPException(400, "Memory text cannot be empty") @@ -88,9 +88,8 @@ def update_memory( memory_id: int, payload: schemas.MemoryUpdate, 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) if memory is None or memory.adventure_id != adventure_id: raise HTTPException(404, "Memory not found") @@ -108,9 +107,8 @@ def delete_memory( adventure_id: int, memory_id: int, 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) if memory is None or memory.adventure_id != adventure_id: raise HTTPException(404, "Memory not found") diff --git a/backend/app/routers/adventures/refresh.py b/backend/app/routers/adventures/refresh.py index 7505a7c..c83fdff 100644 --- a/backend/app/routers/adventures/refresh.py +++ b/backend/app/routers/adventures/refresh.py @@ -14,7 +14,7 @@ from ... import models, schemas, worldstate from ...database import get_db from . import turns -from .deps import CurrentUser, get_adventure_or_404, router +from .deps import CurrentUser, current_adventure, router from .scenario_text import ( CARD_FIELDS, SCENARIO_TEXT_FIELDS, fill_placeholders, scenario_card_specs, scenario_placeholder_names, @@ -113,10 +113,11 @@ def plan_refresh( @router.get("/{adventure_id}/refresh", response_model=schemas.RefreshPlan) 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.""" - adventure = get_adventure_or_404(adventure_id, db, user) scenario = resolve_source_scenario(adventure, db, user) if scenario is None: raise HTTPException(404, "No scenario to update from") @@ -137,6 +138,7 @@ def refresh_from_scenario( payload: schemas.AdventureRefresh = Body(default=schemas.AdventureRefresh()), db: Session = Depends(get_db), user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """Copies the scenario's current plot text, story cards, and stat schema over this adventure's copy. @@ -150,7 +152,6 @@ def refresh_from_scenario( cards; and, through `worldstate.reconcile`, the live value of every stat the schema still defines. """ - adventure = get_adventure_or_404(adventure_id, db, user) scenario = resolve_source_scenario(adventure, db, user) if scenario is None: raise HTTPException(404, "No scenario to update from") diff --git a/backend/app/routers/adventures/scripts.py b/backend/app/routers/adventures/scripts.py index ca2e2a0..5e878ec 100644 --- a/backend/app/routers/adventures/scripts.py +++ b/backend/app/routers/adventures/scripts.py @@ -11,7 +11,7 @@ from sqlalchemy.orm import Session from ... import models, schemas 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 @@ -59,9 +59,10 @@ def _mark_out_of_date( @router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut]) 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] @@ -74,12 +75,12 @@ def sync_adventure_script( adv_script_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """Overwrites this copy's code with the latest from its library script. `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) if script is None or script.adventure_id != adventure_id: raise HTTPException(404, "Script not found") @@ -104,9 +105,8 @@ def update_adventure_script( adv_script_id: int, payload: schemas.AdventureScriptUpdate, 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) if script is None or script.adventure_id != adventure_id: raise HTTPException(404, "Script not found") diff --git a/backend/app/routers/adventures/takes.py b/backend/app/routers/adventures/takes.py index 01e083f..13bfb37 100644 --- a/backend/app/routers/adventures/takes.py +++ b/backend/app/routers/adventures/takes.py @@ -16,12 +16,12 @@ from ...context import cursors from ...context import lineage from ...database import get_db from ...scripting import ScriptPipeline +from ...sse import SSE_HEADERS 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 .paging import action_window, annotate_takes, current_window -from .turns import SSE_HEADERS @router.post("/{adventure_id}/retry") @@ -30,6 +30,7 @@ def retry_action( request: Request, db: Session = Depends(get_db), user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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 is rewritten or deleted. """ - adventure = get_adventure_or_404(adventure_id, db, user) limits.rate_limit("turn", request, user) turns.check_demo_cap(db, user) turns.acquire_turn_lock(adventure_id) @@ -79,7 +79,7 @@ def list_variants( adventure_id: int, action_id: int, db: Session = Depends(get_db), - user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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 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) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -119,7 +118,7 @@ def select_variant( action_id: int, payload: schemas.VariantSelect, 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. @@ -130,7 +129,6 @@ def select_variant( would leave the story contradicting itself. The attempts of earlier turns stay readable through `list_variants`. """ - adventure = get_adventure_or_404(adventure_id, db, user) action = db.get(models.Action, action_id) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -172,7 +170,7 @@ def fork_from_attempt( adventure_id: int, action_id: int, 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. @@ -184,7 +182,6 @@ def fork_from_attempt( * 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. """ - adventure = get_adventure_or_404(adventure_id, db, user) action = db.get(models.Action, action_id) if action is None or action.adventure_id != adventure_id: raise HTTPException(404, "Action not found") @@ -236,6 +233,7 @@ def add_take( request: Request, db: Session = Depends(get_db), user: models.User = CurrentUser, + adventure: models.Adventure = Depends(current_adventure), ): """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, 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.check_row_cap("actions", db, user, adventure=adventure) turns.check_demo_cap(db, user) @@ -317,7 +314,9 @@ def add_take( @router.post("/{adventure_id}/undo", response_model=schemas.ActionPage) 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. @@ -325,7 +324,6 @@ def undo_turn( ran, and it prunes any memory that summarized the removed actions. The turn 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) try: # Only the last turn is removed, so fetch the two actions it can diff --git a/backend/app/routers/adventures/turns.py b/backend/app/routers/adventures/turns.py index b828302..58f4402 100644 --- a/backend/app/routers/adventures/turns.py +++ b/backend/app/routers/adventures/turns.py @@ -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 module, which every caller reads through. """ -import json import threading from fastapi import Depends, HTTPException, Request @@ -20,9 +19,10 @@ from ...context import build_context, cursors from ...database import get_db from ...providers import OpenAICompatibleProvider, PromptParts, ProviderError from ...scripting import ScriptPipeline +from ...sse import SSE_HEADERS, sse, turn_error 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 .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. -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: """Serializes one action for the wire. @@ -425,8 +404,8 @@ def create_action( request: Request, db: Session = Depends(get_db), 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.check_row_cap("actions", db, user, adventure=adventure) check_demo_cap(db, user) diff --git a/backend/app/routers/chat.py b/backend/app/routers/chat.py index 78da113..379db67 100644 --- a/backend/app/routers/chat.py +++ b/backend/app/routers/chat.py @@ -19,7 +19,7 @@ from sqlalchemy.orm import Session from .. import auth, limits, models, schemas from ..database import get_db 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 router = APIRouter(prefix="/api/chat", tags=["chat"]) diff --git a/backend/app/sse.py b/backend/app/sse.py new file mode 100644 index 0000000..21365d4 --- /dev/null +++ b/backend/app/sse.py @@ -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}) diff --git a/backend/app/worldstate/apply.py b/backend/app/worldstate/apply.py index c5dfe53..336ab7a 100644 --- a/backend/app/worldstate/apply.py +++ b/backend/app/worldstate/apply.py @@ -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. """ import copy +from typing import NamedTuple 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..`, 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.` 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.` 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.` and `player.`. + 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..`. 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., " + f"world., npc.., flags. or milestones..", + } + + def _coerce_number(value): if isinstance(value, bool): # `bool` is an `int` subclass, so reject it here. return None @@ -177,10 +272,6 @@ def apply_override(world_state: dict, stat_schema: dict, overrides: dict) -> tup if not isinstance(overrides, dict): 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: if stat_def.get("type") == "text": 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(): path = str(raw_path) - parts = path.split(".") - - 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"}) - continue - if not isinstance(value, bool): - report["rejected"].append({"path": path, "reason": "not a boolean"}) - continue - flags = ws.setdefault("flags", {}) - old = bool(flags.get(fid, False)) - flags[fid] = value - report["applied"].append({"path": path, "old": old, "new": value}) + target, rejection = _resolve(path, stat_schema) + if rejection is not None: + report["rejected"].append(rejection) continue - if parts[0] == "milestones" and len(parts) == 2: - mid = parts[1] - if mid not in milestones: - report["rejected"].append({"path": path, "reason": "unknown milestone"}) - continue + if target.kind == "flag": 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 - reached = ws.setdefault("milestones", {}) - old = bool(reached.get(mid, {}).get("reached")) + flags = target.container(ws) + old = bool(flags.get(target.key, False)) + flags[target.key] = value + report["applied"].append({"path": path, "old": old, "new": value}) + + elif target.kind == "milestone": + # An override toggles a milestone in both directions, so unlike a + # delta it takes false as well as true. + if not isinstance(value, bool): + report["rejected"].append({ + "path": path, "reason": "not a boolean", + "fix": f"`{path}` takes true or false.", + }) + continue + reached = target.container(ws) + old = bool(reached.get(target.key, {}).get("reached")) if value: - reached[mid] = {"reached": True} + reached[target.key] = {"reached": True} else: - reached.pop(mid, None) + reached.pop(target.key, None) report["applied"].append({"path": path, "old": old, "new": value}) - continue - if parts[0] in STAT_SECTIONS and len(parts) == 2: - stat_def = (stat_schema.get(parts[0]) or {}).get(parts[1]) - 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., " - f"world., npc.., flags. or milestones..", - }) + else: + set_stat(target.container(ws), target.key, target.stat_def, value, path) return ws, report @@ -307,46 +360,29 @@ def apply_delta(world_state: dict, stat_schema: dict, delta: dict, if not isinstance(delta, dict): 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(): path = str(raw_path) - parts = path.split(".") + target, rejection = _resolve(path, stat_schema) + if rejection is not None: + report["rejected"].append(rejection) + continue - # `flags.` is a two-way boolean, and either value is accepted. - 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 + if target.kind == "flag": + # A flag goes both ways, so either value is accepted. if not isinstance(change, bool): report["rejected"].append({ "path": path, "reason": "not a boolean", "fix": f"`{path}` takes true or false.", }) continue - flags = ws.setdefault("flags", {}) - old = bool(flags.get(fid, False)) + flags = target.container(ws) + old = bool(flags.get(target.key, False)) if change != old: - flags[fid] = change + flags[target.key] = change report["applied"].append({"path": path, "old": old, "new": change}) - continue - # `milestones.` is a sticky boolean, and only `true` is accepted. - if parts[0] == "milestones" and len(parts) == 2: - 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 + elif target.kind == "milestone": + # A milestone is sticky, so only true is accepted. if change is not True: report["rejected"].append({ "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.", }) continue - reached = ws.setdefault("milestones", {}) - if reached.get(mid, {}).get("reached"): + reached = target.container(ws) + if reached.get(target.key, {}).get("reached"): 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}) - continue - # world. / player. - if parts[0] in STAT_SECTIONS and len(parts) == 2: - stat_def = (stat_schema.get(parts[0]) or {}).get(parts[1]) - 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: - _apply_stat(container, parts[1], stat_def, change, path, - action_index, meta, report) - continue + elif target.stat_def.get("type") == "text": + _apply_text_stat(target.container(ws), target.key, target.stat_def, change, + path, action_index, meta, report) - # `npc..`. 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., " - f"world., npc.., flags. or milestones..", - }) + else: + _apply_stat(target.container(ws), target.key, target.stat_def, change, + path, action_index, meta, report) return ws, report diff --git a/backend/tests/test_state_revert.py b/backend/tests/test_state_revert.py index 5ffd8a6..b1869d2 100644 --- a/backend/tests/test_state_revert.py +++ b/backend/tests/test_state_revert.py @@ -81,7 +81,7 @@ def test_undo_reverts_state_to_before_the_turn(db): _add(db, adv, 2, "ai", state_after={"gold": 10}) 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 [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}) 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 [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") _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} @@ -120,7 +120,7 @@ def test_undo_raises_when_nothing_to_undo(db): _add(db, adv, 0, "start") db.commit() 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 @@ -133,7 +133,7 @@ def test_undo_blocked_by_active_turn_lock(db): adventures.turns.acquire_turn_lock(adv.id) # a turn is "generating" try: 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 # The failed undo must not have released someone else's lock. 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.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} 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 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} # Nothing is written until a replacement actually arrives: the attempt on