diff --git a/backend/app/routers/adventures.py b/backend/app/routers/adventures.py index a58a797..f2d9405 100644 --- a/backend/app/routers/adventures.py +++ b/backend/app/routers/adventures.py @@ -595,7 +595,17 @@ def sse(obj: dict) -> str: SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"} -def action_json(action: models.Action) -> dict: +def action_json(action: models.Action, db: Session | None = None) -> dict: + """One action, as the wire sees it. + + `db` is what gives it the pager numbers, and a turn that has just been + played must pass it: the take it created is often the *second* at its turn, + so the message arrives needing a pager the client cannot draw from a count + of one. Without this a retry showed no pager until the page was reloaded — + the same trap as the adventure GET, which builds its window a third way. + """ + if db is not None: + annotate_takes(db, action.adventure_id, [action]) return schemas.ActionOut.model_validate(action).model_dump(mode="json") @@ -824,7 +834,7 @@ async def _generate_turn( db.commit() db.refresh(ai_action) yield _SAVED - yield sse({"type": "done", "action": action_json(ai_action), "script": pipeline.report()}) + yield sse({"type": "done", "action": action_json(ai_action, db), "script": pipeline.report()}) # Phase 6: fire-and-forget summarization/embedding (opens its own DB # session). Skipped on the demo key — background AI calls would be # unmetered spend on the server-funded key. @@ -845,13 +855,25 @@ async def run_player_turn( db: Session, payload: schemas.ActionCreate, user: models.User, + preformatted: bool = False, ): + """Play a player's turn: their action, then the reply to it. + + `preformatted` says the text already carries the "> You ..." conventions and + must be written as it stands. That is the case for another take of a turn + the player already played (SP9): the editor is seeded with the stored text, + which is the formatted form — the same text plain edit puts in the box and + writes back verbatim. Formatting it again gives "> You > You ...". + """ pipeline = ScriptPipeline(adventure, db) # An empty do/say/story is just a continue. if payload.type != "continue" and payload.text.strip(): # onInput sees the formatted text (as in AI Dungeon: "> You ..."). - formatted = format_player_input(payload.type, payload.text) + formatted = ( + payload.text.strip() if preformatted + else format_player_input(payload.type, payload.text) + ) modified, stop = pipeline.run("input", formatted) if not modified.strip(): yield sse({"type": "error", "detail": "A script's input modifier returned empty text.", @@ -876,7 +898,7 @@ async def run_player_turn( # collection is stale — without this, build_context and next_index for # the AI action would not see the player action just saved. db.expire(adventure, ["actions"]) - yield sse({"type": "player", "action": action_json(player_action)}) + yield sse({"type": "player", "action": action_json(player_action, db)}) if stop: # onInput { stop: true } prevents the AI call. yield sse({"type": "stopped", "script": pipeline.report()}) @@ -1558,6 +1580,9 @@ def add_take( db, schemas.ActionCreate(type=action.type, text=payload.text), user, + # The client seeded its editor from the stored text, which already + # carries the "> You ..." conventions. + preformatted=True, ) return StreamingResponse( with_turn_lock(adventure_id, stream), diff --git a/backend/tests/test_take_parentage.py b/backend/tests/test_take_parentage.py index 2cc9685..764f022 100644 --- a/backend/tests/test_take_parentage.py +++ b/backend/tests/test_take_parentage.py @@ -16,6 +16,7 @@ mean. python -m pytest tests/test_take_parentage.py -v """ +import json import os import tempfile @@ -262,6 +263,34 @@ def test_the_page_carries_the_pager_numbers(client): assert ai[0]["take_index"] == 2, "the newest take is the one being read" +def _done_action(response) -> dict: + """The action carried by the `done` event of a turn's SSE stream.""" + for line in response.text.splitlines(): + if not line.startswith("data: "): + continue + event = json.loads(line[len("data: "):]) + if event.get("type") == "done": + return event["action"] + raise AssertionError("no done event in the stream") + + +def test_the_streamed_action_carries_the_pager_too(client): + """Found by driving it, not by testing it. + + A retry's reply *is* the second take of its turn, so it arrives needing a + pager. The stream builds its own ActionOut and so missed the annotation: + the pager appeared only once the page was reloaded, which is exactly the + moment nobody reloads. + """ + _play(client) + r = client.post(f"/api/adventures/{client.adv_id}/retry") + assert r.status_code == 200, r.text + + action = _done_action(r) + assert action["take_count"] == 2, "the take that just landed knows it has a sibling" + assert action["take_index"] == 1 + + def test_a_turn_nobody_retook_reads_one_of_one(client): _play(client) for action in _page(client): @@ -334,6 +363,27 @@ def test_a_players_own_turn_can_be_played_again(client): assert "press on" not in blob, "what followed the old text stays behind" +def test_a_retaken_player_turn_is_not_formatted_twice(client): + """Found by driving it: "> You > You open the door." + + The editor is seeded from the stored text, which is already the formatted + form — the same text plain edit puts in the box and writes back verbatim. + Running it through the formatter again doubles the prefix. + """ + _play(client, "open the door") + _play(client, "press on") + + first = _user_rows(client.adv_id)[0] + assert first.text.startswith("> You "), "stored formatted, which is the premise" + r = _take(client, first.id, "> You smash the door instead.") + assert r.status_code == 200, r.text + + new_take = [a for a in _user_rows(client.adv_id) + if "smash the door instead" in a.text][0] + assert new_take.text == "> You smash the door instead." + assert "> You > You" not in new_take.text + + def test_the_line_left_behind_keeps_its_whole_story(client): _play(client, "open the door") _play(client, "press on") diff --git a/backend/tests/test_take_state.py b/backend/tests/test_take_state.py new file mode 100644 index 0000000..eee415c --- /dev/null +++ b/backend/tests/test_take_state.py @@ -0,0 +1,215 @@ +"""Phase 14 SP9 — what a take does to the shared state. + +A turn does not only write text. A script mutates `script_state`, the referee +mutates `world_state`, and both are *shared* — they belong to the adventure, not +to the node. So playing a turn again has to put them back to where they were +before that turn ran, or the new take stacks its mutations on top of the one it +replaces and the numbers drift every time the player asks for another take. + +`retry` has done this since SP4 (`attempts.roll_back_before`). These are the +same guarantee for the two roads SP9 opened: a take of a turn the story moved +past, and a take of the player's own turn. Both create a branch, which is the +interesting part — the rollback has to survive leaving the line it was on. + + python -m pytest tests/test_take_state.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 + +SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} + +# Ten gold a turn, every turn. A number that only ever goes up is the clearest +# possible witness to a rollback: if a take stacks instead of replacing, it says +# so in one digit. +GOLD_SCRIPT = """ +const modifier = (text) => { + state.gold = (state.gold || 0) + 10; + return { text }; +}; +modifier(text); +""" + + +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="takestate@example.com") + setup.add(user) + setup.flush() + setup.add(models.Settings(user_id=user.id, api_key="enc:dummy", model="test-model")) + scenario = models.Scenario(user_id=user.id, title="S", stat_schema=SCHEMA) + setup.add(scenario) + setup.flush() + adv = models.Adventure( + user_id=user.id, title="Vault", scenario_id=scenario.id, + script_state={}, world_state={"player": {"hp": 100}}, + ) + setup.add(adv) + setup.flush() + setup.add(models.Action(adventure_id=adv.id, index=0, type="start", text="You begin.")) + 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() + + 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) + + +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 + + +def _take(client, action_id, text=""): + r = client.post( + f"/api/adventures/{client.adv_id}/actions/{action_id}/takes", + json={"text": text}, + ) + assert r.status_code == 200, r.text + return r + + +def _gold(adv_id) -> int: + db = SessionLocal() + try: + return (db.get(models.Adventure, adv_id).script_state or {}).get("gold", 0) + finally: + db.close() + + +def _rows(adv_id, type_): + db = SessionLocal() + try: + return ( + db.query(models.Action) + .filter(models.Action.adventure_id == adv_id, models.Action.type == type_) + .order_by(models.Action.id) + .all() + ) + finally: + db.close() + + +def test_the_script_runs_once_a_turn(client): + """The premise the rest of the file rests on.""" + _play(client) + assert _gold(client.adv_id) == 10 + _play(client, "press on") + assert _gold(client.adv_id) == 20 + + +def test_a_take_of_a_past_ai_turn_does_not_stack_its_script(client): + """Two turns played, then the first one taken again. + + The take leaves the path just before turn one, so the state it starts from + is the state turn one started from — nothing, not the twenty that two turns + had accumulated. Then its own run adds ten. + """ + _play(client) + _play(client, "press on") + assert _gold(client.adv_id) == 20 + + first_ai = _rows(client.adv_id, "ai")[0] + _take(client, first_ai.id) + + assert _gold(client.adv_id) == 10, "rolled back to before that turn, then run once" + + +def test_a_take_of_a_player_turn_does_not_stack_its_script(client): + _play(client) + _play(client, "press on") + assert _gold(client.adv_id) == 20 + + first_player = _rows(client.adv_id, "do")[0] + _take(client, first_player.id, "> You do something else.") + + assert _gold(client.adv_id) == 10 + + +def test_writing_below_a_passed_take_starts_from_that_take_s_state(client): + """The `after_id` road, which forks on the way to writing. + + The take being written under produced the state its own turn left behind — + ten — and the turn played on top of it adds the next ten. The twenty the + abandoned line reached has nothing to do with this branch. + """ + _play(client) + r = client.post(f"/api/adventures/{client.adv_id}/retry") + assert r.status_code == 200, r.text + _play(client, "press on") + assert _gold(client.adv_id) == 20 + + discarded = [a for a in _rows(client.adv_id, "ai") if not a.live][0] + _play(client, "a different way", after_id=discarded.id) + + assert _gold(client.adv_id) == 20, "that take's ten, plus this turn's ten" + + +def test_the_line_left_behind_keeps_the_state_it_reached(client): + """Switching back finds the abandoned line's numbers where it left them.""" + _play(client) + _play(client, "press on") + first_ai = _rows(client.adv_id, "ai")[0] + _take(client, first_ai.id) + assert _gold(client.adv_id) == 10 + + branches = client.get(f"/api/adventures/{client.adv_id}/branches").json() + root = [b for b in branches if b["parent_branch_id"] is None][0] + r = client.post(f"/api/adventures/{client.adv_id}/branches/{root['id']}/switch") + assert r.status_code == 200, r.text + + assert _gold(client.adv_id) == 20, "the first telling still has its two turns"