diff --git a/backend/app/memorybank.py b/backend/app/memorybank.py index c1f9979..4b9928c 100644 --- a/backend/app/memorybank.py +++ b/backend/app/memorybank.py @@ -84,6 +84,21 @@ def story_actions(adventure: models.Adventure) -> list[models.Action]: return [a for a in adventure.actions if a.text.strip()] +def prune_dangling_memories(adventure: models.Adventure, db: Session) -> int: + """Delete memories that summarized actions which no longer exist (e.g. after + undo). source_start/source_end are Action.index values; a memory is dangling + if any covered action is past the current end of the story. Returns the count + removed. Cursors are self-healing in run_post_turn, so this is cleanup only.""" + max_index = max((a.index for a in adventure.actions), default=-1) + dangling = [ + m for m in adventure.memories + if m.source_end is not None and m.source_end > max_index + ] + for m in dangling: + db.delete(m) + return len(dangling) + + # ---------- Retrieval (runs inside the turn, before build_context) ---------- async def retrieve_memories( diff --git a/backend/app/migrations.py b/backend/app/migrations.py index 332e2b2..5cbc461 100644 --- a/backend/app/migrations.py +++ b/backend/app/migrations.py @@ -75,6 +75,10 @@ MIGRATIONS: list[tuple[int, str]] = [ # re-synced on demand. NULL for copies made before this column existed. (24, "ALTER TABLE adventure_scripts ADD COLUMN source_script_id INTEGER " "REFERENCES scripts(id) ON DELETE SET NULL"), + # Per-action snapshot of the shared script_state as it was before that + # action's hooks ran, enabling undo/retry to roll state back. JSON is valid + # on both SQLite and Postgres. + (25, "ALTER TABLE actions ADD COLUMN state_before JSON"), ] LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1) diff --git a/backend/app/models.py b/backend/app/models.py index df6c624..ae0fbe1 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -181,6 +181,10 @@ class Action(Base): # Reasoning-model "thinking" that preceded the text (AI actions only). reasoning: Mapped[str | None] = mapped_column(Text, nullable=True) context_snapshot: Mapped[dict | None] = mapped_column(JSON, nullable=True) + # Copy of Adventure.script_state as it was immediately BEFORE this action's + # script hooks ran, so undo/retry can roll the shared scoreboard back. + # NULL for actions created before this column existed. + state_before: Mapped[dict | None] = mapped_column(JSON, nullable=True) created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow) adventure: Mapped[Adventure] = relationship(back_populates="actions") diff --git a/backend/app/routers/adventures.py b/backend/app/routers/adventures.py index 93ca82b..3d815bd 100644 --- a/backend/app/routers/adventures.py +++ b/backend/app/routers/adventures.py @@ -1,3 +1,4 @@ +import copy import json import re import threading @@ -153,6 +154,13 @@ def get_script_state( return {"state": state} +def snapshot_state(adventure: models.Adventure) -> dict: + """Deep copy of the shared script_state, to staple onto an action so undo/ + retry can restore it. Independent of later hook mutations.""" + state = adventure.script_state if isinstance(adventure.script_state, dict) else {} + return copy.deepcopy(state) + + @router.patch("/{adventure_id}", response_model=schemas.AdventureOut) def update_adventure( adventure_id: int, @@ -258,6 +266,9 @@ async def generate_turn( ) else: memories = await memorybank.retrieve_memories(adventure, settings, update_stats=True) + # Scoreboard as it stands before this AI turn's context/output hooks mutate + # it — stapled onto the AI action so retry can start over from here. + state_before = snapshot_state(adventure) system_text, story_text, snapshot = build_context(adventure, settings, memories) # onModelContext: scripts see (and may rewrite) the whole assembled context. @@ -327,6 +338,7 @@ async def generate_turn( text=text, reasoning="".join(reasoning_chunks).strip() or None, context_snapshot=snapshot, + state_before=state_before, ) db.add(ai_action) adventure.updated_at = models.utcnow() @@ -362,6 +374,9 @@ async def run_player_turn( # An empty do/say/story is just a continue. if payload.type != "continue" and payload.text.strip(): + # Scoreboard before the input hook mutates it — the pre-turn state that + # undo restores to (the AI action keeps its own post-input snapshot). + state_before = snapshot_state(adventure) # onInput sees the formatted text (as in AI Dungeon: "> You ..."). formatted = format_player_input(payload.type, payload.text) modified, stop = pipeline.run("input", formatted) @@ -374,6 +389,7 @@ async def run_player_turn( index=next_index(adventure), type=payload.type, text=modified, + state_before=state_before, ) db.add(player_action) db.commit() @@ -426,7 +442,13 @@ def retry_action( acquire_turn_lock(adventure_id) try: if adventure.actions and adventure.actions[-1].type == "ai": - db.delete(adventure.actions[-1]) + last_ai = adventure.actions[-1] + # Roll the scoreboard back to before this AI turn's hooks ran, so + # regenerating starts fresh instead of stacking output mutations on + # top of the discarded attempt. NULL for pre-migration actions. + if last_ai.state_before is not None: + adventure.script_state = copy.deepcopy(last_ai.state_before) + db.delete(last_ai) db.commit() db.refresh(adventure) except BaseException: @@ -446,18 +468,34 @@ def retry_action( def undo_turn( adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser ): - """Delete the last turn: the trailing AI action plus its player action, if any.""" + """Delete the last turn: the trailing AI action plus its player action, if any. + + Also rolls the shared script_state back to before that turn ran and prunes + any memory that summarized the removed actions. The turn lock guards against + undoing while a turn is still generating.""" adventure = get_adventure_or_404(adventure_id, db, user) - actions = list(adventure.actions) - if not actions or actions[-1].type == "start": - raise HTTPException(400, "Nothing to undo") - last = actions.pop() - db.delete(last) - if last.type == "ai" and actions and actions[-1].type in ("do", "say", "story"): - db.delete(actions.pop()) - db.commit() - db.refresh(adventure) - return adventure.actions + acquire_turn_lock(adventure_id) + try: + actions = list(adventure.actions) + if not actions or actions[-1].type == "start": + raise HTTPException(400, "Nothing to undo") + last = actions.pop() + # The earliest action removed in this turn holds the pre-turn scoreboard. + first_removed = last + db.delete(last) + if last.type == "ai" and actions and actions[-1].type in ("do", "say", "story"): + first_removed = actions.pop() + db.delete(first_removed) + if first_removed.state_before is not None: + adventure.script_state = copy.deepcopy(first_removed.state_before) + db.flush() # apply deletes so pruning sees the shrunken action list + db.expire(adventure, ["actions"]) + memorybank.prune_dangling_memories(adventure, db) + db.commit() + db.refresh(adventure) + return adventure.actions + finally: + _active_turns.discard(adventure_id) # ---------- Import / Export ---------- diff --git a/backend/tests/test_state_revert.py b/backend/tests/test_state_revert.py new file mode 100644 index 0000000..4e75552 --- /dev/null +++ b/backend/tests/test_state_revert.py @@ -0,0 +1,201 @@ +"""Tests for undo/retry rolling back the shared script_state scoreboard +(plan/11-state-revert-and-retry-fix.md). + +Run from the backend dir: python -m pytest tests/test_state_revert.py -v +""" +import os +import tempfile + +# Point the app at a throwaway SQLite file BEFORE importing anything that binds +# the engine at import time (app.database reads AIDND_DB_PATH on import). +_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) +_tmp.close() +os.environ["AIDND_DB_PATH"] = _tmp.name +os.environ.pop("AIDND_DATABASE_URL", None) +os.environ.pop("DATABASE_URL", None) + +import pytest +from fastapi import HTTPException + +from app import memorybank, models +from app.database import Base, SessionLocal, engine +from app.routers import adventures + + +@pytest.fixture() +def db(): + Base.metadata.create_all(bind=engine) + session = SessionLocal() + try: + yield session + finally: + session.close() + Base.metadata.drop_all(bind=engine) + adventures._active_turns.clear() + + +def _make_adventure(db, script_state): + user = models.User(is_guest=False) + db.add(user) + db.flush() + adv = models.Adventure(user_id=user.id, title="T", script_state=script_state) + db.add(adv) + db.flush() + return user, adv + + +def _add(db, adv, index, type_, text="x", state_before=None): + a = models.Action( + adventure_id=adv.id, index=index, type=type_, text=text, + state_before=state_before, + ) + db.add(a) + db.flush() + return a + + +# ---------------------------------------------------------------- undo + +def test_undo_reverts_state_to_before_the_turn(db): + # A turn took the scoreboard from {gold:0} -> {gold:10}. The player action + # carries the pre-turn snapshot; current state is the mutated one. + user, adv = _make_adventure(db, {"gold": 10}) + _add(db, adv, 0, "start", state_before=None) + _add(db, adv, 1, "do", state_before={"gold": 0}) + _add(db, adv, 2, "ai", state_before={"gold": 0}) + db.commit() + + adventures.undo_turn(adv.id, db=db, user=user) + + assert adv.script_state == {"gold": 0} + assert [a.type for a in adv.actions] == ["start"] + + +def test_undo_of_bare_continue_uses_ai_snapshot(db): + # A "continue" turn has no player action; the AI action's own snapshot is + # the pre-turn state. + user, adv = _make_adventure(db, {"gold": 5}) + _add(db, adv, 0, "start") + _add(db, adv, 1, "ai", state_before={"gold": 0}) + db.commit() + + adventures.undo_turn(adv.id, db=db, user=user) + + assert adv.script_state == {"gold": 0} + assert [a.type for a in adv.actions] == ["start"] + + +def test_undo_leaves_state_untouched_when_snapshot_missing(db): + # Pre-migration actions have state_before = NULL: don't clobber the state. + user, adv = _make_adventure(db, {"gold": 10}) + _add(db, adv, 0, "start") + _add(db, adv, 1, "do", state_before=None) + _add(db, adv, 2, "ai", state_before=None) + db.commit() + + adventures.undo_turn(adv.id, db=db, user=user) + + assert adv.script_state == {"gold": 10} + + +def test_undo_raises_when_nothing_to_undo(db): + user, adv = _make_adventure(db, {}) + _add(db, adv, 0, "start") + db.commit() + with pytest.raises(HTTPException) as exc: + adventures.undo_turn(adv.id, db=db, user=user) + assert exc.value.status_code == 400 + + +def test_undo_blocked_by_active_turn_lock(db): + user, adv = _make_adventure(db, {}) + _add(db, adv, 0, "start") + _add(db, adv, 1, "ai", state_before={}) + db.commit() + + adventures.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) + assert exc.value.status_code == 409 + # The failed undo must not have released someone else's lock. + assert adv.id in adventures._active_turns + finally: + adventures._active_turns.discard(adv.id) + + +def test_undo_prunes_memory_covering_removed_actions(db): + user, adv = _make_adventure(db, {}) + for i in range(4): + _add(db, adv, i, "ai" if i % 2 else "do", state_before={}) + # A memory summarizing actions up to index 3, which undo will delete. + covering = models.Memory(adventure_id=adv.id, text="m", source_start=0, source_end=3) + keep = models.Memory(adventure_id=adv.id, text="k", source_start=0, source_end=1) + db.add_all([covering, keep]) + db.commit() + + adventures.undo_turn(adv.id, db=db, user=user) # removes indexes 2 & 3 + + texts = {m.text for m in adv.memories} + assert texts == {"k"} + + +# ---------------------------------------------------------------- prune helper + +def test_prune_dangling_memories_counts_and_removes(db): + user, adv = _make_adventure(db, {}) + _add(db, adv, 0, "do") + _add(db, adv, 1, "ai") + db.add_all([ + models.Memory(adventure_id=adv.id, text="live", source_start=0, source_end=1), + models.Memory(adventure_id=adv.id, text="dead", source_start=2, source_end=5), + ]) + db.commit() + + removed = memorybank.prune_dangling_memories(adv, db) + db.commit() + db.refresh(adv) # expire_on_commit=False: reload the memories collection + + assert removed == 1 + assert {m.text for m in adv.memories} == {"live"} + + +# ---------------------------------------------------------------- snapshot + +def test_snapshot_state_is_an_independent_deep_copy(db): + _, adv = _make_adventure(db, {"nested": {"n": 1}}) + snap = adventures.snapshot_state(adv) + adv.script_state["nested"]["n"] = 99 + assert snap == {"nested": {"n": 1}} # unaffected by later mutation + + +def test_snapshot_state_handles_non_dict(db): + _, adv = _make_adventure(db, {}) + adv.script_state = None + assert adventures.snapshot_state(adv) == {} + + +# ---------------------------------------------------------------- retry + +def test_retry_restores_state_before_regenerating(db, monkeypatch): + # Retry deletes the last AI action and must roll the scoreboard back to that + # action's snapshot so regeneration doesn't stack output mutations. + user, adv = _make_adventure(db, {"gold": 20}) # 20 = double-applied bug value + _add(db, adv, 0, "start") + _add(db, adv, 1, "do", state_before={"gold": 0}) + _add(db, adv, 2, "ai", state_before={"gold": 10}) + db.commit() + + monkeypatch.setattr(adventures.limits, "rate_limit", lambda *a, **k: None) + monkeypatch.setattr(adventures, "check_demo_cap", lambda *a, **k: None) + + async def _noop(*a, **k): + if False: + yield # make it an async generator + monkeypatch.setattr(adventures, "generate_turn", _noop) + + adventures.retry_action(adv.id, request=None, db=db, user=user) + + assert adv.script_state == {"gold": 10} + assert [a.type for a in adv.actions] == ["start", "do"] + adventures._active_turns.discard(adv.id) diff --git a/backend/tests/test_turn_flow_integration.py b/backend/tests/test_turn_flow_integration.py new file mode 100644 index 0000000..c199db2 --- /dev/null +++ b/backend/tests/test_turn_flow_integration.py @@ -0,0 +1,136 @@ +"""End-to-end HTTP tests for undo/retry state revert, driving real turns through +the actual routes + scripting engine with only the LLM provider mocked. + +A script's output hook adds 10 gold each turn; we assert the shared scoreboard +behaves correctly across play / undo / retry. + + python -m pytest tests/test_turn_flow_integration.py -v +""" +import os +import tempfile + +_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) +_tmp.close() +os.environ["AIDND_DB_PATH"] = _tmp.name +os.environ.pop("AIDND_DATABASE_URL", None) +os.environ.pop("DATABASE_URL", None) + +import pytest +from fastapi import Depends +from fastapi.testclient import TestClient + +from app import auth, limits, models +from app.database import Base, SessionLocal, engine, get_db +from app.main import app +from app.providers import PromptParts +from app.routers import adventures + +GOLD_SCRIPT = """ +const modifier = (text) => { + state.gold = (state.gold || 0) + 10; + return { text }; +}; +modifier(text); +""" + + +class FakeProvider: + """Stand-in for OpenAICompatibleProvider: streams one fixed line, no network.""" + def __init__(self, *a, **k): + pass + + async def generate(self, parts: PromptParts, *, temperature, max_tokens): + yield ("text", "The torch flickers as you press onward.") + + +@pytest.fixture() +def client(monkeypatch): + Base.metadata.create_all(bind=engine) + setup = SessionLocal() + user = models.User(is_guest=False, email="tester@example.com") + setup.add(user) + setup.flush() + setup.add(models.Settings(user_id=user.id, api_key="enc:dummy", model="test-model")) + adv = models.Adventure(user_id=user.id, title="Cave", script_state={}) + setup.add(adv) + setup.flush() + setup.add(models.Action(adventure_id=adv.id, index=0, type="start", text="You enter a cave.")) + setup.add(models.AdventureScript( + adventure_id=adv.id, position=0, enabled=True, name="Gold", + output_js=GOLD_SCRIPT, + )) + setup.commit() + adv_id, user_id = adv.id, user.id + setup.close() + + # Force a real (non-demo) turn that uses our fake provider. + monkeypatch.setattr(adventures, "OpenAICompatibleProvider", FakeProvider) + monkeypatch.setattr(auth, "resolve_provider_config", lambda s: auth.ProviderConfig( + "http://fake", "k", "test-model", False)) + monkeypatch.setattr(limits, "rate_limit", lambda *a, **k: None) + monkeypatch.setattr(limits, "check_row_cap", lambda *a, **k: None) + + def _current_user(db=Depends(get_db)): + return db.get(models.User, user_id) + + app.dependency_overrides[auth.get_current_user] = _current_user + + c = TestClient(app) + c.adv_id = adv_id + try: + yield c + finally: + app.dependency_overrides.clear() + adventures._active_turns.clear() + Base.metadata.drop_all(bind=engine) + + +def _state(adv_id): + db = SessionLocal() + try: + return db.get(models.Adventure, adv_id).script_state + finally: + db.close() + + +def _play(client, type_="do", text="look around"): + r = client.post(f"/api/adventures/{client.adv_id}/actions", json={"type": type_, "text": text}) + assert r.status_code == 200, r.text + return r + + +def test_play_then_undo_reverts_gold(client): + assert _state(client.adv_id) == {} + _play(client) + assert _state(client.adv_id) == {"gold": 10} + + r = client.post(f"/api/adventures/{client.adv_id}/undo") + assert r.status_code == 200, r.text + assert _state(client.adv_id) == {} # scoreboard rolled back + + +def test_two_turns_then_undo_reverts_only_last(client): + _play(client) + _play(client) + assert _state(client.adv_id) == {"gold": 20} + + client.post(f"/api/adventures/{client.adv_id}/undo") + assert _state(client.adv_id) == {"gold": 10} # back to after turn 1, not 0 + + +def test_retry_does_not_double_apply_gold(client): + _play(client) + assert _state(client.adv_id) == {"gold": 10} + + # Before the fix this produced 20 (output hook ran twice); now it stays 10. + r = client.post(f"/api/adventures/{client.adv_id}/retry") + assert r.status_code == 200, r.text + assert _state(client.adv_id) == {"gold": 10} + + +def test_retry_then_undo_still_clean(client): + _play(client) + client.post(f"/api/adventures/{client.adv_id}/retry") + assert _state(client.adv_id) == {"gold": 10} + client.post(f"/api/adventures/{client.adv_id}/undo") + assert _state(client.adv_id) == {} diff --git a/plan/11-state-revert-and-retry-fix.md b/plan/11-state-revert-and-retry-fix.md new file mode 100644 index 0000000..18eda1c --- /dev/null +++ b/plan/11-state-revert-and-retry-fix.md @@ -0,0 +1,81 @@ +# Plan: undo/retry state revert (+ concurrency lock) + +Fixes three linked issues around `script_state` (the shared per-adventure +"scoreboard" scripts write to) and undo/retry. + +## Background + +- `script_state` is one shared dict on `Adventure` (`models.py:90`), mutated in + exactly ONE place: `pipeline.py:89` (`self.adventure.script_state = state`). +- Today, undo (`adventures.py:445`) and retry (`:422`) delete *actions* but never + touch `script_state`, so state never rolls back. + +## Issue 1 — retry double-applies state (pre-existing bug) + +Retry deletes the last AI action and regenerates. The output hook already mutated +`script_state` on the first attempt; regenerating runs it again, stacking the change +(e.g. "add 10 gold" → 20 gold after one retry). Same root cause as undo not +reverting. + +## Issue 4 — undo has no concurrency guard + +Turns take `acquire_turn_lock` (`:189`); undo does not, so undo can race a turn +that is still streaming. + +## Issue 2 — Memory Bank leftovers (smaller than expected) + +`run_post_turn` already clamps `memory_cursor`/`summary_cursor` down to the current +action count (`memorybank.py:178-182`), so there is NO cursor stall. The only +remainder: a `Memory` created from a turn that was later undone stays behind, its +`source_start/source_end` now pointing past the end of the story. + +--- + +## The fix + +### 1. Snapshot state per turn (Issue 1 + enables undo revert) + +- Add column `state_before: JSON nullable` to `Action` (`models.py`). +- Migration: append `(25, "ALTER TABLE actions ADD COLUMN state_before JSON")` + to `migrations.py`. `JSON` is valid on both SQLite and Postgres. +- In `run_player_turn` / `generate_turn`, when the FIRST action of a turn is + created, stash `copy.deepcopy(adventure.script_state)` onto it — captured + *before* any hook runs. (Player action for do/say/story; the AI action for a + bare `continue`.) +- Fix retry directly: before regenerating, restore + `adventure.script_state` from the deleted AI action's `state_before` so the + output hook starts from the pre-turn scoreboard instead of the mutated one. + +### 2. Revert on undo (Issue depends on #1) + +- In `undo_turn`, after deleting the popped actions, set + `adventure.script_state = .state_before` (fall back + to `{}` if null, i.e. pre-migration turns), then commit. +- Only wire this into undo + retry — NOT the arbitrary + `delete_action` endpoint (`:830`); mid-history state revert is undefined. + +### 3. Lock undo (Issue 4) + +- Wrap `undo_turn` body in `acquire_turn_lock(adventure_id)` / + `_active_turns.discard(...)` in a `finally`. It's synchronous (not SSE), so no + `with_turn_lock` wrapper needed — just acquire and discard. + +### 4. Clean up dangling memories (Issue 2, optional) + +- In `undo_turn` after deleting actions, delete any `Memory` whose `source_start` + is >= the new story-action count (i.e. summarized a turn that no longer exists). + Cursors already self-heal, so this is polish, not correctness. + +## Known limits (document, don't fix) + +- Story cards a script created (`_apply_cards`, `pipeline.py:87`) are NOT reverted — + only text + script_state roll back. +- Pre-migration turns have `state_before = NULL` → undo falls back to `{}`. +- Demo turn cap is not refunded on undo (intentional). + +## Test checklist + +- Script that increments a counter: play → undo → counter back to prior value. +- Same script: play → retry → counter changes once, not twice. +- Undo during an active stream returns 409, doesn't corrupt state. +- Undo a turn old enough to have been summarized: no orphaned memory left.