From d766b1d66fd109faf6da9789954b103e6c0e33a7 Mon Sep 17 00:00:00 2001 From: Chris Guidry Date: Sat, 28 Mar 2026 23:10:08 -0400 Subject: [PATCH] Add initiative tracking tools for structured combat MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The DM was tracking initiative order, enemy HP, conditions, and effect durations entirely in its head. That works for short fights but drifts badly on longer encounters and breaks on context truncation. Now the DM has a proper initiative tracker — a state machine that handles turn order, HP, conditions with 5e-correct duration expiry (anchored to the source's turn), and auto-syncs player HP to the character sheet. The full combat state gets injected into the system prompt every turn, so the DM always sees ground truth instead of trying to remember it. The MCP server dynamically swaps tool sets: narrative mode gets the usual 10 tools plus enter_initiative, initiative mode gets 7 combat tools (damage, heal, condition, next_turn, etc.) plus roll, recall, and update_character. This keeps the DM's decision space tight — it only sees what's relevant to the current mode. Works for non-combat initiative too (environmental hazards, chases, timed puzzles) since the tracker is just "who goes when" with HP and conditions bolted on. Co-Authored-By: Claude Opus 4.6 (1M context) --- .loq_cache | 1 + prompts/dm-system.md | 64 ++++- src/storied/engine.py | 14 + src/storied/initiative.py | 524 +++++++++++++++++++++++++++++++++++++ src/storied/mcp_server.py | 28 +- src/storied/tools.py | 24 +- tests/test_engine.py | 2 + tests/test_execute_tool.py | 99 +++++++ tests/test_initiative.py | 507 +++++++++++++++++++++++++++++++++++ tests/test_mcp_server.py | 86 +++++- 10 files changed, 1334 insertions(+), 15 deletions(-) create mode 100644 .loq_cache create mode 100644 src/storied/initiative.py create mode 100644 tests/test_initiative.py diff --git a/.loq_cache b/.loq_cache new file mode 100644 index 0000000..74b7605 --- /dev/null +++ b/.loq_cache @@ -0,0 +1 @@ +{"version":1,"config_hash":4557771575092473650,"entries":{"src/storied/initiative.py":{"mtime_secs":1774751731,"mtime_nanos":975235964,"lines":524}}} \ No newline at end of file diff --git a/prompts/dm-system.md b/prompts/dm-system.md index e91b948..6383fd8 100644 --- a/prompts/dm-system.md +++ b/prompts/dm-system.md @@ -4,18 +4,35 @@ This is collaborative storytelling in a fantasy game. Players may explore morall ## Available Tools +### Always Available | Tool | Purpose | |------|---------| -| `set_scene` | **Call after every response.** Logs what happened, advances the clock, updates the scene | | `roll` | Roll dice (e.g., `roll("1d20+5", "attack")`) | | `recall` | Look up rules or world content | +| `update_character` | Modify character stats (HP, gold, equipment) | + +### Narrative Mode (outside initiative) +| Tool | Purpose | +|------|---------| +| `set_scene` | **Call after every response.** Logs what happened, advances the clock, updates the scene | | `establish` | Create or update entities (NPCs, locations, items, threads) | | `mark` | Record what happened to an entity | | `note_discovery` | Record what the player learned | -| `update_character` | Modify character stats (HP, gold, equipment) | | `create_character` | Create a new character | | `tune` | Update your style/personality tuning based on player feedback | | `end_session` | Gracefully end the session | +| `enter_initiative` | Enter initiative mode for combat or turn-based encounters | + +### Initiative Mode (during initiative) +| Tool | Purpose | +|------|---------| +| `next_turn` | Advance to the next combatant's turn | +| `damage` | Deal damage to a combatant (tracks defeat at 0 HP) | +| `heal` | Heal a combatant (clamped to max HP) | +| `condition` | Add or remove a condition (Prone, Stunned, etc.) | +| `add_combatant` | Add reinforcements or late arrivals | +| `remove_combatant` | Remove a combatant who fled or was banished | +| `end_initiative` | End initiative and return to narrative mode | ## After Every Response: Call `set_scene` @@ -198,14 +215,45 @@ Set the DC mentally, roll, then narrate the outcome. Don't announce DCs to the p **In combat**: More mechanical transparency is fine. The player should understand the tactical situation - hits, misses, how wounded enemies look. But still narrate, don't just announce numbers. -## Combat Flow +## Initiative Mode + +When ordered action matters — combat, environmental dangers, timed puzzles, chase sequences — use initiative mode. The tracker handles turn order, HP, and conditions so you don't have to hold it in memory. + +### Entering Initiative + +1. Roll initiative for all participants (use `roll`) +2. Decide turn order (you handle tie-breaking per 5e rules) +3. Call `enter_initiative` with all combatants listed in turn order + +After calling `enter_initiative`, finish your response with narration setting the scene. Initiative tools become available starting next turn. + +### During Initiative + +Your context always includes the full initiative table — current turn, HP, AC, conditions, who's next. You never need to remember combat state; it's always right there. + +Each combatant's turn: +- **Player's turn**: Narrate the tactical situation, then wait for their input. After they act, resolve rolls and effects, then call `next_turn`. +- **Monster's turn**: Decide their action, roll attacks/saves/damage, apply `damage`/`condition`, then call `next_turn`. + +Use `damage` and `heal` to track HP changes on combatants. Use `condition` to apply or remove conditions with optional durations. When you damage or heal the player character, their character sheet is automatically synced — no need to call `update_character` separately for HP during initiative. + +### Ending Initiative + +Call `end_initiative` when combat is resolved — enemies defeated, fled, or surrendered. It reports total rounds and duration. Then call `set_scene` with the aftermath and use the reported duration. + +### Non-Combat Initiative + +Initiative isn't just for combat. Use it for: +- **Environmental hazards**: collapsing dungeon, rising flood, spreading fire +- **Chase sequences**: tracking distance and obstacles turn by turn +- **Timed puzzles**: a mechanism with rounds to solve before something triggers +- **Social confrontations**: tense negotiations where order matters + +Any situation where "who goes when" matters deserves initiative. -In combat, track internally: -- Initiative order -- Enemy HP and conditions -- Active effects and durations +### Narrative During Initiative -Narrate each exchange with dramatic weight. A hit isn't just "8 damage" - describe the impact. Show how wounded enemies are through their behavior, not HP counts. +Narrate each exchange with dramatic weight. A hit isn't just "8 damage" — describe the impact. Show how wounded enemies are through their behavior, not HP counts. ## Player Input diff --git a/src/storied/engine.py b/src/storied/engine.py index 637ee30..93a42dd 100644 --- a/src/storied/engine.py +++ b/src/storied/engine.py @@ -63,6 +63,14 @@ def _tool_notification(name: str) -> str: "note_discovery": "Noting discovery", "tune": "Tuning style", "end_session": "Saving session", + "enter_initiative": "Entering initiative", + "next_turn": "Next turn", + "add_combatant": "Adding combatant", + "remove_combatant": "Removing combatant", + "damage": "Applying damage", + "heal": "Healing", + "condition": "Updating condition", + "end_initiative": "Ending initiative", } label = labels.get(short, short) return f"\n[{label}...]\n" @@ -241,6 +249,12 @@ class DMEngine: parts.append(entity_context) loaded_names.add(name) + # Initiative state (injected when active so the DM never loses track) + if self._mcp.ctx.initiative.active: + initiative_context = self._mcp.ctx.initiative.format_for_context() + self._context_parts["Initiative"] = initiative_context + parts.append(initiative_context) + # Terminal width so the DM can size display blocks import os try: diff --git a/src/storied/initiative.py b/src/storied/initiative.py new file mode 100644 index 0000000..4c5ff93 --- /dev/null +++ b/src/storied/initiative.py @@ -0,0 +1,524 @@ +"""Initiative tracking for structured combat and turn-based encounters.""" + +from dataclasses import dataclass, field + + +@dataclass +class TrackedCondition: + """A condition on a combatant with optional duration. + + Duration is anchored to the source's turn per 5e rules. Duration -1 + means the condition persists until manually removed. + """ + + name: str + source: str + duration: int = -1 + ends_on: str = "start" + + +@dataclass +class Combatant: + """A participant in initiative order.""" + + name: str + initiative: int + hp: int + hp_max: int + ac: int + is_player: bool = False + conditions: list[TrackedCondition] = field(default_factory=list) + defeated: bool = False + + +class InitiativeTracker: + """Pure state machine for tracking initiative order and combat state. + + Starts inactive. Call begin() to enter initiative mode, end() to leave. + The DM passes combatants in their desired turn order (pre-sorted). + """ + + def __init__(self) -> None: + self.active: bool = False + self.combatants: list[Combatant] = [] + self.current_index: int = 0 + self.round: int = 0 + + @property + def current_combatant(self) -> Combatant | None: + if not self.active or not self.combatants: + return None + return self.combatants[self.current_index] + + def begin(self, combatants: list[Combatant]) -> str: + """Enter initiative mode with the given combatants in turn order.""" + self.combatants = list(combatants) + self.current_index = 0 + self.round = 1 + self.active = True + + first = self.combatants[0] + lines = [f"Initiative started — Round 1", ""] + lines.append(self._format_order()) + lines.append("") + lines.append(f"**{first.name}** goes first.") + return "\n".join(lines) + + def next_turn(self) -> str: + """Advance to the next combatant's turn.""" + if not self.active: + return "Initiative is not active." + + # Process end-of-turn effects for the combatant we're leaving + leaving = self.combatants[self.current_index] + expired = self._process_effects(leaving.name, "end") + + # Advance to next living combatant + self.current_index = self._next_living_index(self.current_index) + + current = self.combatants[self.current_index] + started_turn = self._process_effects(current.name, "start") + + parts = [] + if expired: + parts.append(f"Expired: {', '.join(expired)}") + if started_turn: + parts.append(f"Expired: {', '.join(started_turn)}") + + parts.append( + f"Round {self.round} — **{current.name}**'s turn " + f"({current.hp}/{current.hp_max} HP, AC {current.ac})" + ) + + if current.conditions: + cond_str = ", ".join(self._format_condition(c) for c in current.conditions) + parts.append(f"Conditions: {cond_str}") + + # Check if only one side remains + hint = self._one_side_hint() + if hint: + parts.append(hint) + + return "\n".join(parts) + + def apply_damage(self, target: str, amount: int) -> str: + """Deal damage to a combatant.""" + combatant = self._find(target) + if combatant is None: + return f"Combatant '{target}' not found." + + old_hp = combatant.hp + combatant.hp = max(0, combatant.hp - amount) + result = f"{target} takes {amount} damage ({old_hp} \u2192 {combatant.hp}/{combatant.hp_max} HP)" + + if combatant.hp == 0: + combatant.defeated = True + result += " \u2014 DOWN!" + elif combatant.hp <= combatant.hp_max // 2 < old_hp: + result += " \u2014 Bloodied" + + return result + + def apply_heal(self, target: str, amount: int) -> str: + """Heal a combatant.""" + combatant = self._find(target) + if combatant is None: + return f"Combatant '{target}' not found." + + old_hp = combatant.hp + combatant.hp = min(combatant.hp_max, combatant.hp + amount) + actual = combatant.hp - old_hp + if combatant.defeated and combatant.hp > 0: + combatant.defeated = False + return f"{target} heals {actual} HP ({old_hp} \u2192 {combatant.hp}/{combatant.hp_max})" + + def add_condition( + self, + target: str, + condition: str, + duration: int = -1, + ends_on: str = "start", + source: str = "", + ) -> str: + """Apply a condition to a combatant.""" + combatant = self._find(target) + if combatant is None: + return f"Combatant '{target}' not found." + + tc = TrackedCondition( + name=condition, source=source, duration=duration, ends_on=ends_on, + ) + combatant.conditions.append(tc) + + dur_str = f" ({duration} rds)" if duration > 0 else "" + return f"{target} is now {condition}{dur_str}" + + def remove_condition(self, target: str, condition: str) -> str: + """Remove a condition from a combatant.""" + combatant = self._find(target) + if combatant is None: + return f"Combatant '{target}' not found." + + before = len(combatant.conditions) + combatant.conditions = [c for c in combatant.conditions if c.name != condition] + if len(combatant.conditions) == before: + return f"{target} does not have {condition}." + return f"{condition} removed from {target}." + + def add_combatant(self, combatant: Combatant) -> str: + """Add a combatant at the correct initiative position.""" + insert_idx = len(self.combatants) + for i, c in enumerate(self.combatants): + if combatant.initiative > c.initiative: + insert_idx = i + break + + self.combatants.insert(insert_idx, combatant) + + if insert_idx <= self.current_index: + self.current_index += 1 + + return f"{combatant.name} joins initiative (initiative {combatant.initiative})" + + def remove_combatant(self, name: str) -> str: + """Remove a combatant from initiative.""" + idx = None + for i, c in enumerate(self.combatants): + if c.name.lower() == name.lower(): + idx = i + break + + if idx is None: + return f"Combatant '{name}' not found." + + removed = self.combatants.pop(idx) + + if not self.combatants: + self.active = False + return f"{removed.name} removed. No combatants remain — initiative ended." + + if idx < self.current_index: + self.current_index -= 1 + elif idx == self.current_index: + if self.current_index >= len(self.combatants): + self.current_index = 0 + self.round += 1 + + return f"{removed.name} removed from initiative." + + def end(self) -> str: + """End initiative and return a summary.""" + defeated = [c for c in self.combatants if c.defeated] + survivors = [c for c in self.combatants if not c.defeated] + rounds = self.round + duration_sec = rounds * 6 + + lines = [f"Initiative ended after {rounds} rounds ({duration_sec} seconds)."] + + if defeated: + names = ", ".join(c.name for c in defeated) + lines.append(f"Defeated: {names}") + + if survivors: + parts = [] + for c in survivors: + conds = "" + if c.conditions: + conds = f" [{', '.join(co.name for co in c.conditions)}]" + parts.append(f"{c.name} ({c.hp}/{c.hp_max} HP{conds})") + lines.append(f"Survivors: {', '.join(parts)}") + + self.active = False + self.combatants = [] + self.current_index = 0 + self.round = 0 + + return "\n".join(lines) + + def format_for_context(self) -> str: + """Format full initiative state for system prompt injection.""" + if not self.active: + return "" + + current = self.combatants[self.current_index] + + # Find who's next (next living combatant after current) + next_idx = self._peek_next_living(self.current_index) + next_up = self.combatants[next_idx] if next_idx is not None else None + + lines = [f"## Active Initiative \u2014 Round {self.round}", ""] + lines.append("| # | Combatant | Init | HP | AC | Conditions |") + lines.append("|---|-----------|------|----|----|------------|") + + for i, c in enumerate(self.combatants): + marker = " > " if i == self.current_index else " " + name = f"**{c.name}**" if i == self.current_index else c.name + if c.defeated: + name = f"~~{c.name}~~" + hp_str = f"{c.hp}/{c.hp_max}" + conds = ", ".join(self._format_condition(co) for co in c.conditions) + if c.defeated and not conds: + conds = "Defeated" + lines.append(f"|{marker}| {name} | {c.initiative} | {hp_str} | {c.ac} | {conds} |") + + lines.append("") + + cond_str = "" + if current.conditions: + cond_str = ", " + ", ".join(co.name for co in current.conditions) + lines.append( + f"**Current turn:** {current.name} " + f"({current.hp}/{current.hp_max} HP, AC {current.ac}{cond_str})" + ) + + if next_up: + lines.append(f"**Up next:** {next_up.name}") + + lines.append(f"**Round:** {self.round}") + lines.append("") + lines.append( + f"Resolve {current.name}'s turn, then call `next_turn` to advance." + ) + + return "\n".join(lines) + + def _find(self, name: str) -> Combatant | None: + for c in self.combatants: + if c.name.lower() == name.lower(): + return c + return None + + def _next_living_index(self, from_index: int) -> int: + """Find the next living combatant after from_index, wrapping around.""" + n = len(self.combatants) + for offset in range(1, n + 1): + idx = (from_index + offset) % n + if idx == 0 and offset > 0: + self.round += 1 + if not self.combatants[idx].defeated: + return idx + return from_index # all defeated, shouldn't happen + + def _peek_next_living(self, from_index: int) -> int | None: + """Peek at next living combatant without modifying state.""" + n = len(self.combatants) + for offset in range(1, n): + idx = (from_index + offset) % n + if not self.combatants[idx].defeated: + return idx + return None + + def _process_effects(self, source_name: str, phase: str) -> list[str]: + """Process and expire effects anchored to source_name at the given phase.""" + expired: list[str] = [] + for c in self.combatants: + remaining: list[TrackedCondition] = [] + for cond in c.conditions: + if cond.source.lower() != source_name.lower() or cond.ends_on != phase: + remaining.append(cond) + continue + if cond.duration == -1: + remaining.append(cond) + continue + cond.duration -= 1 + if cond.duration <= 0: + expired.append(f"{cond.name} on {c.name}") + else: + remaining.append(cond) + c.conditions = remaining + return expired + + def _one_side_hint(self) -> str | None: + """Check if only one side (player vs non-player) has living combatants.""" + living = [c for c in self.combatants if not c.defeated] + has_player = any(c.is_player for c in living) + has_non_player = any(not c.is_player for c in living) + if has_player and not has_non_player: + return "Only player combatants remaining." + if has_non_player and not has_player: + return "Only non-player combatants remaining." + return None + + def _format_condition(self, cond: TrackedCondition) -> str: + if cond.duration == -1: + return cond.name + return f"{cond.name} ({cond.duration} rds)" + + def _format_order(self) -> str: + parts = [] + for c in self.combatants: + parts.append(f" {c.initiative}: {c.name} ({c.hp}/{c.hp_max} HP, AC {c.ac})") + return "\n".join(parts) + + +# --- Tool functions (thin wrappers around tracker methods) --- + +def enter_initiative(combatants_raw: list[dict], tracker: InitiativeTracker) -> str: + """Enter initiative mode with pre-sorted combatants.""" + if tracker.active: + return "Initiative is already active. Call end_initiative first." + + combatants = [ + Combatant( + name=c["name"], + initiative=c["initiative"], + hp=c["hp"], + hp_max=c["hp_max"], + ac=c["ac"], + is_player=c.get("is_player", False), + ) + for c in combatants_raw + ] + return tracker.begin(combatants) + + +def execute_initiative_tool( + tool_name: str, tool_input: dict, tracker: InitiativeTracker, +) -> str | None: + """Dispatch an initiative tool call. Returns None for unknown tools.""" + if tool_name not in ALL_INITIATIVE_TOOL_NAMES: + return None + + if tool_name == "enter_initiative": + return enter_initiative(tool_input["combatants"], tracker) + + if not tracker.active: + return "Initiative is not active. Call enter_initiative first." + + if tool_name == "next_turn": + return tracker.next_turn() + elif tool_name == "add_combatant": + c = Combatant( + name=tool_input["name"], + initiative=tool_input["initiative"], + hp=tool_input["hp"], + hp_max=tool_input["hp_max"], + ac=tool_input["ac"], + is_player=tool_input.get("is_player", False), + ) + return tracker.add_combatant(c) + elif tool_name == "remove_combatant": + return tracker.remove_combatant(tool_input["name"]) + elif tool_name == "damage": + return tracker.apply_damage(tool_input["target"], tool_input["amount"]) + elif tool_name == "heal": + return tracker.apply_heal(tool_input["target"], tool_input["amount"]) + elif tool_name == "condition": + action = tool_input.get("action", "add") + if action == "remove": + return tracker.remove_condition( + tool_input["target"], tool_input["condition"], + ) + return tracker.add_condition( + target=tool_input["target"], + condition=tool_input["condition"], + duration=tool_input.get("duration", -1), + ends_on=tool_input.get("ends_on", "start"), + source=tool_input.get("source", ""), + ) + elif tool_name == "end_initiative": + return tracker.end() + + return None + + +# --- Tool definitions --- + +_COMBATANT_PROPS: dict = { + "name": {"type": "string", "description": "Combatant name"}, + "initiative": {"type": "integer", "description": "Initiative roll total"}, + "hp": {"type": "integer", "description": "Current hit points"}, + "hp_max": {"type": "integer", "description": "Maximum hit points"}, + "ac": {"type": "integer", "description": "Armor class"}, + "is_player": {"type": "boolean", "description": "True for the player character"}, +} +_COMBATANT_REQUIRED = ["name", "initiative", "hp", "hp_max", "ac"] + +_TARGET_AMOUNT_SCHEMA: dict = { + "type": "object", + "properties": { + "target": {"type": "string", "description": "Combatant name"}, + "amount": {"type": "integer", "description": "Amount"}, + }, + "required": ["target", "amount"], +} + +_NO_INPUT: dict = {"type": "object", "properties": {}} + +ENTER_INITIATIVE_DEFINITION: dict = { + "name": "enter_initiative", + "description": ( + "Enter initiative mode for combat or any turn-based encounter. " + "Provide all participants in their desired turn order (you handle " + "tie-breaking). Roll initiative for everyone first, then call this. " + "Initiative tools become available on the next turn." + ), + "input_schema": { + "type": "object", + "properties": { + "combatants": { + "type": "array", + "description": "Participants in initiative order (first acts first)", + "items": { + "type": "object", + "properties": _COMBATANT_PROPS, + "required": _COMBATANT_REQUIRED, + }, + }, + }, + "required": ["combatants"], + }, +} + +COMBAT_TOOL_DEFINITIONS: list[dict] = [ + {"name": "next_turn", + "description": "Advance to the next combatant's turn. Skips defeated. Call after resolving actions.", + "input_schema": _NO_INPUT}, + {"name": "add_combatant", + "description": "Add a combatant (reinforcements, surprised creatures waking up).", + "input_schema": { + "type": "object", "properties": _COMBATANT_PROPS, + "required": _COMBATANT_REQUIRED}}, + {"name": "remove_combatant", + "description": "Remove a combatant who fled, was banished, or is otherwise out.", + "input_schema": { + "type": "object", + "properties": {"name": {"type": "string", "description": "Combatant name"}}, + "required": ["name"]}}, + {"name": "damage", + "description": "Deal damage. Tracks defeat at 0 HP, reports Bloodied at half. Auto-syncs player character sheet.", + "input_schema": _TARGET_AMOUNT_SCHEMA}, + {"name": "heal", + "description": "Heal a combatant. Clamped to max HP. Revives defeated. Auto-syncs player character sheet.", + "input_schema": _TARGET_AMOUNT_SCHEMA}, + {"name": "condition", + "description": ( + "Add or remove a condition. For timed effects, duration counts down " + "on the source's turn. Duration -1 = until manually removed."), + "input_schema": { + "type": "object", + "properties": { + "target": {"type": "string", "description": "Combatant name"}, + "condition": {"type": "string", "description": "Condition name (Prone, Stunned, etc)"}, + "action": {"type": "string", "enum": ["add", "remove"]}, + "duration": {"type": "integer", "description": "Rounds until expiry (-1 = until removed)"}, + "ends_on": {"type": "string", "enum": ["start", "end"], + "description": "Expires at start or end of source's turn"}, + "source": {"type": "string", "description": "Who caused the condition"}, + }, + "required": ["target", "condition"]}}, + {"name": "end_initiative", + "description": ( + "End initiative and return to narrative. Returns summary with rounds, " + "defeated, and survivor HP. Use the duration in your set_scene call."), + "input_schema": _NO_INPUT}, +] + +ALL_INITIATIVE_TOOL_NAMES: set[str] = ( + {ENTER_INITIATIVE_DEFINITION["name"]} + | {d["name"] for d in COMBAT_TOOL_DEFINITIONS} +) + +INITIATIVE_KEEP_NARRATIVE: frozenset[str] = frozenset({ + "roll", "recall", "update_character", +}) diff --git a/src/storied/mcp_server.py b/src/storied/mcp_server.py index 30bfd6c..c0cdc03 100644 --- a/src/storied/mcp_server.py +++ b/src/storied/mcp_server.py @@ -19,6 +19,11 @@ from mcp.types import TextContent, Tool from storied.log import CampaignLog from storied.search import VectorIndex +from storied.initiative import ( + COMBAT_TOOL_DEFINITIONS, + ENTER_INITIATIVE_DEFINITION, + INITIATIVE_KEEP_NARRATIVE, +) from storied.tools import ( EntityIndex, PLANNER_TOOL_DEFINITIONS, @@ -31,7 +36,6 @@ from storied.tools import ( ) TOOL_SETS: dict[str, list[dict]] = { - "dm": TOOL_DEFINITIONS, "planner": PLANNER_TOOL_DEFINITIONS, "seeder": SEEDER_TOOL_DEFINITIONS, } @@ -43,6 +47,14 @@ EXECUTORS = { } +def _dm_tool_definitions(ctx: ToolContext) -> list[dict]: + """Return DM tools based on whether initiative is active.""" + if ctx.initiative.active: + kept = [d for d in TOOL_DEFINITIONS if d["name"] in INITIATIVE_KEEP_NARRATIVE] + return kept + COMBAT_TOOL_DEFINITIONS + return TOOL_DEFINITIONS + [ENTER_INITIATIVE_DEFINITION] + + def _to_mcp_tool(defn: dict) -> Tool: """Convert an Anthropic-format tool definition to an MCP Tool.""" return Tool( @@ -112,18 +124,26 @@ def start_server( vector_index=VectorIndex(world_dir / "search.db", on_empty=_populate_index), ) - definitions = TOOL_SETS.get(tool_set, TOOL_DEFINITIONS) + static_definitions = TOOL_SETS.get(tool_set) executor = EXECUTORS.get(tool_set, execute_tool) - mcp_tools = [_to_mcp_tool(d) for d in definitions] mcp = Server("storied") @mcp.list_tools() async def list_tools() -> list[Tool]: - return mcp_tools + if tool_set == "dm": + return [_to_mcp_tool(d) for d in _dm_tool_definitions(ctx)] + return [_to_mcp_tool(d) for d in (static_definitions or TOOL_DEFINITIONS)] @mcp.call_tool() async def call_tool(name: str, arguments: dict) -> list[TextContent]: + if tool_set == "dm": + allowed = {d["name"] for d in _dm_tool_definitions(ctx)} + if name not in allowed: + return [TextContent( + type="text", + text=f"Tool '{name}' not available in current mode.", + )] result = executor(name, arguments, ctx) return [TextContent(type="text", text=str(result))] diff --git a/src/storied/tools.py b/src/storied/tools.py index 25c3853..7c82a06 100644 --- a/src/storied/tools.py +++ b/src/storied/tools.py @@ -6,12 +6,17 @@ the tool descriptions that Claude sees. import re import threading -from dataclasses import dataclass +from dataclasses import dataclass, field from pathlib import Path import yaml from storied.character import create_character as char_create +from storied.initiative import ( + ALL_INITIATIVE_TOOL_NAMES, + InitiativeTracker, + execute_initiative_tool, +) from storied.character import update_character as char_update from storied.dice import roll as dice_roll from storied.log import CampaignLog @@ -79,6 +84,16 @@ class ToolContext: campaign_log: CampaignLog entity_index: EntityIndex vector_index: VectorIndex + initiative: InitiativeTracker = field(default_factory=InitiativeTracker) + + +def _sync_player_hp(target: str, ctx: ToolContext, result: str) -> str: + """Auto-sync player character sheet when damage/heal targets a player.""" + combatant = ctx.initiative._find(target) + if combatant and combatant.is_player: + char_update(ctx.player_id, {"hp.current": combatant.hp}, ctx.base_path) + result += f" (character sheet synced to {combatant.hp} HP)" + return result def roll(notation: str, reason: str | None = None) -> dict: @@ -979,6 +994,13 @@ TOOL_DEFINITIONS = [ def execute_tool(tool_name: str, tool_input: dict, ctx: ToolContext) -> str: """Execute a tool by name with the given input.""" + if tool_name in ALL_INITIATIVE_TOOL_NAMES: + result = execute_initiative_tool(tool_name, tool_input, ctx.initiative) + if result is not None: + if tool_name in ("damage", "heal"): + result = _sync_player_hp(tool_input["target"], ctx, result) + return result + if tool_name == "roll": result = roll(tool_input["notation"]) rolls_str = ", ".join(str(r) for r in result["rolls"]) diff --git a/tests/test_engine.py b/tests/test_engine.py index 0208eb7..ed0a5e6 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -75,6 +75,7 @@ class TestDMEngineContext: (prompts_dir / "dm-system.md").write_text("You are a DM.") with patch("storied.engine.start_mcp_server") as mock_mcp: + from storied.initiative import InitiativeTracker from storied.tools import EntityIndex mock_mcp.return_value = type("Handle", (), { @@ -82,6 +83,7 @@ class TestDMEngineContext: "ctx": type("Ctx", (), { "entity_index": EntityIndex(world_dir), "vector_index": None, + "initiative": InitiativeTracker(), })(), })() return DMEngine( diff --git a/tests/test_execute_tool.py b/tests/test_execute_tool.py index 9f1cfcb..5ee90ef 100644 --- a/tests/test_execute_tool.py +++ b/tests/test_execute_tool.py @@ -1,5 +1,6 @@ """Tests for execute_tool dispatch and uncovered tool functions.""" +from storied.initiative import Combatant from storied.tools import ( ToolContext, _auto_mark_present, @@ -234,6 +235,104 @@ class TestAutoMarkPresent: assert marked == [] +class TestInitiativeViaExecuteTool: + """Tests that initiative tools route through execute_tool.""" + + def test_enter_initiative(self, ctx: ToolContext): + result = execute_tool("enter_initiative", { + "combatants": [ + {"name": "Kira", "initiative": 18, "hp": 25, "hp_max": 25, "ac": 16, "is_player": True}, + {"name": "Goblin", "initiative": 10, "hp": 7, "hp_max": 7, "ac": 15}, + ], + }, ctx) + + assert "Initiative started" in result + assert ctx.initiative.active + + def test_damage_via_execute_tool(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Goblin", initiative=10, hp=7, hp_max=7, ac=15), + ]) + + result = execute_tool("damage", {"target": "Goblin", "amount": 3}, ctx) + + assert "3" in result + assert ctx.initiative._find("Goblin").hp == 4 + + def test_damage_syncs_player_hp(self, ctx: ToolContext): + execute_tool("create_character", { + "name": "Kira", "race": "Human", "char_class": "Fighter", + "level": 1, "abilities": { + "strength": 16, "dexterity": 12, "constitution": 14, + "intelligence": 10, "wisdom": 13, "charisma": 8, + }, + "hp_max": 25, "ac": 16, + }, ctx) + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16, is_player=True), + ]) + + result = execute_tool("damage", {"target": "Kira", "amount": 7}, ctx) + + assert "synced" in result + from storied.character import load_character + char = load_character(ctx.player_id, ctx.base_path) + assert char["hp"]["current"] == 18 + + def test_heal_syncs_player_hp(self, ctx: ToolContext): + execute_tool("create_character", { + "name": "Kira", "race": "Human", "char_class": "Fighter", + "level": 1, "abilities": { + "strength": 16, "dexterity": 12, "constitution": 14, + "intelligence": 10, "wisdom": 13, "charisma": 8, + }, + "hp_max": 25, "ac": 16, + }, ctx) + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=20, hp_max=25, ac=16, is_player=True), + ]) + + result = execute_tool("heal", {"target": "Kira", "amount": 3}, ctx) + + assert "synced" in result + from storied.character import load_character + char = load_character(ctx.player_id, ctx.base_path) + assert char["hp"]["current"] == 23 + + def test_damage_no_sync_for_non_player(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Goblin", initiative=10, hp=7, hp_max=7, ac=15), + ]) + + result = execute_tool("damage", {"target": "Goblin", "amount": 3}, ctx) + + assert "synced" not in result + + def test_end_initiative_via_execute_tool(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Goblin", initiative=10, hp=7, hp_max=7, ac=15), + ]) + + result = execute_tool("end_initiative", {}, ctx) + + assert "ended" in result.lower() + assert not ctx.initiative.active + + def test_initiative_tools_not_in_planner(self, ctx: ToolContext): + result = planner_execute_tool("enter_initiative", { + "combatants": [], + }, ctx) + + assert "not available" in result + + def test_initiative_tools_not_in_seeder(self, ctx: ToolContext): + result = seeder_execute_tool("damage", { + "target": "Goblin", "amount": 5, + }, ctx) + + assert "not available" in result + + class TestSeederExecuteTool: """Tests for seeder tool restriction.""" diff --git a/tests/test_initiative.py b/tests/test_initiative.py new file mode 100644 index 0000000..88b3a62 --- /dev/null +++ b/tests/test_initiative.py @@ -0,0 +1,507 @@ +"""Tests for initiative tracking system.""" + +import pytest + +from storied.initiative import ( + ALL_INITIATIVE_TOOL_NAMES, + COMBAT_TOOL_DEFINITIONS, + ENTER_INITIATIVE_DEFINITION, + INITIATIVE_KEEP_NARRATIVE, + Combatant, + InitiativeTracker, + TrackedCondition, + execute_initiative_tool, +) + + +@pytest.fixture +def tracker() -> InitiativeTracker: + return InitiativeTracker() + + +@pytest.fixture +def combatants() -> list[Combatant]: + """Three combatants in initiative order (pre-sorted by DM).""" + return [ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16, is_player=True), + Combatant(name="Goblin 1", initiative=14, hp=7, hp_max=7, ac=15), + Combatant(name="Goblin 2", initiative=10, hp=7, hp_max=7, ac=15), + ] + + +class TestTrackerLifecycle: + def test_starts_inactive(self, tracker: InitiativeTracker): + assert not tracker.active + + def test_begin_activates(self, tracker: InitiativeTracker, combatants: list[Combatant]): + tracker.begin(combatants) + + assert tracker.active + assert tracker.round == 1 + assert tracker.current_index == 0 + + def test_begin_preserves_list_order( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + + assert [c.name for c in tracker.combatants] == ["Kira", "Goblin 1", "Goblin 2"] + + def test_end_deactivates(self, tracker: InitiativeTracker, combatants: list[Combatant]): + tracker.begin(combatants) + summary = tracker.end() + + assert not tracker.active + assert "1" in summary # round count + + def test_end_reports_defeated( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 7) + summary = tracker.end() + + assert "Goblin 1" in summary + assert "defeated" in summary.lower() + + def test_end_reports_duration( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.next_turn() + tracker.next_turn() + tracker.next_turn() # back to round 2 + summary = tracker.end() + + assert "2" in summary # round 2 + + +class TestTurnAdvancement: + def test_next_turn_advances( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + assert tracker.current_combatant.name == "Kira" + + result = tracker.next_turn() + + assert tracker.current_combatant.name == "Goblin 1" + assert "Goblin 1" in result + + def test_round_wraps(self, tracker: InitiativeTracker, combatants: list[Combatant]): + tracker.begin(combatants) + tracker.next_turn() # -> Goblin 1 + tracker.next_turn() # -> Goblin 2 + result = tracker.next_turn() # -> Kira, round 2 + + assert tracker.round == 2 + assert tracker.current_combatant.name == "Kira" + assert "Round 2" in result + + def test_skips_defeated( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 7) # defeat Goblin 1 + tracker.next_turn() # should skip Goblin 1 -> Goblin 2 + + assert tracker.current_combatant.name == "Goblin 2" + + def test_hints_one_side_remaining( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 7) + tracker.apply_damage("Goblin 2", 7) + result = tracker.next_turn() # wraps to Kira + + assert "only" in result.lower() or "remaining" in result.lower() + + +class TestDamageAndHealing: + def test_damage_reduces_hp( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.apply_damage("Goblin 1", 3) + + goblin = tracker._find("Goblin 1") + assert goblin.hp == 4 + assert "3" in result # damage amount + + def test_damage_defeats_at_zero( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.apply_damage("Goblin 1", 7) + + goblin = tracker._find("Goblin 1") + assert goblin.hp == 0 + assert goblin.defeated + assert "down" in result.lower() + + def test_damage_clamps_to_zero( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 100) + + assert tracker._find("Goblin 1").hp == 0 + + def test_damage_reports_bloodied( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.apply_damage("Kira", 13) # 25 -> 12, half is 12 + + assert "bloodied" in result.lower() + + def test_heal_increases_hp( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Kira", 10) + result = tracker.apply_heal("Kira", 5) + + assert tracker._find("Kira").hp == 20 + assert "5" in result + + def test_heal_clamps_to_max( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Kira", 5) + tracker.apply_heal("Kira", 100) + + assert tracker._find("Kira").hp == 25 + + def test_heal_revives_defeated( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 7) + assert tracker._find("Goblin 1").defeated + + tracker.apply_heal("Goblin 1", 3) + + goblin = tracker._find("Goblin 1") + assert goblin.hp == 3 + assert not goblin.defeated + + def test_damage_unknown_target( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.apply_damage("Nobody", 5) + + assert "not found" in result.lower() + + def test_heal_unknown_target( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.apply_heal("Nobody", 5) + + assert "not found" in result.lower() + + +class TestConditions: + def test_add_condition( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.add_condition("Goblin 1", "Prone", duration=-1, source="Kira") + + goblin = tracker._find("Goblin 1") + assert len(goblin.conditions) == 1 + assert goblin.conditions[0].name == "Prone" + assert "Prone" in result + + def test_remove_condition( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.add_condition("Goblin 1", "Prone", duration=-1, source="Kira") + result = tracker.remove_condition("Goblin 1", "Prone") + + goblin = tracker._find("Goblin 1") + assert len(goblin.conditions) == 0 + assert "Prone" in result + + def test_remove_nonexistent_condition( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.remove_condition("Goblin 1", "Invisible") + + assert "not found" in result.lower() or "no" in result.lower() + + def test_condition_unknown_target( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.add_condition("Nobody", "Prone", duration=-1, source="Kira") + + assert "not found" in result.lower() + + def test_effect_expires_end_of_source_turn( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + """Effect with ends_on='end' expires when leaving source's turn.""" + tracker.begin(combatants) + # Kira (idx 0) applies 1-round effect on Goblin 1, ends at end of Kira's turn + tracker.add_condition( + "Goblin 1", "Stunned", duration=1, ends_on="end", source="Kira", + ) + + # Advance from Kira's turn -> processes end-of-Kira effects + tracker.next_turn() # Kira -> Goblin 1 + + goblin = tracker._find("Goblin 1") + assert not any(c.name == "Stunned" for c in goblin.conditions) + + def test_effect_expires_start_of_source_turn( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + """Effect with ends_on='start' expires when arriving at source's turn.""" + tracker.begin(combatants) + # Kira applies 1-round effect, ends at start of Kira's next turn + tracker.add_condition( + "Goblin 1", "Frightened", duration=1, ends_on="start", source="Kira", + ) + + # Full round: Kira -> G1 -> G2 -> Kira (round 2, start of Kira's turn) + tracker.next_turn() # -> Goblin 1 + tracker.next_turn() # -> Goblin 2 + tracker.next_turn() # -> Kira (round 2) + + goblin = tracker._find("Goblin 1") + assert not any(c.name == "Frightened" for c in goblin.conditions) + + def test_indefinite_condition_persists( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + """Duration -1 never auto-expires.""" + tracker.begin(combatants) + tracker.add_condition("Goblin 1", "Grappled", duration=-1, source="Kira") + + # Full round + for _ in range(3): + tracker.next_turn() + + goblin = tracker._find("Goblin 1") + assert any(c.name == "Grappled" for c in goblin.conditions) + + def test_multi_round_duration( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + """2-round effect lasts through 2 full rounds of the source's turns.""" + tracker.begin(combatants) + tracker.add_condition( + "Goblin 1", "Held", duration=2, ends_on="end", source="Kira", + ) + + # Round 1: Kira -> G1 -> G2 (end of Kira's turn, duration 2 -> 1) + tracker.next_turn() # -> G1 + goblin = tracker._find("Goblin 1") + assert any(c.name == "Held" for c in goblin.conditions) + + tracker.next_turn() # -> G2 + tracker.next_turn() # -> Kira (round 2) + + # Round 2: leaving Kira's turn (duration 1 -> 0, expires) + tracker.next_turn() # -> G1 + + goblin = tracker._find("Goblin 1") + assert not any(c.name == "Held" for c in goblin.conditions) + + +class TestAddRemoveCombatant: + def test_add_combatant( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.add_combatant( + Combatant(name="Archer", initiative=12, hp=9, hp_max=9, ac=13), + ) + + assert len(tracker.combatants) == 4 + assert "Archer" in result + + def test_add_inserts_by_initiative( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) # Kira(18), G1(14), G2(10) + tracker.add_combatant( + Combatant(name="Archer", initiative=12, hp=9, hp_max=9, ac=13), + ) + + names = [c.name for c in tracker.combatants] + assert names == ["Kira", "Goblin 1", "Archer", "Goblin 2"] + + def test_add_before_current_adjusts_index( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.next_turn() # -> Goblin 1 (index 1) + + # Add someone with higher initiative (inserted before current) + tracker.add_combatant( + Combatant(name="Archer", initiative=16, hp=9, hp_max=9, ac=13), + ) + + # Current should still be Goblin 1 + assert tracker.current_combatant.name == "Goblin 1" + + def test_remove_combatant( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.remove_combatant("Goblin 2") + + assert len(tracker.combatants) == 2 + assert "Goblin 2" in result + + def test_remove_current_advances( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.next_turn() # -> Goblin 1 + tracker.remove_combatant("Goblin 1") + + # Should advance to Goblin 2 (or adjust so current is valid) + assert tracker.current_combatant.name == "Goblin 2" + + def test_remove_unknown( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + result = tracker.remove_combatant("Nobody") + + assert "not found" in result.lower() + + +class TestFormatForContext: + def test_includes_table( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + context = tracker.format_for_context() + + assert "Kira" in context + assert "Goblin 1" in context + assert "25/25" in context # HP display + + def test_marks_current_turn( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + context = tracker.format_for_context() + + assert "**Kira**" in context # current combatant is bold + assert "Current turn" in context + + def test_shows_defeated( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.apply_damage("Goblin 1", 7) + context = tracker.format_for_context() + + assert "~~Goblin 1~~" in context or "Defeated" in context + + def test_shows_conditions( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + tracker.add_condition("Goblin 1", "Prone", duration=-1, source="Kira") + context = tracker.format_for_context() + + assert "Prone" in context + + def test_shows_up_next( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + context = tracker.format_for_context() + + assert "Up next" in context + assert "Goblin 1" in context.split("Up next")[1] + + def test_shows_round( + self, tracker: InitiativeTracker, combatants: list[Combatant], + ): + tracker.begin(combatants) + context = tracker.format_for_context() + + assert "Round" in context + + +class TestDispatch: + def test_all_tool_names_present(self): + expected = { + "enter_initiative", "next_turn", "add_combatant", + "remove_combatant", "damage", "heal", "condition", + "end_initiative", + } + assert ALL_INITIATIVE_TOOL_NAMES == expected + + def test_enter_initiative_definition_exists(self): + assert ENTER_INITIATIVE_DEFINITION["name"] == "enter_initiative" + + def test_combat_definitions_exclude_enter(self): + names = {d["name"] for d in COMBAT_TOOL_DEFINITIONS} + assert "enter_initiative" not in names + assert "next_turn" in names + + def test_keep_narrative_tools(self): + assert INITIATIVE_KEEP_NARRATIVE == frozenset({"roll", "recall", "update_character"}) + + def test_dispatch_routes_enter(self, tracker: InitiativeTracker): + result = execute_initiative_tool( + "enter_initiative", + {"combatants": [ + {"name": "Kira", "initiative": 18, "hp": 25, "hp_max": 25, "ac": 16, "is_player": True}, + {"name": "Goblin", "initiative": 10, "hp": 7, "hp_max": 7, "ac": 15}, + ]}, + tracker, + ) + + assert tracker.active + assert "Kira" in result + + def test_dispatch_routes_damage(self, tracker: InitiativeTracker): + tracker.begin([ + Combatant(name="Goblin", initiative=10, hp=7, hp_max=7, ac=15), + ]) + + result = execute_initiative_tool( + "damage", {"target": "Goblin", "amount": 3}, tracker, + ) + + assert "3" in result + + def test_dispatch_unknown_tool(self, tracker: InitiativeTracker): + result = execute_initiative_tool("fake_tool", {}, tracker) + + assert result is None + + def test_guard_when_inactive(self, tracker: InitiativeTracker): + result = execute_initiative_tool( + "next_turn", {}, tracker, + ) + + assert "not active" in result.lower() + + def test_enter_errors_when_active(self, tracker: InitiativeTracker): + tracker.begin([ + Combatant(name="Goblin", initiative=10, hp=7, hp_max=7, ac=15), + ]) + + result = execute_initiative_tool( + "enter_initiative", + {"combatants": [{"name": "X", "initiative": 1, "hp": 1, "hp_max": 1, "ac": 10}]}, + tracker, + ) + + assert "already active" in result.lower() diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index a1a9d53..d1c6e0a 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -2,8 +2,14 @@ import pytest -from storied.mcp_server import _to_mcp_tool -from storied.tools import TOOL_DEFINITIONS +from storied.initiative import ( + COMBAT_TOOL_DEFINITIONS, + ENTER_INITIATIVE_DEFINITION, + INITIATIVE_KEEP_NARRATIVE, + Combatant, +) +from storied.mcp_server import _dm_tool_definitions, _to_mcp_tool +from storied.tools import TOOL_DEFINITIONS, ToolContext class TestToMcpTool: @@ -31,3 +37,79 @@ class TestToMcpTool: tool = _to_mcp_tool(defn) assert tool.name == defn["name"] assert tool.inputSchema is not None + + +class TestDynamicDmTools: + """Tests that the DM tool list changes based on initiative state.""" + + def test_narrative_mode_includes_enter_initiative(self, ctx: ToolContext): + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + assert "enter_initiative" in names + + def test_narrative_mode_includes_all_narrative_tools(self, ctx: ToolContext): + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + assert "set_scene" in names + assert "establish" in names + assert "end_session" in names + + def test_narrative_mode_excludes_combat_tools(self, ctx: ToolContext): + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + assert "next_turn" not in names + assert "damage" not in names + assert "end_initiative" not in names + + def test_narrative_mode_count(self, ctx: ToolContext): + defs = _dm_tool_definitions(ctx) + + assert len(defs) == len(TOOL_DEFINITIONS) + 1 # +1 for enter_initiative + + def test_initiative_mode_includes_combat_tools(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16), + ]) + + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + assert "next_turn" in names + assert "damage" in names + assert "end_initiative" in names + + def test_initiative_mode_keeps_narrative_subset(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16), + ]) + + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + for tool_name in INITIATIVE_KEEP_NARRATIVE: + assert tool_name in names + + def test_initiative_mode_excludes_narrative_only_tools(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16), + ]) + + defs = _dm_tool_definitions(ctx) + names = {d["name"] for d in defs} + + assert "set_scene" not in names + assert "establish" not in names + assert "enter_initiative" not in names + + def test_initiative_mode_count(self, ctx: ToolContext): + ctx.initiative.begin([ + Combatant(name="Kira", initiative=18, hp=25, hp_max=25, ac=16), + ]) + + defs = _dm_tool_definitions(ctx) + expected = len(INITIATIVE_KEEP_NARRATIVE) + len(COMBAT_TOOL_DEFINITIONS) + + assert len(defs) == expected -- 2.51.2