Roll back script_state on undo/retry; fix retry double-apply

The shared per-adventure script_state ("scoreboard" scripts write to) was
never reverted by undo, and retry re-ran the output hook on top of the already-
mutated state, double-applying its changes (e.g. "+10 gold" became +20).

Each action now snapshots script_state as it was immediately before its own
hooks ran (new Action.state_before column, migration 25):
- undo restores the turn's first-action snapshot, prunes memories that
  summarized the removed actions, and takes the turn lock against races.
- retry restores the AI action's snapshot before regenerating.

Story-card mutations are not reverted (documented limit). Adds the project's
first test suite: unit + full HTTP integration through the real scripting
engine (14 tests). See plan/11-state-revert-and-retry-fix.md.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
parththakkar106
2026-07-21 02:58:38 +05:30
co-authored by Claude Opus 4.8
parent 92044a53c4
commit dab1807118
7 changed files with 491 additions and 12 deletions
+15
View File
@@ -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(
+4
View File
@@ -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)
+4
View File
@@ -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")
+50 -12
View File
@@ -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 ----------
+201
View File
@@ -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)
+136
View File
@@ -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) == {}