main
py 71 lines 2.16 KB
Raw
1 from __future__ import annotations
2
3 import pytest
4
5 from app.contracts import (
6 McpErrorCode,
7 McpGatewayError,
8 McpRiskClass,
9 McpToolDefinition,
10 McpToolExecution,
11 McpToolKind,
12 StrictContract,
13 )
14 from app.policy import McpToolPolicy
15 from app.registry import DuplicateMcpToolError, McpToolRegistry
16
17
18 class EmptyInput(StrictContract):
19 pass
20
21
22 async def handler(_arguments, _auth, _request_id):
23 return McpToolExecution(data={})
24
25
26 def definition(name: str, risk: McpRiskClass = McpRiskClass.SAFE_READ) -> McpToolDefinition:
27 return McpToolDefinition(
28 name=name,
29 description="test",
30 input_model=EmptyInput,
31 risk_class=risk,
32 kind=McpToolKind.INTERNAL,
33 handler=handler,
34 )
35
36
37 def test_known_tools_and_schemas_are_discoverable(container) -> None:
38 assert container.registry.get("get_research_readiness") is not None
39 discovered = {item["name"]: item for item in container.registry.discover()}
40 assert len(discovered) == 9
41 assert {item["riskClass"] for item in discovered.values()} == {"SAFE_READ"}
42 assert {item["kind"] for item in discovered.values()} == {"INTERNAL"}
43 assert discovered["get_company_analysis"]["riskClass"] == "SAFE_READ"
44 assert discovered["get_company_analysis"]["kind"] == "INTERNAL"
45 assert discovered["get_company_analysis"]["inputSchema"]["required"] == ["globalInstrumentId"]
46
47
48 def test_duplicate_registration_is_rejected() -> None:
49 registry = McpToolRegistry()
50 registry.register(definition("same"))
51 with pytest.raises(DuplicateMcpToolError, match="already registered"):
52 registry.register(definition("same"))
53
54
55 def test_safe_read_is_allowed(auth) -> None:
56 McpToolPolicy().authorize(definition("read"), auth)
57
58
59 @pytest.mark.parametrize(
60 "risk",
61 [
62 McpRiskClass.SENSITIVE_READ,
63 McpRiskClass.WRITE_NON_FINANCIAL,
64 McpRiskClass.FINANCIAL_ACTION,
65 McpRiskClass.ADMIN_ACTION,
66 ],
67 )
68 def test_every_non_safe_risk_class_is_denied(auth, risk) -> None:
69 with pytest.raises(McpGatewayError) as error:
70 McpToolPolicy().authorize(definition("unsafe", risk), auth)
71 assert error.value.code == McpErrorCode.MCP_TOOL_DENIED