diff --git a/backend/app/context/__init__.py b/backend/app/context/__init__.py index 3cfd2f2..674f15f 100644 --- a/backend/app/context/__init__.py +++ b/backend/app/context/__init__.py @@ -1,3 +1,11 @@ -from .builder import build_context, count_tokens, story_actions, truncate_to_last_tokens +from . import history +from .builder import build_context, count_tokens, truncate_to_last_tokens +from .history import story_actions -__all__ = ["build_context", "count_tokens", "story_actions", "truncate_to_last_tokens"] +__all__ = [ + "build_context", + "count_tokens", + "history", + "story_actions", + "truncate_to_last_tokens", +] diff --git a/backend/app/context/builder.py b/backend/app/context/builder.py index 43d4d1d..6bd1bbc 100644 --- a/backend/app/context/builder.py +++ b/backend/app/context/builder.py @@ -17,9 +17,11 @@ from dataclasses import dataclass import tiktoken from .. import models, worldstate +from . import history AUTHORS_NOTE_DEPTH = 3 # actions from the end of history CARD_BUDGET_SHARE = 0.4 # max share of non-reserved budget that story cards may take +NPC_WINDOW = 6 # actions of story searched for NPC trigger words ("in scene") SEPARATOR = "\n\n" @@ -76,25 +78,12 @@ def _history_text(action: models.Action) -> str: return text -def story_actions( - adventure: models.Adventure, exclude_action_id: int | None = None -) -> list[models.Action]: - """The adventure's non-empty actions, as the model should see them. - - `exclude_action_id` drops one action from the story — used by retry, where - the row being regenerated is still attached to the adventure (it holds the - variant history) but must not appear in the context assembled to replace it. - """ - return [ - a for a in adventure.actions - if a.text.strip() and (exclude_action_id is None or a.id != exclude_action_id) - ] - - def _visible_npcs(actions: list[models.Action], stat_schema: dict) -> dict[str, str]: """Defined NPCs whose trigger words appear in the recent story — the ones - "in scene", so only their stats get injected. Maps npc id -> display name.""" - recent = SEPARATOR.join(a.text for a in actions[-6:]).lower() + "in scene", so only their stats get injected. Maps npc id -> display name. + + `actions` is already the last handful (see NPC_WINDOW).""" + recent = SEPARATOR.join(a.text for a in actions).lower() visible: dict[str, str] = {} for npc_key, ndef in (stat_schema.get("npcs") or {}).items(): if not isinstance(ndef, dict): @@ -127,9 +116,8 @@ def build_context( ) -> tuple[str, str, dict]: """Returns (system_text, story_text, context_report). `memory_bank` is the result of memorybank.retrieve_memories (None when the bank is off); - `exclude_action_id` omits one action from the story (see story_actions).""" + `exclude_action_id` omits one action from the story (see history.py).""" script_mem = _script_memory(adventure) - actions = story_actions(adventure, exclude_action_id) # ----- Always-included components ----- system_sections: list[Section] = [Section("narrator", settings.narrator_prompt.strip())] @@ -142,7 +130,10 @@ def build_context( if guide: system_sections.append(Section("world_state_guide", guide)) block = worldstate.render_state_section( - adventure.world_state, stat_schema, _visible_npcs(actions, stat_schema) + adventure.world_state, stat_schema, + _visible_npcs( + history.tail(adventure, NPC_WINDOW, exclude_action_id), stat_schema + ), ) if block: system_sections.append(Section("world_state", block)) @@ -181,6 +172,14 @@ def build_context( ) available = max(256, settings.context_token_budget - reserved) + # Only the newest actions can reach the prompt: everything below is either + # truncated to `available` tokens or stops at the budget. Fetch a window + # that is provably larger than that and no more — a long adventure would + # otherwise read its entire history every turn to use the tail of it. + actions = history.window_covering( + adventure, available, count_tokens, exclude_action_id + ) + # ----- Story cards: triggered by recent story text (the window history could fill) ----- trigger_window = truncate_to_last_tokens(SEPARATOR.join(a.text for a in actions), available) triggered = _match_cards(adventure.story_cards, trigger_window) @@ -268,7 +267,9 @@ def build_context( "memories": memory_bank, "history": { "included": len(included_actions), - "total": len(actions), + # The whole story, not just the window fetched above — Insights + # reports "N of M actions included" and M is the real total. + "total": history.count(adventure, exclude_action_id), "oldest_truncated": oldest_truncated, }, "settings": { diff --git a/backend/app/context/history.py b/backend/app/context/history.py new file mode 100644 index 0000000..18300e6 --- /dev/null +++ b/backend/app/context/history.py @@ -0,0 +1,303 @@ +"""Reading the story without reading all of it. + +`story_actions()` walked `adventure.actions`, which loads every row of the +adventure — then every caller threw almost all of it away. The context builder +concatenates the story and immediately cuts it back to the token budget; the +NPC-in-scene check looks at the last 6; memory retrieval looks at the last 4; +the post-turn cursor clamp only wants a count. So a turn on a 200-action +adventure read ~840 KB to use maybe 70 KB of it, and the cost grew with every +turn played. + +This module serves those shapes directly from SQL — a tail, a slice, a count — +so the read is bounded by the context budget instead of by the length of the +story. + +Two rules hold everything together: + +* **One definition of "story action".** The cursors in memorybank are + *positions* in this filtered, index-ordered list, so SQL and Python must + agree on membership exactly or a cursor silently points at a different + action. `_STORY_TEXT` and `is_story_text()` are that one definition, written + twice; keep them in step. +* **Never load twice.** If `adventure.actions` is already in memory (the + scripting pipeline hands the whole history to user scripts, as AI Dungeon + does), every helper here slices that instead of issuing a query, so a + scripted adventure pays what it always paid and nothing more. +""" + +from sqlalchemy import func, inspect as sa_inspect +from sqlalchemy.orm import Session, defer, object_session + +from .. import models + +# How many of the newest actions to read before checking whether the token +# budget is covered. When it isn't, the next size is worked out from the +# average action length just measured rather than by blind doubling — guessing +# high means reading hundreds of actions to use sixty of them. +WINDOW_START = 32 +WINDOW_MARGIN = 0.15 # aim this far past the budget, so one more round is rare +WINDOW_STEP = 8 # ...and at least this many more actions each round + + +def _sql_stripped(column): + """`column` with leading/trailing whitespace removed, portably. + + SQLite and Postgres both accept single-argument `trim()`, but it strips + spaces only — Python's `.strip()` also drops newlines and tabs, and an + action of nothing but a newline would otherwise count as story text here + and not in Python. `replace()` and `trim()` are the two string functions + both dialects spell identically, so fold the other whitespace into spaces + first. (Form feed and vertical tab are not covered; nothing produces them.) + """ + folded = column + for char in ("\n", "\r", "\t"): + folded = func.replace(folded, char, " ") + return func.trim(folded) + + +_STORY_TEXT = _sql_stripped(models.Action.text) != "" + + +def is_story_text(text: str) -> bool: + """The Python half of `_STORY_TEXT` — keep the two in step.""" + return bool(text.strip()) + + +def _loaded_actions(adventure: models.Adventure) -> list[models.Action] | None: + """The adventure's actions if they are already in memory, else None. + + Slicing an already-loaded collection is free; issuing a query beside it + would mean paying for the same rows twice. + """ + state = sa_inspect(adventure) + if state.detached or "actions" in state.unloaded: + return None + return list(adventure.actions) + + +def _from_memory( + adventure: models.Adventure, exclude_action_id: int | None +) -> list[models.Action] | None: + loaded = _loaded_actions(adventure) + if loaded is None: + return None + return [ + a for a in loaded + if is_story_text(a.text) and (exclude_action_id is None or a.id != exclude_action_id) + ] + + +def _filters(adventure: models.Adventure, exclude_action_id: int | None) -> list: + conditions = [models.Action.adventure_id == adventure.id, _STORY_TEXT] + if exclude_action_id is not None: + conditions.append(models.Action.id != exclude_action_id) + return conditions + + +def _query(db: Session, adventure: models.Adventure, exclude_action_id: int | None): + # Reasoning traces are never read from replayed history and can be larger + # than the narration itself on a reasoning model. + return ( + db.query(models.Action) + .filter(*_filters(adventure, exclude_action_id)) + .options(defer(models.Action.reasoning)) + ) + + +def _count_query(db: Session, adventure: models.Adventure, exclude_action_id: int | None): + """A real `SELECT count(...)`. + + Deliberately not `_query(...).count()`: that wraps the entity select in a + subquery, so the emitted SQL names every column — including the deferred + ones this whole design exists to keep off the wire. No bytes come back + either way, but the database still has to read them, and an egress guard + that greps the SQL cannot tell the two apart. + """ + return db.query(func.count(models.Action.id)).filter( + *_filters(adventure, exclude_action_id) + ) + + +def _session(adventure: models.Adventure) -> Session | None: + return object_session(adventure) + + +# ------------------------------------------------------------------ the API + +def story_actions( + adventure: models.Adventure, exclude_action_id: int | None = None +) -> list[models.Action]: + """Every story action, oldest first. + + Still the right call where the whole story is genuinely wanted — user + scripts receive it, per AI Dungeon's scripting API. Prefer `tail`, `slice_` + or `count` anywhere the caller only needs part of it. + + `exclude_action_id` drops one action from the story — used by retry, where + the row being regenerated is still attached to the adventure (it holds the + variant history) but must not appear in the context assembled to replace it. + """ + in_memory = _from_memory(adventure, exclude_action_id) + if in_memory is not None: + return in_memory + db = _session(adventure) + if db is None: + return [] + return _query(db, adventure, exclude_action_id).order_by(models.Action.index).all() + + +def count(adventure: models.Adventure, exclude_action_id: int | None = None) -> int: + """How many story actions there are, without fetching any of them.""" + in_memory = _from_memory(adventure, exclude_action_id) + if in_memory is not None: + return len(in_memory) + db = _session(adventure) + if db is None: + return 0 + return _count_query(db, adventure, exclude_action_id).scalar() or 0 + + +def tail_range( + adventure: models.Adventure, + skip: int, + limit: int, + exclude_action_id: int | None = None, +) -> list[models.Action]: + """`limit` story actions ending `skip` actions before the end, oldest first. + + `skip=0` is the newest slice; `skip=32, limit=16` is the 16 actions just + older than the newest 32. Lets a growing window fetch only the part it + doesn't already have. + """ + if limit <= 0 or skip < 0: + return [] + in_memory = _from_memory(adventure, exclude_action_id) + if in_memory is not None: + stop = len(in_memory) - skip + return in_memory[max(stop - limit, 0):stop] if stop > 0 else [] + db = _session(adventure) + if db is None: + return [] + rows = ( + _query(db, adventure, exclude_action_id) + .order_by(models.Action.index.desc()) + .offset(skip) + .limit(limit) + .all() + ) + rows.reverse() + return rows + + +def tail( + adventure: models.Adventure, limit: int, exclude_action_id: int | None = None +) -> list[models.Action]: + """The newest `limit` story actions, returned oldest first.""" + return tail_range(adventure, 0, limit, exclude_action_id) + + +def slice_( + adventure: models.Adventure, + start: int, + length: int, + exclude_action_id: int | None = None, +) -> list[models.Action]: + """Story actions at positions [start, start + length), oldest first. + + Positions are into the same filtered, index-ordered list the memory cursors + count in, which is why the filter has to match Python's exactly. + """ + if length <= 0 or start < 0: + return [] + in_memory = _from_memory(adventure, exclude_action_id) + if in_memory is not None: + return in_memory[start:start + length] + db = _session(adventure) + if db is None: + return [] + return ( + _query(db, adventure, exclude_action_id) + .order_by(models.Action.index) + .offset(start) + .limit(length) + .all() + ) + + +def position_of_index(adventure: models.Adventure, index: int) -> int: + """The position the story action with `Action.index == index` occupies — + i.e. how many story actions come before it. + + Translates between the two coordinate systems that keep tripping this code + up: cursors are positions, `Memory.source_start/_end` are `Action.index` + values, and the two diverge the moment anything is deleted. + """ + in_memory = _from_memory(adventure, None) + if in_memory is not None: + return next( + (i for i, a in enumerate(in_memory) if a.index >= index), len(in_memory) + ) + db = _session(adventure) + if db is None: + return 0 + return ( + _count_query(db, adventure, None) + .filter(models.Action.index < index) + .scalar() + or 0 + ) + + +def max_action_index(adventure: models.Adventure) -> int: + """Highest `Action.index` in the adventure, story text or not. -1 if empty.""" + loaded = _loaded_actions(adventure) + if loaded is not None: + return max((a.index for a in loaded), default=-1) + db = _session(adventure) + if db is None: + return -1 + highest = ( + db.query(func.max(models.Action.index)) + .filter(models.Action.adventure_id == adventure.id) + .scalar() + ) + return -1 if highest is None else highest + + +def window_covering( + adventure: models.Adventure, + budget_tokens: int, + token_counter, + exclude_action_id: int | None = None, +) -> list[models.Action]: + """The newest story actions whose combined text exceeds `budget_tokens` — + i.e. more than the context builder can possibly include, and never less. + + Measures rather than guesses a chars-per-token ratio, so the prompt is + byte-for-byte what loading the whole story would have produced. Budgets on + the raw text, which is never longer than the rendered history text, so + erring here can only mean fetching slightly too much. + + Each round fetches only the actions it doesn't already hold, so no row is + ever read twice however many rounds it takes. + """ + actions: list[models.Action] = [] + tokens = 0 + size = WINDOW_START + while True: + older = tail_range( + adventure, len(actions), size - len(actions), exclude_action_id + ) + if not older: + return actions # already holding the whole story + actions = older + actions + tokens += sum(token_counter(a.text) for a in older) + if len(actions) < size: + return actions # that was the whole story + if tokens > budget_tokens: + return actions + # Short. Project how many actions the budget takes at the length these + # ones turned out to be, and go straight there. + average = tokens / len(actions) + projected = int(budget_tokens / average * (1 + WINDOW_MARGIN)) + WINDOW_STEP + size = max(projected, size + WINDOW_STEP) diff --git a/backend/app/memorybank.py b/backend/app/memorybank.py index ad9f81a..7df9219 100644 --- a/backend/app/memorybank.py +++ b/backend/app/memorybank.py @@ -25,7 +25,7 @@ import math from sqlalchemy.orm import Session from . import models -from .context import story_actions, truncate_to_last_tokens +from .context import history, story_actions, truncate_to_last_tokens from .database import SessionLocal from .providers import OpenAICompatibleProvider, ProviderError @@ -35,6 +35,7 @@ SUMMARY_INTERVAL = 15 # actions between Story Summary updates MAX_MEMORIES_PER_RUN = 5 # cap catch-up work (e.g. imported adventures) per turn MAX_EMBED_BATCH = 32 RETRIEVAL_WINDOW_TOKENS = 600 # recent story text used as the similarity query +RETRIEVAL_WINDOW_ACTIONS = 4 # ...taken from this many of the newest actions SUMMARY_MAX_WORDS = 250 MEMORY_SYSTEM_PROMPT = ( @@ -88,9 +89,31 @@ def cosine(a: list[float], b: list[float]) -> float: return dot / norm if norm else 0.0 +def settled_count(adventure: models.Adventure) -> int: + """How many story actions are old enough to summarize: all but the newest. + + See settled_story_actions for why one action is held back. Counting rather + than listing keeps the post-turn pass off the whole story. + """ + return max(history.count(adventure) - 1, 0) + + +def settled_slice(adventure: models.Adventure, start: int, length: int) -> list[models.Action]: + """Settled story actions at positions [start, start + length). + + Callers must already have checked against `settled_count()`; this only + fetches, it does not re-clamp. + """ + return history.slice_(adventure, start, length) + + def settled_story_actions(adventure: models.Adventure) -> list[models.Action]: """Story actions old enough to summarize: everything but the newest one. + The plain-list form of the rule. The passes below use `settled_count` and + `settled_slice` instead, which express the same thing without reading the + whole story; this stays as the statement of what they must agree with. + Only the *last* action can be retried, so once an action has another action after it, its text is final. Summarizing right up to the newest action meant a memory could describe an attempt the player then retried away — the @@ -112,8 +135,7 @@ def _rewind_cursors_to_index(adventure: models.Adventure, index: int) -> None: Action.index values, so the two spaces have to be translated between (they diverge as soon as any action is deleted). """ - actions = story_actions(adventure) - position = next((i for i, a in enumerate(actions) if a.index >= index), len(actions)) + position = history.position_of_index(adventure, index) adventure.memory_cursor = min(adventure.memory_cursor, position) adventure.summary_cursor = min(adventure.summary_cursor, position) @@ -127,10 +149,11 @@ def note_action_removed(adventure: models.Adventure, action: models.Action) -> N that was never summarized shifts into the "already covered" range and is skipped forever. """ - actions = story_actions(adventure) - position = next((i for i, a in enumerate(actions) if a.id == action.id), None) - if position is None: - return + if not history.is_story_text(action.text): + return # not in the list the cursors count, so nothing shifts + # Actions are ordered by index, so "how many come before it" is exactly + # "how many have a lower index" — no need to walk the list to find it. + position = history.position_of_index(adventure, action.index) if position < adventure.memory_cursor: adventure.memory_cursor -= 1 if position < adventure.summary_cursor: @@ -148,7 +171,7 @@ def prune_dangling_memories(adventure: models.Adventure, db: Session) -> int: describing them. Rewind to where the earliest discarded memory began, so those actions are summarized again. """ - max_index = max((a.index for a in adventure.actions), default=-1) + max_index = history.max_action_index(adventure) dangling = [ m for m in adventure.memories if m.source_end is not None and m.source_end > max_index @@ -186,9 +209,9 @@ async def retrieve_memories( if not candidates: return {"used": [], "error": None} - actions = story_actions(adventure, exclude_action_id) + recent = history.tail(adventure, RETRIEVAL_WINDOW_ACTIONS, exclude_action_id) query = truncate_to_last_tokens( - "\n\n".join(a.text for a in actions[-4:]), RETRIEVAL_WINDOW_TOKENS + "\n\n".join(a.text for a in recent), RETRIEVAL_WINDOW_TOKENS ) if not query.strip(): return {"used": [], "error": None} @@ -264,9 +287,9 @@ async def run_post_turn(adventure_id: int) -> None: # an already-summarized action in the next block. Both consumers below # read settled actions and bail on a negative remainder, so a cursor # briefly sitting one past the settled end is harmless. - count = len(story_actions(adventure)) - adventure.memory_cursor = min(adventure.memory_cursor, count) - adventure.summary_cursor = min(adventure.summary_cursor, count) + total = history.count(adventure) + adventure.memory_cursor = min(adventure.memory_cursor, total) + adventure.summary_cursor = min(adventure.summary_cursor, total) if adventure.auto_summarize: await _create_due_memories(adventure, settings, db) await _update_story_summary(adventure, settings, db) @@ -281,13 +304,18 @@ async def run_post_turn(adventure_id: int) -> None: async def _create_due_memories( adventure: models.Adventure, settings: models.Settings, db: Session ) -> None: - actions = settled_story_actions(adventure) provider = summary_provider(settings) for _ in range(MAX_MEMORIES_PER_RUN): + # Re-counted each pass: a memory just committed doesn't change the + # count, but this loop is the only thing that moves the cursor, so the + # comparison has to be against a total that is still current. + settled = settled_count(adventure) cursor = adventure.memory_cursor - if len(actions) < MEMORY_START or len(actions) - cursor < MEMORY_INTERVAL: + if settled < MEMORY_START or settled - cursor < MEMORY_INTERVAL: + return + block = settled_slice(adventure, cursor, MEMORY_INTERVAL) + if len(block) < MEMORY_INTERVAL: return - block = actions[cursor:cursor + MEMORY_INTERVAL] excerpt = truncate_to_last_tokens("\n\n".join(a.text for a in block), 2000) try: text = await provider.complete( @@ -312,8 +340,8 @@ async def _create_due_memories( async def _update_story_summary( adventure: models.Adventure, settings: models.Settings, db: Session ) -> None: - actions = settled_story_actions(adventure) - if len(actions) - adventure.summary_cursor < SUMMARY_INTERVAL: + settled = settled_count(adventure) + if settled - adventure.summary_cursor < SUMMARY_INTERVAL: return # Fold in memories covering the uncovered stretch; fall back to raw story @@ -321,10 +349,12 @@ async def _update_story_summary( # summary_cursor is a position into story_actions(); Memory.source_end is # an Action.index. Translate the cursor to an index boundary before # comparing — the two spaces diverge once actions are deleted or empty. - if adventure.summary_cursor < len(actions): - boundary = actions[adventure.summary_cursor].index + if adventure.summary_cursor < settled: + [first_uncovered] = settled_slice(adventure, adventure.summary_cursor, 1) + boundary = first_uncovered.index else: - boundary = actions[-1].index + 1 if actions else 0 + last = settled_slice(adventure, settled - 1, 1) if settled else [] + boundary = last[0].index + 1 if last else 0 new_events = [ m.text for m in adventure.memories @@ -333,7 +363,9 @@ async def _update_story_summary( if new_events: events_text = "\n".join(f"- {t}" for t in new_events) else: - block = actions[adventure.summary_cursor:] + block = settled_slice( + adventure, adventure.summary_cursor, settled - adventure.summary_cursor + ) events_text = truncate_to_last_tokens("\n\n".join(a.text for a in block), 2000) current = adventure.story_summary.strip() @@ -351,7 +383,7 @@ async def _update_story_summary( if not text: return adventure.story_summary = text - adventure.summary_cursor = len(actions) + adventure.summary_cursor = settled db.commit() diff --git a/backend/app/migrations.py b/backend/app/migrations.py index c05dd6e..34522ad 100644 --- a/backend/app/migrations.py +++ b/backend/app/migrations.py @@ -115,12 +115,19 @@ MIGRATIONS: list[tuple[int, str]] = [ # the emit block replayed into history. Lift just that slice into its own # column so the snapshot can be deferred. Backfilled by _backfill_world_delta. (36, "ALTER TABLE actions ADD COLUMN world_delta JSON"), + # Egress, part two: `variants` holds every discarded retry attempt, but a + # list response only needs how many there are. Keep the count beside it so + # the column itself can be deferred — otherwise each retry permanently adds + # ~5 KB to every later load of that adventure. Backfilled by + # _backfill_variant_count. + (37, "ALTER TABLE actions ADD COLUMN variant_count INTEGER NOT NULL DEFAULT 0"), ] LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1) # Migrations that need a data pass after their DDL, keyed by version. WORLD_DELTA_VERSION = 36 +VARIANT_COUNT_VERSION = 37 def _backfill_world_delta(conn) -> None: @@ -154,6 +161,28 @@ def _backfill_world_delta(conn) -> None: conn.execute(text(sql)) +def _backfill_variant_count(conn) -> None: + """Populate actions.variant_count from the existing variants list. + + Server-side for the same reason as _backfill_world_delta: `variants` is the + column being taken off the wire, so counting it in Python would mean + dragging every stored attempt across the network once to avoid dragging it + across forever. + """ + if conn.dialect.name == "sqlite": + sql = """ + UPDATE actions SET variant_count = json_array_length(variants) + WHERE variants IS NOT NULL AND json_valid(variants) + """ + else: + sql = """ + UPDATE actions SET variant_count = jsonb_array_length(variants::jsonb) + WHERE variants IS NOT NULL + AND jsonb_typeof(variants::jsonb) = 'array' + """ + conn.execute(text(sql)) + + def _get_version(conn) -> int: if conn.dialect.name == "sqlite": return conn.execute(text("PRAGMA user_version")).scalar() or 1 @@ -194,6 +223,8 @@ def bootstrap(engine: Engine) -> None: conn.execute(text(sql)) if version == WORLD_DELTA_VERSION: _backfill_world_delta(conn) + if version == VARIANT_COUNT_VERSION: + _backfill_variant_count(conn) current = version _set_version(conn, current) _encrypt_plaintext_api_keys(conn) diff --git a/backend/app/models.py b/backend/app/models.py index ddb57ac..1c427b4 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -234,7 +234,15 @@ class Action(Base): # is its own only version. `variant_index` says which entry `text`, # `reasoning` and `context_snapshot` currently mirror; retry appends and # points here instead of deleting the row, so nothing is lost. - variants: Mapped[list | None] = mapped_column(JSON, nullable=True) + # + # Deferred for the same reason as context_snapshot: a list response only + # ever needs the *count* (see variant_count below), but the column holds + # every discarded attempt's full narration, so loading it in bulk made each + # retry a permanent tax on every later page load of that adventure. + variants: Mapped[list | None] = mapped_column(JSON, nullable=True, deferred=True) + # len(variants), maintained on write by set_variants() so the deferred + # column above never has to be fetched just to count it. 0 = never retried. + variant_count: Mapped[int] = mapped_column(Integer, default=0) variant_index: Mapped[int] = mapped_column(Integer, default=0) created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow) @@ -268,12 +276,6 @@ class Action(Base): out.append({"kind": "stat", "label": label, "delta": delta, "value": new}) return out - @property - def variant_count(self) -> int: - """How many attempts exist for this turn. 0 (not 1) when the action was - never retried — the UI shows its pager only above 1 either way.""" - return len(self.variants) if isinstance(self.variants, list) else 0 - class Script(Base): __tablename__ = "scripts" diff --git a/backend/app/routers/adventures.py b/backend/app/routers/adventures.py index 83ea652..649c63c 100644 --- a/backend/app/routers/adventures.py +++ b/backend/app/routers/adventures.py @@ -6,10 +6,11 @@ import threading from fastapi import APIRouter, Body, Depends, HTTPException, Request from fastapi.responses import StreamingResponse from sqlalchemy import func -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, undefer from .. import auth, images, limits, memorybank, models, schemas, worldstate from ..context import build_context +from ..context import history as context_history from ..database import get_db from ..providers import OpenAICompatibleProvider, PromptParts, ProviderError from ..scripting import ScriptPipeline @@ -349,6 +350,18 @@ def world_delta_of(snapshot: dict | None) -> dict | None: } +def set_variants(action: models.Action, entries: list[dict]) -> None: + """The ONLY way to write Action.variants. + + `variants` is deferred (it holds every discarded attempt's narration), so + `variant_count` exists to answer "how many attempts?" without fetching it. + Writing the list anywhere else would let the two drift and the pager would + lie about how many takes a turn has. + """ + action.variants = entries + action.variant_count = len(entries) + + def variant_of(action: models.Action, adventure: models.Adventure) -> dict: """Freeze an action's *current* content as a variant entry. @@ -475,7 +488,19 @@ def action_json(action: models.Action) -> dict: def next_index(adventure: models.Adventure) -> int: - return max((a.index for a in adventure.actions), default=-1) + 1 + return context_history.max_action_index(adventure) + 1 + + +def last_action(adventure: models.Adventure, db: Session) -> models.Action | None: + """The newest action of any kind, or None. A query rather than + `adventure.actions[-1]`, which would load the entire story to look at + one row.""" + return ( + db.query(models.Action) + .filter(models.Action.adventure_id == adventure.id) + .order_by(models.Action.index.desc()) + .first() + ) async def generate_turn( @@ -644,7 +669,7 @@ async def _generate_turn( "created_at": models.utcnow().isoformat(), **{k: copy.deepcopy(snapshot[k]) for k in VARIANT_SNAPSHOT_KEYS if k in snapshot}, }) - ai_action.variants = history + set_variants(ai_action, history) apply_variant(ai_action, adventure, len(history) - 1) else: ai_action = models.Action( @@ -768,13 +793,14 @@ def retry_action( acquire_turn_lock(adventure_id) last_ai = None try: - if adventure.actions and adventure.actions[-1].type == "ai": - last_ai = adventure.actions[-1] + newest = last_action(adventure, db) + if newest is not None and newest.type == "ai": + last_ai = newest # First retry: the row has no history yet, so record what's on # screen as variant 0 before anything is rolled back — the live # script/world state is precisely that attempt's outcome. - if not last_ai.variants: - last_ai.variants = [variant_of(last_ai, adventure)] + if not last_ai.variant_count: + set_variants(last_ai, [variant_of(last_ai, adventure)]) last_ai.variant_index = 0 # Roll the scoreboard back to before this AI turn's hooks ran, so # regenerating starts fresh instead of stacking output mutations on @@ -855,7 +881,8 @@ def select_variant( variants = action.variants if isinstance(action.variants, list) else [] if not 0 <= payload.index < len(variants): raise HTTPException(400, "No such attempt for this action") - if not adventure.actions or adventure.actions[-1].id != action.id: + newest = last_action(adventure, db) + if newest is None or newest.id != action.id: raise HTTPException( 400, "Only the latest message can be switched — the story has already " @@ -884,16 +911,25 @@ def undo_turn( adventure = get_adventure_or_404(adventure_id, db, user) acquire_turn_lock(adventure_id) try: - actions = list(adventure.actions) - if not actions or actions[-1].type == "start": + # Only the last turn is ever removed, so fetch the two actions it can + # consist of rather than the whole story. + newest = ( + db.query(models.Action) + .filter(models.Action.adventure_id == adventure.id) + .order_by(models.Action.index.desc()) + .limit(2) + .all() + ) + if not newest or newest[0].type == "start": raise HTTPException(400, "Nothing to undo") - last = actions.pop() + last = newest[0] + preceding = newest[1] if len(newest) > 1 else None # The earliest action removed in this turn holds the pre-turn scoreboard. first_removed = last memorybank.note_action_removed(adventure, last) db.delete(last) - if last.type == "ai" and actions and actions[-1].type in ("do", "say", "story"): - first_removed = actions.pop() + if last.type == "ai" and preceding is not None and preceding.type in ("do", "say", "story"): + first_removed = preceding memorybank.note_action_removed(adventure, first_removed) db.delete(first_removed) if first_removed.state_before is not None: @@ -918,6 +954,16 @@ def export_adventure( ): """Full backup: plot components, story cards, scripts (+state), every action.""" adv = get_adventure_or_404(adventure_id, db, user) + # Export is the one read that genuinely wants every attempt, so it asks for + # the deferred `variants` column up front — iterating adv.actions instead + # would lazy-load it one row at a time. + exported_actions = ( + db.query(models.Action) + .filter(models.Action.adventure_id == adv.id) + .options(undefer(models.Action.variants)) + .order_by(models.Action.index) + .all() + ) return { "format": "ai-dnd-adventure-v1", "title": adv.title, @@ -967,7 +1013,7 @@ def export_adventure( "variantIndex": a.variant_index, "createdAt": a.created_at.isoformat(), } - for a in adv.actions + for a in exported_actions ], } @@ -1060,6 +1106,7 @@ def import_adventure( text=str(a["text"]), reasoning=str(a["reasoning"]) if a.get("reasoning") else None, variants=variants or None, + variant_count=len(variants), # Clamped: a bundle could name an index its variant list # doesn't have, which would make the pager point at nothing. variant_index=min(max(int(a.get("variantIndex", 0)), 0), max(len(variants) - 1, 0)), @@ -1493,11 +1540,12 @@ def update_action( action.text = payload.text # Keep the live variant in step, or paging away and back would silently # revert the edit. - variants = action.variants if isinstance(action.variants, list) else [] - if 0 <= action.variant_index < len(variants): - history = copy.deepcopy(variants) - history[action.variant_index]["text"] = payload.text - action.variants = history + if action.variant_count: + variants = action.variants if isinstance(action.variants, list) else [] + if 0 <= action.variant_index < len(variants): + history = copy.deepcopy(variants) + history[action.variant_index]["text"] = payload.text + set_variants(action, history) db.commit() return action diff --git a/backend/tests/test_egress.py b/backend/tests/test_egress.py index 8b169bb..ef73e93 100644 --- a/backend/tests/test_egress.py +++ b/backend/tests/test_egress.py @@ -22,6 +22,7 @@ from fastapi.testclient import TestClient from sqlalchemy import event, text from app import auth, limits, migrations, models +from app.context import history from app.database import Base, SessionLocal, engine, get_db from app.main import app @@ -37,6 +38,14 @@ BIG_SNAPSHOT = { }, } +# Retry history: each discarded attempt keeps its full narration, so an action +# retried a few times carries several KB that a list response only ever counts. +BIG_VARIANTS = [ + {"text": "z" * 4_000, "reasoning": None, "script_state": {}, + "created_at": "2026-01-01T00:00:00"} + for _ in range(3) +] + @pytest.fixture() def sql_log(): @@ -71,6 +80,10 @@ def client(monkeypatch): context_snapshot=BIG_SNAPSHOT, world_delta={"delta": {"player.hp": -15}, "applied": [{"path": "player.hp", "old": 100, "new": 85}]}, + # Every AI action has been retried twice, so `variants` is carrying + # weight the list response must not pay for. + variants=BIG_VARIANTS if i % 2 else None, + variant_count=len(BIG_VARIANTS) if i % 2 else 0, )) setup.commit() adv_id, user_id = adventure.id, user.id @@ -129,6 +142,48 @@ def test_world_changes_still_works_without_the_snapshot(client): ] +def test_loading_an_adventure_does_not_fetch_variants(client, sql_log): + """Same failure as context_snapshot, one size down: the payload carries + only `variant_count`, but loading the column to compute it made every retry + a permanent tax on every later load of that adventure.""" + r = client.get(f"/api/adventures/{client.adv_id}") + assert r.status_code == 200, r.text + + selects = action_selects(sql_log) + assert selects, "expected at least one SELECT against actions" + offenders = [s for s in selects if "variants" in s] + assert offenders == [], f"variants was fetched in bulk:\n{offenders[0][:400]}" + + +def test_variant_count_survives_variants_being_deferred(client): + """The pager reads this number; it has to be right without the column.""" + r = client.get(f"/api/adventures/{client.adv_id}") + by_type = {} + for action in r.json()["actions"]: + by_type.setdefault(action["type"], []).append(action) + assert all(a["variant_count"] == len(BIG_VARIANTS) for a in by_type["ai"]) + assert all(a["variant_count"] == 0 for a in by_type["do"]) + + +def test_counting_actions_does_not_name_the_deferred_columns(client, sql_log): + """A count that wraps the entity select in a subquery names every column in + the emitted SQL — no bytes come back, but the database still reads them and + the guard above cannot tell it apart from a real bulk fetch.""" + db = SessionLocal() + try: + adventure = db.get(models.Adventure, client.adv_id) + sql_log.clear() + assert history.count(adventure) == 12 + counts = [s for s in sql_log if "count" in s.lower()] + assert counts, "expected a COUNT to be emitted" + for column in ("context_snapshot", "state_before", "world_state_before", "variants"): + assert not any(column in s for s in counts), ( + f"{column} is named by the count query:\n{counts[0][:400]}" + ) + finally: + db.close() + + def test_snapshot_is_still_reachable_on_demand(client): """Deferred means lazy, not gone — Insights still gets the full thing.""" r = client.get(f"/api/adventures/{client.adv_id}") @@ -163,6 +218,27 @@ def test_backfill_populates_world_delta_from_existing_snapshots(client): db.close() +def test_backfill_populates_variant_count_from_existing_variants(client): + """Migration 37 counts the lists server-side — reading them into Python to + count them would mean pulling the column across the wire once to stop + pulling it across forever.""" + db = SessionLocal() + try: + db.execute(text("UPDATE actions SET variant_count = 0")) + db.commit() + + with engine.begin() as conn: + migrations._backfill_variant_count(conn) + + db.expire_all() + actions = db.query(models.Action).order_by(models.Action.index).all() + for action in actions: + expected = len(BIG_VARIANTS) if action.type == "ai" else 0 + assert action.variant_count == expected, f"action {action.index}" + finally: + db.close() + + def test_backfill_leaves_actions_without_world_state_alone(client): db = SessionLocal() try: diff --git a/backend/tests/test_history_window.py b/backend/tests/test_history_window.py new file mode 100644 index 0000000..a7b123c --- /dev/null +++ b/backend/tests/test_history_window.py @@ -0,0 +1,249 @@ +"""The context builder reads a window of the story, not all of it. + +Walking `adventure.actions` every turn made a turn cost O(story length), so a +long adventure read hundreds of KB to use the tail of it — and the cost grew +with every turn played. `app.context.history` serves tails, slices and counts +from SQL instead. + +Two things have to hold, and both are easy to break by accident: + +* the window must produce **exactly** the prompt the full story produced, or + this is a behaviour change wearing an optimization's clothes; +* the helpers must agree with the old list arithmetic, because memorybank's + cursors are *positions* in that list and a cursor off by one silently + summarizes the wrong actions. + + python -m pytest tests/test_history_window.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 sqlalchemy import event + +from app import memorybank, models +from app.context import builder, history +from app.database import Base, SessionLocal, engine + +# Long enough that a window is much smaller than the whole story. +ACTION_COUNT = 200 +NARRATION = ( + "The scrub gives way to a shallow bowl of land where woodsmoke hangs in " + "flat grey layers, and somewhere behind the largest tent a woman is " + "arguing, low and fast. " +) * 3 + +SCHEMA = { + "player": {"hp": {"min": 0, "max": 100, "initial": 100, "desc": "Health"}}, + "npcs": { + "gwen": {"name": "Gwen", "keys": ["gwen"], "desc": "A scout.", + "stats": {"trust": {"min": 0, "max": 100, "initial": 30}}}, + }, +} + + +@pytest.fixture() +def story(): + """An adventure with ACTION_COUNT actions, plus its settings.""" + Base.metadata.create_all(bind=engine) + db = SessionLocal() + user = models.User(is_guest=False, email="window@example.com") + db.add(user) + db.flush() + settings = models.Settings(user_id=user.id, api_key="enc:dummy", model="m") + db.add(settings) + scenario = models.Scenario(user_id=user.id, title="S", stat_schema=SCHEMA, + prompt="A long road." * 50) + db.add(scenario) + db.flush() + adventure = models.Adventure( + user_id=user.id, title="Long", scenario_id=scenario.id, script_state={}, + memory="The hero is hunting bandits. " * 20, + world_state={"player": {"hp": 100}, "npc": {"gwen": {"trust": 30}}, + "milestones": {}, "flags": {}, "_meta": {"last_changed": {}}}, + ) + db.add(adventure) + db.flush() + db.add(models.StoryCard(adventure_id=adventure.id, name="Gwen", keys="gwen", + entry="A scout with sharp eyes.", type="lore")) + for i in range(ACTION_COUNT): + db.add(models.Action( + adventure_id=adventure.id, index=i, + type="ai" if i % 2 else "do", + text=f"[{i}] {NARRATION}", + world_delta={"delta": {"player.hp": -1}, + "applied": [{"path": "player.hp", "old": 100, "new": 99}]}, + )) + db.commit() + db.expire_all() + adventure = db.get(models.Adventure, adventure.id) + settings = db.get(models.Settings, settings.id) + try: + yield db, adventure, settings + finally: + db.close() + Base.metadata.drop_all(bind=engine) + + +def full_window(adventure, budget_tokens, token_counter, exclude_action_id=None): + """Stand-in for window_covering that hands back the entire story, i.e. the + behaviour this module replaced.""" + return history.story_actions(adventure, exclude_action_id) + + +@pytest.fixture() +def actions_loaded(): + """Counts Action rows the ORM materializes, i.e. how much of the story was + actually fetched. rowcount is meaningless for SELECT on SQLite, so count + the objects the mapper builds instead.""" + loaded = {"n": 0} + + def on_load(target, context): + loaded["n"] += 1 + + event.listen(models.Action, "load", on_load) + try: + yield loaded + finally: + event.remove(models.Action, "load", on_load) + + +# ------------------------------------------------------- the prompt is equal + +@pytest.mark.parametrize("budget", [1024, 4096, 8192, 16384, 65536]) +def test_window_builds_the_same_prompt_as_the_whole_story(story, budget, monkeypatch): + db, adventure, settings = story + settings.context_token_budget = budget + + windowed = builder.build_context(adventure, settings) + monkeypatch.setattr(builder.history, "window_covering", full_window) + db.expire(adventure) + everything = builder.build_context(adventure, settings) + + assert windowed[0] == everything[0], "system prompt differs" + assert windowed[1] == everything[1], "story prompt differs" + assert windowed[2]["history"] == everything[2]["history"] + assert windowed[2]["cards"] == everything[2]["cards"] + + +def test_window_matches_on_the_retry_shape(story, monkeypatch): + """Retry excludes the action being regenerated; the exclusion has to reach + the window query, not just the in-memory filter.""" + db, adventure, settings = story + last = history.tail(adventure, 1)[0] + + windowed = builder.build_context(adventure, settings, exclude_action_id=last.id) + assert f"[{last.index}]" not in windowed[1] + + monkeypatch.setattr(builder.history, "window_covering", full_window) + db.expire(adventure) + everything = builder.build_context(adventure, settings, exclude_action_id=last.id) + assert windowed[1] == everything[1] + + +def test_reported_total_is_the_whole_story_not_the_window(story): + """Insights says "N of M actions included"; M must not become the window.""" + db, adventure, settings = story + settings.context_token_budget = 4096 + report = builder.build_context(adventure, settings)[2] + assert report["history"]["total"] == ACTION_COUNT + assert report["history"]["included"] < ACTION_COUNT + + +# ------------------------------------------------------------ it is bounded + +def test_building_context_reads_far_less_than_the_whole_story(story, actions_loaded): + db, adventure, settings = story + # Expire first: expiring afterwards would discard the unflushed change and + # silently put the budget back to its default. + db.expire_all() + # Small enough that the budget, not the length of the story, decides. + settings.context_token_budget = 4096 + actions_loaded["n"] = 0 + + report = builder.build_context(adventure, settings)[2] + included = report["history"]["included"] + + assert included < ACTION_COUNT, "fixture is too short to prove anything" + # The window aims a margin past the budget and re-asks if it fell short, so + # it reads somewhat more than it includes. What matters is that the read is + # a function of the token budget, not of how long the story has got. + assert actions_loaded["n"] < ACTION_COUNT // 2, ( + f"read {actions_loaded['n']} action rows out of {ACTION_COUNT} to " + f"include {included} — the window is not bounding the read" + ) + + +def test_window_is_ordered_and_free_of_duplicates(story): + """The window grows by fetching only what it does not already hold, so an + off-by-one in the offset would show up as a repeated or missing action.""" + db, adventure, settings = story + window = history.window_covering(adventure, 16384, builder.count_tokens) + ids = [a.id for a in window] + assert len(ids) == len(set(ids)), "the same action appeared twice in the window" + assert ids == sorted(ids), "window must be oldest-first" + + +# --------------------------------------------- the cursor arithmetic agrees + +def test_helpers_agree_with_the_full_list(story): + db, adventure, settings = story + actions = history.story_actions(adventure) + assert len(actions) == ACTION_COUNT + + assert history.count(adventure) == len(actions) + assert history.max_action_index(adventure) == max(a.index for a in actions) + assert [a.id for a in history.tail(adventure, 4)] == [a.id for a in actions[-4:]] + assert [a.id for a in history.slice_(adventure, 10, 6)] == [a.id for a in actions[10:16]] + assert [a.id for a in history.tail_range(adventure, 5, 3)] == \ + [a.id for a in actions[-8:-5]] + assert memorybank.settled_count(adventure) == len(actions) - 1 + + for probe in (0, 1, ACTION_COUNT // 2, ACTION_COUNT - 1): + target = actions[probe] + expected = next(i for i, a in enumerate(actions) if a.index >= target.index) + assert history.position_of_index(adventure, target.index) == expected + + +def test_positions_still_line_up_after_a_middle_action_is_deleted(story): + """The gap in Action.index is exactly what makes positions and indexes + diverge — the case that has broken the cursors twice before.""" + db, adventure, settings = story + actions = history.story_actions(adventure) + victim = actions[50] + db.delete(victim) + db.commit() + db.expire(adventure) + + remaining = history.story_actions(adventure) + assert len(remaining) == ACTION_COUNT - 1 + assert history.count(adventure) == ACTION_COUNT - 1 + for probe in (0, 49, 50, 51, ACTION_COUNT - 2): + target = remaining[probe] + expected = next(i for i, a in enumerate(remaining) if a.index >= target.index) + assert history.position_of_index(adventure, target.index) == expected, probe + + +def test_blank_actions_are_excluded_the_same_way_in_sql_and_python(story): + """SQL and Python must agree on membership or a cursor points elsewhere.""" + db, adventure, settings = story + for blank in ("", " ", "\n", "\t\n "): + db.add(models.Action(adventure_id=adventure.id, + index=history.max_action_index(adventure) + 1, + type="story", text=blank)) + db.commit() + db.expire(adventure) + + # SQL path (relationship not loaded) + from_sql = history.count(adventure) + # Python path (relationship loaded) + adventure.actions # noqa: B018 — force the collection into memory + from_python = history.count(adventure) + + assert from_sql == from_python == ACTION_COUNT