rebase
This commit is contained in:
231
app/engine/tools/base.py
Normal file
231
app/engine/tools/base.py
Normal file
@@ -0,0 +1,231 @@
|
||||
"""Base types for tool system: Tool, ToolContext, ToolResult, ToolRegistry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.logging import get_logger
|
||||
from app.core.state_validator import apply_patch
|
||||
from app.models import Entity, StepToolCall, World
|
||||
|
||||
_logger = get_logger(__name__)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Context & Result
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class ToolContext:
|
||||
"""Per-iteration context passed to every tool call."""
|
||||
|
||||
db: AsyncSession
|
||||
world: World
|
||||
user_id: uuid.UUID | None = None
|
||||
step_id: uuid.UUID | None = None
|
||||
stage: str = ""
|
||||
sse_emitter: Any = None # callable: async (event, data) -> None
|
||||
pending_state_changes: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolResult:
|
||||
ok: bool
|
||||
data: dict[str, Any] = field(default_factory=dict)
|
||||
message: str = ""
|
||||
error_code: str = ""
|
||||
error_message: str = ""
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
if self.ok:
|
||||
return {"ok": True, "data": self.data, "message": self.message}
|
||||
return {
|
||||
"ok": False,
|
||||
"error": {"code": self.error_code, "message": self.error_message},
|
||||
}
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Tool base class
|
||||
# --------------------------------------------------------------------------- #
|
||||
class Tool(abc.ABC):
|
||||
"""Abstract base for all tools."""
|
||||
|
||||
name: str = ""
|
||||
category: str = "game" # game | interaction | schema
|
||||
stages: set[str] = set() # which stages can use this tool
|
||||
description: str = ""
|
||||
parameters_schema: dict[str, Any] = {}
|
||||
|
||||
@abc.abstractmethod
|
||||
async def execute(self, arguments: dict[str, Any], ctx: ToolContext) -> ToolResult:
|
||||
"""Run the tool. Must be idempotent within a single transaction."""
|
||||
|
||||
def to_openai_format(self) -> dict[str, Any]:
|
||||
"""Serialize to OpenAI tools format."""
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.parameters_schema,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Registry
|
||||
# --------------------------------------------------------------------------- #
|
||||
class ToolRegistry:
|
||||
"""Holds all registered tools and dispatches calls."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._tools: dict[str, Tool] = {}
|
||||
|
||||
def register(self, tool: Tool) -> None:
|
||||
if not tool.name:
|
||||
raise ValueError("Tool name is required")
|
||||
if tool.name in self._tools:
|
||||
raise ValueError(f"Tool {tool.name} already registered")
|
||||
self._tools[tool.name] = tool
|
||||
|
||||
def get(self, name: str) -> Tool | None:
|
||||
return self._tools.get(name)
|
||||
|
||||
def list_for_stage(self, stage: str) -> list[Tool]:
|
||||
"""Return all tools available at the given stage."""
|
||||
return [t for t in self._tools.values() if stage in t.stages or "*" in t.stages]
|
||||
|
||||
def to_openai_format(self, stage: str) -> list[dict[str, Any]]:
|
||||
return [t.to_openai_format() for t in self.list_for_stage(stage)]
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
name: str,
|
||||
arguments: dict[str, Any],
|
||||
ctx: ToolContext,
|
||||
) -> ToolResult:
|
||||
"""Execute a tool by name. Logs to step_tool_calls and emits SSE."""
|
||||
from sqlalchemy import select as sa_select
|
||||
|
||||
tool = self.get(name)
|
||||
executed_at = datetime.now(timezone.utc)
|
||||
if tool is None:
|
||||
result = ToolResult(
|
||||
ok=False,
|
||||
error_code="unknown_tool",
|
||||
error_message=f"Tool {name!r} is not registered",
|
||||
)
|
||||
else:
|
||||
try:
|
||||
result = await tool.execute(arguments, ctx)
|
||||
except Exception as e: # noqa: BLE001
|
||||
_logger.exception("tool_execution_failed", tool=name, error=str(e))
|
||||
result = ToolResult(
|
||||
ok=False,
|
||||
error_code="tool_exception",
|
||||
error_message=str(e),
|
||||
)
|
||||
|
||||
# Log to step_tool_calls (if we have a step_id)
|
||||
if ctx.step_id is not None:
|
||||
try:
|
||||
ctx.db.add(
|
||||
StepToolCall(
|
||||
step_id=ctx.step_id,
|
||||
tool_name=name,
|
||||
arguments=arguments,
|
||||
result=result.to_dict(),
|
||||
is_success=result.ok,
|
||||
executed_at=executed_at,
|
||||
)
|
||||
)
|
||||
await ctx.db.flush()
|
||||
except Exception as e: # noqa: BLE001
|
||||
_logger.error("tool_log_failed", tool=name, error=str(e))
|
||||
|
||||
# Emit SSE
|
||||
if ctx.sse_emitter is not None:
|
||||
try:
|
||||
await ctx.sse_emitter(
|
||||
"tool_call",
|
||||
{
|
||||
"tool": name,
|
||||
"arguments": arguments,
|
||||
"result": result.to_dict(),
|
||||
"is_success": result.ok,
|
||||
},
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
_logger.warning("sse_tool_call_failed", tool=name, error=str(e))
|
||||
|
||||
return result
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Helpers used by entity_* tools
|
||||
# --------------------------------------------------------------------------- #
|
||||
async def get_entity_by_query(
|
||||
db: AsyncSession, world_id: uuid.UUID, query: Any
|
||||
) -> Entity | None:
|
||||
"""Resolve entity by UUID string or by {entity_type, name}."""
|
||||
from sqlalchemy import select as sa_select
|
||||
|
||||
if isinstance(query, str):
|
||||
try:
|
||||
eid = uuid.UUID(query)
|
||||
except ValueError:
|
||||
return None
|
||||
return (
|
||||
await db.execute(
|
||||
sa_select(Entity).where(
|
||||
Entity.id == eid, Entity.world_id == world_id
|
||||
)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
elif isinstance(query, dict):
|
||||
et = query.get("entity_type")
|
||||
nm = query.get("name")
|
||||
if not et or not nm:
|
||||
return None
|
||||
return (
|
||||
await db.execute(
|
||||
sa_select(Entity).where(
|
||||
Entity.world_id == world_id,
|
||||
Entity.entity_type == et,
|
||||
Entity.name == nm,
|
||||
Entity.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
return None
|
||||
|
||||
|
||||
def apply_env_patch(environment: dict, patch: dict) -> tuple[dict, list[str]]:
|
||||
"""Wrapper around state_validator.apply_patch for environment dicts."""
|
||||
return apply_patch(environment, patch)
|
||||
|
||||
|
||||
# Singleton registry (instantiated in `app.engine.tools.__init__`)
|
||||
_registry: ToolRegistry | None = None
|
||||
|
||||
|
||||
def get_registry() -> ToolRegistry:
|
||||
global _registry
|
||||
if _registry is None:
|
||||
from app.engine.tools.register_all import build_default_registry
|
||||
|
||||
_registry = build_default_registry()
|
||||
return _registry
|
||||
|
||||
|
||||
def reset_registry() -> None:
|
||||
"""Reset the cached registry — used in tests."""
|
||||
global _registry
|
||||
_registry = None
|
||||
Reference in New Issue
Block a user