feat: compose native research workflows in QuantFlow
This commit is contained in:
@@ -310,8 +310,10 @@ class Experiments:
|
||||
preserve_source=True,
|
||||
)
|
||||
|
||||
async def setting_variants(self, body):
|
||||
parents = await self.parents([body.alpha_id], [])
|
||||
async def setting_variants(self, body, *, parent_snapshot=None, extra_evidence=None, kind="variant"):
|
||||
parents = (
|
||||
[parent_snapshot] if parent_snapshot is not None else await self.parents([body.alpha_id], [])
|
||||
)
|
||||
original = parents[0]
|
||||
base = SimulationSettings.model_validate(original["settings"])
|
||||
expression = original["expression"]
|
||||
@@ -362,12 +364,13 @@ class Experiments:
|
||||
)
|
||||
return await self.save(
|
||||
f"{body.alpha_id} · 设置变体",
|
||||
"variant",
|
||||
kind,
|
||||
body.hypothesis,
|
||||
snapshots,
|
||||
parents,
|
||||
candidates,
|
||||
{
|
||||
**(extra_evidence or {}),
|
||||
"method": "settings",
|
||||
"rejected": rejected,
|
||||
"operators_snapshot": operators_snapshot,
|
||||
|
||||
@@ -22,6 +22,7 @@ from .workspace_contracts import (
|
||||
ImportCommit,
|
||||
ImportPreview,
|
||||
SettingVariants,
|
||||
WorkflowSpec,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/v1/research", tags=["research"], dependencies=[Depends(require_auth)])
|
||||
@@ -41,7 +42,7 @@ async def assets(
|
||||
limit: int = Query(25, ge=1, le=100),
|
||||
offset: int = Query(0, ge=0),
|
||||
):
|
||||
if kind not in ("template", "feature", "view"):
|
||||
if kind not in ("template", "feature", "view", "workflow"):
|
||||
raise HTTPException(422, "当前素材类型尚未开放")
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Assets(db).list(kind, q, limit, offset)
|
||||
@@ -49,7 +50,7 @@ async def assets(
|
||||
|
||||
@router.post("/assets", status_code=201)
|
||||
async def save_asset(body: AssetWrite, request: Request):
|
||||
if body.kind not in ("template", "feature", "view"):
|
||||
if body.kind not in ("template", "feature", "view", "workflow"):
|
||||
raise HTTPException(422, "当前素材类型尚未开放")
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Assets(db).save(body)
|
||||
@@ -63,7 +64,7 @@ async def asset(asset_id: str, request: Request, version: int | None = Query(Non
|
||||
|
||||
@router.put("/assets/{asset_id}")
|
||||
async def update_asset(asset_id: str, body: AssetWrite, request: Request):
|
||||
if body.kind not in ("template", "feature", "view"):
|
||||
if body.kind not in ("template", "feature", "view", "workflow"):
|
||||
raise HTTPException(422, "当前素材类型尚未开放")
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Assets(db).save(body, asset_id)
|
||||
@@ -240,21 +241,42 @@ async def fixed_recipe():
|
||||
async def start_flow(body: FlowStart, request: Request):
|
||||
from .workflows import Workflows
|
||||
|
||||
if body.workflow_id:
|
||||
raise HTTPException(422, "自定义流程将在 QuantFlow 阶段开放")
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
config = await request.app.state.ai.config(db)
|
||||
result = await Workflows(db).start(body, config.revision)
|
||||
from .workflows import fixed_workflow
|
||||
from .workspace_contracts import WorkflowSpec
|
||||
|
||||
graph = (
|
||||
WorkflowSpec.model_validate(
|
||||
(await Assets(db).get(body.workflow_id, body.workflow_version, "workflow"))["content"]
|
||||
)
|
||||
if body.workflow_id
|
||||
else fixed_workflow(body.budget.max_rounds)
|
||||
)
|
||||
needs_model = any(
|
||||
n.type == "generate"
|
||||
or (n.type == "feature" and not n.config.get("asset_id"))
|
||||
or (n.type == "variant" and n.config.get("method", "structure") == "structure")
|
||||
for n in graph.nodes
|
||||
)
|
||||
config = await request.app.state.ai.config(db) if needs_model else None
|
||||
result = await Workflows(db).start(body, config.revision if config else None)
|
||||
request.app.state.research.wake.set()
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/flows/runs")
|
||||
async def flow_runs(request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)):
|
||||
async def flow_runs(
|
||||
request: Request,
|
||||
limit: int = Query(25, ge=1, le=100),
|
||||
offset: int = Query(0, ge=0),
|
||||
kind: str | None = None,
|
||||
):
|
||||
from .workflows import Workflows
|
||||
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Workflows(db).list(limit, offset)
|
||||
if kind not in (None, "pipeline", "quantflow"):
|
||||
raise HTTPException(422, "未知研究运行类型")
|
||||
return await Workflows(db).list(limit, offset, kind)
|
||||
|
||||
|
||||
@router.get("/flows/runs/{run_id}")
|
||||
@@ -276,3 +298,17 @@ async def control_flow(run_id: str, body: FlowControl, request: Request):
|
||||
request.app.state.research.wake.set()
|
||||
request.app.state.runner.backtests.wake.set()
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/flows/nodes")
|
||||
async def flow_nodes():
|
||||
from .workflows import NODE_TYPES
|
||||
|
||||
return {"items": [{"type": key, **value} for key, value in NODE_TYPES.items()]}
|
||||
|
||||
|
||||
@router.post("/flows/validate")
|
||||
async def validate_flow(body: WorkflowSpec):
|
||||
from .workflows import validate_graph
|
||||
|
||||
return {"valid": True, "order": validate_graph(body)}
|
||||
|
||||
+214
-56
@@ -6,21 +6,32 @@ Backtest source fields are provenance, never a grant to execute automatically.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
from collections import defaultdict
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from ..backtests.contracts import ControlInput, StartInput
|
||||
from ..backtests.contracts import ControlInput, SimulationSettings, StartInput
|
||||
from ..backtests.service import Backtests, uid
|
||||
from ..models import Account, ResearchFlowRun, ResearchStepRun, now
|
||||
from .assets import Assets
|
||||
from .evaluations import Evaluations
|
||||
from .experiments import Experiments
|
||||
from .experiments import Experiments, scope_of
|
||||
from .features import Features
|
||||
from .model import request_model
|
||||
from .workflows import validate_graph
|
||||
from .workspace_contracts import AssetWrite, EvaluateInput, Expansion, Generation, TemplateSpec, WorkflowSpec
|
||||
from .workspace_contracts import (
|
||||
AssetWrite,
|
||||
EvaluateInput,
|
||||
Expansion,
|
||||
FeatureSpec,
|
||||
Generation,
|
||||
SettingVariants,
|
||||
TemplateSpec,
|
||||
WorkflowSpec,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
ACTIVE = ("queued", "running")
|
||||
@@ -162,7 +173,8 @@ class ResearchRuntime:
|
||||
model_work = await self.prepare(run_id)
|
||||
if model_work:
|
||||
step_id, context, revision = model_work
|
||||
result, evidence = await request_model(self.ai, context, TemplateSpec, revision)
|
||||
output_type = FeatureSpec if context.get("method") == "feature" else TemplateSpec
|
||||
result, evidence = await request_model(self.ai, context, output_type, revision)
|
||||
await self.finish_model(run_id, step_id, result, evidence)
|
||||
except HTTPException as exc:
|
||||
await self.fail(run_id, str(exc.detail))
|
||||
@@ -271,7 +283,7 @@ class ResearchRuntime:
|
||||
else None
|
||||
)
|
||||
step.output = {
|
||||
**(previous.output if previous else {}),
|
||||
**(previous.output if previous else {"template": run.authorization.get("template")}),
|
||||
"type": "context",
|
||||
"input_ids": run.authorization["input_ids"],
|
||||
}
|
||||
@@ -287,50 +299,45 @@ class ResearchRuntime:
|
||||
}
|
||||
else:
|
||||
return await self.reserve_model(db, run, step, data, node)
|
||||
elif node.type == "feature":
|
||||
reference = run.authorization.get("node_assets", {}).get(key)
|
||||
if step.output.get("feature"):
|
||||
ref = step.output["feature"]
|
||||
reference = await Assets(db).get(ref["id"], ref["version"], "feature")
|
||||
if reference:
|
||||
template = await Features(db).to_template(reference["id"], reference["version"])
|
||||
step.output = {
|
||||
**data,
|
||||
**step.output,
|
||||
"type": "context",
|
||||
"feature": {"id": reference["id"], "version": reference["version"]},
|
||||
"input_ids": reference["content"]["input_ids"],
|
||||
"template": {"id": template["id"], "version": template["version"]},
|
||||
}
|
||||
else:
|
||||
return await self.reserve_model(db, run, step, data, node)
|
||||
elif node.type == "expand":
|
||||
if not data.get("template"):
|
||||
raise HTTPException(422, "展开节点没有固定模板版本")
|
||||
body = Expansion(
|
||||
asset_id=data["template"]["id"],
|
||||
version=data["template"]["version"],
|
||||
input_ids=run.authorization["input_ids"],
|
||||
hypothesis=run.authorization["hypothesis"],
|
||||
settings=run.authorization["settings"],
|
||||
mode="random",
|
||||
limit=run.authorization["batch_candidates"],
|
||||
seed=run.authorization["seed"] + run.round,
|
||||
parent_alpha_ids=run.authorization["parent_alpha_ids"],
|
||||
parent_experiment_ids=data.get("parent_experiment_ids", []),
|
||||
reference = run.authorization.get("node_assets", {}).get(key) or data.get("template")
|
||||
await self.expand(db, run, step, {**data, "template": reference})
|
||||
return
|
||||
elif node.type == "variant":
|
||||
if node.config.get("method", "structure") == "structure":
|
||||
if step.output.get("template"):
|
||||
await self.expand(db, run, step, step.output)
|
||||
return
|
||||
return await self.reserve_model(db, run, step, data, node)
|
||||
experiment = await Experiments(db).setting_variants(
|
||||
SettingVariants(
|
||||
alpha_id=run.authorization["parent_alpha_ids"][0],
|
||||
input_ids=run.authorization["input_ids"],
|
||||
hypothesis=run.authorization["hypothesis"],
|
||||
),
|
||||
parent_snapshot=run.authorization["parents"][0],
|
||||
extra_evidence={"flow_run_id": run.id, "node_id": key, "round": run.round},
|
||||
kind=run.authorization["kind"],
|
||||
)
|
||||
frozen_parents = run.authorization["parents"] + await Experiments(db).parents(
|
||||
[], body.parent_experiment_ids
|
||||
)
|
||||
experiment = await Experiments(db).create(
|
||||
body,
|
||||
"pipeline",
|
||||
{
|
||||
"flow_run_id": run.id,
|
||||
"node_id": key,
|
||||
"round": run.round,
|
||||
"method": "pipeline",
|
||||
"authorized_seed_snapshots": run.authorization["parents"],
|
||||
},
|
||||
parent_snapshots=frozen_parents,
|
||||
)
|
||||
step.output = {
|
||||
"type": "candidates",
|
||||
"experiment_id": experiment["id"],
|
||||
"candidate_ids": [
|
||||
c["client_item_id"]
|
||||
for c in experiment["candidates"]
|
||||
if c["validation"]["status"] == "valid"
|
||||
],
|
||||
"template": data["template"],
|
||||
}
|
||||
if not step.output["candidate_ids"]:
|
||||
step.status, step.error = "blocked", "候选均未通过本地校验,请核实字段、算子和设置"
|
||||
halt(run, "needs_review", step.error)
|
||||
return
|
||||
self.candidates(run, step, experiment, {})
|
||||
return
|
||||
elif node.type == "backtest":
|
||||
await self.backtest(db, run, step, data)
|
||||
return
|
||||
@@ -349,6 +356,55 @@ class ResearchRuntime:
|
||||
"experiment_id": data["experiment_id"],
|
||||
"backtest_run_id": data["backtest_run_id"],
|
||||
}
|
||||
elif node.type == "condition":
|
||||
step.output = {**data, "type": "evaluation"}
|
||||
elif node.type == "filter":
|
||||
report = (await Evaluations(db).get(data["evaluation_id"]))["report"]
|
||||
ids = [
|
||||
r["client_item_id"]
|
||||
for r in report["records"]
|
||||
if r["verdict"] in node.config.get("verdicts", ["pass"])
|
||||
]
|
||||
experiment = await Experiments(db).get(data["experiment_id"])
|
||||
template = experiment["evidence"].get("template")
|
||||
step.output = {
|
||||
**data,
|
||||
"type": "candidates",
|
||||
"candidate_ids": ids,
|
||||
"template": {"id": template["id"], "version": template["version"]}
|
||||
if template
|
||||
else None,
|
||||
}
|
||||
if not ids:
|
||||
step.status, step.updated_at = "skipped", now()
|
||||
return
|
||||
elif node.type == "summarize":
|
||||
# References keep joins bounded, without recursively copying the upstream graph.
|
||||
step.output = {
|
||||
"type": "summary",
|
||||
"artifacts": [
|
||||
{
|
||||
"step_id": steps[e.source].id,
|
||||
"node_id": e.source,
|
||||
**{
|
||||
k: steps[e.source].output[k]
|
||||
for k in (
|
||||
"type",
|
||||
"template",
|
||||
"feature",
|
||||
"experiment_id",
|
||||
"evaluation_id",
|
||||
"backtest_run_id",
|
||||
"verdict",
|
||||
)
|
||||
if k in steps[e.source].output
|
||||
},
|
||||
}
|
||||
for e in upstream_edges
|
||||
if steps[e.source].status != "skipped"
|
||||
and (not e.branch or steps[e.source].output.get("verdict") == e.branch)
|
||||
],
|
||||
}
|
||||
elif node.type == "iterate":
|
||||
step.output = {
|
||||
**data,
|
||||
@@ -370,10 +426,16 @@ class ResearchRuntime:
|
||||
generation = Generation(
|
||||
name=run.name,
|
||||
hypothesis=run.authorization["hypothesis"],
|
||||
input_ids=run.authorization["input_ids"],
|
||||
input_ids=data.get("input_ids", run.authorization["input_ids"]),
|
||||
method="feature"
|
||||
if node.type == "feature"
|
||||
else "structure"
|
||||
if node.type == "variant"
|
||||
else "template",
|
||||
parent_experiment_ids=[data["experiment_id"]] if data.get("experiment_id") else [],
|
||||
)
|
||||
context = await Experiments(db).generation_context(generation)
|
||||
context["method"] = generation.method
|
||||
context["parents"] = run.authorization["parents"] + context["parents"]
|
||||
context["operators"] = run.authorization["operators_snapshot"]["content"]["items"][:100]
|
||||
context["node_prompt"] = node.config.get("prompt", "")
|
||||
@@ -414,9 +476,16 @@ class ResearchRuntime:
|
||||
step = await db.get(ResearchStepRun, step_id)
|
||||
if not step or step.status != "running":
|
||||
return
|
||||
if isinstance(result, FeatureSpec) and set(result.input_ids) != set(
|
||||
[i["id"] for i in step.output["context"]["inputs"]]
|
||||
):
|
||||
raise HTTPException(422, "模型不能改变已固定的输入范围")
|
||||
# A paused/stopped run may collect this already-issued model output, but cannot advance.
|
||||
asset = await Assets(db).save(
|
||||
AssetWrite(kind="template", content=result.model_dump(mode="json")),
|
||||
AssetWrite(
|
||||
kind="feature" if isinstance(result, FeatureSpec) else "template",
|
||||
content=result.model_dump(mode="json"),
|
||||
),
|
||||
provenance={
|
||||
"flow_run_id": run_id,
|
||||
"step_id": step_id,
|
||||
@@ -424,18 +493,92 @@ class ResearchRuntime:
|
||||
"context": step.output["context"],
|
||||
},
|
||||
)
|
||||
feature = (
|
||||
{"id": asset["id"], "version": asset["version"]} if isinstance(result, FeatureSpec) else None
|
||||
)
|
||||
mark_model_attempt(step, "completed")
|
||||
step.output = {
|
||||
"model_attempts": step.output.get("model_attempts", []),
|
||||
"type": "template",
|
||||
"template": {"id": asset["id"], "version": asset["version"]},
|
||||
"type": "context" if feature else "template",
|
||||
"feature": feature,
|
||||
"input_ids": result.input_ids
|
||||
if feature
|
||||
else [i["id"] for i in step.output["context"]["inputs"]],
|
||||
"template": None if feature else {"id": asset["id"], "version": asset["version"]},
|
||||
"parent_experiment_ids": step.output.get("parent_experiment_ids", []),
|
||||
"generation": evidence,
|
||||
"reserved_call": step.output["reserved_call"],
|
||||
}
|
||||
step.status, step.updated_at = "completed", now()
|
||||
node = next(n for n in run.definition["nodes"] if n["id"] == step.node_id)
|
||||
# Keep post-processing durable and separate: pause/account changes are checked again
|
||||
# before converting features or expanding variants on the next active tick.
|
||||
step.status = "generated" if node["type"] in ("variant", "feature") else "completed"
|
||||
step.updated_at = now()
|
||||
changed(run)
|
||||
|
||||
def candidates(self, run, step, experiment, data):
|
||||
step.output = {
|
||||
**data,
|
||||
"type": "candidates",
|
||||
"experiment_id": experiment["id"],
|
||||
"candidate_ids": [
|
||||
c["client_item_id"] for c in experiment["candidates"] if c["validation"]["status"] == "valid"
|
||||
],
|
||||
}
|
||||
ids = step.output["candidate_ids"]
|
||||
maximum = run.authorization["batch_candidates"]
|
||||
if len(ids) > maximum:
|
||||
step.output = {
|
||||
**step.output,
|
||||
"candidate_ids": random.Random(run.authorization["seed"] + run.round).sample(ids, maximum),
|
||||
}
|
||||
step.status, step.updated_at = "completed", now()
|
||||
if not step.output["candidate_ids"]:
|
||||
step.status, step.error = "blocked", "候选均未通过本地校验,请核实字段、算子和设置"
|
||||
if run.status in ACTIVE:
|
||||
halt(run, "needs_review", step.error)
|
||||
|
||||
async def expand(self, db, run, step, data):
|
||||
reference = data.get("template")
|
||||
if not reference:
|
||||
raise HTTPException(422, "展开节点没有固定模板版本")
|
||||
scope = scope_of(SimulationSettings.model_validate(run.authorization["settings"]))
|
||||
ids = [
|
||||
i["id"]
|
||||
for i in run.authorization["inputs"]
|
||||
if i["scope"] == scope and i["id"] in data.get("input_ids", run.authorization["input_ids"])
|
||||
]
|
||||
if not ids:
|
||||
raise HTTPException(422, "模板展开需要与基础设置匹配的固定输入")
|
||||
body = Expansion(
|
||||
asset_id=reference["id"],
|
||||
version=reference["version"],
|
||||
input_ids=ids,
|
||||
hypothesis=run.authorization["hypothesis"],
|
||||
settings=run.authorization["settings"],
|
||||
mode="random",
|
||||
limit=run.authorization["batch_candidates"],
|
||||
seed=run.authorization["seed"] + run.round,
|
||||
parent_alpha_ids=run.authorization["parent_alpha_ids"],
|
||||
parent_experiment_ids=data.get("parent_experiment_ids", []),
|
||||
)
|
||||
frozen_parents = run.authorization["parents"] + await Experiments(db).parents(
|
||||
[], body.parent_experiment_ids
|
||||
)
|
||||
experiment = await Experiments(db).create(
|
||||
body,
|
||||
run.authorization["kind"],
|
||||
{
|
||||
"flow_run_id": run.id,
|
||||
"node_id": step.node_id,
|
||||
"round": run.round,
|
||||
"method": run.authorization["kind"],
|
||||
"authorized_seed_snapshots": run.authorization["parents"],
|
||||
},
|
||||
parent_snapshots=frozen_parents,
|
||||
)
|
||||
self.candidates(run, step, experiment, data)
|
||||
|
||||
async def backtest(self, db, run, step, data):
|
||||
service = Backtests(db)
|
||||
if step.backtest_run_id:
|
||||
@@ -470,12 +613,27 @@ class ResearchRuntime:
|
||||
):
|
||||
raise HTTPException(409, "研究候选预览不再匹配保存的授权步骤")
|
||||
experiment = await Experiments(db).get(step.output["experiment_id"])
|
||||
if experiment["evidence"].get("flow_run_id") != run.id or {
|
||||
s["id"] for s in experiment["inputs"]
|
||||
} != set(run.authorization["input_ids"]):
|
||||
if (
|
||||
experiment["evidence"].get("flow_run_id") != run.id
|
||||
or {s["id"] for s in experiment["inputs"]}.issubset(set(run.authorization["input_ids"])) is False
|
||||
):
|
||||
raise HTTPException(403, "候选不属于此研究运行的固定输入范围")
|
||||
if "backtest" not in run.authorization["methods"] or any(
|
||||
c["settings"] != run.authorization["settings"]
|
||||
c["settings"]
|
||||
not in (
|
||||
run.authorization.get("allowed_settings", [run.authorization["settings"]])
|
||||
if experiment["evidence"].get("method") == "settings"
|
||||
else [run.authorization["settings"]]
|
||||
)
|
||||
or not c.get("input_ids")
|
||||
or any(
|
||||
not any(
|
||||
i["id"] == input_id
|
||||
and i["scope"] == scope_of(SimulationSettings.model_validate(c["settings"]))
|
||||
for i in run.authorization["inputs"]
|
||||
)
|
||||
for input_id in c["input_ids"]
|
||||
)
|
||||
for c in experiment["candidates"]
|
||||
if c["client_item_id"] in step.output["candidate_ids"]
|
||||
):
|
||||
|
||||
@@ -5,7 +5,7 @@ from collections import defaultdict
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from ..backtests.contracts import fingerprint
|
||||
from ..backtests.contracts import SimulationSettings, fingerprint
|
||||
from ..backtests.service import uid
|
||||
from ..models import Account, ResearchFlowRun, ResearchStepRun
|
||||
from .assets import Assets
|
||||
@@ -64,8 +64,12 @@ def validate_graph(graph):
|
||||
if node.type != "summarize" and len(incoming[node.id]) > 1:
|
||||
raise HTTPException(422, "仅汇总节点接受多个上游;其他节点需要唯一输入")
|
||||
allowed = (
|
||||
{"prompt"}
|
||||
if node.type in ("feature", "generate")
|
||||
{"prompt", "asset_id", "version"}
|
||||
if node.type == "feature"
|
||||
else {"asset_id", "version"}
|
||||
if node.type == "expand"
|
||||
else {"prompt"}
|
||||
if node.type == "generate"
|
||||
else {"method"}
|
||||
if node.type == "variant"
|
||||
else {"verdicts"}
|
||||
@@ -76,6 +80,15 @@ def validate_graph(graph):
|
||||
)
|
||||
if set(node.config) - allowed:
|
||||
raise HTTPException(422, f"节点 {node.id} 包含不支持的配置")
|
||||
if "asset_id" in node.config or "version" in node.config:
|
||||
if (
|
||||
not isinstance(node.config.get("asset_id"), str)
|
||||
or not node.config["asset_id"]
|
||||
or len(node.config["asset_id"]) > 36
|
||||
or type(node.config.get("version")) is not int
|
||||
or node.config["version"] < 1
|
||||
):
|
||||
raise HTTPException(422, "节点素材引用需要 ID 和正整数版本")
|
||||
if "prompt" in node.config and (
|
||||
not isinstance(node.config["prompt"], str) or len(node.config["prompt"]) > 10000
|
||||
):
|
||||
@@ -84,7 +97,7 @@ def validate_graph(graph):
|
||||
raise HTTPException(422, "未知变体方法")
|
||||
if node.type == "filter" and (
|
||||
not isinstance(node.config.get("verdicts", ["pass"]), list)
|
||||
or not set(node.config.get("verdicts", ["pass"])).issubset({"pass", "review", "block"})
|
||||
or any(v not in ("pass", "review", "block") for v in node.config.get("verdicts", ["pass"]))
|
||||
):
|
||||
raise HTTPException(422, "筛选结果必须为 pass/review/block")
|
||||
if node.type == "iterate" and (
|
||||
@@ -169,8 +182,48 @@ class Workflows:
|
||||
if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds:
|
||||
raise HTTPException(422, "流程迭代上限超过本次授权轮数")
|
||||
experiments = Experiments(self.db)
|
||||
inputs, _ = await experiments.inputs(body.input_ids, scope_of(body.settings))
|
||||
settings_variant = any(
|
||||
n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes
|
||||
)
|
||||
if settings_variant and len(body.parent_alpha_ids) != 1:
|
||||
raise HTTPException(422, "设置变体需要且只能选择一个种子 Alpha")
|
||||
if any(n.type == "variant" for n in graph.nodes) and not body.parent_alpha_ids:
|
||||
raise HTTPException(422, "变体节点需要种子 Alpha")
|
||||
inputs, _ = await experiments.inputs(
|
||||
body.input_ids, None if settings_variant else scope_of(body.settings)
|
||||
)
|
||||
parents = await experiments.parents(body.parent_alpha_ids, [])
|
||||
allowed_settings = [body.settings.model_dump(mode="json")]
|
||||
if settings_variant:
|
||||
base = SimulationSettings.model_validate(parents[0]["settings"])
|
||||
for snapshot in inputs:
|
||||
scope = snapshot["scope"]
|
||||
target = SimulationSettings.model_validate(
|
||||
{
|
||||
**base.model_dump(),
|
||||
"instrumentType": scope["instrument_type"],
|
||||
**{key: scope[key] for key in ("region", "universe", "delay")},
|
||||
}
|
||||
)
|
||||
errors, _ = await experiments.settings_check(target)
|
||||
if errors:
|
||||
raise HTTPException(422, "目标范围:" + ";".join(errors))
|
||||
if target.model_dump(mode="json") not in allowed_settings:
|
||||
allowed_settings.append(target.model_dump(mode="json"))
|
||||
node_assets = {}
|
||||
for node in graph.nodes:
|
||||
if node.config.get("asset_id"):
|
||||
ref = await Assets(self.db).get(
|
||||
node.config["asset_id"],
|
||||
node.config["version"],
|
||||
"feature" if node.type == "feature" else "template",
|
||||
)
|
||||
if node.type == "feature":
|
||||
if not set(ref["content"]["input_ids"]).issubset(body.input_ids):
|
||||
raise HTTPException(422, "特征方案输入超出本次固定范围")
|
||||
if not ref["content"].get("template"):
|
||||
raise HTTPException(422, "特征方案需要输出模板")
|
||||
node_assets[node.id] = ref
|
||||
errors, settings_snapshot = await experiments.settings_check(body.settings)
|
||||
if errors:
|
||||
raise HTTPException(422, ";".join(errors))
|
||||
@@ -187,6 +240,7 @@ class Workflows:
|
||||
if (
|
||||
any(
|
||||
node.type == "expand"
|
||||
and not node.config.get("asset_id")
|
||||
and any(
|
||||
e.target == node.id and next(n for n in graph.nodes if n.id == e.source).type == "input"
|
||||
for e in graph.edges
|
||||
@@ -212,6 +266,8 @@ class Workflows:
|
||||
"template": template,
|
||||
"workflow": asset,
|
||||
"methods": methods,
|
||||
"node_assets": node_assets,
|
||||
"allowed_settings": allowed_settings,
|
||||
"settings_snapshot": settings_snapshot,
|
||||
"operators_snapshot": operators_snapshot,
|
||||
"kind": "quantflow" if asset else "pipeline",
|
||||
@@ -272,13 +328,16 @@ class Workflows:
|
||||
}
|
||||
)
|
||||
|
||||
async def list(self, limit=25, offset=0):
|
||||
async def list(self, limit=25, offset=0, kind=None):
|
||||
query = select(ResearchFlowRun)
|
||||
if kind:
|
||||
query = query.where(ResearchFlowRun.authorization["kind"].as_string() == kind)
|
||||
rows = await self.db.scalars(
|
||||
select(ResearchFlowRun).order_by(ResearchFlowRun.created_at.desc()).limit(limit).offset(offset)
|
||||
query.order_by(ResearchFlowRun.created_at.desc()).limit(limit).offset(offset)
|
||||
)
|
||||
return {
|
||||
"items": [await self.get(row.id) for row in rows],
|
||||
"total": await self.db.scalar(select(func.count()).select_from(ResearchFlowRun)),
|
||||
"total": await self.db.scalar(select(func.count()).select_from(query.subquery())),
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
@@ -184,6 +184,24 @@ async def list_flows(ctx, args):
|
||||
|
||||
|
||||
CAPABILITIES += (
|
||||
Capability(
|
||||
name="search_research_workflows",
|
||||
schema=AssetQuery,
|
||||
description="分页查阅 QuantFlow 原生流程与版本,不启动运行。",
|
||||
label="搜索研究流程",
|
||||
renderer="research",
|
||||
effect="query",
|
||||
handler=lambda ctx, args: Assets(ctx.business.db).list("workflow", **args.model_dump()),
|
||||
),
|
||||
Capability(
|
||||
name="get_research_workflow",
|
||||
schema=FixedAssetReference,
|
||||
description="读取指定研究流程版本及原生节点连接,配合 get_research_run 解释执行产物。",
|
||||
label="读取流程版本",
|
||||
renderer="research",
|
||||
effect="query",
|
||||
handler=lambda ctx, args: Assets(ctx.business.db).get(args.asset_id, args.version, "workflow"),
|
||||
),
|
||||
Capability(
|
||||
name="get_research_run",
|
||||
schema=FlowReference,
|
||||
|
||||
Reference in New Issue
Block a user