Files
RollCalcPython/tests/test_mcp_adapter.py
T

257 lines
9.7 KiB
Python

import inspect
import sys
import types
import unittest
from unittest.mock import patch
import mcp_server
from roll_calculation import ArticleRepository, calculate_roll, get_article
from rollcalc_mcp_tools import (
analyze_transport_capacity_result,
calculate_roll_diameter_result,
get_article_result,
search_articles_result,
)
from transport_calculation import analyze_transport, transport_presets
class _FakeFastMCP:
def __init__(self, name):
self.name = name
self.tools = []
def tool(self):
def register(function):
self.tools.append(function)
return function
return register
class McpAdapterTests(unittest.TestCase):
def test_get_article_delegates_to_the_domain_for_success_not_found_and_conflict(self):
for arguments in (
{"article_number": "146900"},
{"article_number": "does-not-exist"},
{"article_number": "146900", "article_name_hint": "Other product"},
):
with self.subTest(arguments=arguments):
self.assertEqual(
get_article_result(**arguments),
get_article(**arguments),
)
def test_search_articles_delegates_to_domain_and_preserves_candidates(self):
query = "Bfix NSP 4900, 5,00 x 40 m"
result = search_articles_result(query)
self.assertEqual(result, ArticleRepository.load().search(query))
self.assertEqual(result["status"], "search_results")
self.assertEqual(
[candidate["article_number"] for candidate in result["candidates"]],
["180205", "8180205", "182815", "206900"],
)
self.assertEqual(
result["candidates"][1]["production_site"], "Malaysia"
)
def test_search_articles_preserves_ambiguous_and_not_found_domain_results(self):
ambiguous = search_articles_result("Bentofix NSP 4900")
missing = search_articles_result("not a RollCalc product")
self.assertGreater(ambiguous["total_matches"], 1)
self.assertGreater(len(ambiguous["candidates"]), 1)
self.assertEqual(missing, {
"status": "search_not_found",
"query": "not a RollCalc product",
"total_matches": 0,
"candidates": [],
})
def test_calculate_roll_diameter_delegates_for_normal_weighted_and_invalid_cases(self):
for arguments in (
{
"article_number": "146900",
"roll_length_m": 50.0,
"core_diameter_mm": 150.0,
},
{
"article_number": "146900",
"roll_length_m": 50.0,
"core_diameter_mm": 150.0,
"width_m": 5.8,
"area_weight_g_m2": 1767.2,
},
{"roll_length_m": -1.0, "core_diameter_mm": 150.0, "thickness_mm": 2.0},
):
with self.subTest(arguments=arguments):
actual = calculate_roll_diameter_result(**arguments)
expected = calculate_roll({**arguments, "include_roll_weight": True})
if expected["status"] == "success":
self.assertEqual(
actual["transport_roll_inputs"]["roll_diameter_mm"],
845.9,
)
del actual["transport_roll_inputs"]
self.assertEqual(actual, expected)
def test_roll_calculation_adapter_always_requests_weight_from_the_domain(self):
result = calculate_roll_diameter_result(
article_number="114030",
roll_length_m=15.0,
core_diameter_mm=100.0,
)
self.assertEqual(result["status"], "success")
self.assertIsNotNone(result["calculation"]["roll_weight_kg"])
def test_roll_calculation_adapter_uses_shared_transport_input_helper(self):
expected_inputs = {"roll_diameter_mm": 123.4}
with patch(
"rollcalc_mcp_tools.transport_roll_inputs_from_roll_calculation",
return_value=expected_inputs,
) as helper:
result = calculate_roll_diameter_result(
article_number="114030",
roll_length_m=15.0,
core_diameter_mm=100.0,
)
helper.assert_called_once()
self.assertEqual(result["transport_roll_inputs"], expected_inputs)
def test_transport_tool_delegates_to_domain_with_lkw_sattelzug_preset(self):
arguments = {
"transport_preset": "lkw_sattelzug",
"roll_diameter_mm": 1000.0,
"core_diameter_mm": 150.0,
"roll_width_m": 2.0,
"roll_weight_kg": 1000.0,
"product_length_m": 50.0,
}
self.assertEqual(
analyze_transport_capacity_result(**arguments),
analyze_transport(arguments),
)
result = analyze_transport_capacity_result(**arguments)
self.assertEqual(result["status"], "success")
self.assertEqual(result["analysis"]["transport_preset"], "lkw_sattelzug")
self.assertEqual(result["analysis"]["final_rolls"], 24)
self.assertEqual(result["analysis"]["limiting"]["name"], "Weight")
def test_server_registers_the_existing_tools_and_article_search(self):
mcp = types.ModuleType("mcp")
server = types.ModuleType("mcp.server")
fastmcp = types.ModuleType("mcp.server.fastmcp")
fastmcp.FastMCP = _FakeFastMCP
with patch.dict(sys.modules, {
"mcp": mcp,
"mcp.server": server,
"mcp.server.fastmcp": fastmcp,
}):
instance = mcp_server.create_server()
self.assertEqual(instance.name, "RollCalc")
self.assertEqual(
[tool.__name__ for tool in instance.tools],
[
"get_article",
"search_articles",
"calculate_roll_diameter",
"analyze_transport_capacity",
],
)
def test_server_schema_exposes_weight_and_transport_chaining_contract(self):
instance = mcp_server.create_server()
tools = instance._tool_manager._tools
search_tool = tools["search_articles"]
roll_schema = tools["calculate_roll_diameter"].parameters
transport_tool = tools["analyze_transport_capacity"]
transport_schema = transport_tool.parameters
self.assertEqual(search_tool.parameters["required"], ["query"])
self.assertEqual(
search_tool.parameters["properties"]["query"]["type"], "string"
)
self.assertNotIn("include_roll_weight", roll_schema["properties"])
preset_schema = transport_schema["properties"]["transport_preset"]
self.assertEqual(
preset_schema["enum"],
[preset.key for preset in transport_presets()],
)
self.assertEqual(
transport_schema["required"],
[
"transport_preset",
"roll_diameter_mm",
"core_diameter_mm",
"roll_width_m",
"roll_weight_kg",
"product_length_m",
],
)
for field in transport_schema["required"]:
if field == "transport_preset":
continue
self.assertEqual(
transport_schema["properties"][field]["type"], "number"
)
for field in (
"length_m",
"width_m",
"height_m",
"max_weight_kg",
"margin_side_m",
"margin_ceiling_m",
):
self.assertNotIn(field, transport_schema["properties"])
self.assertNotIn(
field,
inspect.signature(analyze_transport_capacity_result).parameters,
)
description = transport_tool.description
for text in (
"Use this tool when the user asks how many rolls fit on a known RollCalc transport type",
"do not ask the user for vehicle dimensions or payload",
"LKW-Sattelzug -> lkw_sattelzug",
"LKW-Tandem -> lkw_tandem",
"20ft Container -> container_20ft",
"40ft Container -> container_40ft",
"40ft High Cube -> container_40ft_hc",
'transport_preset=\"lkw_sattelzug\"',
"first call get_article and calculate_roll_diameter",
"transport_roll_inputs bundle unchanged",
"Do not choose among minimum, average, or maximum diameter",
"product_length_m",
"do not estimate or invent",
"stateless tool does not recalculate a roll",
):
self.assertIn(text, description)
search_description = search_tool.description
for text in (
"name, family, designation, or descriptive article text",
"Do not guess or invent an article number",
"deterministic candidates from the RollCalc article master data",
"total_matches is exactly 1",
"use that candidate's exact article_number",
"do NOT select the first or highest-ranked candidate",
"ask the user which article number is intended",
"do not calculate yet",
"no matching RollCalc article was found",
"more specific designation or article number",
"do not invent an article",
"production_site=Malaysia is descriptive metadata only",
"does not imply Bentofix, Bento 2, or any production machine",
"Ranking/order expresses relevance only",
"get_article -> calculate_roll_diameter -> optionally analyze_transport_capacity",
):
self.assertIn(text, search_description)
if __name__ == "__main__":
unittest.main()