210 lines
8.8 KiB
Python
210 lines
8.8 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Local mock test for the MCP endpoint (POST /mcp JSON-RPC).
|
|
|
|
Covers: initialize, tools/list, tools/call for all 9 tools, plus
|
|
error paths (parse error, unknown method, missing required code) and
|
|
notification 202. Run with local python (3.6-compatible stdlib only).
|
|
"""
|
|
import sys, io, os, json, time, urllib.request
|
|
|
|
if sys.stdout is None:
|
|
sys.stdout = io.StringIO()
|
|
if sys.stderr is None:
|
|
sys.stderr = io.StringIO()
|
|
|
|
SRC = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src")
|
|
sys.path.insert(0, SRC)
|
|
|
|
import bridge_util, bridge_http_server, bridge_data_adapter
|
|
|
|
class FakeDF:
|
|
"""Minimal pandas.DataFrame stand-in (columns/shape/iat/index) so the
|
|
test runs on stdlib-only Python (no pandas needed locally)."""
|
|
def __init__(self, columns, rows):
|
|
self.columns = columns
|
|
self._rows = rows
|
|
@property
|
|
def shape(self):
|
|
return (len(self._rows), len(self.columns))
|
|
def __len__(self):
|
|
return len(self._rows)
|
|
def __iter__(self):
|
|
return iter(self.columns)
|
|
@property
|
|
def index(self):
|
|
return ["20260818", "20260819"]
|
|
def iat(self, i, j):
|
|
return self._rows[i][j]
|
|
|
|
class FakeCtx:
|
|
def get_market_data_ex(self, fields, codes, **kw):
|
|
cols = ["open", "high", "low", "close", "volume", "amount"]
|
|
rows = [[10.0, 10.5, 9.8, 10.2, 100.0, 1000.0],
|
|
[11.0, 11.5, 10.6, 11.2, 200.0, 2200.0]]
|
|
return {codes[0]: FakeDF(cols, rows)}
|
|
def get_full_tick(self, codes):
|
|
return {c: {"lastPrice": 10.2, "volume": 500,
|
|
"askPrice": [10.3], "bidPrice": [10.1]} for c in codes}
|
|
def get_instrument_detail(self, code):
|
|
return {"InstrumentName": "Test", "UpStopPrice": 11.0,
|
|
"DownStopPrice": 9.0, "OpenDate": "20010827"}
|
|
def get_trading_dates(self, *a, **kw):
|
|
return ["20260818", "20260819", "20260820"]
|
|
|
|
bridge_util.CTX = FakeCtx()
|
|
bridge_util.ACCOUNT = "TESTACC"
|
|
bridge_util.ACCOUNT_TYPE = "STOCK"
|
|
|
|
def fake_trade(acct, atype, dtype, strategy_name=None):
|
|
if dtype == "position":
|
|
return [type("Pos", (), {"m_strInstrumentID": "600519", "m_strInstrumentName": "Moutai",
|
|
"m_nVolume": 100, "m_nCanUseVolume": 50,
|
|
"m_nFrozenVolume": 50, "m_dOpenPrice": 1500.0,
|
|
"m_dFloatProfit": 1000.0})()]
|
|
if dtype == "order":
|
|
return [type("O", (), {"m_strInstrumentID": "600519", "m_nOrderStatus": 48})()]
|
|
return [type("D", (), {"m_strInstrumentID": "600519", "m_nVolume": 100,
|
|
"m_dPrice": 1500.0})()]
|
|
|
|
bridge_util.QMT_API["get_trade_detail_data"] = fake_trade
|
|
bridge_http_server.PORT = 18692
|
|
bridge_http_server.TOKEN = ""
|
|
bridge_http_server.ACCOUNT = "TESTACC"
|
|
bridge_util._refresh_trade_cache("TESTACC", "STOCK")
|
|
bridge_http_server.start_server()
|
|
time.sleep(0.8)
|
|
|
|
BASE = "http://127.0.0.1:18692/mcp"
|
|
|
|
def post(payload):
|
|
data = json.dumps(payload).encode()
|
|
req = urllib.request.Request(BASE, data=data,
|
|
headers={"Content-Type": "application/json"})
|
|
with urllib.request.urlopen(req, timeout=5) as r:
|
|
return r.status, json.loads(r.read().decode())
|
|
|
|
def check(name, cond, detail=""):
|
|
if not cond:
|
|
raise AssertionError("%s FAILED %s" % (name, detail))
|
|
print("OK %s" % name)
|
|
|
|
# 1. initialize
|
|
st, r = post({"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
|
"params": {"protocolVersion": "2025-06-18",
|
|
"capabilities": {}, "clientInfo": {"name": "t", "version": "1"}}})
|
|
check("initialize", st == 200 and r["result"]["protocolVersion"] == "2025-06-18"
|
|
and r["result"]["capabilities"].get("tools") is not None,
|
|
str(r))
|
|
|
|
# 2. tools/list -> 9 tools
|
|
st, r = post({"jsonrpc": "2.0", "id": 2, "method": "tools/list"})
|
|
tools = r["result"]["tools"]
|
|
names = [t["name"] for t in tools]
|
|
check("tools/list count", len(tools) == 9, str(names))
|
|
for need in ["qmt_kline", "qmt_quote", "qmt_tick", "qmt_instrument",
|
|
"qmt_trading_dates", "qmt_positions", "qmt_asset",
|
|
"qmt_orders", "qmt_trades"]:
|
|
check("tool %s" % need, need in names)
|
|
|
|
# 3. tools/call qmt_kline
|
|
st, r = post({"jsonrpc": "2.0", "id": 3, "method": "tools/call",
|
|
"params": {"name": "qmt_kline",
|
|
"arguments": {"code": "600519", "period": "1d", "count": 5}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_kline", text["code"] == "600519.SH" and text["count"] == 2, str(text)[:120])
|
|
|
|
# 4. tools/call qmt_quote
|
|
st, r = post({"jsonrpc": "2.0", "id": 4, "method": "tools/call",
|
|
"params": {"name": "qmt_quote", "arguments": {"code": "600519.SH"}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_quote", text["ok"] is True and text["data"]["lastPrice"] == 10.2, str(text)[:120])
|
|
|
|
# 5. tools/call qmt_tick
|
|
st, r = post({"jsonrpc": "2.0", "id": 5, "method": "tools/call",
|
|
"params": {"name": "qmt_tick",
|
|
"arguments": {"codes": "600000.SH,000001.SZ"}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_tick", text["ok"] is True and len(text["data"]) == 2, str(text)[:120])
|
|
|
|
# 6. tools/call qmt_instrument
|
|
st, r = post({"jsonrpc": "2.0", "id": 6, "method": "tools/call",
|
|
"params": {"name": "qmt_instrument", "arguments": {"code": "600519"}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_instrument", text["data"]["UpStopPrice"] == 11.0, str(text)[:120])
|
|
|
|
# 7. tools/call qmt_trading_dates
|
|
st, r = post({"jsonrpc": "2.0", "id": 7, "method": "tools/call",
|
|
"params": {"name": "qmt_trading_dates",
|
|
"arguments": {"start": "20260818", "end": "20260820"}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_trading_dates", text["count"] == 3, str(text)[:120])
|
|
|
|
# 8. tools/call qmt_positions
|
|
st, r = post({"jsonrpc": "2.0", "id": 8, "method": "tools/call",
|
|
"params": {"name": "qmt_positions", "arguments": {}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_positions", text["data"]["summary"]["count"] == 1, str(text)[:120])
|
|
|
|
# 9. tools/call qmt_asset
|
|
st, r = post({"jsonrpc": "2.0", "id": 9, "method": "tools/call",
|
|
"params": {"name": "qmt_asset", "arguments": {}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_asset", "data" in text, str(text)[:120])
|
|
|
|
# 10. tools/call qmt_orders (with code filter)
|
|
st, r = post({"jsonrpc": "2.0", "id": 10, "method": "tools/call",
|
|
"params": {"name": "qmt_orders",
|
|
"arguments": {"code": "600519", "status": "active"}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_orders", text["ok"] is True and len(text["data"]) == 1, str(text)[:120])
|
|
|
|
# 11. tools/call qmt_trades
|
|
st, r = post({"jsonrpc": "2.0", "id": 11, "method": "tools/call",
|
|
"params": {"name": "qmt_trades", "arguments": {}}})
|
|
text = json.loads(r["result"]["content"][0]["text"])
|
|
check("qmt_trades", text["ok"] is True, str(text)[:120])
|
|
|
|
# 12. error: parse error (bad json)
|
|
data = b"{not json"
|
|
req = urllib.request.Request(BASE, data=data,
|
|
headers={"Content-Type": "application/json"})
|
|
with urllib.request.urlopen(req, timeout=5) as rr:
|
|
err = json.loads(rr.read().decode())
|
|
check("parse error", err["error"]["code"] == -32700, str(err))
|
|
|
|
# 13. error: unknown method
|
|
st, r = post({"jsonrpc": "2.0", "id": 13, "method": "bogus/method"})
|
|
check("unknown method", r["error"]["code"] == -32601, str(r))
|
|
|
|
# 14. error: missing required code
|
|
st, r = post({"jsonrpc": "2.0", "id": 14, "method": "tools/call",
|
|
"params": {"name": "qmt_kline", "arguments": {}}})
|
|
check("missing code", r["error"]["code"] == -32602, str(r))
|
|
|
|
# 15. error: unknown tool
|
|
st, r = post({"jsonrpc": "2.0", "id": 15, "method": "tools/call",
|
|
"params": {"name": "qmt_nope", "arguments": {}}})
|
|
check("unknown tool", r["error"]["code"] == -32602, str(r))
|
|
|
|
# 16. notification -> 202 no body
|
|
data = json.dumps({"jsonrpc": "2.0", "method": "notifications/initialized"}).encode()
|
|
req = urllib.request.Request(BASE, data=data,
|
|
headers={"Content-Type": "application/json"})
|
|
with urllib.request.urlopen(req, timeout=5) as rr:
|
|
check("notification 202", rr.status == 202 and rr.read() == b"", "status=%s" % rr.status)
|
|
|
|
# 17. OPTIONS preflight -> 204 with CORS headers
|
|
req = urllib.request.Request("http://127.0.0.1:18692/mcp", method="OPTIONS")
|
|
try:
|
|
with urllib.request.urlopen(req, timeout=5) as rr:
|
|
check("options 204", rr.status == 204
|
|
and rr.headers.get("Access-Control-Allow-Origin") == "*",
|
|
"status=%s" % rr.status)
|
|
except urllib.error.HTTPError as e:
|
|
check("options 204", e.code == 204
|
|
and e.headers.get("Access-Control-Allow-Origin") == "*",
|
|
"status=%s" % e.code)
|
|
|
|
print("=== MCP ENDPOINT OK ===")
|
|
bridge_http_server.stop_server()
|