"""Tool definitions and handlers for the orchestrator's tool-calling loop.""" from __future__ import annotations import json import random import uuid from typing import Any, Awaitable, Callable, Dict, List, Optional from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.core.llm import build_tool_schema from app.core.rag import get_rag from app.core.state_validator import apply_patch, validate_state from app.logging_setup import get_logger from app.models import DeferredTrigger, GlossaryEntry, World log = get_logger("tools") # === Tool schemas (OpenAI function-calling format) === DICE_ROLL_SCHEMA = build_tool_schema( name="dice_roll", description="Roll dice. Use 'sides' (e.g. 20 for d20) and optional 'count' (default 1) and 'modifier'. Returns the rolls and total.", params={ "type": "object", "properties": { "sides": {"type": "integer", "description": "Number of sides on the die, e.g. 20 for d20"}, "count": {"type": "integer", "description": "Number of dice to roll", "default": 1}, "modifier": {"type": "integer", "description": "Modifier to add to total", "default": 0}, "label": {"type": "string", "description": "What this roll represents, e.g. 'attack' or 'perception'"}, }, "required": ["sides"], }, ) UPDATE_STATE_SCHEMA = build_tool_schema( name="update_state", description="Apply a patch to world state. Paths use dot notation. ops: set, unset, append, increment, remove.", params={ "type": "object", "properties": { "patch": { "type": "object", "description": "JSON-patch object with optional keys: set, unset, append, increment, remove. Each is a dict of path->value (or list of paths for unset).", "properties": { "set": {"type": "object"}, "unset": {"type": "array", "items": {"type": "string"}}, "append": {"type": "object"}, "increment": {"type": "object"}, "remove": {"type": "object"}, }, } }, "required": ["patch"], }, ) RAG_QUERY_SCHEMA = build_tool_schema( name="rag_query", description="Search the glossary (NPCs, locations, items, lore) for relevant facts.", params={ "type": "object", "properties": { "query": {"type": "string", "description": "Free-text search query"}, "limit": {"type": "integer", "description": "Max results", "default": 5}, }, "required": ["query"], }, ) RAG_ADD_SCHEMA = build_tool_schema( name="rag_add", description="Add a new entry to the glossary (NPC, location, item, lore, event).", params={ "type": "object", "properties": { "kind": {"type": "string", "enum": ["npc", "location", "item", "lore", "event", "rule"]}, "name": {"type": "string"}, "description": {"type": "string"}, "payload": {"type": "object", "description": "Optional extra fields"}, }, "required": ["kind", "name", "description"], }, ) SCHEDULE_TRIGGER_SCHEMA = build_tool_schema( name="schedule_trigger", description="Schedule a deferred event tied to world time. When world time reaches fire_at, the system will fire it.", params={ "type": "object", "properties": { "fire_at": {"type": "string", "description": "World time string in same format as world.current_time, e.g. 'day_3_hour_14'"}, "description": {"type": "string", "description": "What should happen"}, "payload": {"type": "object", "description": "Arbitrary structured payload for the trigger runner"}, }, "required": ["fire_at", "description"], }, ) ADVANCE_TIME_SCHEMA = build_tool_schema( name="advance_time", description="Advance the world's internal clock by days/hours/minutes. Use this when the action takes time.", params={ "type": "object", "properties": { "days": {"type": "integer", "default": 0}, "hours": {"type": "integer", "default": 0}, "minutes": {"type": "integer", "default": 0}, "reason": {"type": "string", "description": "Why time advances"}, }, }, ) RUN_SUBAGENT_SCHEMA = build_tool_schema( name="run_subagent", description="Spawn a sub-agent with clean context for a focused sub-task (e.g. generate NPC backstory, room description).", params={ "type": "object", "properties": { "task": {"type": "string", "description": "The specific task for the sub-agent"}, "context": {"type": "string", "description": "Minimal context needed (max 200 words)"}, }, "required": ["task"], }, ) ALL_TOOL_SCHEMAS = [ DICE_ROLL_SCHEMA, UPDATE_STATE_SCHEMA, RAG_QUERY_SCHEMA, RAG_ADD_SCHEMA, SCHEDULE_TRIGGER_SCHEMA, ADVANCE_TIME_SCHEMA, RUN_SUBAGENT_SCHEMA, ] # === Tool handlers === class ToolContext: """Holds everything tools need to execute.""" def __init__( self, db: AsyncSession, world: World, session_id: uuid.UUID, user_id: uuid.UUID, subagent_runner: Optional[Callable[[str, str], Awaitable[str]]] = None, settings_map: Optional[Dict[str, Any]] = None, ): self.db = db self.world = world self.session_id = session_id self.user_id = user_id self.subagent_runner = subagent_runner self.settings_map = settings_map or {} # Track time advancement during this iteration self.time_advance: Dict[str, int] = {"days": 0, "hours": 0, "minutes": 0} # Track scheduled triggers self.scheduled_triggers: List[Dict[str, Any]] = [] # Track rag facts added self.rag_added: List[Dict[str, Any]] = [] async def handle_tool_call(name: str, args: Dict[str, Any], ctx: ToolContext) -> Dict[str, Any]: if name == "dice_roll": sides = int(args.get("sides", 20)) count = int(args.get("count", 1)) modifier = int(args.get("modifier", 0)) label = args.get("label", "") rolls = [random.randint(1, sides) for _ in range(max(1, count))] total = sum(rolls) + modifier return {"rolls": rolls, "modifier": modifier, "total": total, "label": label} if name == "update_state": patch = args.get("patch", {}) new_state = apply_patch(ctx.world.state, patch) schema = ctx.world.definition.get("world_schema", {}) ok, errors = validate_state(new_state, schema) if not ok: return {"ok": False, "errors": errors, "state_unchanged": True} ctx.world.state = new_state return {"ok": True, "new_state_summary": _summarize_state(new_state)} if name == "rag_query": query = args.get("query", "") limit = int(args.get("limit", 5)) rag = await get_rag(ctx.settings_map) results = await rag.search_glossary(ctx.world.id, query, limit=limit, settings_map=ctx.settings_map) return {"results": results} if name == "rag_add": kind = args.get("kind", "lore") entry_name = args.get("name", "") desc = args.get("description", "") extra = args.get("payload", {}) or {} entry = GlossaryEntry( world_id=ctx.world.id, session_id=ctx.session_id, kind=kind, name=entry_name, description=desc, payload=extra, ) ctx.db.add(entry) await ctx.db.flush() rag = await get_rag(ctx.settings_map) await rag.upsert_glossary( world_id=ctx.world.id, entry_id=entry.id, kind=kind, name=entry_name, description=desc, payload=extra, settings_map=ctx.settings_map, ) ctx.rag_added.append({"kind": kind, "name": entry_name, "description": desc}) return {"ok": True, "entry_id": str(entry.id)} if name == "schedule_trigger": fire_at = args.get("fire_at", "") description = args.get("description", "") payload = args.get("payload", {}) or {} trigger = DeferredTrigger( session_id=ctx.session_id, fire_at=fire_at, description=description, payload=payload, ) ctx.db.add(trigger) await ctx.db.flush() ctx.scheduled_triggers.append({ "id": str(trigger.id), "fire_at": fire_at, "description": description, }) return {"ok": True, "trigger_id": str(trigger.id)} if name == "advance_time": days = int(args.get("days", 0)) hours = int(args.get("hours", 0)) minutes = int(args.get("minutes", 0)) ctx.time_advance["days"] += days ctx.time_advance["hours"] += hours ctx.time_advance["minutes"] += minutes return { "ok": True, "advance": {"days": days, "hours": hours, "minutes": minutes}, "reason": args.get("reason", ""), } if name == "run_subagent": if ctx.subagent_runner is None: return {"error": "subagent_runner_not_available"} task = args.get("task", "") context = args.get("context", "") try: result = await ctx.subagent_runner(task, context) return {"result": result} except Exception as e: return {"error": str(e)} return {"error": f"unknown_tool: {name}"} def _summarize_state(state: Dict[str, Any]) -> str: """Quick human-readable summary of state for the LLM.""" if not state: return "(empty)" parts: List[str] = [] player = state.get("player", {}) if player: name = player.get("name", "?") stats = player.get("stats", {}) location = player.get("location", "?") hp = stats.get("health", "?") hp_max = stats.get("health_max", "?") mp = stats.get("mana", "?") parts.append(f"player={name} hp={hp}/{hp_max} mp={mp} loc={location}") inv = player.get("inventory", []) if isinstance(player, dict) else [] if inv: parts.append("inv=" + ", ".join(f"{i.get('name','?')}x{i.get('qty',1)}" for i in inv[:8])) npcs = state.get("npcs", []) if npcs: parts.append(f"npcs={len(npcs)}") return " | ".join(parts)