Make the branch happen when you write, not when you look

Stepping between takes used to be two controls and two meanings. At the tip a
chip switched; above the tip it only *previewed*, and taking that line needed a
second button next to it. Which one you got depended on where you were standing,
which is the thing that made the tree unusable when it was driven by hand.

So the server stops caring that anyone is looking. Reading a take the story
moved past changes nothing and creates nothing. `ActionCreate.after_id` names
the node a turn is played after, and naming a take the story left is the first
moment the player has said which line they mean -- so that is where the fork
happens, and only there.

`stand_on` is the old fork endpoint's body, lifted out whole. It already knew
the two cases and got them right: at the tip the takes are leaves nobody built
on, so it is a switch and no branch is made; past the tip the line being left
keeps every turn it has, so the take needs a branch. Both callers now share it,
which is the point -- a fork asked for and a fork arrived at are the same move.

Six tests. Five fail with the grouping reverted to the coordinate, and the one
that does not is deliberate: naming the tip in `after_id` must stay an ordinary
turn that forks nothing, which guards against over-correcting rather than
against the original bug. The nesting test passes both ways too and is kept for
what it says, not for what it catches -- a coordinate separates C1's takes from
C2's by accident, because the fork has already put them on different branches.

415 tests.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_017Dvvqn9ZDR4ixeFPHNbww7
This commit is contained in:
parththakkar106
2026-08-18 21:34:12 +05:30
committed by Parth
co-authored by Claude Opus 5
parent 6fa213f6db
commit e126bda387
3 changed files with 331 additions and 17 deletions
+66 -17
View File
@@ -889,6 +889,11 @@ def create_action(
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)
acquire_turn_lock(adventure_id) acquire_turn_lock(adventure_id)
try:
_move_to_after(db, adventure, payload.after_id)
except BaseException:
_active_turns.discard(adventure_id)
raise
return StreamingResponse( return StreamingResponse(
with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)), with_turn_lock(adventure_id, run_player_turn(adventure, db, payload, user)),
media_type="text/event-stream", media_type="text/event-stream",
@@ -896,6 +901,33 @@ def create_action(
) )
def _move_to_after(
db: Session, adventure: models.Adventure, after_id: int | None
) -> None:
"""Put the story where `after_id` says before the turn is played.
This is the moment a branch is born (SP9). Reading a take the story moved
past changes nothing on the server — the player is looking, and looking is
free. Writing below one is the first time they have said which line they
mean, and that is when the fork happens.
A take already on the path needs nothing: it is where the story is.
"""
if after_id is None:
return
node = db.get(models.Action, after_id)
if node is None or node.adventure_id != adventure.id:
raise HTTPException(404, "Action not found")
if node.live and lineage.path_of(db, adventure).contains(node):
return
if not node.live and len(attempts.group(db, node)) < 2:
# Not reachable through the pager, so nothing put the player here.
raise HTTPException(400, "That take is not one of this turn's.")
stand_on(db, adventure, node)
db.commit()
db.refresh(adventure)
@router.post("/{adventure_id}/retry") @router.post("/{adventure_id}/retry")
def retry_action( def retry_action(
adventure_id: int, adventure_id: int,
@@ -1381,23 +1413,7 @@ def fork_from_attempt(
) )
acquire_turn_lock(adventure_id) acquire_turn_lock(adventure_id)
try: try:
newest = last_action(adventure, db) stand_on(db, adventure, action)
at_the_tip = (
newest is not None
and newest.branch_id == action.branch_id
and newest.depth == action.depth
)
if at_the_tip:
# The story at this coordinate is about to say something else, so
# what was derived from it is withdrawn — the same move retry
# makes. A fork needs none of that: it leaves the coordinate, and
# its memory, exactly where they are (see `tree.fork`).
memorybank.forget_node(db, adventure, action)
cursors.rewind_all(adventure, action.branch_id, (action.depth or 0) - 1)
attempts.make_live(db, adventure, action)
else:
tree.fork(db, adventure, action)
attempts.restore_state(adventure, action)
adventure.updated_at = models.utcnow() adventure.updated_at = models.utcnow()
db.commit() db.commit()
db.refresh(adventure) db.refresh(adventure)
@@ -1406,6 +1422,39 @@ def fork_from_attempt(
_active_turns.discard(adventure_id) _active_turns.discard(adventure_id)
def stand_on(
db: Session, adventure: models.Adventure, action: models.Action
) -> None:
"""Make `action` the take the story tells, forking only if it has to.
Two cases, and the caller does not have to know which. While the turn is
still the tip its takes are leaves nobody has built on, so this is a switch
and no branch is created. Once the story has moved past, the line being left
keeps every turn it has, so the take needs a branch of its own.
Called from the fork endpoint and from a turn played below a take the story
moved past — the same move, once as a request and once as the thing that
happens on the way to writing (SP9).
"""
newest = last_action(adventure, db)
at_the_tip = (
newest is not None
and newest.branch_id == action.branch_id
and newest.depth == action.depth
)
if at_the_tip:
# The story at this coordinate is about to say something else, so what
# was derived from it is withdrawn — the same move retry makes. A fork
# needs none of that: it leaves the coordinate, and its memory, exactly
# where they are (see `tree.fork`).
memorybank.forget_node(db, adventure, action)
cursors.rewind_all(adventure, action.branch_id, (action.depth or 0) - 1)
attempts.make_live(db, adventure, action)
else:
tree.fork(db, adventure, action)
attempts.restore_state(adventure, action)
@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), user: models.User = CurrentUser
+8
View File
@@ -244,6 +244,14 @@ class ActionUpdate(BaseModel):
class ActionCreate(BaseModel): class ActionCreate(BaseModel):
type: Literal["do", "say", "story", "continue"] type: Literal["do", "say", "story", "continue"]
text: ActionText = "" text: ActionText = ""
# The node this action is played after (SP9). Omitted means "the tip",
# which is every ordinary turn.
#
# Naming a take the story moved past is how a branch gets made: stepping
# between takes costs nothing and creates nothing, and the fork happens on
# the first thing written below one. That is the only moment the player has
# said which line they mean — before it, they were reading.
after_id: int | None = None
class AdventureOut(ORMModel): class AdventureOut(ORMModel):
+257
View File
@@ -0,0 +1,257 @@
"""Phase 14 SP9 — a turn's takes are grouped by their parent, not by where they sit.
SP4 made every take a node at the same (branch, depth). That coordinate answers
"which takes belong to this turn" right up until one of them is forked onto its
own branch — at which point it *leaves* the coordinate and reads as the only
take of its turn, with its siblings unreachable from the line it was taken on.
The parent does not move when a branch does, which is the whole of the fix. It
also gets the nesting right for free: takes under C1 and takes under C2 share a
depth and, until one forks, a branch. Only the parent separates them.
And the branch itself is lazy now. Stepping between takes creates nothing —
looking is free. The fork happens on the first thing *written* below a take the
story moved past, which is the first moment the player has said which line they
mean.
python -m pytest tests/test_take_parentage.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 attempts, 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
class ScriptedProvider:
replies: list = []
calls = 0
def __init__(self, *a, **k):
pass
async def generate(self, parts: PromptParts, *, temperature, max_tokens):
index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1)
ScriptedProvider.calls += 1
yield ("text", ScriptedProvider.replies[index])
@pytest.fixture()
def client(monkeypatch):
Base.metadata.create_all(bind=engine)
setup = SessionLocal()
user = models.User(is_guest=False, email="takes@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={}, world_state={}
)
setup.add(adv)
setup.flush()
setup.add(
models.Action(adventure_id=adv.id, index=0, type="start", text="You enter.")
)
setup.commit()
adv_id, user_id = adv.id, user.id
setup.close()
# Distinct replies so a take can be told apart from its siblings by text.
ScriptedProvider.replies = [f"Take {n}." for n in range(1, 40)]
ScriptedProvider.calls = 0
monkeypatch.setattr(adventures, "OpenAICompatibleProvider", ScriptedProvider)
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)
# ------------------------------------------------------------------ helpers
def _play(client, text="look around", after_id=None):
body = {"type": "do", "text": text}
if after_id is not None:
body["after_id"] = after_id
r = client.post(f"/api/adventures/{client.adv_id}/actions", json=body)
assert r.status_code == 200, r.text
return r
def _retry(client):
r = client.post(f"/api/adventures/{client.adv_id}/retry")
assert r.status_code == 200, r.text
def _branch_count(adv_id) -> int:
db = SessionLocal()
try:
return (
db.query(models.Branch).filter(models.Branch.adventure_id == adv_id).count()
)
finally:
db.close()
def _ai_rows(adv_id) -> list[models.Action]:
"""Every AI node ever written, oldest first — live or not, any branch."""
db = SessionLocal()
try:
return (
db.query(models.Action)
.filter(
models.Action.adventure_id == adv_id,
models.Action.type == "ai",
)
.order_by(models.Action.id)
.all()
)
finally:
db.close()
def _group_size(action_id: int) -> int:
db = SessionLocal()
try:
node = db.get(models.Action, action_id)
return len(attempts.group(db, node))
finally:
db.close()
# ------------------------------------------------------------------ tests
def test_retaken_turn_groups_all_its_takes(client):
"""The baseline the rest of the file leans on: three takes, one turn."""
_play(client)
_retry(client)
_retry(client)
takes = _ai_rows(client.adv_id)
assert len(takes) == 3
for take in takes:
assert _group_size(take.id) == 3, "every take sees the whole turn"
def test_a_forked_take_keeps_its_siblings(client):
"""The bug this subphase exists for.
Under coordinate grouping the forked take reads 1/1: it has left the
(branch, depth) the others are still at. The parent does not move with it.
"""
_play(client)
_retry(client)
_retry(client)
# The story moves past the turn, so taking a different take is now a fork.
_play(client, "press on")
first_take = _ai_rows(client.adv_id)[0]
assert first_take.live is False
# Writing below it is what forks -- see the next test.
_play(client, "go back and try this instead", after_id=first_take.id)
assert _group_size(first_take.id) == 3, (
"a take forked onto its own branch is still a take of the same turn"
)
def test_stepping_between_takes_creates_no_branch(client):
"""Looking is free. Only writing commits to a line."""
_play(client)
_retry(client)
_play(client, "press on")
before = _branch_count(client.adv_id)
first_take = _ai_rows(client.adv_id)[0]
# Reading a take: the pager fetches the turn's takes and shows one.
r = client.get(
f"/api/adventures/{client.adv_id}/actions/{first_take.id}/variants"
)
assert r.status_code == 200, r.text
assert len(r.json()) == 2
assert _branch_count(client.adv_id) == before, "reading forked nothing"
def test_writing_below_a_passed_take_forks_exactly_once(client):
_play(client)
_retry(client)
_play(client, "press on")
before = _branch_count(client.adv_id)
first_take = _ai_rows(client.adv_id)[0]
_play(client, "a different way", after_id=first_take.id)
assert _branch_count(client.adv_id) == before + 1, "one write, one branch"
def test_takes_under_one_parent_do_not_count_takes_under_its_sibling(client):
"""The player's own example: 3/3 on one line, 2/2 on the other.
C1 and C2 are takes of the same turn. What is played *below* each of them
is a different turn, and the two must not pool -- they share a depth, and
until the fork they share a branch too.
"""
_play(client)
_retry(client) # two takes at this turn: C1, C2 (C2 live)
c1, c2 = _ai_rows(client.adv_id)[:2]
# Below C2, the live one: a turn with three takes.
_play(client, "down the second path")
_retry(client)
_retry(client)
under_c2 = [a for a in _ai_rows(client.adv_id) if a.parent_id is not None]
under_c2 = [a for a in under_c2 if a.id not in (c1.id, c2.id)]
assert len(under_c2) == 3
# Now take C1 instead, and play a turn with two takes below it.
_play(client, "down the first path", after_id=c1.id)
_retry(client)
everything = _ai_rows(client.adv_id)
under_c1 = [a for a in everything if a.parent_id not in (None,)
and a.id not in (c1.id, c2.id) and a.id not in {x.id for x in under_c2}]
assert len(under_c1) == 2
assert _group_size(under_c2[0].id) == 3, "C2's line keeps its three takes"
assert _group_size(under_c1[0].id) == 2, "C1's line counts only its own two"
def test_naming_a_take_that_is_already_the_story_just_plays_on(client):
"""`after_id` pointing at the tip is an ordinary turn, and forks nothing."""
_play(client)
before = _branch_count(client.adv_id)
live = [a for a in _ai_rows(client.adv_id) if a.live][-1]
_play(client, "carry on", after_id=live.id)
assert _branch_count(client.adv_id) == before