488 lines
27 KiB
Python
488 lines
27 KiB
Python
"""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) == 29
|
|
assert any(tool.name == "get_pyramid_distribution" for tool in listed.tools)
|
|
assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools}
|
|
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"
|
|
|
|
|
|
async def test_self_correlation_sdk_workflow_cache_and_staleness(mcp_app, monkeypatch):
|
|
import httpx2
|
|
from mcp import ClientSession
|
|
from mcp.client.streamable_http import streamable_http_client
|
|
|
|
from app.alphas import upsert_alpha
|
|
from app.models import Job, SelfCorrelation
|
|
from tests.conftest import alpha
|
|
from tests.test_alpha_management import points
|
|
|
|
data = points([1, -2, 4, 3] * 20)
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
for raw in (alpha("target"), alpha("peer", status="ACTIVE"), alpha("pending")):
|
|
await upsert_alpha(db, raw)
|
|
db.add(Pnl(alpha_id="target", raw={}, points=data))
|
|
calls = []
|
|
|
|
async def pnl(alpha_id):
|
|
calls.append(alpha_id)
|
|
return {"records": [{"date": p["date"], "pnl": p["value"]} for p in data]}
|
|
|
|
monkeypatch.setattr(mcp_app.state.runner.client, "pnl", pnl)
|
|
principal, secret = await credentials(mcp_app, {"research:read", "research:refresh"})
|
|
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 = {t.name: t for t in (await client.list_tools()).tools}
|
|
assert listed["get_self_correlation"].annotations.read_only_hint
|
|
assert not listed["check_self_correlation"].annotations.read_only_hint
|
|
assert listed["check_self_correlation"].annotations.open_world_hint
|
|
before = await client.call_tool("get_self_correlation", {"alpha_id": "target"})
|
|
assert before.structured_content["status"] == "not_cached"
|
|
assert before.structured_content["result"] is None and calls == []
|
|
mcp_app.state.runner.wake.clear()
|
|
started = await client.call_tool("check_self_correlation", {"alpha_ids": ["target", "target"]})
|
|
assert not started.is_error, started
|
|
job_id = started.structured_content["job_id"]
|
|
assert mcp_app.state.runner.wake.is_set() and calls == []
|
|
assert started.structured_content["alpha_ids"] == ["target"]
|
|
again = await client.call_tool("check_self_correlation", {"alpha_ids": ["target"]})
|
|
assert again.structured_content["job_id"] == job_id
|
|
queued = await client.call_tool("get_refresh_job", {"job_id": job_id})
|
|
assert queued.structured_content["status"] == "queued"
|
|
await mcp_app.state.runner.execute(job_id)
|
|
done = await client.call_tool("get_refresh_job", {"job_id": job_id})
|
|
assert done.structured_content["status"] == "completed"
|
|
assert done.structured_content["processed"] == 1
|
|
result = await client.call_tool("get_self_correlation", {"alpha_id": "target"})
|
|
body = result.structured_content
|
|
assert body["source"] == "local" and body["status"] == "available"
|
|
assert body["result"]["max_correlation"] == pytest.approx(1)
|
|
assert body["result"]["compared_count"] == 1
|
|
assert body["result"]["matches"][0]["alpha_id"] == "peer"
|
|
assert calls == ["peer"]
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
row = await db.get(SelfCorrelation, "target")
|
|
row.stale = True
|
|
audit = await db.scalar(select(MCPAudit).where(MCPAudit.tool == "check_self_correlation"))
|
|
assert audit.business_id == job_id and audit.result_code == "OK"
|
|
assert await db.scalar(select(func.count()).select_from(Job).where(Job.kind == "self_correlation")) == 1
|
|
assert await db.get(Pnl, "peer") is not None
|
|
stale = await invoke(mcp_app, principal, "get_self_correlation", {"alpha_id": "target"})
|
|
assert stale["status"] == "stale" and stale["result"]["stale"] and calls == ["peer"]
|
|
|
|
|
|
async def test_self_correlation_read_only_scope_and_invalid_inputs(mcp_app):
|
|
from app.alphas import upsert_alpha
|
|
from app.models import Job
|
|
from tests.conftest import alpha
|
|
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
await upsert_alpha(db, alpha("known"))
|
|
_, secret = await credentials(mcp_app, {"research:read"})
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mcp_app), base_url="http://testserver",
|
|
headers={"Authorization": f"Bearer {secret}", "Accept": "application/json, text/event-stream"}) as http:
|
|
listed = await http.post(ENDPOINT, json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"})
|
|
names = {t["name"] for t in listed.json()["result"]["tools"]}
|
|
assert "get_self_correlation" in names and "check_self_correlation" not in names
|
|
for name, args, status in (
|
|
("check_self_correlation", {"alpha_ids": ["known"]}, 403),
|
|
("get_self_correlation", {"alpha_id": "known"}, 200),
|
|
):
|
|
response = await http.post(ENDPOINT, json={"jsonrpc": "2.0", "id": 2, "method": "tools/call",
|
|
"params": {"name": name, "arguments": args}})
|
|
assert response.status_code == status
|
|
principal, _ = await credentials(mcp_app)
|
|
for args, code in (
|
|
({"alpha_ids": []}, "INVALID_INPUT"),
|
|
({"alpha_ids": ["known"] * 101}, "INVALID_INPUT"),
|
|
({"alpha_ids": ["../secret"]}, "INVALID_INPUT"),
|
|
({"alpha_ids": ["known"], "force": True}, "INVALID_INPUT"),
|
|
({"alpha_ids": ["known", "missing"]}, "NOT_FOUND"),
|
|
):
|
|
failure = await mcp_app.state.mcp.invoke(principal, "check_self_correlation", args)
|
|
assert failure.is_error and failure.structured_content["error"]["code"] == code
|
|
if code == "NOT_FOUND":
|
|
assert failure.structured_content["error"]["affected_items"] == ["missing"]
|
|
missing = await mcp_app.state.mcp.invoke(principal, "get_self_correlation", {"alpha_id": "missing"})
|
|
assert missing.is_error and missing.structured_content["error"]["code"] == "NOT_FOUND"
|
|
async with mcp_app.state.sessions() as db:
|
|
assert await db.scalar(select(func.count()).select_from(Job)) == 0
|
|
|
|
|
|
async def test_worldquant_authentication_queue_and_recovery(mcp_app):
|
|
from app.models import Account, Job
|
|
|
|
principal, _ = await credentials(mcp_app)
|
|
runner = mcp_app.state.runner
|
|
runner.client.disconnect()
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
account = await db.get(Account, 1)
|
|
account.connection_status = "disconnected"
|
|
before = await invoke(mcp_app, principal, "get_worldquant_connection")
|
|
assert before["credentials_configured"] and not before["session_authenticated"]
|
|
runner.wake.clear()
|
|
started = await invoke(mcp_app, principal, "authenticate_worldquant")
|
|
assert started["status"] == "queued" and runner.wake.is_set()
|
|
again = await invoke(mcp_app, principal, "authenticate_worldquant")
|
|
assert again["job_id"] == started["job_id"]
|
|
await runner.execute(started["job_id"])
|
|
done = await invoke(mcp_app, principal, "get_worldquant_connection", {"job_id": started["job_id"]})
|
|
assert done["job"]["status"] == "completed"
|
|
assert done["connection_status"] == "connected" and done["session_authenticated"]
|
|
async with mcp_app.state.sessions() as db:
|
|
assert await db.scalar(select(func.count()).select_from(Job)) == 1
|
|
audit = await db.scalar(select(MCPAudit).where(MCPAudit.tool == "authenticate_worldquant"))
|
|
assert audit.business_id == started["job_id"]
|
|
assert "synthetic-platform-secret" not in str(done)
|
|
assert "password" not in str(done) and "verification_url" not in str(done)
|
|
|
|
|
|
async def test_worldquant_authentication_permissions_and_challenge(mcp_app):
|
|
from fastapi import HTTPException
|
|
|
|
from app.models import Account, Job
|
|
from app.worldquant import VerificationRequired
|
|
|
|
readonly, _ = await credentials(mcp_app, {"research:read"})
|
|
await invoke(mcp_app, readonly, "get_worldquant_connection")
|
|
with pytest.raises(HTTPException) as denied:
|
|
await mcp_app.state.mcp.invoke(readonly, "authenticate_worldquant", {})
|
|
assert denied.value.status_code == 403
|
|
principal, _ = await credentials(mcp_app)
|
|
invalid = await mcp_app.state.mcp.invoke(principal, "authenticate_worldquant", {"password": "untrusted"})
|
|
assert invalid.is_error and invalid.structured_content["error"]["code"] == "INVALID_INPUT"
|
|
runner = mcp_app.state.runner
|
|
original = runner.client.authenticate
|
|
|
|
async def challenge(*args, **kwargs):
|
|
runner.client.verification_url = "https://api.worldquantbrain.com/authentication/test-challenge"
|
|
raise VerificationRequired(runner.client.verification_url)
|
|
|
|
runner.client.authenticate = challenge
|
|
started = await invoke(mcp_app, principal, "authenticate_worldquant")
|
|
await runner.execute(started["job_id"])
|
|
waiting = await invoke(mcp_app, principal, "get_worldquant_connection", {"job_id": started["job_id"]})
|
|
assert waiting["requires_human_verification"] and waiting["job"]["status"] == "waiting_auth"
|
|
blocked = await mcp_app.state.mcp.invoke(principal, "authenticate_worldquant", {})
|
|
assert blocked.structured_content["error"]["code"] == "VERIFICATION_REQUIRED"
|
|
runner.client.authenticate = original
|
|
# Simulate completion of the human challenge; keep normal profile/identity verification.
|
|
async def verified():
|
|
runner.client.verification_url = None
|
|
await original("synthetic@example.com", "synthetic-platform-secret", force=True)
|
|
runner.client.verify = verified
|
|
resumed = await invoke(mcp_app, principal, "authenticate_worldquant", {"action": "verify"})
|
|
await runner.execute(resumed["job_id"])
|
|
done = await invoke(mcp_app, principal, "get_worldquant_connection", {"job_id": resumed["job_id"]})
|
|
assert done["connection_status"] == "connected" and done["job"]["status"] == "completed"
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
account = await db.get(Account, 1)
|
|
account.password_encrypted = None
|
|
db.add(Job(id="unrelated", kind="pnl_refresh"))
|
|
missing = await mcp_app.state.mcp.invoke(principal, "authenticate_worldquant", {})
|
|
assert missing.structured_content["error"]["code"] == "CREDENTIALS_NOT_CONFIGURED"
|
|
wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"})
|
|
assert wrong.structured_content["error"]["code"] == "NOT_FOUND"
|
|
|
|
async def test_mcp_preparations_freeze_at_submit_and_survive_deletion(mcp_app):
|
|
from app.catalog.contracts import Scope
|
|
from app.preparations.contracts import PreparationReference
|
|
from app.preparations.service import Preparations
|
|
principal, _ = await credentials(mcp_app)
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
collection = await Preparations(db).create("MCP prepared", "", Scope(region="USA", universe="TOP3000", delay=1),
|
|
[{"id": "close", "field_id": "close", "name": "Close", "dataset_id": "pv1", "dataset_name": "Price",
|
|
"description": "Synthetic close", "field_type": "MATRIX", "source": "local", "fetched_at": now().isoformat()}])
|
|
refs = [{"id": collection["id"], "version": 1}]
|
|
found = await invoke(mcp_app, principal, "search_data_preparations", {"q": "MCP prepared"})
|
|
assert found["items"][0]["id"] == collection["id"]
|
|
detail = await invoke(mcp_app, principal, "get_data_preparation", {**refs[0], "limit": 1})
|
|
assert detail["fields"]["items"][0]["dataset_id"] == "pv1"
|
|
result = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
|
async with mcp_app.state.sessions.begin() as db:
|
|
run = await db.get(BacktestRun, result["backtest_run_id"])
|
|
snapshot_id = run.source["input_snapshot_ids"][0]
|
|
await Preparations(db).remove([PreparationReference(**refs[0])])
|
|
assert (await Preparations(db).snapshot(snapshot_id))["fields"][0]["description"] == "Synthetic close"
|
|
# Idempotent replay uses the already fixed run even after the collection is gone.
|
|
replay = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
|
assert replay["backtest_run_id"] == result["backtest_run_id"]
|