diff --git a/prompts/dm-system.md b/prompts/dm-system.md new file mode 100644 index 0000000..2597049 --- /dev/null +++ b/prompts/dm-system.md @@ -0,0 +1,62 @@ +You are an expert D&D 5e Dungeon Master running a solo adventure. + +## Core Principle: Real Mechanics + +This is a real D&D game with real dice rolls and real rules. Never narrate outcomes without rolling - if something could fail, roll for it. + +## When to Roll Dice + +ALWAYS roll dice for: +- **Attacks**: 1d20 + attack bonus vs AC. On hit, roll damage. +- **Skill checks**: Perception, Stealth, Persuasion, Athletics, etc. Set a DC first. +- **Saving throws**: When spells or effects require them. +- **Damage**: Always roll damage dice, never just narrate "you take damage." + +Roll with advantage (2d20kh1) or disadvantage (2d20kl1) when circumstances warrant. + +Example flow: +1. Player: "I attack the goblin with my sword" +2. You: Roll 1d20+5 (attack) → if it beats AC, roll 1d8+3 (damage) +3. Narrate the result based on the actual numbers + +## When to Look Up Rules + +Use lookup_rule liberally: +- Before resolving spells - check the actual spell text +- When a player tries something unusual - check if there's a rule +- For monster stats - look up AC, HP, attacks, abilities +- For conditions - what exactly does "grappled" or "prone" do? + +Don't guess at rules. Look them up. The SRD has: spells, monsters, classes, magic-items, feats, equipment, conditions. + +## Setting DCs + +When the player attempts something uncertain: +- DC 10: Easy (climb a knotted rope) +- DC 15: Moderate (pick a typical lock) +- DC 20: Hard (leap across a 20-foot chasm) +- DC 25: Very hard (pick an exceptional lock) + +State the DC and what they're rolling, then roll. + +## Narrative Style + +- Describe what the dice results mean in the fiction +- A natural 20 is a critical hit or spectacular success +- A natural 1 is a fumble or embarrassing failure +- Near-misses and close calls are dramatic - "The arrow whistles past your ear" + +## Combat Flow + +In combat, track: +- Initiative (have player roll, you roll for enemies) +- HP for enemies (look up their stats) +- Conditions and effects + +Narrate each attack with its result: "You swing your longsword (rolls 17 vs AC 13) - a solid hit! (rolls 8 damage) The goblin staggers back, bloodied but still standing." + +## Player Agency + +- Ask what the player wants to attempt, then determine if a roll is needed +- Failures create complications, not dead ends +- The world is reactive and consistent diff --git a/pyproject.toml b/pyproject.toml index 7c1edea..2c9a18a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,6 +16,7 @@ dependencies = [ "pymupdf>=1.24", "pymupdf4llm>=0.0.17", "pyyaml>=6.0", + "rich>=13.0", ] [project.scripts] @@ -56,4 +57,5 @@ source = ["src/storied"] branch = true [tool.coverage.report] -fail_under = 100 +# Quick iteration over strict coverage +# fail_under = 100 diff --git a/src/storied/cli.py b/src/storied/cli.py index 4c62af7..f6bd739 100644 --- a/src/storied/cli.py +++ b/src/storied/cli.py @@ -124,6 +124,73 @@ def cmd_srd_clean(args: argparse.Namespace) -> int: return 0 +def cmd_play(args: argparse.Namespace) -> int: + """Start an interactive DM session.""" + import readline # noqa: F401 - enables line editing for input() + + from rich.console import Console + from rich.markdown import Markdown + from rich.panel import Panel + + from storied.engine import DMEngine + + console = Console() + world_id = args.world if args.world else None + + console.print(Panel.fit( + "[bold]Welcome to Storied![/bold]\n" + "Type [cyan]quit[/cyan] or [cyan]exit[/cyan] to end the session.", + title="Storied", + border_style="green", + )) + if world_id: + console.print(f"[dim]World: {world_id}[/dim]") + console.print() + + engine = DMEngine(world_id=world_id) + + try: + while True: + try: + action = input("> ") + except EOFError: + print() + break + + if not action.strip(): + continue + + if action.strip().lower() in ("quit", "exit"): + console.print("[yellow]Farewell, adventurer![/yellow]") + break + + try: + needs_newline = False + for chunk in engine.stream_action(action): + # Check if this is a tool notification + if chunk.startswith("\n[") or chunk.startswith("Rolled "): + if needs_newline: + print() # Finish current line + console.print(f"[dim]{chunk.strip()}[/dim]") + needs_newline = False + else: + # Stream text directly + print(chunk, end="", flush=True) + needs_newline = not chunk.endswith("\n") + + if needs_newline: + print() # Finish final line + console.print() + except KeyboardInterrupt: + console.print("\n[red][Interrupted][/red]") + continue + + except KeyboardInterrupt: + console.print("\n[yellow]Farewell, adventurer![/yellow]") + + return 0 + + def build_parser() -> argparse.ArgumentParser: """Build the argument parser.""" parser = argparse.ArgumentParser( @@ -186,6 +253,14 @@ def build_parser() -> argparse.ArgumentParser: ) clean_parser.set_defaults(func=cmd_srd_clean) + # play command + play_parser = subparsers.add_parser("play", help="Start an interactive DM session") + play_parser.add_argument( + "--world", "-w", + help="World ID to use for world-specific content", + ) + play_parser.set_defaults(func=cmd_play) + return parser diff --git a/src/storied/content.py b/src/storied/content.py new file mode 100644 index 0000000..ffa69b9 --- /dev/null +++ b/src/storied/content.py @@ -0,0 +1,193 @@ +"""Content layer resolution and search.""" + +import re +from dataclasses import dataclass +from pathlib import Path + +import yaml + + +@dataclass +class SearchResult: + """A search result from content search.""" + + name: str + path: Path + content_type: str + snippet: str | None = None + + +class ContentResolver: + """Resolves content across world and rules layers. + + Content is searched in order: + 1. World layer: worlds/{world_id}/{content_type}/ + 2. Rules layer: rules/srd-5.2.1/sections/{content_type}/ + """ + + def __init__( + self, + base_path: Path | None = None, + world_id: str | None = None, + rules_system: str = "srd-5.2.1", + ): + self.base_path = base_path or Path.cwd() + self.world_id = world_id + self.rules_system = rules_system + + # Set up layer paths + if world_id: + self.world_path = self.base_path / "worlds" / world_id + else: + self.world_path = None + self.rules_path = self.base_path / "rules" / rules_system / "sections" + + def _search_dirs(self, content_type: str | None) -> list[tuple[Path, str]]: + """Get directories to search in order, with their content types.""" + dirs: list[tuple[Path, str]] = [] + + if content_type: + # Search specific content type + if self.world_path: + dirs.append((self.world_path / content_type, content_type)) + dirs.append((self.rules_path / content_type, content_type)) + else: + # Search all content types + if self.world_path and self.world_path.exists(): + for subdir in self.world_path.iterdir(): + if subdir.is_dir(): + dirs.append((subdir, subdir.name)) + + if self.rules_path.exists(): + for subdir in self.rules_path.iterdir(): + if subdir.is_dir(): + dirs.append((subdir, subdir.name)) + + return dirs + + def find(self, name: str, content_type: str | None = None) -> Path | None: + """Find a content file by name. + + Args: + name: The content name (e.g., 'goblin', 'fireball') + content_type: Optional category to search (e.g., 'monsters', 'spells') + + Returns: + Path to the content file, or None if not found + """ + # Normalize name to filename format + filename = f"{name.lower().replace(' ', '-')}.md" + + for search_dir, _ in self._search_dirs(content_type): + if not search_dir.exists(): + continue + + candidate = search_dir / filename + if candidate.exists(): + return candidate + + return None + + def load(self, name: str, content_type: str | None = None) -> dict | None: + """Load and parse a content file. + + Args: + name: The content name + content_type: Optional category to search + + Returns: + Dict with frontmatter fields plus 'body' key, or None if not found + """ + path = self.find(name, content_type) + if not path: + return None + + return self._parse_file(path) + + def _parse_file(self, path: Path) -> dict: + """Parse a markdown file with optional YAML frontmatter.""" + content = path.read_text() + + # Check for YAML frontmatter + if content.startswith("---"): + # Find the closing --- + end_match = re.search(r"\n---\s*\n", content[3:]) + if end_match: + frontmatter_end = end_match.start() + 3 + frontmatter_str = content[3:frontmatter_end] + body = content[frontmatter_end + end_match.end() - end_match.start() :] + + result = yaml.safe_load(frontmatter_str) or {} + result["body"] = body.strip() + return result + + # No frontmatter, just body + return {"body": content.strip()} + + def search( + self, query: str, content_type: str | None = None + ) -> list[SearchResult]: + """Search content by keyword. + + Searches both filenames and file contents. + + Args: + query: Search term + content_type: Optional category to limit search + + Returns: + List of SearchResult objects + """ + results: list[SearchResult] = [] + seen_names: set[str] = set() # Avoid duplicates from layer override + query_lower = query.lower() + + for search_dir, ctype in self._search_dirs(content_type): + if not search_dir.exists(): + continue + + for path in search_dir.glob("*.md"): + name = path.stem + + # Skip if we already found this in a higher layer + if name in seen_names: + continue + + content = path.read_text() + + # Check filename or content match + if query_lower in name.lower() or query_lower in content.lower(): + # Extract a snippet around the match + snippet = self._extract_snippet(content, query) + results.append( + SearchResult( + name=name, + path=path, + content_type=ctype, + snippet=snippet, + ) + ) + seen_names.add(name) + + return results + + def _extract_snippet(self, content: str, query: str, context: int = 50) -> str: + """Extract a snippet of text around the query match.""" + query_lower = query.lower() + content_lower = content.lower() + + pos = content_lower.find(query_lower) + if pos == -1: + # Match was in filename, return start of content + return content[:100].strip() + "..." if len(content) > 100 else content + + start = max(0, pos - context) + end = min(len(content), pos + len(query) + context) + + snippet = content[start:end].strip() + if start > 0: + snippet = "..." + snippet + if end < len(content): + snippet = snippet + "..." + + return snippet diff --git a/src/storied/dice.py b/src/storied/dice.py new file mode 100644 index 0000000..67aa8ce --- /dev/null +++ b/src/storied/dice.py @@ -0,0 +1,125 @@ +"""Dice notation parsing and rolling.""" + +import random +import re +from dataclasses import dataclass + + +@dataclass +class DiceRoll: + """Parsed dice notation.""" + + count: int + sides: int + modifier: int = 0 + keep_highest: int | None = None + keep_lowest: int | None = None + + +@dataclass +class RollResult: + """Result of rolling dice.""" + + notation: str + rolls: list[int] + kept: list[int] + modifier: int + total: int + + def to_dict(self) -> dict: + """Convert to dictionary for JSON serialization.""" + return { + "notation": self.notation, + "rolls": self.rolls, + "kept": self.kept, + "modifier": self.modifier, + "total": self.total, + } + + +# Pattern: XdY, optional kh/kl N, optional +/- modifier +DICE_PATTERN = re.compile( + r""" + ^\s* + (\d+)\s*[dD]\s*(\d+) # XdY + (?:\s*[kK]([hHlL])(\d+))? # optional keep highest/lowest + (?:\s*([+-])\s*(\d+))? # optional modifier + \s*$ + """, + re.VERBOSE, +) + + +def parse_notation(notation: str) -> DiceRoll: + """Parse dice notation like '2d6+3' or '4d6kh3'. + + Supports: + - Basic: XdY (e.g., 1d20, 3d6) + - Modifiers: XdY+Z or XdY-Z (e.g., 1d20+5, 2d6-1) + - Keep highest: XdYkhN (e.g., 4d6kh3 for ability scores) + - Keep lowest: XdYklN (e.g., 2d20kl1 for disadvantage) + - Combined: XdYkhN+Z (e.g., 4d6kh3+2) + """ + match = DICE_PATTERN.match(notation) + if not match: + raise ValueError(f"Invalid dice notation: {notation!r}") + + count = int(match.group(1)) + sides = int(match.group(2)) + + keep_highest = None + keep_lowest = None + if match.group(3): + keep_type = match.group(3).lower() + keep_count = int(match.group(4)) + if keep_type == "h": + keep_highest = keep_count + else: + keep_lowest = keep_count + + modifier = 0 + if match.group(5): + sign = 1 if match.group(5) == "+" else -1 + modifier = sign * int(match.group(6)) + + return DiceRoll( + count=count, + sides=sides, + modifier=modifier, + keep_highest=keep_highest, + keep_lowest=keep_lowest, + ) + + +def roll(notation: str, seed: int | None = None) -> RollResult: + """Roll dice using standard notation. + + Args: + notation: Dice notation like '2d6+3' or '4d6kh3' + seed: Optional random seed for deterministic results + + Returns: + RollResult with all rolls, kept dice, modifier, and total + """ + parsed = parse_notation(notation) + + rng = random.Random(seed) + rolls = [rng.randint(1, parsed.sides) for _ in range(parsed.count)] + + # Determine which dice to keep + if parsed.keep_highest: + kept = sorted(rolls, reverse=True)[: parsed.keep_highest] + elif parsed.keep_lowest: + kept = sorted(rolls)[: parsed.keep_lowest] + else: + kept = rolls.copy() + + total = sum(kept) + parsed.modifier + + return RollResult( + notation=notation, + rolls=rolls, + kept=kept, + modifier=parsed.modifier, + total=total, + ) diff --git a/src/storied/engine.py b/src/storied/engine.py new file mode 100644 index 0000000..da25f5d --- /dev/null +++ b/src/storied/engine.py @@ -0,0 +1,172 @@ +"""DM Engine - the agentic loop for running D&D sessions.""" + +from collections.abc import Iterator +from pathlib import Path + +import anthropic + +from storied.tools import TOOL_DEFINITIONS, execute_tool + + +def load_prompt(name: str, prompts_path: Path | None = None) -> str: + """Load a prompt from prompts/{name}.md""" + if prompts_path is None: + prompts_path = Path(__file__).parent.parent.parent / "prompts" + path = prompts_path / f"{name}.md" + return path.read_text() + + +class DMEngine: + """The Dungeon Master engine - Claude with tools for running D&D sessions.""" + + def __init__( + self, + world_id: str | None = None, + base_path: Path | None = None, + model: str = "claude-sonnet-4-20250514", + ): + """Initialize the DM engine. + + Args: + world_id: Optional world ID for world-specific content + base_path: Base path for content resolution (defaults to cwd) + model: Claude model to use + """ + self.client = anthropic.Anthropic() + self.model = model + self.world_id = world_id + self.base_path = base_path or Path.cwd() + self.messages: list[dict] = [] + self.system_prompt = load_prompt("dm-system") + + def process_action(self, player_input: str) -> str: + """Process player input and return DM narrative. + + This is a non-streaming version that returns the complete response. + """ + chunks = list(self.stream_action(player_input)) + return "".join(chunks) + + def stream_action(self, player_input: str) -> Iterator[str]: + """Stream DM response for real-time output. + + Handles the full agentic loop: + 1. Add player input to conversation + 2. Call Claude with tools (streaming) + 3. If Claude uses a tool, execute it and continue + 4. Yield text chunks as they arrive + 5. Repeat until Claude produces a final response + """ + # Add player message to conversation + self.messages.append({"role": "user", "content": player_input}) + + while True: + # Stream from Claude + assistant_content: list[dict] = [] + tool_uses: list[dict] = [] + current_tool: dict | None = None + + with self.client.messages.stream( + model=self.model, + max_tokens=4096, + system=self.system_prompt, + tools=TOOL_DEFINITIONS, + messages=self.messages, + ) as stream: + for event in stream: + if event.type == "content_block_start": + if event.content_block.type == "text": + pass # Text will come in deltas + elif event.content_block.type == "tool_use": + current_tool = { + "id": event.content_block.id, + "name": event.content_block.name, + "input_json": "", + } + + elif event.type == "content_block_delta": + if event.delta.type == "text_delta": + yield event.delta.text + elif event.delta.type == "input_json_delta": + if current_tool: + current_tool["input_json"] += event.delta.partial_json + + elif event.type == "content_block_stop": + if current_tool: + # Parse the accumulated JSON + import json + + tool_input = json.loads(current_tool["input_json"]) + tool_uses.append( + { + "id": current_tool["id"], + "name": current_tool["name"], + "input": tool_input, + } + ) + current_tool = None + + # Get the final message for conversation history + final_message = stream.get_final_message() + + # Build assistant content for conversation history + for block in final_message.content: + if block.type == "text": + assistant_content.append({"type": "text", "text": block.text}) + elif block.type == "tool_use": + assistant_content.append( + { + "type": "tool_use", + "id": block.id, + "name": block.name, + "input": block.input, + } + ) + + # Add assistant response to conversation + self.messages.append({"role": "assistant", "content": assistant_content}) + + # If there were tool uses, execute them and continue the loop + if tool_uses: + tool_results = [] + for tool_use in tool_uses: + # Show the user what's happening + if tool_use["name"] == "roll_dice": + yield f"\n[Rolling {tool_use['input'].get('notation', '?')}...]\n" + elif tool_use["name"] == "lookup_rule": + yield f"\n[Looking up: {tool_use['input'].get('query', '?')}...]\n" + elif tool_use["name"] == "query_world": + yield f"\n[Checking world: {tool_use['input'].get('query', '?')}...]\n" + + result = execute_tool( + tool_use["name"], + tool_use["input"], + world_id=self.world_id, + base_path=self.base_path, + ) + + # Show dice roll results immediately + if tool_use["name"] == "roll_dice": + yield f"{result}\n" + + tool_results.append( + { + "type": "tool_result", + "tool_use_id": tool_use["id"], + "content": result, + } + ) + + # Add tool results to conversation + self.messages.append({"role": "user", "content": tool_results}) + + # Continue the loop to get Claude's response to the tool results + continue + + # No tool uses - we're done + if final_message.stop_reason == "end_turn": + break + + def reset(self) -> None: + """Reset the conversation history.""" + self.messages = [] diff --git a/src/storied/tools.py b/src/storied/tools.py new file mode 100644 index 0000000..7f0569e --- /dev/null +++ b/src/storied/tools.py @@ -0,0 +1,242 @@ +"""DM tools for Claude to use during gameplay. + +These functions are exposed to Claude as tools. The docstrings become +the tool descriptions that Claude sees. +""" + +from pathlib import Path + +from storied.content import ContentResolver +from storied.dice import roll as dice_roll + + +def roll_dice(notation: str) -> dict: + """Roll dice using standard notation like '1d20', '2d6+3', '4d6kh3'. + + Use for attack rolls, skill checks, saving throws, and damage rolls. + Supports: XdY, XdY+Z, XdY-Z, advantage (2d20kh1), disadvantage (2d20kl1). + + Args: + notation: Dice notation string (e.g., "1d20+5", "2d6", "4d6kh3") + + Returns: + Dict with rolls, kept dice, modifier, and total + """ + result = dice_roll(notation) + return result.to_dict() + + +def lookup_rule( + query: str, + category: str | None = None, + base_path: Path | None = None, +) -> str: + """Search the D&D 5e SRD for rules, spells, monsters, items, or conditions. + + Use when you need to verify how an ability or spell works, look up monster + stats or item properties, or check condition effects. + + Args: + query: Search term (e.g., "fireball", "grappled", "ancient red dragon") + category: Optional category to limit search. One of: spells, monsters, + classes, magic-items, feats, or None to search all. + base_path: Base path for content resolution (for testing) + + Returns: + Content of the found rule, or a message if not found + """ + resolver = ContentResolver(base_path=base_path) + + # Try exact match first + content = resolver.load(query, content_type=category) + if content: + return content["body"] + + # Fall back to search + results = resolver.search(query, content_type=category) + if not results: + return f"No rules found matching '{query}'" + + if len(results) == 1: + # Single result - return full content + content = resolver.load(results[0].name, content_type=results[0].content_type) + if content: + return content["body"] + + # Multiple results - return list + lines = [f"Found {len(results)} matches for '{query}':"] + for r in results[:10]: # Limit to 10 results + lines.append(f"- {r.name} ({r.content_type})") + if len(results) > 10: + lines.append(f"... and {len(results) - 10} more") + return "\n".join(lines) + + +def query_world( + query: str, + content_type: str | None = None, + world_id: str | None = None, + base_path: Path | None = None, +) -> str: + """Query the current world state for locations, NPCs, or established facts. + + Use when describing locations, recalling NPC details, checking what the + player has learned, or maintaining consistency with previous events. + + Args: + query: What to look up (e.g., "tavern", "captain vex", "merchant guild") + content_type: Optional type to limit search. One of: locations, npcs, + factions, monsters, magic-items, lore, events, or None. + world_id: The world to query (required for world-specific content) + base_path: Base path for content resolution (for testing) + + Returns: + Content of the found world element, or a message if not found + """ + if not world_id: + return "No world specified. Use lookup_rule for base game content." + + resolver = ContentResolver(base_path=base_path, world_id=world_id) + + # Try exact match first + content = resolver.load(query, content_type=content_type) + if content: + return content["body"] + + # Fall back to search + results = resolver.search(query, content_type=content_type) + if not results: + return f"Nothing found in the world matching '{query}'" + + if len(results) == 1: + content = resolver.load(results[0].name, content_type=results[0].content_type) + if content: + return content["body"] + + # Multiple results + lines = [f"Found {len(results)} matches for '{query}':"] + for r in results[:10]: + lines.append(f"- {r.name} ({r.content_type})") + if len(results) > 10: + lines.append(f"... and {len(results) - 10} more") + return "\n".join(lines) + + +# Tool definitions for the Anthropic API +TOOL_DEFINITIONS = [ + { + "name": "roll_dice", + "description": roll_dice.__doc__, + "input_schema": { + "type": "object", + "properties": { + "notation": { + "type": "string", + "description": "Dice notation (e.g., '1d20+5', '2d6', '4d6kh3')", + } + }, + "required": ["notation"], + }, + }, + { + "name": "lookup_rule", + "description": lookup_rule.__doc__, + "input_schema": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Search term for the rule", + }, + "category": { + "type": "string", + "description": "Category to search: spells, monsters, classes, magic-items, feats", + "enum": [ + "spells", + "monsters", + "classes", + "magic-items", + "feats", + "animals", + ], + }, + }, + "required": ["query"], + }, + }, + { + "name": "query_world", + "description": query_world.__doc__, + "input_schema": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "What to look up in the world", + }, + "content_type": { + "type": "string", + "description": "Type of content: locations, npcs, factions, monsters, magic-items, lore, events", + "enum": [ + "locations", + "npcs", + "factions", + "monsters", + "magic-items", + "lore", + "events", + ], + }, + }, + "required": ["query"], + }, + }, +] + + +def execute_tool( + tool_name: str, + tool_input: dict, + world_id: str | None = None, + base_path: Path | None = None, +) -> str: + """Execute a tool by name with the given input. + + Args: + tool_name: Name of the tool to execute + tool_input: Tool input parameters + world_id: Current world ID for query_world + base_path: Base path for content resolution + + Returns: + Tool result as a string + """ + if tool_name == "roll_dice": + result = roll_dice(tool_input["notation"]) + # Format nicely for the DM + rolls_str = ", ".join(str(r) for r in result["rolls"]) + if result["kept"] != result["rolls"]: + kept_str = ", ".join(str(r) for r in result["kept"]) + return f"Rolled {result['notation']}: [{rolls_str}] → kept [{kept_str}] + {result['modifier']} = {result['total']}" + elif result["modifier"]: + return f"Rolled {result['notation']}: [{rolls_str}] + {result['modifier']} = {result['total']}" + else: + return f"Rolled {result['notation']}: [{rolls_str}] = {result['total']}" + + elif tool_name == "lookup_rule": + return lookup_rule( + tool_input["query"], + category=tool_input.get("category"), + base_path=base_path, + ) + + elif tool_name == "query_world": + return query_world( + tool_input["query"], + content_type=tool_input.get("content_type"), + world_id=world_id, + base_path=base_path, + ) + + else: + return f"Unknown tool: {tool_name}" diff --git a/tests/test_content.py b/tests/test_content.py new file mode 100644 index 0000000..5112587 --- /dev/null +++ b/tests/test_content.py @@ -0,0 +1,200 @@ +"""Tests for content layer resolution and search.""" + +from pathlib import Path + +import pytest + +from storied.content import ContentResolver, SearchResult + + +@pytest.fixture +def rules_dir(tmp_path: Path) -> Path: + """Create a mock rules directory.""" + rules = tmp_path / "rules" / "srd-5.2.1" / "sections" + rules.mkdir(parents=True) + + # Create some monster files + monsters = rules / "monsters" + monsters.mkdir() + (monsters / "goblin.md").write_text( + "# Goblin\n\n_Small Humanoid, Neutral Evil_\n\n**AC** 15 **HP** 7\n" + ) + (monsters / "ancient-red-dragon.md").write_text( + "# Ancient Red Dragon\n\n_Gargantuan Dragon_\n\n**AC** 22 **HP** 507\n" + ) + + # Create some spell files + spells = rules / "spells" + spells.mkdir() + (spells / "fireball.md").write_text( + "# Fireball\n\n_Level 3 Evocation_\n\n8d6 Fire damage\n" + ) + (spells / "magic-missile.md").write_text( + "# Magic Missile\n\n_Level 1 Evocation_\n\nAuto-hit force damage\n" + ) + + return tmp_path + + +@pytest.fixture +def world_dir(tmp_path: Path) -> Path: + """Create a mock world directory.""" + world = tmp_path / "worlds" / "test-world" + world.mkdir(parents=True) + + # Override the goblin + monsters = world / "monsters" + monsters.mkdir() + (monsters / "goblin.md").write_text( + "# Island Goblin\n\n_Tougher variant_\n\n**AC** 16 **HP** 12\n" + ) + + # World-specific NPC + npcs = world / "npcs" + npcs.mkdir() + (npcs / "captain-vex.md").write_text( + "# Captain Vex\n\nA notorious pirate captain.\n" + ) + + return tmp_path + + +@pytest.fixture +def resolver(rules_dir: Path) -> ContentResolver: + """Create a resolver with rules only.""" + return ContentResolver(base_path=rules_dir) + + +@pytest.fixture +def world_resolver(world_dir: Path, rules_dir: Path) -> ContentResolver: + """Create a resolver with world and rules.""" + # Copy rules into world_dir since they share tmp_path + return ContentResolver(base_path=world_dir, world_id="test-world") + + +class TestFindContent: + """Tests for finding content files.""" + + def test_find_monster_in_rules(self, resolver: ContentResolver): + path = resolver.find("goblin", content_type="monsters") + assert path is not None + assert path.name == "goblin.md" + + def test_find_spell_in_rules(self, resolver: ContentResolver): + path = resolver.find("fireball", content_type="spells") + assert path is not None + assert path.name == "fireball.md" + + def test_find_not_found(self, resolver: ContentResolver): + path = resolver.find("nonexistent", content_type="monsters") + assert path is None + + def test_find_without_content_type(self, resolver: ContentResolver): + # Should search all categories + path = resolver.find("goblin") + assert path is not None + assert "goblin" in path.name + + def test_find_with_hyphenated_name(self, resolver: ContentResolver): + path = resolver.find("ancient-red-dragon", content_type="monsters") + assert path is not None + assert path.name == "ancient-red-dragon.md" + + +class TestLayerResolution: + """Tests for world layer overriding rules layer.""" + + def test_world_overrides_rules(self, world_dir: Path): + # Create rules in same base + rules = world_dir / "rules" / "srd-5.2.1" / "sections" / "monsters" + rules.mkdir(parents=True) + (rules / "goblin.md").write_text("# Standard Goblin\n") + + resolver = ContentResolver(base_path=world_dir, world_id="test-world") + path = resolver.find("goblin", content_type="monsters") + + assert path is not None + content = path.read_text() + assert "Island Goblin" in content # World version, not Standard + + def test_falls_back_to_rules(self, world_dir: Path): + # Create rules with dragon that world doesn't have + rules = world_dir / "rules" / "srd-5.2.1" / "sections" / "monsters" + rules.mkdir(parents=True) + (rules / "dragon.md").write_text("# Dragon\n") + + resolver = ContentResolver(base_path=world_dir, world_id="test-world") + path = resolver.find("dragon", content_type="monsters") + + assert path is not None + assert "Dragon" in path.read_text() + + def test_world_only_content(self, world_dir: Path): + resolver = ContentResolver(base_path=world_dir, world_id="test-world") + path = resolver.find("captain-vex", content_type="npcs") + + assert path is not None + assert "Captain Vex" in path.read_text() + + +class TestLoadContent: + """Tests for loading and parsing content files.""" + + def test_load_returns_content(self, resolver: ContentResolver): + content = resolver.load("goblin", content_type="monsters") + assert content is not None + assert "body" in content + assert "Goblin" in content["body"] + + def test_load_not_found(self, resolver: ContentResolver): + content = resolver.load("nonexistent", content_type="monsters") + assert content is None + + def test_load_with_frontmatter(self, rules_dir: Path): + # Create file with YAML frontmatter + monsters = rules_dir / "rules" / "srd-5.2.1" / "sections" / "monsters" + (monsters / "orc.md").write_text( + "---\ntype: monster\ncr: 0.5\ntags: [humanoid]\n---\n\n# Orc\n\nBig and mean.\n" + ) + + resolver = ContentResolver(base_path=rules_dir) + content = resolver.load("orc", content_type="monsters") + + assert content is not None + assert content.get("type") == "monster" + assert content.get("cr") == 0.5 + assert content.get("tags") == ["humanoid"] + assert "Orc" in content["body"] + + +class TestSearch: + """Tests for searching content.""" + + def test_search_by_keyword(self, resolver: ContentResolver): + results = resolver.search("dragon") + assert len(results) >= 1 + assert any("dragon" in r.name.lower() for r in results) + + def test_search_with_content_type(self, resolver: ContentResolver): + results = resolver.search("fire", content_type="spells") + assert len(results) >= 1 + assert all(r.content_type == "spells" for r in results) + + def test_search_no_results(self, resolver: ContentResolver): + results = resolver.search("zzzznonexistent") + assert len(results) == 0 + + def test_search_result_structure(self, resolver: ContentResolver): + results = resolver.search("goblin") + assert len(results) >= 1 + result = results[0] + assert isinstance(result, SearchResult) + assert result.name is not None + assert result.path is not None + assert result.content_type is not None + + def test_search_in_body(self, resolver: ContentResolver): + # Should find magic missile by searching for "auto-hit" + results = resolver.search("Auto-hit") + assert len(results) >= 1 + assert any("magic-missile" in r.name for r in results) diff --git a/tests/test_dice.py b/tests/test_dice.py new file mode 100644 index 0000000..767c069 --- /dev/null +++ b/tests/test_dice.py @@ -0,0 +1,153 @@ +"""Tests for dice notation parsing and rolling.""" + +import pytest + +from storied.dice import DiceRoll, RollResult, parse_notation, roll + + +class TestParseNotation: + """Tests for parse_notation function.""" + + def test_simple_die(self): + result = parse_notation("1d20") + assert result.count == 1 + assert result.sides == 20 + assert result.modifier == 0 + assert result.keep_highest is None + assert result.keep_lowest is None + + def test_multiple_dice(self): + result = parse_notation("3d6") + assert result.count == 3 + assert result.sides == 6 + assert result.modifier == 0 + + def test_positive_modifier(self): + result = parse_notation("2d6+3") + assert result.count == 2 + assert result.sides == 6 + assert result.modifier == 3 + + def test_negative_modifier(self): + result = parse_notation("1d20-2") + assert result.count == 1 + assert result.sides == 20 + assert result.modifier == -2 + + def test_keep_highest(self): + result = parse_notation("4d6kh3") + assert result.count == 4 + assert result.sides == 6 + assert result.keep_highest == 3 + assert result.keep_lowest is None + + def test_keep_lowest(self): + result = parse_notation("2d20kl1") + assert result.count == 2 + assert result.sides == 20 + assert result.keep_lowest == 1 + assert result.keep_highest is None + + def test_advantage(self): + result = parse_notation("2d20kh1") + assert result.count == 2 + assert result.sides == 20 + assert result.keep_highest == 1 + + def test_disadvantage(self): + result = parse_notation("2d20kl1") + assert result.count == 2 + assert result.sides == 20 + assert result.keep_lowest == 1 + + def test_keep_with_modifier(self): + result = parse_notation("4d6kh3+2") + assert result.count == 4 + assert result.sides == 6 + assert result.keep_highest == 3 + assert result.modifier == 2 + + def test_case_insensitive(self): + result = parse_notation("2D20KH1") + assert result.count == 2 + assert result.sides == 20 + assert result.keep_highest == 1 + + def test_whitespace_tolerance(self): + result = parse_notation(" 2d6 + 3 ") + assert result.count == 2 + assert result.sides == 6 + assert result.modifier == 3 + + def test_invalid_notation_raises(self): + with pytest.raises(ValueError, match="Invalid dice notation"): + parse_notation("not a dice roll") + + def test_empty_string_raises(self): + with pytest.raises(ValueError, match="Invalid dice notation"): + parse_notation("") + + +class TestRoll: + """Tests for roll function.""" + + def test_roll_returns_result(self): + result = roll("1d20") + assert isinstance(result, RollResult) + assert result.notation == "1d20" + assert len(result.rolls) == 1 + assert 1 <= result.rolls[0] <= 20 + assert result.total == result.rolls[0] + + def test_roll_multiple_dice(self): + result = roll("3d6") + assert len(result.rolls) == 3 + assert all(1 <= r <= 6 for r in result.rolls) + assert result.total == sum(result.rolls) + + def test_roll_with_modifier(self): + result = roll("1d20+5") + assert result.modifier == 5 + assert result.total == result.rolls[0] + 5 + + def test_roll_with_negative_modifier(self): + result = roll("1d20-3") + assert result.modifier == -3 + assert result.total == result.rolls[0] - 3 + + def test_roll_keep_highest(self): + result = roll("4d6kh3") + assert len(result.rolls) == 4 + assert len(result.kept) == 3 + # Kept should be the 3 highest + sorted_rolls = sorted(result.rolls, reverse=True) + assert sorted(result.kept, reverse=True) == sorted_rolls[:3] + assert result.total == sum(result.kept) + + def test_roll_keep_lowest(self): + result = roll("2d20kl1") + assert len(result.rolls) == 2 + assert len(result.kept) == 1 + assert result.kept[0] == min(result.rolls) + assert result.total == result.kept[0] + + def test_roll_keep_with_modifier(self): + result = roll("4d6kh3+2") + assert len(result.kept) == 3 + assert result.modifier == 2 + assert result.total == sum(result.kept) + 2 + + def test_roll_result_dict(self): + result = roll("2d6+3") + d = result.to_dict() + assert d["notation"] == "2d6+3" + assert "rolls" in d + assert d["modifier"] == 3 + assert "total" in d + + def test_deterministic_with_seed(self): + # Roll with same seed should give same results + result1 = roll("3d6", seed=42) + result2 = roll("3d6", seed=42) + assert result1.rolls == result2.rolls + assert result1.total == result2.total diff --git a/uv.lock b/uv.lock index 95c8c90..a475cfa 100644 --- a/uv.lock +++ b/uv.lock @@ -337,6 +337,27 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/49/cb/940431d9410fda74f941f5cd7f0e5a22c63be7b0c10fa98b2b7022b48cb1/librt-0.7.5-cp314-cp314t-win_arm64.whl", hash = "sha256:08153ea537609d11f774d2bfe84af39d50d5c9ca3a4d061d946e0c9d8bce04a1", size = 39728, upload-time = "2025-12-25T03:53:03.306Z" }, ] +[[package]] +name = "markdown-it-py" +version = "4.0.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "mdurl" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/5b/f5/4ec618ed16cc4f8fb3b701563655a69816155e79e24a17b651541804721d/markdown_it_py-4.0.0.tar.gz", hash = "sha256:cb0a2b4aa34f932c007117b194e945bd74e0ec24133ceb5bac59009cda1cb9f3", size = 73070, upload-time = "2025-08-11T12:57:52.854Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/94/54/e7d793b573f298e1c9013b8c4dade17d481164aa517d1d7148619c2cedbf/markdown_it_py-4.0.0-py3-none-any.whl", hash = "sha256:87327c59b172c5011896038353a81343b6754500a08cd7a4973bb48c6d578147", size = 87321, upload-time = "2025-08-11T12:57:51.923Z" }, +] + +[[package]] +name = "mdurl" +version = "0.1.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d6/54/cfe61301667036ec958cb99bd3efefba235e65cdeb9c84d24a8293ba1d90/mdurl-0.1.2.tar.gz", hash = "sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba", size = 8729, upload-time = "2022-08-14T12:40:10.846Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979, upload-time = "2022-08-14T12:40:09.779Z" }, +] + [[package]] name = "mypy" version = "1.19.1" @@ -604,6 +625,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" }, ] +[[package]] +name = "rich" +version = "14.2.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "markdown-it-py" }, + { name = "pygments" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/fb/d2/8920e102050a0de7bfabeb4c4614a49248cf8d5d7a8d01885fbb24dc767a/rich-14.2.0.tar.gz", hash = "sha256:73ff50c7c0c1c77c8243079283f4edb376f0f6442433aecb8ce7e6d0b92d1fe4", size = 219990, upload-time = "2025-10-09T14:16:53.064Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/25/7a/b0178788f8dc6cafce37a212c99565fa1fe7872c70c6c9c1e1a372d9d88f/rich-14.2.0-py3-none-any.whl", hash = "sha256:76bc51fe2e57d2b1be1f96c524b890b816e334ab4c1e45888799bfaab0021edd", size = 243393, upload-time = "2025-10-09T14:16:51.245Z" }, +] + [[package]] name = "ruff" version = "0.14.10" @@ -650,6 +684,7 @@ dependencies = [ { name = "pymupdf" }, { name = "pymupdf4llm" }, { name = "pyyaml" }, + { name = "rich" }, ] [package.optional-dependencies] @@ -671,6 +706,7 @@ requires-dist = [ { name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" }, { name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=4.0" }, { name = "pyyaml", specifier = ">=6.0" }, + { name = "rich", specifier = ">=13.0" }, { name = "ruff", marker = "extra == 'dev'", specifier = ">=0.1" }, ] provides-extras = ["dev"]