Files
worldquant-alpha-system/backend/tests/test_mcp.py
T

281 lines
15 KiB
Python
Raw Normal View History

"""MCP integration uses real database/transport and a synthetic WorldQuant only."""
import asyncio
from datetime import timedelta
import httpx
import pytest
from sqlalchemy import func, select
from app.mcp_api.auth import SCOPES, authenticate, create_token
from app.models import (
BacktestPreview,
BacktestResult,
BacktestRun,
MCPAudit,
MCPToken,
Pnl,
now,
)
from tests.test_backtests import candidate, execute, setup
ENDPOINT = "/api/v1/mcp/"
@pytest.fixture
async def mcp_app(settings):
from app.main import create_app
from app.models import Base
from app.worldquant import WqClient
settings.mcp_enabled = True
def no_network(request):
raise AssertionError("No real platform")
application = create_app(settings, WqClient(settings, transport=httpx.MockTransport(no_network)))
async with application.state.engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
ready, stop = asyncio.Event(), asyncio.Event()
async def lifespan():
async with application.router.lifespan_context(application):
ready.set()
await stop.wait()
task = asyncio.create_task(lifespan())
await ready.wait()
try:
await setup(application)
yield application
finally:
stop.set()
await task
async def credentials(app, scopes=SCOPES):
async with app.state.sessions.begin() as db:
row, secret = await create_token(db, "synthetic test", scopes)
principal = await authenticate(db, secret)
return principal, secret
def submission(key="batch-1", items=None, **extra):
return {"name": "MCP batch", "candidates": items or [candidate()], "idempotency_key": key, **extra}
async def invoke(app, principal, name, arguments=None):
result = await app.state.mcp.invoke(principal, name, arguments or {})
assert not result.is_error, result.structured_content
return result.structured_content
async def test_http_auth_and_discovery(mcp_app):
principal, secret = await credentials(mcp_app, {"research:read"})
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mcp_app), base_url="http://testserver") as client:
payload = {"jsonrpc": "2.0", "id": 1, "method": "tools/list"}
assert (await client.post(ENDPOINT, json=payload)).status_code == 401
headers = {"Authorization": f"Bearer {secret}", "Accept": "application/json, text/event-stream"}
response = await client.post(ENDPOINT, json=payload, headers=headers)
assert response.status_code == 200, response.text
names = {t["name"] for t in response.json()["result"]["tools"]}
assert "search_backtests" in names and "submit_backtests" not in names
denied = {"jsonrpc": "2.0", "id": 2, "method": "tools/call", "params": {"name": "submit_backtests", "arguments": submission()}}
assert (await client.post(ENDPOINT, json=denied, headers=headers)).status_code == 403
assert (await client.get("/api/v1/account", headers=headers)).status_code == 401
assert (await client.post(ENDPOINT, json=payload, headers=headers | {"Origin": "https://bad.test"})).status_code == 403
assert (await client.post(ENDPOINT, json=payload, headers=headers | {"Host": "bad.test"})).status_code == 403
async with mcp_app.state.sessions.begin() as db:
row = await db.get(MCPToken, principal.token_id)
row.revoked_at = now()
assert (await client.post(ENDPOINT, json=payload, headers=headers)).status_code == 401
async def test_submit_replay_duplicates_and_rotation(mcp_app):
principal, _ = await credentials(mcp_app)
result = await invoke(mcp_app, principal, "submit_backtests", submission())
second, _ = await credentials(mcp_app)
replay = await invoke(mcp_app, second, "submit_backtests", submission())
assert replay == result
conflict = await mcp_app.state.mcp.invoke(second, "submit_backtests", submission(items=[candidate(1)]))
assert conflict.is_error and conflict.structured_content["error"]["code"] == "IDEMPOTENCY_CONFLICT"
duplicate = await mcp_app.state.mcp.invoke(second, "submit_backtests", submission("different"))
assert duplicate.structured_content["error"]["code"] == "DUPLICATE_INPUT"
rerun = await invoke(mcp_app, second, "submit_backtests", submission("different", duplicate_policy="rerun"))
assert rerun["backtest_run_id"] != result["backtest_run_id"]
async with mcp_app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 2
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 2
audits = list(await db.scalars(select(MCPAudit)))
assert len(audits) == 5
assert result["source"]["kind"] == "mcp"
assert "#backtests?run_id=" in result["web_url"]
async def test_validation_and_within_batch_atomic(mcp_app):
principal, _ = await credentials(mcp_app)
incomplete = candidate()
incomplete["settings"] = {"region": "USA", "universe": "TOP3000", "delay": 1}
result = await mcp_app.state.mcp.invoke(principal, "submit_backtests", submission(items=[incomplete]))
assert result.structured_content["error"]["code"] == "INVALID_INPUT"
result = await mcp_app.state.mcp.invoke(principal, "submit_backtests", submission(items=[candidate(), candidate() | {"client_item_id": "second"}]))
assert result.structured_content["error"]["code"] == "DUPLICATE_INPUT"
# Failed requests did not consume the key or leave previews.
await invoke(mcp_app, principal, "submit_backtests", submission())
async with mcp_app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestPreview)) == 1
async def test_results_control_and_evidence(mcp_app):
principal, _ = await credentials(mcp_app)
result = await invoke(mcp_app, principal, "submit_backtests", submission(items=[candidate(0), candidate(1)]))
rid = result["backtest_run_id"]
control = {"run_id": rid, "action": "pause", "expected_version": 1, "idempotency_key": "pause"}
paused = await invoke(mcp_app, principal, "control_backtest", control)
assert paused == await invoke(mcp_app, principal, "control_backtest", control)
await invoke(mcp_app, principal, "control_backtest", control | {"action": "resume", "expected_version": 2, "idempotency_key": "resume"})
await execute(mcp_app, mcp_app.state.runner.backtests, rid)
data = await invoke(mcp_app, principal, "get_backtest_results", {"run_id": rid, "limit": 1})
assert data["has_more"] and data["items"][0]["metrics"]["is"]["sharpe"] is None
item = data["items"][0]
async with mcp_app.state.sessions.begin() as db:
saved = await db.get(BacktestResult, item["id"])
saved.snapshot = {**saved.snapshot, "is": {"checks": [{"name": "a", "result": "FAIL"}, {"name": "b", "result": "NEW_STATUS"}]}}
data = await invoke(mcp_app, principal, "get_backtest_results", {"run_id": rid})
assert data["items"][0]["checks"]["counts"]["FAIL"] == 1
assert data["items"][0]["checks"]["counts"]["UNKNOWN"] == 1
pnl = await invoke(mcp_app, principal, "get_backtest_artifact", {"item_id": item["id"], "kind": "pnl"})
assert pnl["status"] == "not_cached"
snapshot = await invoke(mcp_app, principal, "get_backtest_artifact", {"item_id": item["id"], "kind": "snapshot", "limit": 1})
assert snapshot["has_more"]
history = await invoke(mcp_app, principal, "search_backtests", {"candidates": [candidate()], "source": "mcp"})
assert history["total"] == 1 and history["items"][0]["match_type"] == "exact_input"
progress = await invoke(mcp_app, principal, "get_backtest", {"run_id": rid, "after": 0, "event_limit": 1})
assert progress["events"]["has_more"] and progress["submission_counts"]["actual_platform_consumption"] is None
async def test_official_sdk_client_and_error_contract(mcp_app):
import httpx2
from mcp import ClientSession
from mcp.client.streamable_http import streamable_http_client
principal, secret = await credentials(mcp_app)
async with httpx2.AsyncClient(transport=httpx2.ASGITransport(app=mcp_app),
headers={"Authorization": f"Bearer {secret}"}) as http:
async with streamable_http_client("http://testserver/api/v1/mcp/", http_client=http) as streams:
async with ClientSession(streams[0], streams[1]) as client:
await client.initialize()
listed = await client.list_tools()
assert len(listed.tools) == 11
caps = await client.call_tool("get_research_capabilities", {})
assert caps.structured_content["max_candidates"] == 100
result = await client.call_tool("submit_backtests", submission())
assert not result.is_error, result
rid = result.structured_content["backtest_run_id"]
await execute(mcp_app, mcp_app.state.runner.backtests, rid)
results = await client.call_tool("get_backtest_results", {"run_id": rid})
assert results.structured_content["items"][0]["persistence_status"] == "saved"
bad = await client.call_tool("submit_backtests", submission("duplicate"))
assert bad.is_error and bad.structured_content["error"]["code"] == "DUPLICATE_INPUT"
await client.call_tool("get_backtest", {"run_id": rid})
async def test_expiry_binding_disabled_and_audit_redaction(app, mcp_app):
from fastapi import HTTPException
from app.models import Account
principal, secret = await credentials(mcp_app)
await mcp_app.state.mcp.invoke(principal, "get_research_capabilities", {}, request_id=secret)
async with mcp_app.state.sessions.begin() as db:
audit = await db.scalar(select(MCPAudit))
assert secret not in audit.request_id
row = await db.get(MCPToken, principal.token_id)
row.expires_at = now() - timedelta(seconds=1)
async with mcp_app.state.sessions() as db:
with pytest.raises(HTTPException):
await authenticate(db, secret)
_, another = await credentials(mcp_app)
async with mcp_app.state.sessions.begin() as db:
account = await db.get(Account, 1)
account.wq_user_id = "CHANGED"
async with mcp_app.state.sessions() as db:
with pytest.raises(HTTPException):
await authenticate(db, another)
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client:
assert (await client.post(ENDPOINT, json={})).status_code == 404
async def test_metadata_refresh_and_pnl_pagination(mcp_app):
from app.worldquant import WqClient
from tests.research_metadata_fake import response
principal, _ = await credentials(mcp_app)
caps = await invoke(mcp_app, principal, "get_research_metadata", {"query": {"kind": "settings"}})
assert caps["status"] == "not_cached"
await mcp_app.state.runner.client.close()
from tests.backtest_fake import Platform
platform = Platform()
client = WqClient(mcp_app.state.settings, transport=httpx.MockTransport(lambda request: response(request) or platform(request)))
await client.authenticate("synthetic@example.com", "synthetic-password")
mcp_app.state.runner.client = client
mcp_app.state.runner.backtests.client = client
operators = await invoke(mcp_app, principal, "refresh_research_data", {"query": {"kind": "operators"}})
assert operators["status"] == "completed"
page = await invoke(mcp_app, principal, "get_research_metadata", {"query": {"kind": "operators", "limit": 1}})
assert page["total"] == 2 and page["has_more"]
await invoke(mcp_app, principal, "refresh_research_data", {"query": {"kind": "settings"}})
await invoke(mcp_app, principal, "refresh_research_data", {"query": {"kind": "field_availability", "scope": {"region": "USA", "universe": "TOP3000", "delay": 1}, "field_id": "close"}})
invalid = await mcp_app.state.mcp.invoke(principal, "submit_backtests", submission(items=[candidate(region="INVALID")]))
assert invalid.structured_content["error"]["code"] == "UNSUPPORTED_SETTINGS"
run = await invoke(mcp_app, principal, "submit_backtests", submission())
await execute(mcp_app, mcp_app.state.runner.backtests, run["backtest_run_id"])
result = await invoke(mcp_app, principal, "get_backtest_results", {"run_id": run["backtest_run_id"]})
item = result["items"][0]
async with mcp_app.state.sessions.begin() as db:
db.add(Pnl(alpha_id=item["alpha_id"], raw={}, points=[{"date": f"2025-01-0{i}", "value": i} for i in range(1, 4)]))
pnl = await invoke(mcp_app, principal, "get_backtest_artifact", {"item_id": item["id"], "kind": "pnl", "limit": 1, "date_from": "2025-01-02"})
assert pnl["total"] == 2 and pnl["items"][0]["value"] == 2 and pnl["has_more"]
async def test_lost_http_response_replays_without_new_run(mcp_app):
import json
_, secret = await credentials(mcp_app)
class LoseResponse(httpx.ASGITransport):
dropped = False
async def handle_async_request(self, request):
response = await super().handle_async_request(request)
if not self.dropped:
self.dropped = True
await response.aread()
raise httpx.ReadError("synthetic lost response", request=request)
return response
transport = LoseResponse(app=mcp_app)
async with httpx.AsyncClient(transport=transport, base_url="http://testserver",
headers={"Authorization": f"Bearer {secret}", "Accept": "application/json, text/event-stream"}) as client:
payload = {"jsonrpc": "2.0", "id": "lost", "method": "tools/call",
"params": {"name": "submit_backtests", "arguments": submission()}}
with pytest.raises(httpx.ReadError):
await client.post(ENDPOINT, json=payload)
response = await client.post(ENDPOINT, json=payload)
assert response.status_code == 200
assert not response.json()["result"]["isError"]
async with mcp_app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 1
audits = list(await db.scalars(select(MCPAudit)))
assert all("wqmcp_" not in json.dumps(a.__dict__, default=str) for a in audits)
async def test_refresh_job_error_pages_and_kind_isolation(mcp_app):
from app.models import Job, JobItem
principal, _ = await credentials(mcp_app)
async with mcp_app.state.sessions.begin() as db:
db.add(Job(id="refresh", kind="pnl_refresh", failed=2))
db.add(Job(id="auth", kind="connect"))
await db.flush()
db.add_all([JobItem(job_id="refresh", alpha_id=str(i), error="synthetic error") for i in range(2)])
result = await invoke(mcp_app, principal, "get_refresh_job", {"job_id": "refresh", "limit": 1})
assert result["errors"]["total"] == 2 and result["errors"]["has_more"]
result = await invoke(mcp_app, principal, "get_refresh_job", {"job_id": "refresh", "offset": 1})
assert result["errors"]["items"][0]["alpha_id"] == "1"
error = await mcp_app.state.mcp.invoke(principal, "get_refresh_job", {"job_id": "auth"})
assert error.structured_content["error"]["code"] == "NOT_FOUND"