rebase
This commit is contained in:
114
tests/unit/test_tool_registry.py
Normal file
114
tests/unit/test_tool_registry.py
Normal file
@@ -0,0 +1,114 @@
|
||||
"""Tests for `app.engine.tools.base.ToolRegistry` and a few game tools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from app.engine.tools.base import Tool, ToolContext, ToolRegistry, ToolResult
|
||||
|
||||
|
||||
class _StubTool(Tool):
|
||||
name = "stub"
|
||||
category = "game"
|
||||
stages = {"test_stage"}
|
||||
description = "Stub tool for tests"
|
||||
parameters_schema = {"type": "object", "properties": {}}
|
||||
|
||||
def __init__(self, return_ok: bool = True):
|
||||
self.return_ok = return_ok
|
||||
self.calls = []
|
||||
|
||||
async def execute(self, arguments, ctx):
|
||||
self.calls.append(arguments)
|
||||
if self.return_ok:
|
||||
return ToolResult(ok=True, data=arguments, message="ok")
|
||||
return ToolResult(ok=False, error_code="stub_error", error_message="intentional")
|
||||
|
||||
|
||||
def test_registry_register_and_get():
|
||||
reg = ToolRegistry()
|
||||
t = _StubTool()
|
||||
reg.register(t)
|
||||
assert reg.get("stub") is t
|
||||
|
||||
|
||||
def test_registry_duplicate_register_raises():
|
||||
reg = ToolRegistry()
|
||||
reg.register(_StubTool())
|
||||
with pytest.raises(ValueError):
|
||||
reg.register(_StubTool())
|
||||
|
||||
|
||||
def test_registry_list_for_stage_filters():
|
||||
reg = ToolRegistry()
|
||||
t1 = _StubTool()
|
||||
t1.stages = {"a", "b"}
|
||||
t2 = _StubTool()
|
||||
t2.name = "stub2"
|
||||
t2.stages = {"b", "c"}
|
||||
reg.register(t1)
|
||||
reg.register(t2)
|
||||
a_tools = reg.list_for_stage("a")
|
||||
assert a_tools == [t1]
|
||||
b_tools = reg.list_for_stage("b")
|
||||
assert set(b_tools) == {t1, t2}
|
||||
c_tools = reg.list_for_stage("c")
|
||||
assert c_tools == [t2]
|
||||
|
||||
|
||||
def test_registry_to_openai_format():
|
||||
reg = ToolRegistry()
|
||||
reg.register(_StubTool())
|
||||
fmt = reg.to_openai_format("test_stage")
|
||||
assert len(fmt) == 1
|
||||
assert fmt[0]["type"] == "function"
|
||||
assert fmt[0]["function"]["name"] == "stub"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registry_execute_unknown_tool():
|
||||
reg = ToolRegistry()
|
||||
ctx = ToolContext(db=None, world=None)
|
||||
result = await reg.execute("nonexistent", {}, ctx)
|
||||
assert not result.ok
|
||||
assert result.error_code == "unknown_tool"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registry_execute_dispatches():
|
||||
reg = ToolRegistry()
|
||||
t = _StubTool()
|
||||
reg.register(t)
|
||||
ctx = ToolContext(db=None, world=None)
|
||||
result = await reg.execute("stub", {"x": 1}, ctx)
|
||||
assert result.ok
|
||||
assert t.calls == [{"x": 1}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registry_execute_catches_exceptions():
|
||||
class _BoomTool(Tool):
|
||||
name = "boom"
|
||||
stages = {"*"}
|
||||
async def execute(self, arguments, ctx):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
reg = ToolRegistry()
|
||||
reg.register(_BoomTool())
|
||||
ctx = ToolContext(db=None, world=None)
|
||||
result = await reg.execute("boom", {}, ctx)
|
||||
assert not result.ok
|
||||
assert result.error_code == "tool_exception"
|
||||
|
||||
|
||||
def test_tool_result_to_dict_ok():
|
||||
r = ToolResult(ok=True, data={"a": 1}, message="hello")
|
||||
assert r.to_dict() == {"ok": True, "data": {"a": 1}, "message": "hello"}
|
||||
|
||||
|
||||
def test_tool_result_to_dict_error():
|
||||
r = ToolResult(ok=False, error_code="x", error_message="boom")
|
||||
d = r.to_dict()
|
||||
assert d["ok"] is False
|
||||
assert d["error"]["code"] == "x"
|
||||
assert d["error"]["message"] == "boom"
|
||||
Reference in New Issue
Block a user