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