170 lines
6.0 KiB
Python
170 lines
6.0 KiB
Python
"""tools.py + MCP surface — small-model-friendly tool schemas (plan §6, §13).
|
|
|
|
Checks the tool *surface* (names, short descriptions, small input schemas)
|
|
over the real FastMCP server using the in-process client, plus a golden
|
|
result for the headline workflow.
|
|
|
|
Note: the plan's "ten tools" are the reads + aggregates; the scaffold adds
|
|
four admin/setup tools (scan/map/export/import), so the registered surface
|
|
is fourteen functions.
|
|
"""
|
|
|
|
import json
|
|
|
|
import pytest
|
|
from mcp.client import Client
|
|
|
|
from mcp_server.tools import TOOL_FUNCTIONS, register_tools
|
|
|
|
EXPECTED_TOOLS = {
|
|
"find_experiments",
|
|
"get_experiment_context",
|
|
"get_next_protocol_step",
|
|
"list_experiment_templates",
|
|
"validate_protocol_template",
|
|
"complete_next_protocol_step",
|
|
"complete_protocol_step",
|
|
"record_protocol_observation",
|
|
"create_experiment_from_template",
|
|
"adjust_inventory",
|
|
"scan_instance",
|
|
"map_resource",
|
|
"export_mapping",
|
|
"import_mapping",
|
|
}
|
|
|
|
|
|
def result_json(call_result):
|
|
"""Extract the JSON payload from a CallToolResult (structured or text)."""
|
|
structured = getattr(call_result, "structured_content", None)
|
|
if structured is not None:
|
|
return structured
|
|
for block in call_result.content:
|
|
text = getattr(block, "text", None)
|
|
if text:
|
|
return json.loads(text)
|
|
raise AssertionError(f"No usable content in tool result: {call_result!r}")
|
|
|
|
|
|
@pytest.fixture
|
|
async def mcp_server(patch_build_services):
|
|
from mcp.server.fastmcp import FastMCP
|
|
|
|
patch_build_services()
|
|
mcp = FastMCP("labvoice-test")
|
|
register_tools(mcp)
|
|
return mcp
|
|
|
|
|
|
@pytest.fixture
|
|
async def mcp_client(mcp_server):
|
|
async with Client(mcp_server, raise_exceptions=True) as client:
|
|
yield client
|
|
|
|
|
|
# --- tool surface ------------------------------------------------------------
|
|
|
|
|
|
async def test_all_tools_registered_and_listed(mcp_client):
|
|
tools = await mcp_client.list_tools()
|
|
names = {t.name for t in tools}
|
|
assert names == EXPECTED_TOOLS
|
|
assert len(TOOL_FUNCTIONS) == len(EXPECTED_TOOLS)
|
|
|
|
|
|
async def test_every_tool_has_a_short_description(mcp_client):
|
|
"""Plan §6: tool descriptions are one or two short sentences."""
|
|
for tool in await mcp_client.list_tools():
|
|
assert tool.description, f"{tool.name} needs a description"
|
|
assert len(tool.description) <= 300, f"{tool.name} description too long for a small model"
|
|
|
|
|
|
async def test_input_schemas_are_small(mcp_client):
|
|
"""Plan §13 lint: small input schemas — no kitchen-sink parameter lists."""
|
|
for tool in await mcp_client.list_tools():
|
|
properties = tool.inputSchema.get("properties", {})
|
|
assert len(properties) <= 5, f"{tool.name} takes {len(properties)} parameters"
|
|
assert tool.inputSchema.get("type") == "object"
|
|
|
|
|
|
async def test_input_schemas_match_function_signatures(mcp_client):
|
|
import inspect
|
|
|
|
for tool in await mcp_client.list_tools():
|
|
fn = next(f for f in TOOL_FUNCTIONS if f.__name__ == tool.name)
|
|
params = inspect.signature(fn).parameters
|
|
properties = tool.inputSchema.get("properties", {})
|
|
assert set(properties) == set(params), f"{tool.name}: schema/args mismatch"
|
|
required = set(tool.inputSchema.get("required", []))
|
|
expected_required = {
|
|
name for name, p in params.items() if p.default is inspect.Parameter.empty
|
|
}
|
|
assert required == expected_required, f"{tool.name}: required set mismatch"
|
|
|
|
|
|
async def test_mutating_tools_are_clearly_named(mcp_client):
|
|
"""A small model must not confuse reads with mutations (plan §2)."""
|
|
mutating = {
|
|
"complete_next_protocol_step",
|
|
"complete_protocol_step",
|
|
"record_protocol_observation",
|
|
"create_experiment_from_template",
|
|
"adjust_inventory",
|
|
"map_resource",
|
|
"scan_instance",
|
|
"import_mapping",
|
|
}
|
|
tools = {t.name: t for t in await mcp_client.list_tools()}
|
|
for name in mutating:
|
|
assert name in tools
|
|
|
|
|
|
# --- golden workflow result over the MCP transport -----------------------------
|
|
|
|
|
|
async def test_headline_workflow_golden_result_over_mcp(mcp_client, client):
|
|
"""Plan §13: golden result for complete_next_protocol_step (§7)."""
|
|
call_result = await mcp_client.call_tool(
|
|
"complete_next_protocol_step",
|
|
{"experiment_id": 123, "comment": "Done at the bench"},
|
|
)
|
|
payload = result_json(call_result)
|
|
|
|
assert payload["ok"] is True
|
|
assert payload["status"] == "completed"
|
|
assert payload["experiment_id"] == 123
|
|
assert payload["step"]["id"] == 9
|
|
assert payload["step"]["finished"] is True
|
|
assert payload["consumed"][0]["resource_key"] == "ethanol_absolute"
|
|
assert payload["consumed"][0]["container_id"] == 31
|
|
assert payload["consumed"][0]["amount"] == "2.0 mL"
|
|
assert payload["consumed"][0]["remaining"] == "48.0 mL"
|
|
assert payload["next_step"]["id"] == 10
|
|
# stock really moved in eLabFTW
|
|
assert client.items[12].containers[0].qty_stored == 48.0
|
|
|
|
|
|
async def test_read_tool_over_mcp_returns_compact_json(mcp_client):
|
|
call_result = await mcp_client.call_tool("get_next_protocol_step", {"experiment_id": 123})
|
|
payload = result_json(call_result)
|
|
assert payload["step"]["id"] == 9
|
|
assert payload["consumables"][0]["resource_key"] == "ethanol_absolute"
|
|
|
|
|
|
async def test_clarification_over_mcp_is_a_typed_error(mcp_client, resolver_stub):
|
|
resolver_stub.mappings.clear()
|
|
call_result = await mcp_client.call_tool(
|
|
"complete_next_protocol_step", {"experiment_id": 123}
|
|
)
|
|
payload = result_json(call_result)
|
|
# Either the transport surfaces the error or the payload carries it —
|
|
# both are acceptable, but it must be the clarification kind either way.
|
|
if getattr(call_result, "is_error", False):
|
|
text = " ".join(getattr(b, "text", "") for b in call_result.content)
|
|
assert "clarification" in text or "ethanol" in text
|
|
else:
|
|
assert payload.get("error") == "clarification" or payload.get("status") in (
|
|
"clarification",
|
|
"reverted",
|
|
)
|