main
py 156 lines 5.25 KB
Raw
1 #!/usr/bin/env python3
2 """End-to-end LOCAL MCP smoke test over the official STDIO transport."""
3 from __future__ import annotations
4
5 import asyncio
6 import json
7 import sys
8 import threading
9 from contextlib import contextmanager
10 from datetime import datetime, timezone
11 from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
12 from pathlib import Path
13 from urllib.parse import urlparse
14
15 from mcp import Client, StdioServerParameters
16
17
18 ROOT = Path(__file__).resolve().parents[1]
19 INSTRUMENT_ID = "11111111-1111-4111-8111-111111111111"
20 NOW = datetime(2026, 9, 11, 10, 0, tzinfo=timezone.utc).isoformat()
21
22
23 class PersistedReadFixture(BaseHTTPRequestHandler):
24 def log_message(self, _format: str, *_args) -> None:
25 return
26
27 def do_GET(self) -> None: # noqa: N802 - stdlib callback name
28 path = urlparse(self.path).path
29 if path == f"/api/v1/research/readiness/{INSTRUMENT_ID}":
30 self._json(
31 {
32 "globalInstrumentId": INSTRUMENT_ID,
33 "overallStatus": "READY",
34 "requirements": [],
35 "generatedAt": NOW,
36 }
37 )
38 return
39 if path == "/api/v1/research/sector-performance":
40 self._json(
41 {
42 "region": "INDIA",
43 "sector": "FINANCIAL_SERVICES",
44 "period": "MONTH",
45 "leaders": [],
46 "asOf": NOW,
47 }
48 )
49 return
50 self._json({"message": "not found"}, status=404)
51
52 def do_POST(self) -> None: # noqa: N802 - stdlib callback name
53 path = urlparse(self.path).path
54 length = int(self.headers.get("Content-Length", "0"))
55 if length:
56 self.rfile.read(length)
57 if path == f"/api/v1/research/analysis/{INSTRUMENT_ID}":
58 self._json(
59 {
60 "globalInstrumentId": INSTRUMENT_ID,
61 "ruleEngineVersion": "STOCK_RULE_ENGINE_V1",
62 "score": 78,
63 "generatedAt": NOW,
64 }
65 )
66 return
67 self._json({"message": "not found"}, status=404)
68
69 def _json(self, value: dict, *, status: int = 200) -> None:
70 body = json.dumps(value).encode("utf-8")
71 self.send_response(status)
72 self.send_header("Content-Type", "application/json")
73 self.send_header("Content-Length", str(len(body)))
74 self.end_headers()
75 self.wfile.write(body)
76
77
78 @contextmanager
79 def persisted_read_server():
80 server = ThreadingHTTPServer(("127.0.0.1", 0), PersistedReadFixture)
81 thread = threading.Thread(target=server.serve_forever, daemon=True)
82 thread.start()
83 try:
84 yield f"http://127.0.0.1:{server.server_port}"
85 finally:
86 server.shutdown()
87 server.server_close()
88 thread.join(timeout=2)
89
90
91 async def rejected(client: Client, tool: str) -> bool:
92 try:
93 result = await client.call_tool(tool, {})
94 except Exception:
95 return True
96 return bool(result.is_error)
97
98
99 async def run_smoke(research_base_url: str) -> None:
100 environment = {
101 "AIP_ENVIRONMENT": "LOCAL",
102 "AIP_FEATURE_MCP_ENABLED": "true",
103 "AIP_MCP_TRANSPORT": "stdio",
104 "AIP_MCP_RESEARCH_BASE_URL": research_base_url,
105 "AIP_MCP_SERVICE_IDENTITY": "local-mcp-smoke",
106 "AIP_MCP_SCOPES": "mcp:read",
107 "AIP_MCP_EXTERNAL_PROVIDERS_ENABLED": "false",
108 "PYTHONUNBUFFERED": "1",
109 }
110 parameters = StdioServerParameters(
111 command=sys.executable,
112 args=["-m", "app.main", "--transport", "stdio"],
113 env=environment,
114 cwd=ROOT,
115 )
116 async with Client(parameters) as client:
117 print("initialize_mcp=PASS")
118 listing = await client.list_tools()
119 names = {tool.name for tool in listing.tools}
120 assert len(names) == 9 and "get_research_readiness" in names
121 print(f"list_tools=PASS count={len(names)}")
122
123 readiness = await client.call_tool(
124 "get_research_readiness", {"globalInstrumentId": INSTRUMENT_ID}
125 )
126 assert readiness.structured_content["data"]["overallStatus"] == "READY"
127 print("get_research_readiness=PASS provider_calls=0")
128
129 analysis = await client.call_tool(
130 "get_company_analysis", {"globalInstrumentId": INSTRUMENT_ID}
131 )
132 assert analysis.structured_content["provenance"]["ruleEngineVersion"] == "STOCK_RULE_ENGINE_V1"
133 print("get_company_analysis=PASS rule_engine=STOCK_RULE_ENGINE_V1 provider_calls=0")
134
135 sector = await client.call_tool(
136 "get_sector_performance",
137 {"region": "INDIA", "sector": "FINANCIAL_SERVICES", "period": "MONTH"},
138 )
139 assert sector.structured_content["data"]["sector"] == "FINANCIAL_SERVICES"
140 print("get_sector_performance=PASS provider_calls=0")
141
142 assert await rejected(client, "not_a_tool")
143 print("invalid_tool_rejected=PASS")
144 assert await rejected(client, "place_order")
145 print("unsafe_tool_unavailable=PASS")
146 print("clean_shutdown=PASS")
147
148
149 def main() -> None:
150 with persisted_read_server() as research_base_url:
151 asyncio.run(run_smoke(research_base_url))
152 print("local_mcp_smoke=PASS")
153
154
155 if __name__ == "__main__":
156 main()