feat: 增加 MCP 研究模板保存工具
This commit is contained in:
@@ -11,7 +11,7 @@ from sqlalchemy import select
|
||||
from ..models import Account, Admin, MCPToken, now
|
||||
from ..security import token_hash
|
||||
|
||||
SCOPES = frozenset({"research:read", "research:refresh", "backtests:execute", "backtests:control"})
|
||||
SCOPES = frozenset({"research:read", "research:refresh", "research:write", "backtests:execute", "backtests:control"})
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -21,6 +21,7 @@ from ..research_access.service import ResearchAccess, ResearchError
|
||||
|
||||
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
|
||||
TOOLS = {
|
||||
"create_research_template": (c.CreateTemplate, "create_template", "research:write", "将调用方大模型研究后自行总结的参数化模板保存到模板工坊,供用户后续批量回测。先用 get_backtest_results 阅读实际指标和检查,选择 1–20 个已完成采集的 source_item_ids,并说明 hypothesis;不要把 completed 当作检查通过。template 使用 {name} 占位符及逐一对应的 variables,字段变量须声明 MATRIX/VECTOR/GROUP,VECTOR 聚合须明确写入表达式。提供唯一名称和 idempotency_key,可附 reference。返回模板 ID、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"),
|
||||
"get_submission_check": (c.SelfCorrelationReference, "submission_check_context", "research:read", "读取已导入 Alpha 的表达式、Description、snapshot 和缓存检查结果;不发起检查。先核对或生成三段 Description,再调用 check_submission。"),
|
||||
"check_submission": (c.SubmissionCheck, "check_submission", "research:refresh", "对单个待提交 Alpha 写回已获用户授权的 Description 并调用平台 GET /check,返回 job_id。须先用 get_submission_check 获取 snapshot;保留本地自相关门槛和冲突保护。通过 get_refresh_job 查进度、get_submission_check 读结果。无论检查结果如何,都不会调用 /submit 或正式提交 Alpha。"),
|
||||
"get_worldquant_connection": (c.ConnectionReference, "connection", "research:read", "读取 WorldQuant 连接状态及可选认证 job_id 的进度,不发起认证;人工验证在网页完成。"),
|
||||
@@ -69,7 +70,7 @@ class MCPResearchServer:
|
||||
return types.ListToolsResult(tools=[types.Tool(name=name, description=description,
|
||||
inputSchema=schema.model_json_schema(), annotations=types.ToolAnnotations(
|
||||
readOnlyHint=scope == "research:read", destructiveHint=method == "control",
|
||||
idempotentHint=method in {"submit", "control"} or scope == "research:read",
|
||||
idempotentHint=method in {"submit", "control", "create_template"} or scope == "research:read",
|
||||
openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate"}))
|
||||
for name, (schema, method, scope, description) in TOOLS.items()
|
||||
if scope in principal.scopes and "research:read" in principal.scopes])
|
||||
@@ -123,7 +124,7 @@ class MCPResearchServer:
|
||||
data = {"error": ResearchError(code, "研究操作失败;可使用原幂等键重试或查询历史", retryable=True).data}
|
||||
db.add(MCPAudit(id=str(uuid4()), token_id=principal.token_id, tool=name,
|
||||
request_id=fingerprint({"request_id": request_id}), input_digest=digest,
|
||||
business_id=data.get("backtest_run_id", data.get("job_id")),
|
||||
business_id=data.get("backtest_run_id", data.get("job_id", data.get("template_id"))),
|
||||
result_code=code, elapsed_ms=int((time.monotonic()-started)*1000)))
|
||||
if not error:
|
||||
if access.wake == "backtests":
|
||||
|
||||
@@ -8,7 +8,7 @@ from sqlalchemy import func, select
|
||||
|
||||
from ..models import Account, MCPToken, now
|
||||
from ..security import require_auth
|
||||
from .auth import create_token
|
||||
from .auth import SCOPES, create_token
|
||||
|
||||
router = APIRouter(prefix="/api/v1/mcp-tokens", tags=["mcp-tokens"], dependencies=[Depends(require_auth)])
|
||||
|
||||
@@ -17,7 +17,7 @@ class TokenInput(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
name: str = Field(min_length=1, max_length=100)
|
||||
days: int = Field(default=90, ge=1, le=365, strict=True)
|
||||
scopes: list[str] = Field(default_factory=lambda: ["research:read"], max_length=4)
|
||||
scopes: list[str] = Field(default_factory=lambda: ["research:read"], max_length=len(SCOPES))
|
||||
|
||||
|
||||
def token_output(row, account):
|
||||
|
||||
@@ -7,6 +7,7 @@ from pydantic import Field, model_validator
|
||||
|
||||
from ..backtests.contracts import Candidate, SimulationSettings
|
||||
from ..catalog.contracts import CatalogFilters, Scope
|
||||
from ..research.workspace_contracts import TemplateSpec
|
||||
from ..schemas import Contract
|
||||
|
||||
Identifier = Annotated[str, Field(min_length=1, max_length=100)]
|
||||
@@ -74,6 +75,26 @@ class Control(Contract):
|
||||
idempotency_key: Identifier
|
||||
|
||||
|
||||
class CreateTemplate(Contract):
|
||||
template: TemplateSpec
|
||||
hypothesis: str = Field(min_length=1, max_length=10000)
|
||||
source_item_ids: list[RunId] = Field(min_length=1, max_length=20)
|
||||
reference: str | None = Field(default=None, max_length=200)
|
||||
idempotency_key: Identifier
|
||||
|
||||
@model_validator(mode="after")
|
||||
def research_template(self):
|
||||
self.template.name = self.template.name.strip()
|
||||
self.hypothesis = self.hypothesis.strip()
|
||||
if not self.template.name or not self.hypothesis or not self.template.expression.strip():
|
||||
raise ValueError("模板名称、表达式和研究假设不能为空")
|
||||
if self.template.category != "template":
|
||||
raise ValueError("此工具仅创建完整模板,不保存表达式片段")
|
||||
if len(set(self.source_item_ids)) != len(self.source_item_ids):
|
||||
raise ValueError("来源候选 ID 不能重复")
|
||||
return self
|
||||
|
||||
|
||||
class CatalogSearch(Contract):
|
||||
filters: CatalogFilters
|
||||
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
|
||||
|
||||
@@ -53,6 +53,12 @@ class ResearchAccess:
|
||||
"confirmation": "调用者须已获本批执行授权;直接提交后返回稳定运行 ID",
|
||||
"duplicate_policies": ["reject", "rerun"], "permissions": sorted(self.principal.scopes),
|
||||
"metadata_only": True, "actual_platform_allowance": None,
|
||||
"templates": {
|
||||
"create_with": "create_research_template", "required_scope": "research:write",
|
||||
"authored_by": "caller", "max_source_items": 20,
|
||||
"source_items_with": "get_backtest_results", "starts_backtests": False,
|
||||
"web_url": f"{self.public_origin}/#templates",
|
||||
},
|
||||
"submission_check": {
|
||||
"check_with": "check_submission", "read_with": "get_submission_check",
|
||||
"job_with": "get_refresh_job", "max_targets": 1,
|
||||
@@ -220,6 +226,18 @@ class ResearchAccess:
|
||||
async def history(self, args):
|
||||
return await self.evidence.history(args)
|
||||
|
||||
async def create_template(self, args):
|
||||
"""Persist the caller's template and evidence without model or queue execution."""
|
||||
from .templates import create_template
|
||||
|
||||
previous, digest = await self.previous("create_research_template", args)
|
||||
if previous:
|
||||
return previous.response
|
||||
result = await create_template(self.db, args, self.principal)
|
||||
result["web_url"] = f"{self.public_origin}/#templates"
|
||||
return await self.remember("create_research_template", args, digest, result,
|
||||
business_id=result["template_id"])
|
||||
|
||||
async def previous(self, operation, args):
|
||||
# PostgreSQL row lock is shared with HTTP start and catalog/job creation.
|
||||
account = await self.db.scalar(select(Account).where(Account.id == self.principal.account_id).with_for_update())
|
||||
@@ -233,13 +251,13 @@ class ResearchAccess:
|
||||
raise ResearchError("IDEMPOTENCY_CONFLICT", "幂等键已用于不同内容")
|
||||
return row, digest
|
||||
|
||||
async def remember(self, operation, args, digest, result):
|
||||
async def remember(self, operation, args, digest, result, *, business_id, wake=None):
|
||||
result["_meta"] = {"schema_version": 1, "observed_at": now().isoformat(), "source": "system"}
|
||||
self.db.add(ResearchRequest(id=str(uuid4()), account_id=self.principal.account_id,
|
||||
operation=operation, idempotency_key=args.idempotency_key, digest=digest,
|
||||
business_id=result["backtest_run_id"], response=encode_snapshot(result)))
|
||||
business_id=business_id, response=encode_snapshot(result)))
|
||||
await self.db.flush()
|
||||
self.wake = "backtests"
|
||||
self.wake = wake
|
||||
return result
|
||||
|
||||
async def validate_settings(self, candidates):
|
||||
@@ -286,7 +304,8 @@ class ResearchAccess:
|
||||
result = {**result, "input_digest": digest, "batch_count": preview["batch_count"],
|
||||
"duplicates": {"within_batch": within, "historical_matches": history["total"]},
|
||||
"validation": validation, "web_url": self.run_url(result["backtest_run_id"])}
|
||||
return await self.remember("submit_backtests", args, digest, result)
|
||||
return await self.remember("submit_backtests", args, digest, result,
|
||||
business_id=result["backtest_run_id"], wake="backtests")
|
||||
|
||||
async def run(self, args):
|
||||
result = await self.backtests.run(args.run_id)
|
||||
@@ -325,4 +344,5 @@ class ResearchAccess:
|
||||
"indefinite_account_block_cleared": args.action == "resume" and bool(before["scheduler"]["blocked_reason"])
|
||||
and before["scheduler"]["blocked_until"] is None,
|
||||
"note": "暂停/停止仅阻止后续提交;已提交模拟继续采集。recover 不重新提交。"}
|
||||
return await self.remember("control_backtest", args, digest, result)
|
||||
return await self.remember("control_backtest", args, digest, result,
|
||||
business_id=result["backtest_run_id"], wake="backtests")
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Persist caller-authored templates with frozen local research evidence."""
|
||||
|
||||
import math
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from ..models import BacktestItem, BacktestResult, ResearchAsset
|
||||
from ..research.assets import Assets
|
||||
from ..research.workspace_contracts import AssetWrite
|
||||
from .queries import item_summary
|
||||
|
||||
|
||||
async def create_template(db, args, principal):
|
||||
"""Save a new asset inside the caller's account-locked, idempotent transaction.
|
||||
|
||||
Args contain the external model's TemplateSpec and local source item IDs.
|
||||
Return the versioned asset and theoretical combination count. Raise
|
||||
ResearchError for name conflicts or missing/incomplete research evidence.
|
||||
Stored evidence proves provenance, not profitability or platform eligibility;
|
||||
schema validation does not validate every expanded FASTEXPR combination.
|
||||
"""
|
||||
from .service import ResearchError
|
||||
|
||||
existing = await db.scalar(select(ResearchAsset).where(
|
||||
ResearchAsset.kind == "template", ResearchAsset.name == args.template.name,
|
||||
).order_by(ResearchAsset.id).limit(1))
|
||||
if existing:
|
||||
raise ResearchError("TEMPLATE_NAME_CONFLICT", "模板名称已存在,请使用新名称;此工具不覆盖已有模板",
|
||||
affected_items=[{"template_id": existing.id, "version": existing.version}])
|
||||
rows = (await db.execute(select(BacktestItem, BacktestResult).outerjoin(
|
||||
BacktestResult, BacktestResult.item_id == BacktestItem.id,
|
||||
).where(BacktestItem.id.in_(args.source_item_ids)))).all()
|
||||
found = {item.id: (item, result) for item, result in rows}
|
||||
missing = [item_id for item_id in args.source_item_ids if item_id not in found]
|
||||
if missing:
|
||||
raise ResearchError("NOT_FOUND", "部分来源回测候选不存在", affected_items=missing)
|
||||
incomplete = [item.id for item, result in rows if (
|
||||
item.platform_status != "completed" or item.collection_status != "complete"
|
||||
or item.persistence_status != "saved" or not result or not result.complete
|
||||
)]
|
||||
if incomplete:
|
||||
raise ResearchError("SOURCE_NOT_READY", "来源候选须完成回测、结果采集和持久化,请先读取 get_backtest_results",
|
||||
affected_items=sorted(incomplete))
|
||||
provenance = {
|
||||
"source": {"kind": "mcp", "reference": args.reference},
|
||||
"hypothesis": args.hypothesis,
|
||||
"mcp_token_id": principal.token_id,
|
||||
"admin_id": principal.admin_id,
|
||||
"source_items": [item_summary(*found[item_id]) for item_id in args.source_item_ids],
|
||||
}
|
||||
asset = await Assets(db).save(AssetWrite(
|
||||
kind="template", content=args.template.model_dump(mode="json"),
|
||||
), provenance=provenance)
|
||||
return {
|
||||
**asset, "template_id": asset["id"],
|
||||
"combination_count": str(math.prod(len(v.values) for v in args.template.variables.values())),
|
||||
"validation": {"structure": "valid", "source_evidence": "recorded",
|
||||
"expanded_candidates": "not_validated", "platform_semantics": "unknown"},
|
||||
"next_step": "在模板工坊选择固定输入及模拟设置,展开并核验候选,再确认批量回测。",
|
||||
}
|
||||
@@ -162,7 +162,7 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
|
||||
async with ClientSession(streams[0], streams[1]) as client:
|
||||
await client.initialize()
|
||||
listed = await client.list_tools()
|
||||
assert len(listed.tools) == 17
|
||||
assert len(listed.tools) == 18
|
||||
caps = await client.call_tool("get_research_capabilities", {})
|
||||
assert caps.structured_content["max_candidates"] == 100
|
||||
result = await client.call_tool("submit_backtests", submission())
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
"""External-model templates share the real asset/expansion path, using synthetic research."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import httpx2
|
||||
import pytest
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from app.mcp_api.auth import SCOPES, authenticate
|
||||
from app.mcp_api.server import MCPResearchServer
|
||||
from app.models import (
|
||||
Account,
|
||||
BacktestItem,
|
||||
BacktestResult,
|
||||
BacktestRun,
|
||||
Job,
|
||||
MCPAudit,
|
||||
ResearchAsset,
|
||||
ResearchRequest,
|
||||
ResearchRevision,
|
||||
)
|
||||
from tests.test_backtests import candidate, execute
|
||||
from tests.test_mcp import ENDPOINT, credentials, invoke, submission
|
||||
from tests.test_mcp import mcp_app as mcp_app_fixture
|
||||
from tests.test_research_workspace import catalog as catalog_fixture
|
||||
from tests.test_research_workspace import research_input as research_input_fixture
|
||||
|
||||
mcp_test_app = mcp_app_fixture
|
||||
catalog = catalog_fixture
|
||||
research_input = research_input_fixture
|
||||
TOOL = "create_research_template"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def app(mcp_test_app):
|
||||
# The catalog fixture connects this same synthetic account; keep its email consistent.
|
||||
async with mcp_test_app.state.sessions.begin() as db:
|
||||
(await db.get(Account, 1)).email = "test@example.com"
|
||||
return mcp_test_app
|
||||
|
||||
|
||||
def template_request(item_id):
|
||||
return {
|
||||
"template": {
|
||||
"name": "研究后的字段排序模板",
|
||||
"description": "比较相同类型字段及偏移参数,检验研究假设的稳定性。",
|
||||
"expression": "rank({field}) + {offset}",
|
||||
"variables": {
|
||||
"field": {"kind": "field", "field_type": "MATRIX", "values": ["TEST_FIN_001", "TEST_FIN_002"]},
|
||||
"offset": {"kind": "integer", "values": [0, 1, 5]},
|
||||
},
|
||||
},
|
||||
"hypothesis": "来源结果值得进一步检验,字段替换和参数组合仍需独立回测。",
|
||||
"source_item_ids": [item_id],
|
||||
"reference": "synthetic-research-round-1",
|
||||
"idempotency_key": "template-1",
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def completed_source(app):
|
||||
principal, _ = await credentials(app)
|
||||
run = await invoke(app, principal, "submit_backtests", submission(items=[
|
||||
candidate() | {"expression": "rank(TEST_FIN_001) + 0"},
|
||||
]))
|
||||
await execute(app, app.state.runner.backtests, run["backtest_run_id"])
|
||||
result = await invoke(app, principal, "get_backtest_results", {"run_id": run["backtest_run_id"]})
|
||||
item_id = result["items"][0]["id"]
|
||||
async with app.state.sessions.begin() as db:
|
||||
row = await db.get(BacktestResult, item_id)
|
||||
row.snapshot = {**row.snapshot, "is": {"sharpe": 1.9, "fitness": 0.9, "checks": [
|
||||
{"name": "LOW_FITNESS", "result": "FAIL"}, {"name": "FUTURE_CHECK", "result": "NEW_STATUS"},
|
||||
]}}
|
||||
result = await invoke(app, principal, "get_backtest_results", {"run_id": run["backtest_run_id"]})
|
||||
return result["items"][0]
|
||||
|
||||
|
||||
async def test_sdk_template_creation_frozen_evidence_and_web_expansion(app, logged_in, completed_source, research_input, monkeypatch):
|
||||
principal, secret = await credentials(app, {"research:read", "research:write"})
|
||||
body = template_request(completed_source["id"])
|
||||
async with app.state.sessions() as db:
|
||||
job_count = await db.scalar(select(func.count()).select_from(Job))
|
||||
factory = Mock(side_effect=AssertionError("Template persistence must not call a model"))
|
||||
monkeypatch.setattr(app.state.ai, "model_factory", factory)
|
||||
app.state.runner.wake.clear()
|
||||
app.state.runner.backtests.wake.clear()
|
||||
async with httpx2.AsyncClient(transport=httpx2.ASGITransport(app=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 = {tool.name: tool for tool in (await client.list_tools()).tools}
|
||||
tool = listed[TOOL]
|
||||
assert "submit_backtests" not in listed
|
||||
assert not tool.annotations.read_only_hint and not tool.annotations.destructive_hint
|
||||
assert tool.annotations.idempotent_hint and not tool.annotations.open_world_hint
|
||||
assert {"template", "hypothesis", "source_item_ids", "idempotency_key"} <= set(tool.input_schema["required"])
|
||||
assert tool.input_schema["additionalProperties"] is False
|
||||
caps = await client.call_tool("get_research_capabilities", {})
|
||||
assert caps.structured_content["templates"]["create_with"] == TOOL
|
||||
result = await client.call_tool(TOOL, body)
|
||||
assert not result.is_error, result
|
||||
saved = result.structured_content
|
||||
assert json.loads(result.content[0].text) == saved
|
||||
assert saved["id"] == saved["template_id"] and saved["version"] == 1
|
||||
assert saved["combination_count"] == "6" and saved["web_url"] == "http://testserver/#templates"
|
||||
assert saved["validation"]["expanded_candidates"] == "not_validated"
|
||||
assert saved["provenance"]["source"] == {"kind": "mcp", "reference": body["reference"]}
|
||||
assert saved["provenance"]["hypothesis"] == body["hypothesis"]
|
||||
assert saved["provenance"]["source_items"] == [completed_source]
|
||||
factory.assert_not_called()
|
||||
assert not app.state.runner.wake.is_set() and not app.state.runner.backtests.wake.is_set()
|
||||
async with app.state.sessions.begin() as db:
|
||||
audit = await db.scalar(select(MCPAudit).where(MCPAudit.tool == TOOL))
|
||||
assert audit.business_id == saved["id"] and audit.result_code == "OK"
|
||||
assert secret not in json.dumps(audit.__dict__, default=str)
|
||||
row = await db.get(BacktestResult, completed_source["id"])
|
||||
row.snapshot = {**row.snapshot, "is": {"sharpe": 999}}
|
||||
# The browser reads the same library; later source changes do not rewrite the saved evidence.
|
||||
listing = (await logged_in.get("/api/v1/research/assets?kind=template")).json()
|
||||
assert listing["items"][0]["id"] == saved["id"]
|
||||
stored = (await logged_in.get(f'/api/v1/research/assets/{saved["id"]}?version=1')).json()
|
||||
assert stored["provenance"]["source_items"][0]["metrics"]["is"]["sharpe"] == 1.9
|
||||
assert stored["provenance"]["source_items"][0]["checks"]["counts"]["FAIL"] == 1
|
||||
assert stored["provenance"]["source_items"][0]["checks"]["counts"]["UNKNOWN"] == 1
|
||||
expanded = await logged_in.post("/api/v1/research/experiments", json={
|
||||
"asset_id": saved["id"], "version": saved["version"], "input_ids": [research_input["id"]],
|
||||
"hypothesis": body["hypothesis"], "settings": completed_source["settings"], "limit": 100,
|
||||
})
|
||||
assert expanded.status_code == 201, expanded.text
|
||||
experiment = expanded.json()
|
||||
assert {c["expression"] for c in experiment["candidates"]} == {
|
||||
f"rank({field}) + {offset}" for field in ["TEST_FIN_001", "TEST_FIN_002"] for offset in [0, 1, 5]
|
||||
}
|
||||
assert all(c["validation"]["status"] == "valid" for c in experiment["candidates"])
|
||||
assert experiment["evidence"]["template"]["provenance"] == stored["provenance"]
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(Job)) == job_count
|
||||
|
||||
|
||||
async def test_template_idempotency_concurrency_rotation_and_name_conflict(app, completed_source):
|
||||
principal, _ = await credentials(app, {"research:read", "research:write"})
|
||||
body = template_request(completed_source["id"])
|
||||
first, repeated = await asyncio.gather(invoke(app, principal, TOOL, body), invoke(app, principal, TOOL, body))
|
||||
assert first == repeated
|
||||
rotated, _ = await credentials(app, {"research:read", "research:write"})
|
||||
restarted = MCPResearchServer(app.state.sessions, app.state.runner, app.state.settings)
|
||||
replay = await restarted.invoke(rotated, TOOL, body)
|
||||
assert not replay.is_error and replay.structured_content == first
|
||||
conflict = await app.state.mcp.invoke(principal, TOOL, body | {"hypothesis": "另一假设"})
|
||||
assert conflict.structured_content["error"]["code"] == "IDEMPOTENCY_CONFLICT"
|
||||
same_name = await app.state.mcp.invoke(principal, TOOL, body | {"idempotency_key": "new-key"})
|
||||
assert same_name.structured_content["error"]["code"] == "TEMPLATE_NAME_CONFLICT"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchAsset)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchRevision)) == 1
|
||||
requests = list(await db.scalars(select(ResearchRequest).where(ResearchRequest.operation == TOOL)))
|
||||
assert len(requests) == 1 and requests[0].business_id == first["id"]
|
||||
|
||||
|
||||
async def test_template_invalid_inputs_and_missing_sources_are_atomic(app, completed_source):
|
||||
principal, _ = await credentials(app)
|
||||
valid = template_request(completed_source["id"])
|
||||
variants = [
|
||||
{"source_item_ids": []}, {"source_item_ids": [completed_source["id"]] * 21},
|
||||
{"source_item_ids": [completed_source["id"]] * 2}, {"hypothesis": " "},
|
||||
{"force": True}, {"idempotency_key": ""},
|
||||
{"template": valid["template"] | {"expression": "rank({missing})"}},
|
||||
{"template": valid["template"] | {"category": "fragment"}},
|
||||
{"template": valid["template"] | {"name": " "}},
|
||||
{"template": valid["template"] | {"expression": " ", "variables": {}}},
|
||||
]
|
||||
missing_type = deepcopy(valid["template"])
|
||||
missing_type["variables"]["field"].pop("field_type")
|
||||
variants.append({"template": missing_type})
|
||||
for changed in variants:
|
||||
failed = await app.state.mcp.invoke(principal, TOOL, valid | changed)
|
||||
assert failed.is_error and failed.structured_content["error"]["code"] == "INVALID_INPUT"
|
||||
missing = await app.state.mcp.invoke(principal, TOOL, valid | {"source_item_ids": [completed_source["id"], "missing"]})
|
||||
assert missing.structured_content["error"]["code"] == "NOT_FOUND"
|
||||
assert missing.structured_content["error"]["affected_items"] == ["missing"]
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchAsset)) == 0
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchRequest).where(ResearchRequest.operation == TOOL)) == 0
|
||||
await invoke(app, principal, TOOL, valid)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("state", ["pending", "collection_failed", "unsaved", "incomplete_snapshot", "missing_snapshot"])
|
||||
async def test_template_rejects_incomplete_source_without_consuming_key(app, completed_source, state):
|
||||
principal, _ = await credentials(app)
|
||||
body = template_request(completed_source["id"])
|
||||
async with app.state.sessions.begin() as db:
|
||||
item = await db.get(BacktestItem, completed_source["id"])
|
||||
result = await db.get(BacktestResult, item.id)
|
||||
if state == "pending":
|
||||
item.platform_status = "pending"
|
||||
elif state == "collection_failed":
|
||||
item.collection_status = "failed"
|
||||
elif state == "unsaved":
|
||||
item.persistence_status = "pending"
|
||||
elif state == "incomplete_snapshot":
|
||||
result.complete = False
|
||||
else:
|
||||
await db.delete(result)
|
||||
failed = await app.state.mcp.invoke(principal, TOOL, body)
|
||||
assert failed.is_error and failed.structured_content["error"]["code"] == "SOURCE_NOT_READY"
|
||||
assert failed.structured_content["error"]["affected_items"] == body["source_item_ids"]
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchAsset)) == 0
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchRequest).where(ResearchRequest.operation == TOOL)) == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("scopes", [{"research:read"}, SCOPES - {"research:write"}])
|
||||
async def test_template_write_permission_not_granted_to_old_keys(app, scopes):
|
||||
_, secret = await credentials(app, scopes)
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=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"})
|
||||
assert TOOL not in {tool["name"] for tool in listed.json()["result"]["tools"]}
|
||||
denied = await http.post(ENDPOINT, json={"jsonrpc": "2.0", "id": 2, "method": "tools/call",
|
||||
"params": {"name": TOOL, "arguments": template_request("missing")}})
|
||||
assert denied.status_code == 403
|
||||
|
||||
|
||||
async def test_browser_can_issue_template_only_and_all_permissions(app, logged_in):
|
||||
for scopes in [{"research:read", "research:write"}, SCOPES]:
|
||||
response = await logged_in.post("/api/v1/mcp-tokens", json={"name": "template writer", "scopes": sorted(scopes)})
|
||||
assert response.status_code == 201, response.text
|
||||
async with app.state.sessions() as db:
|
||||
principal = await authenticate(db, response.json()["token"])
|
||||
assert principal.scopes == scopes
|
||||
Reference in New Issue
Block a user