feat(research): streamline template details and enable bot versioning
Deploy production / deploy (push) Successful in 1m35s
Deploy production / deploy (push) Successful in 1m35s
This commit is contained in:
@@ -33,6 +33,9 @@ TOOLS = {
|
||||
"get_pyramid_distribution": (c.PyramidQuery, "pyramid_distribution", "research:read", "实时读取指定 region(如 USA、GLB)和 delay(0/1)的个人 Pyramid Alpha 分布;必传 current_date(YYYY-MM-DD),自动按自然年四季度取完整起止日(如2026-09-13对应2026-07-01至2026-09-30),传给平台 startDate/endDate,不使用默认周期。按用户约定 alphaCount>=3 为 lit(已点亮),1–2 为 in_progress,0 为 unlit;每项含 category、alpha_count、距3条的 remaining。复用平台认证,未连接时先调用 authenticate_worldquant;缺失数据不当作0。不回测、不提交。"),
|
||||
"search_data_preparations": (c.PreparationSearch, "preparations", "research:read", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"),
|
||||
"get_data_preparation": (c.PreparationRead, "preparation", "research:read", "按集合 ID 与版本分页预览字段、类型、描述和数据集归属。提交回测时携带 preparation_refs,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"),
|
||||
"search_research_templates": (c.TemplateSearch, "templates", "research:read", "分页搜索模板工坊的模板与最新版本,不执行研究。"),
|
||||
"get_research_template": (c.TemplateRead, "template", "research:read", "读取模板内容、字段定义和来源;新增版本前核对最新版本。"),
|
||||
"create_research_template_version": (c.CreateTemplateVersion, "create_template_version", "research:write", "为已有模板新增不可变版本。先读取模板,携带 template_id、expected_version、完整 template、hypothesis、1–20 个已完成采集的 source_item_ids 及 idempotency_key。沿用创建模板的结构和来源校验;版本冲突须重新读取,不覆盖历史、不调用模型、不回测。"),
|
||||
"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 和缓存检查结果;check_summary 分离 Alpha 检查和 REGULAR_SUBMISSION 提交限制,原始 checks 保留;限制不代表当前额度。不发起检查。先核对或生成三段 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。"),
|
||||
@@ -82,7 +85,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", "create_template", "save_super_plan", "build_super_candidates"} or scope == "research:read",
|
||||
idempotentHint=method in {"submit", "control", "create_template", "create_template_version", "save_super_plan", "build_super_candidates"} or scope == "research:read",
|
||||
openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate", "pyramid_distribution", "preview_super_selection"}))
|
||||
for name, (schema, method, scope, description) in TOOLS.items()
|
||||
if scope in principal.scopes and "research:read" in principal.scopes])
|
||||
|
||||
@@ -148,16 +148,25 @@ class Experiments:
|
||||
)
|
||||
variables = {}
|
||||
for name, variable in template.variables.items():
|
||||
values = variable.values
|
||||
# Empty field definitions bind only to the selected immutable input scope.
|
||||
# Existing explicit domains remain restrictions and are never silently widened.
|
||||
if variable.kind == "field":
|
||||
for value in variable.values:
|
||||
if not values:
|
||||
values = sorted(field for field, kind in fields.items() if kind == variable.field_type)
|
||||
if not values:
|
||||
raise HTTPException(422, f"变量 {name} 没有匹配的 {variable.field_type} 字段,请调整数据准备")
|
||||
for value in values:
|
||||
if fields.get(str(value)) != variable.field_type:
|
||||
raise HTTPException(422, f"变量 {name} 的字段 {value} 不在固定输入中或类型不符")
|
||||
if variable.kind == "group" and any(
|
||||
str(v) not in GROUPS and fields.get(str(v)) != "GROUP" for v in variable.values
|
||||
):
|
||||
raise HTTPException(422, f"分组变量 {name} 未在固定输入中核实")
|
||||
if not values:
|
||||
raise HTTPException(422, f"变量 {name} 缺少候选取值,请通过模板接口补充,或将固定参数直接写入表达式")
|
||||
variables[name] = [
|
||||
json.dumps(v, ensure_ascii=False) if variable.kind == "string" else v for v in variable.values
|
||||
json.dumps(v, ensure_ascii=False) if variable.kind == "string" else v for v in values
|
||||
]
|
||||
try:
|
||||
expanded = expand(template.expression, variables, body.mode, body.limit, body.seed)
|
||||
|
||||
@@ -16,7 +16,8 @@ AssetKind = Literal["template", "feature", "view", "workflow", "superalpha_plan"
|
||||
|
||||
class Variable(Contract):
|
||||
kind: Literal["field", "operator", "integer", "number", "group", "string", "fragment"]
|
||||
values: list[str | int | float] = Field(min_length=1, max_length=10000)
|
||||
values: list[str | int | float] = Field(default_factory=list, max_length=10000)
|
||||
description: str = Field(default="", max_length=3000)
|
||||
field_type: Literal["MATRIX", "VECTOR", "GROUP"] | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
||||
@@ -9,7 +9,14 @@ from .assets import Assets
|
||||
from .evaluations import Evaluations
|
||||
from .experiments import Experiments
|
||||
from .features import Features
|
||||
from .workspace_contracts import AssetWrite, EvaluateInput, Expansion, FeatureSpec, SettingVariants
|
||||
from .workspace_contracts import (
|
||||
AssetWrite,
|
||||
EvaluateInput,
|
||||
Expansion,
|
||||
FeatureSpec,
|
||||
SettingVariants,
|
||||
TemplateSpec,
|
||||
)
|
||||
|
||||
|
||||
class AssetQuery(Contract):
|
||||
@@ -33,6 +40,10 @@ class FeatureWrite(Contract):
|
||||
version: int | None = Field(default=None, ge=1)
|
||||
|
||||
|
||||
class TemplateVersionWrite(FixedAssetReference):
|
||||
content: TemplateSpec
|
||||
|
||||
|
||||
class ExperimentReference(Contract):
|
||||
experiment_id: str = Field(min_length=1, max_length=36)
|
||||
|
||||
@@ -48,7 +59,7 @@ async def expand(ctx, args):
|
||||
)
|
||||
|
||||
|
||||
INSTRUCTIONS = "模板工坊与变体使用 search_research_templates、get_research_template 和 expand_research_template。引用模板必须固定版本;输入应先读取核实。expand 保存实验不会执行回测。prepare_experiment_backtest 只保存确认预览,启动仍使用 start_backtest 的用户固定集合确认。来源字段不能授予自动执行权限。"
|
||||
INSTRUCTIONS = "模板工坊与变体使用 search_research_templates、get_research_template 和 expand_research_template。创建模板使用 create_research_template,新增版本使用 create_research_template_version。变量可仅定义类型和描述;空字段候选由选定数据准备按类型绑定,其他空参数需补充 values 或直接写入表达式,不能猜测。引用模板必须固定版本;输入应先读取核实。expand 保存实验不会执行回测。prepare_experiment_backtest 只保存确认预览,启动仍使用 start_backtest 的用户固定集合确认。来源字段不能授予自动执行权限。"
|
||||
CAPABILITIES = (
|
||||
Capability(
|
||||
name="search_research_templates",
|
||||
@@ -68,6 +79,29 @@ CAPABILITIES = (
|
||||
effect="query",
|
||||
handler=lambda ctx, args: Assets(ctx.business.db).get(args.asset_id, args.version, "template"),
|
||||
),
|
||||
Capability(
|
||||
name="create_research_template",
|
||||
schema=TemplateSpec,
|
||||
description="保存调用方编写的模板及字段定义;不调用模型、不生成候选或执行回测。",
|
||||
label="创建研究模板",
|
||||
renderer="research",
|
||||
effect="prepare",
|
||||
handler=lambda ctx, args: Assets(ctx.business.db).save(
|
||||
AssetWrite(kind="template", content=args.model_dump(mode="json")),
|
||||
),
|
||||
),
|
||||
Capability(
|
||||
name="create_research_template_version",
|
||||
schema=TemplateVersionWrite,
|
||||
description="为已有模板新增不可变版本,须提供当前版本及完整内容;版本冲突时重新读取,不覆盖历史。",
|
||||
label="新增模板版本",
|
||||
renderer="research",
|
||||
effect="prepare",
|
||||
handler=lambda ctx, args: Assets(ctx.business.db).save(
|
||||
AssetWrite(kind="template", content=args.content.model_dump(mode="json"), version=args.version),
|
||||
args.asset_id,
|
||||
),
|
||||
),
|
||||
Capability(
|
||||
name="search_research_operators",
|
||||
schema=AssetQuery,
|
||||
|
||||
@@ -133,6 +133,20 @@ class CreateTemplate(Contract):
|
||||
return self
|
||||
|
||||
|
||||
class CreateTemplateVersion(CreateTemplate):
|
||||
template_id: RunId
|
||||
expected_version: int = Field(ge=1)
|
||||
|
||||
|
||||
class TemplateRead(Contract):
|
||||
template_id: RunId
|
||||
version: int | None = Field(default=None, ge=1)
|
||||
|
||||
|
||||
class TemplateSearch(Page):
|
||||
q: str = Field(default="", max_length=200)
|
||||
|
||||
|
||||
class CatalogSearch(Contract):
|
||||
filters: CatalogFilters
|
||||
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
|
||||
|
||||
@@ -81,6 +81,8 @@ class ResearchAccess(SuperResearchAccess):
|
||||
"metadata_only": True, "actual_platform_allowance": None,
|
||||
"templates": {
|
||||
"create_with": "create_research_template", "required_scope": "research:write",
|
||||
"version_with": "create_research_template_version",
|
||||
"read_with": "get_research_template", "search_with": "search_research_templates",
|
||||
"authored_by": "caller", "max_source_items": 20,
|
||||
"source_items_with": "get_backtest_results", "starts_backtests": False,
|
||||
"web_url": f"{self.public_origin}/#templates",
|
||||
@@ -285,6 +287,30 @@ class ResearchAccess(SuperResearchAccess):
|
||||
return await self.remember("create_research_template", args, digest, result,
|
||||
business_id=result["template_id"])
|
||||
|
||||
async def template(self, args):
|
||||
"""Read the exact saved template revision without executing research."""
|
||||
from ..research.assets import Assets
|
||||
|
||||
return await Assets(self.db).get(args.template_id, args.version, "template")
|
||||
|
||||
async def templates(self, args):
|
||||
"""Search the same template library used by the browser."""
|
||||
from ..research.assets import Assets
|
||||
|
||||
return await Assets(self.db).list("template", **args.model_dump())
|
||||
|
||||
async def create_template_version(self, args):
|
||||
"""Append an idempotent, optimistic revision with refreshed source evidence."""
|
||||
from .templates import create_template
|
||||
|
||||
operation = "create_research_template_version"
|
||||
previous, digest = await self.previous(operation, args)
|
||||
if previous:
|
||||
return previous.response
|
||||
result = await create_template(self.db, args, self.principal, asset_id=args.template_id)
|
||||
result["web_url"] = f"{self.public_origin}/#templates"
|
||||
return await self.remember(operation, args, digest, result, business_id=args.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())
|
||||
|
||||
@@ -10,23 +10,27 @@ 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.
|
||||
async def create_template(db, args, principal, *, asset_id=None):
|
||||
"""Save a new asset or revision 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.
|
||||
Optional asset_id selects a version update guarded by args.expected_version;
|
||||
a stale version raises HTTP 409 and cannot overwrite a historical revision.
|
||||
"""
|
||||
from .service import ResearchError
|
||||
|
||||
existing = await db.scalar(select(ResearchAsset).where(
|
||||
ResearchAsset.kind == "template", ResearchAsset.name == args.template.name,
|
||||
ResearchAsset.id != asset_id if asset_id else True,
|
||||
).order_by(ResearchAsset.id).limit(1))
|
||||
if existing:
|
||||
raise ResearchError("TEMPLATE_NAME_CONFLICT", "模板名称已存在,请使用新名称;此工具不覆盖已有模板",
|
||||
affected_items=[{"template_id": existing.id, "version": existing.version}])
|
||||
previous = await Assets(db).get(asset_id, expected_kind="template") if asset_id else None
|
||||
rows = (await db.execute(select(BacktestItem, BacktestResult).outerjoin(
|
||||
BacktestResult, BacktestResult.item_id == BacktestItem.id,
|
||||
).where(BacktestItem.id.in_(args.source_item_ids)))).all()
|
||||
@@ -48,12 +52,16 @@ async def create_template(db, args, principal):
|
||||
"admin_id": principal.admin_id,
|
||||
"source_items": [item_summary(*found[item_id]) for item_id in args.source_item_ids],
|
||||
}
|
||||
if previous:
|
||||
provenance["parent_template"] = {"id": asset_id, "version": args.expected_version}
|
||||
asset = await Assets(db).save(AssetWrite(
|
||||
kind="template", content=args.template.model_dump(mode="json"),
|
||||
), provenance=provenance)
|
||||
version=args.expected_version if asset_id else None,
|
||||
), asset_id=asset_id, provenance=provenance)
|
||||
return {
|
||||
**asset, "template_id": asset["id"],
|
||||
"combination_count": str(math.prod(len(v.values) for v in args.template.variables.values())),
|
||||
"combination_count": (str(math.prod(len(v.values) for v in args.template.variables.values()))
|
||||
if all(v.values for v in args.template.variables.values()) else None),
|
||||
"validation": {"structure": "valid", "source_evidence": "recorded",
|
||||
"expanded_candidates": "not_validated", "platform_semantics": "unknown"},
|
||||
"next_step": "在模板工坊选择固定输入及模拟设置,展开并核验候选,再确认批量回测。",
|
||||
|
||||
@@ -162,7 +162,8 @@ 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) == 29
|
||||
assert len(listed.tools) == 32
|
||||
assert {"search_research_templates", "get_research_template", "create_research_template_version"} <= {t.name for t in listed.tools}
|
||||
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", {})
|
||||
|
||||
@@ -236,3 +236,32 @@ async def test_browser_can_issue_template_only_and_all_permissions(app, logged_i
|
||||
async with app.state.sessions() as db:
|
||||
principal = await authenticate(db, response.json()["token"])
|
||||
assert principal.scopes == scopes
|
||||
|
||||
|
||||
async def test_mcp_template_version_is_idempotent_and_preserves_history(app, completed_source):
|
||||
principal, _ = await credentials(app, {"research:read", "research:write"})
|
||||
body = template_request(completed_source["id"])
|
||||
created = await invoke(app, principal, TOOL, body)
|
||||
content = deepcopy(body["template"])
|
||||
content["variables"]["field"] = {"kind": "field", "field_type": "MATRIX", "description": "数据准备中的矩阵字段"}
|
||||
update = {**body, "template": content, "template_id": created["id"],
|
||||
"expected_version": 1, "idempotency_key": "version-2"}
|
||||
tool = "create_research_template_version"
|
||||
first = await invoke(app, principal, tool, update)
|
||||
replay = await invoke(app, principal, tool, update)
|
||||
assert first == replay and first["version"] == 2
|
||||
assert first["combination_count"] is None
|
||||
assert first["provenance"]["parent_template"] == {"id": created["id"], "version": 1}
|
||||
read = await invoke(app, principal, "get_research_template", {"template_id": created["id"], "version": 1})
|
||||
assert read["content"]["variables"]["field"]["values"] == ["TEST_FIN_001", "TEST_FIN_002"]
|
||||
listed = await invoke(app, principal, "search_research_templates", {"q": content["name"]})
|
||||
assert listed["items"][0]["version"] == 2
|
||||
stale = await app.state.mcp.invoke(principal, tool, update | {"idempotency_key": "stale-version"})
|
||||
assert stale.is_error and stale.structured_content["error"]["code"] == "CONFLICT"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchRevision)) == 2
|
||||
denied, _ = await credentials(app, {"research:read"})
|
||||
from fastapi import HTTPException
|
||||
with pytest.raises(HTTPException) as forbidden:
|
||||
await app.state.mcp.invoke(denied, tool, update)
|
||||
assert forbidden.value.status_code == 403
|
||||
|
||||
@@ -526,3 +526,62 @@ async def test_model_cannot_append_preparation_references_to_fixed_inputs(app, l
|
||||
assert response.status_code == 422 and "不能改变" in response.text
|
||||
async with app.state.sessions() as db:
|
||||
assert not await db.scalar(select(ResearchAsset).where(ResearchAsset.name == "untrusted"))
|
||||
|
||||
|
||||
async def test_template_definitions_bind_selected_fields_without_mutating_asset(logged_in, research_input):
|
||||
content = template()
|
||||
content["expression"] = "rank({field}) + rank({field})"
|
||||
content["variables"]["field"] = {"kind": "field", "field_type": "MATRIX", "description": "横截面字段"}
|
||||
saved = await logged_in.post("/api/v1/research/assets", json={"kind": "template", "content": content})
|
||||
assert saved.status_code == 201, saved.text
|
||||
asset = saved.json()
|
||||
assert asset["content"]["variables"]["field"] == {
|
||||
"kind": "field", "field_type": "MATRIX", "description": "横截面字段", "values": [],
|
||||
}
|
||||
body = expansion(research_input["id"], asset_id=asset["id"], version=1, mode="random", limit=2, seed=19)
|
||||
body.pop("template")
|
||||
first = await logged_in.post("/api/v1/research/experiments", json=body)
|
||||
second = await logged_in.post("/api/v1/research/experiments", json=body)
|
||||
assert first.status_code == second.status_code == 201, first.text
|
||||
assert first.json()["candidates"] == second.json()["candidates"]
|
||||
assert len(first.json()["candidates"]) == 2
|
||||
for candidate in first.json()["candidates"]:
|
||||
field = candidate["bindings"]["field"]
|
||||
assert research_input["field_types"][field] == "MATRIX"
|
||||
assert candidate["expression"] == f"rank({field}) + rank({field})"
|
||||
stored = (await logged_in.get(f"/api/v1/research/assets/{asset['id']}")).json()
|
||||
assert stored["content"] == asset["content"]
|
||||
|
||||
|
||||
async def test_empty_template_domains_fail_expansion_with_actionable_errors(logged_in, research_input):
|
||||
body = expansion(research_input["id"])
|
||||
body["template"]["variables"]["field"] = {"kind": "field", "field_type": "GROUP"}
|
||||
response = await logged_in.post("/api/v1/research/experiments", json=body)
|
||||
assert response.status_code == 422 and "没有匹配的 GROUP 字段" in response.text
|
||||
body["template"]["expression"] = "ts_mean(TEST_FIN_001, {window})"
|
||||
body["template"]["variables"] = {"window": {"kind": "integer", "description": "时间窗口"}}
|
||||
response = await logged_in.post("/api/v1/research/experiments", json=body)
|
||||
assert response.status_code == 422 and "变量 window 缺少候选取值" in response.text
|
||||
|
||||
|
||||
async def test_native_bot_creates_template_and_immutable_version(app, logged_in):
|
||||
from fastapi import HTTPException
|
||||
|
||||
from app.ai.capabilities import ToolContext
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.business import Business
|
||||
|
||||
async with app.state.sessions.begin() as db:
|
||||
ctx = ToolContext(Business(db))
|
||||
content = template()
|
||||
content["variables"]["field"].pop("values")
|
||||
created = await CAPABILITIES["create_research_template"].invoke(ctx, content)
|
||||
content["description"] = "修订后的研究解释"
|
||||
args = {"asset_id": created["id"], "version": 1, "content": content}
|
||||
updated = await CAPABILITIES["create_research_template_version"].invoke(ctx, args)
|
||||
assert updated["version"] == 2
|
||||
with pytest.raises(HTTPException) as conflict:
|
||||
await CAPABILITIES["create_research_template_version"].invoke(ctx, args)
|
||||
assert conflict.value.status_code == 409
|
||||
old = (await logged_in.get(f"/api/v1/research/assets/{created['id']}?version=1")).json()
|
||||
assert old["content"]["description"] == "测试经济假设"
|
||||
|
||||
Reference in New Issue
Block a user