diff --git a/src/storied/advancement.py b/src/storied/advancement.py index 1240c73..47f040c 100644 --- a/src/storied/advancement.py +++ b/src/storied/advancement.py @@ -10,10 +10,10 @@ from storied import notifications from storied.character import format_character_context, load_character from storied.claude import run_with_tools from storied.engine import load_prompt -from storied.log import CampaignLog from storied.mcp_server import start_server as start_mcp_server from storied.paths import data_home from storied.session import load_session +from storied.tools._context import get_or_create_ctx @dataclass @@ -46,7 +46,7 @@ def build_advancement_context( parts.append(char_context) # Campaign log — entries since last level-up - log = CampaignLog(world_id) + log = get_or_create_ctx(world_id, player_id).campaign_log parts.append(f"## Campaign Time: {log.get_current_time()}") entries_since_level = log.get_entries_since_tag("level") @@ -116,10 +116,7 @@ def evaluate_advancement( system_prompt = load_prompt("xp-evaluator") - campaign_log = CampaignLog(world_id) - mcp = start_mcp_server( - world_id, player_id, "advancement", campaign_log, - ) + mcp = start_mcp_server(world_id, player_id, "advancement") progress(f"Evaluating advancement with {model}...") diff --git a/src/storied/engine.py b/src/storied/engine.py index 7cc1254..00bbdd9 100644 --- a/src/storied/engine.py +++ b/src/storied/engine.py @@ -19,7 +19,7 @@ from storied.claude import ( stream_with_tools, ) from storied.content import ContentResolver -from storied.log import CampaignLog, TranscriptLog +from storied.log import TranscriptLog from storied.mcp_server import start_server as start_mcp_server from storied.notification_formatters import ( DEFERRED_FORMATTERS, @@ -101,20 +101,24 @@ class DMEngine: if transcript_path: transcript_path.parent.mkdir(parents=True, exist_ok=True) - # Campaign log for time tracking (world-scoped, shared with MCP server) - self._campaign_log = CampaignLog(self.world_id) - # Transcript log for conversation history self._transcript = TranscriptLog(self.world_id) - # Start in-process MCP server (shares CampaignLog with engine) + # Start in-process MCP server. start_server resolves the singleton + # ToolContext (constructing it on first call), so seed_world, + # plot_arc, the engine, and the background ticker/advancement + # threads all share the same CampaignLog, EntityIndex, VectorIndex, + # and InitiativeTracker. self._mcp = start_mcp_server( world_id=self.world_id, player_id=self.player_id, tool_set="dm", - campaign_log=self._campaign_log, ) + # Read the campaign log straight off the singleton so display + # reads always see the same instance the tools mutate. + self._campaign_log = self._mcp.ctx.campaign_log + # Build system prompt with full context self._prompt_name = prompt_name self._base_prompt = load_prompt(prompt_name) diff --git a/src/storied/initiative.py b/src/storied/initiative.py index fa98548..a5d02d1 100644 --- a/src/storied/initiative.py +++ b/src/storied/initiative.py @@ -7,6 +7,7 @@ storied.tools._context, hence the split — keeping it here would create a circular import). """ +import threading from dataclasses import dataclass, field @@ -50,94 +51,114 @@ class InitiativeTracker: self.combatants: list[Combatant] = [] self.current_index: int = 0 self.round: int = 0 + # Single per-instance lock guarding all combatant-state mutations. + # The tracker is a process-wide singleton accessed from the engine + # main thread (display reads) and from MCP server threads (tool + # mutations), so naive iteration over self.combatants can race + # with insert/remove. RLock so format_for_context can be called + # from inside other locked methods if that ever comes up. + self._lock = threading.RLock() @property def current_combatant(self) -> Combatant | None: - if not self.active or not self.combatants: - return None - return self.combatants[self.current_index] + with self._lock: + 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) + with self._lock: + 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." + with self._lock: + 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") + # 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) + # 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") + 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 = [] + 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})" - ) + 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}") + 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) + # Check if only one side remains + hint = self._one_side_hint() + if hint: + parts.append(hint) - return "\n".join(parts) + 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 + with self._lock: + 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 " + f"({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})" + with self._lock: + 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 " + f"({old_hp} \u2192 {combatant.hp}/{combatant.hp_max})" + ) def add_condition( self, @@ -148,146 +169,168 @@ class InitiativeTracker: source: str = "", ) -> str: """Apply a condition to a combatant.""" - combatant = self._find(target) - if combatant is None: - return f"Combatant '{target}' not found." + with self._lock: + 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) + 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}" + 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}." + with self._lock: + 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 + with self._lock: + 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) + self.combatants.insert(insert_idx, combatant) - if insert_idx <= self.current_index: - self.current_index += 1 + if insert_idx <= self.current_index: + self.current_index += 1 - return f"{combatant.name} joins initiative (initiative {combatant.initiative})" + return ( + f"{combatant.name} joins initiative " + f"(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." + with self._lock: + 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 — " + f"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 + with self._lock: + 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 " + f"({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)}") - 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 + self.active = False + self.combatants = [] + self.current_index = 0 + self.round = 0 - return "\n".join(lines) + 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) + with self._lock: + 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} " + f"| {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: diff --git a/src/storied/log.py b/src/storied/log.py index 2d2edcb..596857c 100644 --- a/src/storied/log.py +++ b/src/storied/log.py @@ -3,6 +3,7 @@ from __future__ import annotations import re +import threading from dataclasses import dataclass, field from pathlib import Path @@ -183,6 +184,14 @@ class CampaignLog: self.base_path = base_path or data_home() self.log_dir = self.base_path / "worlds" / world_id / "log" + # Single per-instance lock guarding mutation paths. The campaign + # log is now a process-wide singleton shared across the engine + # thread, the engine's MCP uvicorn thread, and any background + # agents (advancement, ticker), so concurrent set_scene calls + # could otherwise race the read-modify-write on current_time and + # stomp each other's day-file saves. + self._lock = threading.RLock() + # Load or initialize state self._load_state() @@ -337,26 +346,27 @@ class CampaignLog: if isinstance(duration, str): duration = Duration.parse(duration) - anchor = self.current_time.to_anchor() - entry = LogEntry( - anchor=anchor, - event=event, - duration=duration, - tags=tags or [], - ) - self.current_entries.append(entry) + with self._lock: + anchor = self.current_time.to_anchor() + entry = LogEntry( + anchor=anchor, + event=event, + duration=duration, + tags=tags or [], + ) + self.current_entries.append(entry) - if advance_time: - self.current_time = self.current_time.add_duration(duration) + if advance_time: + self.current_time = self.current_time.add_duration(duration) - # Check if we crossed into a new day - if self.current_time.day > self.current_day: - self._roll_day() + # Check if we crossed into a new day + if self.current_time.day > self.current_day: + self._roll_day() - # Save current day entries - self._save_day_file(self.current_day, self.current_entries) - self._save_index() - return anchor + # Save current day entries + self._save_day_file(self.current_day, self.current_entries) + self._save_index() + return anchor def _roll_day(self) -> None: """Archive current day and start a new one.""" @@ -386,11 +396,12 @@ class CampaignLog: days, this returns every entry — useful for scanning the log for casually mentioned entities. """ - entries: list[LogEntry] = [] - start_day = max(1, self.current_day - days + 1) - for day in range(start_day, self.current_day + 1): - entries.extend(self._load_day_entries(day)) - return entries + with self._lock: + entries: list[LogEntry] = [] + start_day = max(1, self.current_day - days + 1) + for day in range(start_day, self.current_day + 1): + entries.extend(self._load_day_entries(day)) + return entries def format_for_context(self) -> str: """Format the log for inclusion in system prompt. @@ -399,33 +410,37 @@ class CampaignLog: top of the DM's context — see ``DMEngine._format_time_header``. This block is just the recent event history. """ - lines: list[str] = [] - - if self.previous_summaries: - lines.append("## Campaign Log") - lines.append("") - lines.append("**Previous Days:**") - for summary in self.previous_summaries[-3:]: # Last 3 days - lines.append(f"- {summary}") + with self._lock: + lines: list[str] = [] - if self.current_entries: - if not lines: + if self.previous_summaries: lines.append("## Campaign Log") - lines.append("") - lines.append(f"**Today (Day {self.current_day}):**") - if len(self.current_entries) > 10: - lines.append(f"({len(self.current_entries) - 10} earlier entries today)") - for entry in self.current_entries[-10:]: - lines.append(f"- {entry.event}") + lines.append("") + lines.append("**Previous Days:**") + for summary in self.previous_summaries[-3:]: # Last 3 days + lines.append(f"- {summary}") - return "\n".join(lines) + if self.current_entries: + if not lines: + lines.append("## Campaign Log") + lines.append("") + lines.append(f"**Today (Day {self.current_day}):**") + if len(self.current_entries) > 10: + lines.append( + f"({len(self.current_entries) - 10} earlier entries today)" + ) + for entry in self.current_entries[-10:]: + lines.append(f"- {entry.event}") + + return "\n".join(lines) def get_all_entries(self) -> list[LogEntry]: """Get every log entry from day 1 through the current day.""" - entries: list[LogEntry] = [] - for day in range(1, self.current_day + 1): - entries.extend(self._load_day_entries(day)) - return entries + with self._lock: + entries: list[LogEntry] = [] + for day in range(1, self.current_day + 1): + entries.extend(self._load_day_entries(day)) + return entries def get_entries_since_tag(self, tag: str) -> list[LogEntry]: """Get all entries after the last occurrence of a tag. @@ -455,15 +470,16 @@ class CampaignLog: """Calculate time since last rest of given type.""" tag = f"rest:{rest_type}" - # Search backwards through entries - total_minutes = 0 - for entry in reversed(self.current_entries): - if tag in entry.tags: - return Duration(minutes=total_minutes) - total_minutes += entry.duration.total_minutes + with self._lock: + # Search backwards through entries + total_minutes = 0 + for entry in reversed(self.current_entries): + if tag in entry.tags: + return Duration(minutes=total_minutes) + total_minutes += entry.duration.total_minutes - # If not found in current day, it's been longer - return Duration(minutes=total_minutes + 8 * 60) # Add 8 hours as estimate + # If not found in current day, it's been longer + return Duration(minutes=total_minutes + 8 * 60) # +8h estimate def load_log(world_id: str = "default", base_path: Path | None = None) -> CampaignLog: diff --git a/src/storied/mcp_server.py b/src/storied/mcp_server.py index 79f75e7..e278fac 100644 --- a/src/storied/mcp_server.py +++ b/src/storied/mcp_server.py @@ -3,8 +3,10 @@ start_server() launches a FastMCP server (SSE transport) on a free localhost port in a background thread. Each call composes a per-role top-level server by mounting the tools/*.py module-level FastMCP instances and applying -tag-based visibility filters. ToolContext is process-global and accessed by -tools via the Dependency subclasses in storied.tools._context. +tag-based visibility filters. The ToolContext is a process-wide singleton +fetched via :func:`get_or_create_ctx`, so every server in the process — DM, +planner, ticker, advancement, seeder — reads and writes the same in-memory +game state. """ import asyncio @@ -17,13 +19,18 @@ import uvicorn from fastmcp import FastMCP from storied import paths -from storied.log import CampaignLog from storied.search import VectorIndex -from storied.tools import character, combat, entities, mechanics, run_code, scene +from storied.tools import ( + character, + combat, + entities, + mechanics, + run_code, + scene, +) from storied.tools._context import ( - EntityIndex, ToolContext, - init_ctx, + get_or_create_ctx, ) ALL_ROLES = {"dm", "planner", "seeder", "advancement", "arc_architect"} @@ -174,13 +181,15 @@ def start_server( # pragma: no cover world_id: str, player_id: str, tool_set: str = "dm", - campaign_log: CampaignLog | None = None, ) -> MCPServerHandle: """Start an in-process FastMCP HTTP server on a free localhost port. Returns an MCPServerHandle with the URL to pass to --mcp-config. - The server runs in a daemon thread and shares the process-global - ToolContext (set via init_ctx) with the caller. + The server runs in a daemon thread and reads from the singleton + ToolContext (constructed lazily by ``get_or_create_ctx`` on the + first call). Subsequent calls — engine, planner, ticker, advancement, + seeder — all bind to the same context, so set_scene mutations made + via any role are visible to every reader. Paths are resolved via :mod:`storied.paths` (the data home is set once at CLI startup via ``configure()``), so this function takes @@ -192,24 +201,13 @@ def start_server( # pragma: no cover in tests/test_mcp_server.py; mocking out uvicorn here would test the mock, not the launcher. """ - if campaign_log is None: - campaign_log = CampaignLog(world_id) - - world_dir = paths.world_path(world_id) + ctx = get_or_create_ctx(world_id, player_id) - vector_index = VectorIndex(world_dir / "search.db") # Populate eagerly so the first recall never races the transcript # upsert at turn end. `_populate_index` is idempotent — subsequent # calls skip the SRD reseed and mtime-check the user/world layers. - _populate_index(world_dir, vector_index) - - ctx = init_ctx( - world_id=world_id, - player_id=player_id, - campaign_log=campaign_log, - entity_index=EntityIndex(world_dir), - vector_index=vector_index, - ) + world_dir = paths.world_path(world_id) + _populate_index(world_dir, ctx.vector_index) server = asyncio.run(_compose_server(tool_set)) diff --git a/src/storied/planner.py b/src/storied/planner.py index df59dce..f4fbf1d 100644 --- a/src/storied/planner.py +++ b/src/storied/planner.py @@ -12,6 +12,7 @@ from storied.character import format_character_context, load_character from storied.claude import run_prompt, run_with_tools from storied.engine import load_prompt from storied.log import CampaignLog +from storied.tools._context import get_or_create_ctx from storied.mcp_server import start_server as start_mcp_server from storied.paths import data_home, world_path from storied.session import ( @@ -249,7 +250,7 @@ def build_planning_context( parts.append(body) # Campaign log — full recent entries so the planner can spot casual mentions - log = CampaignLog(world_id) + log = get_or_create_ctx(world_id, player_id).campaign_log parts.append(f"## Campaign Time: {log.get_current_time()}") recent = log.get_recent_entries(days=2) @@ -344,10 +345,7 @@ def plan_world( context = build_planning_context(world_id, player_id, candidate_pairs) system_prompt = load_prompt("planner-system") - campaign_log = CampaignLog(world_id) - mcp = start_mcp_server( - world_id, player_id, "planner", campaign_log, - ) + mcp = start_mcp_server(world_id, player_id, "planner") progress(f"Planning with {model}...") @@ -442,10 +440,7 @@ def seed_world( ) system_prompt = load_prompt("world-seed") - campaign_log = CampaignLog(world_id) - mcp = start_mcp_server( - world_id, player_id, "seeder", campaign_log, - ) + mcp = start_mcp_server(world_id, player_id, "seeder") progress(f"Seeding with {model}...") @@ -562,10 +557,7 @@ def plot_arc( + char_block ) - campaign_log = CampaignLog(world_id) - mcp = start_mcp_server( - world_id, player_id, "arc_architect", campaign_log, - ) + mcp = start_mcp_server(world_id, player_id, "arc_architect") def on_tool(name: str) -> None: if on_progress: @@ -636,7 +628,7 @@ def build_tick_context( parts.append(prefs.rstrip()) # Campaign log and time - log = CampaignLog(world_id) + log = get_or_create_ctx(world_id, player_id).campaign_log current_time = log.get_current_time() parts.append(f"## Current Game Time: {current_time}") @@ -704,10 +696,7 @@ def tick_world( # pragma: no cover context = build_tick_context(world_id, player_id, entities) system_prompt = load_prompt("world-tick") - campaign_log = CampaignLog(world_id) - mcp = start_mcp_server( - world_id, player_id, "planner", campaign_log, - ) + mcp = start_mcp_server(world_id, player_id, "planner") progress(f"Ticking with {model}...") diff --git a/src/storied/search.py b/src/storied/search.py index 2ddf91a..d0e10d9 100644 --- a/src/storied/search.py +++ b/src/storied/search.py @@ -7,6 +7,7 @@ The index lives as a single search.db file in the world directory. import re import shutil import struct +import threading from collections.abc import Callable from dataclasses import dataclass from pathlib import Path @@ -154,6 +155,13 @@ class VectorIndex: ): self._db_path = db_path self._embed_fn: Callable[[list[str]], list[list[float]]] = _default_embed + # The shared sqlite connection is opened with check_same_thread=False + # so it can be used from any uvicorn worker thread, but a single + # connection still serializes its statements through one cursor — + # concurrent execute calls from different MCP server threads (DM, + # planner, advancement) can race on commit boundaries. RLock so + # ``reindex_directory`` can call ``upsert`` re-entrantly. + self._lock = threading.RLock() self._conn = self._open_or_recreate() @staticmethod @@ -168,21 +176,24 @@ class VectorIndex: def reseed(self, seed_path: Path) -> None: """Replace this index's DB with a copy of the seed and reconnect.""" - self._conn.close() - shutil.copy2(seed_path, self._db_path) - self._conn = self._open_or_recreate() + with self._lock: + self._conn.close() + shutil.copy2(seed_path, self._db_path) + self._conn = self._open_or_recreate() def has_source(self, source: str) -> bool: """True if at least one document is already indexed from ``source``.""" - row = self._conn.execute( - "SELECT 1 FROM documents WHERE source = ? LIMIT 1", - (source,), - ).fetchone() - return row is not None + with self._lock: + row = self._conn.execute( + "SELECT 1 FROM documents WHERE source = ? LIMIT 1", + (source,), + ).fetchone() + return row is not None def close(self) -> None: """Close the database connection.""" - self._conn.close() + with self._lock: + self._conn.close() def _open_or_recreate(self) -> sqlite3.Connection: """Open the database, recreating if corrupt.""" @@ -232,45 +243,47 @@ class VectorIndex: blob = _serialize_f32(vec) preview = text[:200].strip() - self._conn.execute( - "DELETE FROM vec_documents WHERE doc_id = ?", (doc_id,) - ) - self._conn.execute( - "DELETE FROM documents WHERE doc_id = ?", (doc_id,) - ) - - self._conn.execute( - """INSERT INTO documents - (doc_id, path, source, content_type, chunk_index, - title, body_preview, game_day, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""", - ( - doc_id, - metadata.get("path", ""), - metadata["source"], - metadata.get("content_type"), - metadata.get("chunk_index", 0), - metadata.get("title"), - preview, - metadata.get("game_day"), - metadata.get("updated_at", 0.0), - ), - ) - self._conn.execute( - "INSERT INTO vec_documents (doc_id, embedding) VALUES (?, ?)", - (doc_id, blob), - ) - self._conn.commit() + with self._lock: + self._conn.execute( + "DELETE FROM vec_documents WHERE doc_id = ?", (doc_id,) + ) + self._conn.execute( + "DELETE FROM documents WHERE doc_id = ?", (doc_id,) + ) + + self._conn.execute( + """INSERT INTO documents + (doc_id, path, source, content_type, chunk_index, + title, body_preview, game_day, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ( + doc_id, + metadata.get("path", ""), + metadata["source"], + metadata.get("content_type"), + metadata.get("chunk_index", 0), + metadata.get("title"), + preview, + metadata.get("game_day"), + metadata.get("updated_at", 0.0), + ), + ) + self._conn.execute( + "INSERT INTO vec_documents (doc_id, embedding) VALUES (?, ?)", + (doc_id, blob), + ) + self._conn.commit() def delete(self, doc_id: str) -> None: """Remove a document and its embedding.""" - self._conn.execute( - "DELETE FROM vec_documents WHERE doc_id = ?", (doc_id,) - ) - self._conn.execute( - "DELETE FROM documents WHERE doc_id = ?", (doc_id,) - ) - self._conn.commit() + with self._lock: + self._conn.execute( + "DELETE FROM vec_documents WHERE doc_id = ?", (doc_id,) + ) + self._conn.execute( + "DELETE FROM documents WHERE doc_id = ?", (doc_id,) + ) + self._conn.commit() def search( self, @@ -309,16 +322,17 @@ class VectorIndex: fetch_limit = limit * 3 - rows = self._conn.execute( - """SELECT v.doc_id, v.distance, d.path, d.source, - d.content_type, d.body_preview, d.game_day - FROM vec_documents v - JOIN documents d ON v.doc_id = d.doc_id - WHERE v.embedding MATCH ? - AND k = ? - ORDER BY v.distance""", - (blob, fetch_limit), - ).fetchall() + with self._lock: + rows = self._conn.execute( + """SELECT v.doc_id, v.distance, d.path, d.source, + d.content_type, d.body_preview, d.game_day + FROM vec_documents v + JOIN documents d ON v.doc_id = d.doc_id + WHERE v.embedding MATCH ? + AND k = ? + ORDER BY v.distance""", + (blob, fetch_limit), + ).fetchall() hits: list[SearchHit] = [] for doc_id, distance, path, source, ctype, preview, game_day in rows: @@ -364,71 +378,75 @@ class VectorIndex: ``transcripts/`` — those files are indexed separately by the engine under ``source="transcript"``. """ - existing = {} - for row in self._conn.execute( - "SELECT doc_id, updated_at FROM documents WHERE source = ?", - (source,), - ): - existing[row[0]] = row[1] - - seen_doc_ids: set[str] = set() - count = 0 - - for md_file in sorted(directory.rglob("*.md")): - rel = md_file.relative_to(directory) - if skip_subdirs and rel.parts and rel.parts[0] in skip_subdirs: - continue - content = md_file.read_text() - mtime = md_file.stat().st_mtime - content_type = rel.parts[0] if len(rel.parts) > 1 else "" - - chunks = chunk_document(md_file, content) + with self._lock: + existing = {} + for row in self._conn.execute( + "SELECT doc_id, updated_at FROM documents WHERE source = ?", + (source,), + ): + existing[row[0]] = row[1] - for chunk_idx, chunk_text in chunks: - doc_id = f"{source}:{rel}:{chunk_idx}" - seen_doc_ids.add(doc_id) + seen_doc_ids: set[str] = set() + count = 0 - if doc_id in existing and existing[doc_id] == mtime: - count += 1 + for md_file in sorted(directory.rglob("*.md")): + rel = md_file.relative_to(directory) + if skip_subdirs and rel.parts and rel.parts[0] in skip_subdirs: continue + content = md_file.read_text() + mtime = md_file.stat().st_mtime + content_type = rel.parts[0] if len(rel.parts) > 1 else "" + + chunks = chunk_document(md_file, content) + + for chunk_idx, chunk_text in chunks: + doc_id = f"{source}:{rel}:{chunk_idx}" + seen_doc_ids.add(doc_id) + + if doc_id in existing and existing[doc_id] == mtime: + count += 1 + continue + + title_match = re.match(r"^#\s+(.+)", content) + title = ( + title_match.group(1).strip() + if title_match + else md_file.stem + ) + + game_day = None + day_match = re.match(r"day([+-]\d+)", md_file.stem) + if day_match: + game_day = int(day_match.group(1)) + + self.upsert(doc_id, chunk_text, { + "source": source, + "content_type": content_type, + "path": str(md_file), + "title": title, + "chunk_index": chunk_idx, + "game_day": game_day, + "updated_at": mtime, + }) + count += 1 - title_match = re.match(r"^#\s+(.+)", content) - title = ( - title_match.group(1).strip() if title_match else md_file.stem - ) - - game_day = None - day_match = re.match(r"day([+-]\d+)", md_file.stem) - if day_match: - game_day = int(day_match.group(1)) - - self.upsert(doc_id, chunk_text, { - "source": source, - "content_type": content_type, - "path": str(md_file), - "title": title, - "chunk_index": chunk_idx, - "game_day": game_day, - "updated_at": mtime, - }) - count += 1 - - stale = set(existing.keys()) - seen_doc_ids - for doc_id in stale: - self.delete(doc_id) + stale = set(existing.keys()) - seen_doc_ids + for doc_id in stale: + self.delete(doc_id) - return count + return count def stats(self) -> dict: """Return index statistics.""" - total = self._conn.execute( - "SELECT count(*) FROM documents" - ).fetchone()[0] - - by_source: dict[str, int] = {} - for source, cnt in self._conn.execute( - "SELECT source, count(*) FROM documents GROUP BY source" - ): - by_source[source] = cnt + with self._lock: + total = self._conn.execute( + "SELECT count(*) FROM documents" + ).fetchone()[0] + + by_source: dict[str, int] = {} + for source, cnt in self._conn.execute( + "SELECT source, count(*) FROM documents GROUP BY source" + ): + by_source[source] = cnt - return {"total_documents": total, "by_source": by_source} + return {"total_documents": total, "by_source": by_source} diff --git a/src/storied/tools/__init__.py b/src/storied/tools/__init__.py index 35b5cb5..c81f176 100644 --- a/src/storied/tools/__init__.py +++ b/src/storied/tools/__init__.py @@ -18,6 +18,8 @@ from storied.tools._context import ( World, _get_file_lock, _sync_player_hp, + current_ctx, + get_or_create_ctx, init_ctx, reset_ctx, ) @@ -35,6 +37,8 @@ __all__ = [ "_get_file_lock", "_load_entity", "_sync_player_hp", + "current_ctx", + "get_or_create_ctx", "init_ctx", "reset_ctx", ] diff --git a/src/storied/tools/_context.py b/src/storied/tools/_context.py index 91c0205..195bfd0 100644 --- a/src/storied/tools/_context.py +++ b/src/storied/tools/_context.py @@ -62,9 +62,12 @@ class EntityIndex: class ToolContext: """Shared infrastructure for all tool calls. - Created once per process via init_ctx() and exposed to tools through - the Dependency subclasses below. Tools should never reach for the - full ToolContext — they ask for the specific slices they need. + Genuine singleton — one per process, shared by every FastMCP server + (DM, planner, seeder, advancement, arc_architect) so they all read + and write the same in-memory game state. The DMEngine keeps a + reference to this same instance for its display reads. Tools never + reach for the full ToolContext; they ask for the slice they need + via the Dependency subclasses below. Filesystem paths live in :mod:`storied.paths` (module globals, configured at CLI startup), not on the ToolContext. The context @@ -80,8 +83,17 @@ class ToolContext: # --- Process-global ToolContext --------------------------------------------- +# +# Storied serves exactly one campaign at a time, so we collapse the per-MCP +# ToolContext into a true process-wide singleton. Multiple FastMCP servers +# (engine, planner, ticker, advancement, seeder) all retrieve the same +# instance via ``get_or_create_ctx``, so set_scene mutations made in any +# server are visible to every reader. The earlier "init on each start_server" +# pattern caused background agents to silently swap in fresh CampaignLog / +# EntityIndex / VectorIndex copies, leaving the engine reading a stale view. _ctx: ToolContext | None = None +_ctx_lock = threading.Lock() def init_ctx( @@ -91,30 +103,78 @@ def init_ctx( entity_index: EntityIndex, vector_index: VectorIndex, ) -> ToolContext: - """Initialize the process-global ToolContext. + """Set the process-global ToolContext directly. - Idempotent — last writer wins. Tests use this to reset state between cases. + Test-only entry point. Production code should call + :func:`get_or_create_ctx` instead — it constructs the slices itself + and refuses to clobber an existing context. Tests pair this with + :func:`reset_ctx` in fixture teardown so each test starts clean. + """ + global _ctx + with _ctx_lock: + _ctx = ToolContext( + world_id=world_id, + player_id=player_id, + campaign_log=campaign_log, + entity_index=entity_index, + vector_index=vector_index, + ) + return _ctx + + +def get_or_create_ctx(world_id: str, player_id: str) -> ToolContext: + """Return the singleton ToolContext, constructing it on first call. + + Used by every production code path that needs to attach an MCP + server or background agent to the live game state. The first call + builds the CampaignLog, EntityIndex, and VectorIndex from disk; + subsequent callers get the same instance back. Mismatched + world_id/player_id raises rather than silently rebinding, since the + process is committed to one campaign. """ + from storied import paths + from storied.search import VectorIndex as _VectorIndex + global _ctx - _ctx = ToolContext( - world_id=world_id, - player_id=player_id, - campaign_log=campaign_log, - entity_index=entity_index, - vector_index=vector_index, - ) + with _ctx_lock: + if _ctx is not None: + if _ctx.world_id != world_id or _ctx.player_id != player_id: + raise RuntimeError( + f"ToolContext already bound to " + f"world={_ctx.world_id!r} player={_ctx.player_id!r}; " + f"refusing to rebind to " + f"world={world_id!r} player={player_id!r}" + ) + return _ctx + + world_dir = paths.world_path(world_id) + _ctx = ToolContext( + world_id=world_id, + player_id=player_id, + campaign_log=CampaignLog(world_id), + entity_index=EntityIndex(world_dir), + vector_index=_VectorIndex(world_dir / "search.db"), + ) + return _ctx + + +def current_ctx() -> ToolContext | None: + """Return the current ToolContext without constructing one.""" return _ctx def reset_ctx() -> None: """Clear the process-global ToolContext (for test teardown).""" global _ctx - _ctx = None + with _ctx_lock: + _ctx = None def _require() -> ToolContext: if _ctx is None: - raise RuntimeError("ToolContext not initialized; call init_ctx() first") + raise RuntimeError( + "ToolContext not initialized; call get_or_create_ctx() first" + ) return _ctx diff --git a/tests/test_engine.py b/tests/test_engine.py index 1f1d30d..98f2341 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -76,11 +76,13 @@ class TestDMEngineContext: with patch("storied.engine.start_mcp_server") as mock_mcp: from storied.initiative import InitiativeTracker + from storied.log import CampaignLog from storied.tools import EntityIndex mock_mcp.return_value = type("Handle", (), { "url": "http://localhost:0/sse", "ctx": type("Ctx", (), { + "campaign_log": CampaignLog("test"), "entity_index": EntityIndex(world_dir), "vector_index": None, "initiative": InitiativeTracker(), @@ -214,9 +216,12 @@ class TestDMEngineContext: transcript_path = tmp_path / "transcripts" / "session.jsonl" with patch("storied.engine.start_mcp_server") as mock_mcp: + from storied.log import CampaignLog + mock_mcp.return_value = type("Handle", (), { "url": "http://localhost:0/sse", "ctx": type("Ctx", (), { + "campaign_log": CampaignLog("test"), "entity_index": EntityIndex(tmp_path / "worlds" / "test"), "vector_index": None, "initiative": InitiativeTracker(),