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

480 lines
22 KiB
Python
Raw Normal View History

"""Opt-in WorldQuant integration: at most six simulations, never official submission.
Run from backend with --execute and an explicitly authorized credentials file.
Uses an isolated SQLite database and a local deterministic model for orchestration;
only catalog/authentication/simulation/PnL requests reach the official platform.
"""
import argparse
import asyncio
import json
import os
import time
from datetime import datetime, timezone
from pathlib import Path
from unittest.mock import patch
import httpx
from cryptography.fernet import Fernet
from sqlalchemy import select
from app.config import Settings
from app.main import create_app
from app.models import Base, SimulationAttempt
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
SCOPE = {"instrument_type": "EQUITY", "region": "USA", "universe": "TOP3000", "delay": 1}
def spec(window):
return {
"name": f"真实联调反转 {window}",
"description": "接口验收,不代表投资结论",
"expression": f"-rank(ts_delta({{field}}, {window}))",
"variables": {"field": {"kind": "field", "field_type": "MATRIX", "values": ["close"]}},
}
async def main(args):
out = Path(args.output).resolve()
out.mkdir(mode=0o700, parents=True, exist_ok=True)
database = out / "research.sqlite"
if database.exists() and not args.resume:
raise RuntimeError("验收数据库已存在:保留原始运行,不重新发送;请选择新目录或只读核验")
secret_file = out / "encryption.key"
key = secret_file.read_text() if args.resume else Fernet.generate_key().decode()
if not args.resume:
secret_file.write_text(key)
secret_file.chmod(0o600)
settings = Settings(
_env_file=None,
database_url=f"sqlite+aiosqlite:///{database}",
admin_password="isolated-live-research-only",
encryption_key=key,
enable_runner=False,
public_origin="http://testserver",
retry_attempts=1,
)
app = create_app(settings)
async with app.state.engine.begin() as db:
await db.run_sync(Base.metadata.create_all)
database.chmod(0o600)
evidence = {
"platform": "https://api.worldquantbrain.com",
"model": "local deterministic fixture, no external model",
"simulation_cap": 6,
"simulations_sent": 0,
"stages": [],
}
if args.resume:
evidence = json.loads((out / "evidence.json").read_text())
stages = [
row
for row in evidence["stages"]
if row["stage"]
in (
"stage1_template",
"stage1_structure",
"stage1_settings",
"stage2_feature",
"stage3_pipeline",
"stage4_quantflow",
)
]
async with app.state.sessions() as db:
attempts = list(await db.scalars(select(SimulationAttempt)))
if any(a.state != "completed" for a in attempts) or len(stages) != len(attempts):
raise RuntimeError("存在未核实或未完整记录的模拟,请恢复原运行,不重新发送")
def record(stage, **data):
row = {"stage": stage, **data}
evidence["stages"].append(row)
(out / "evidence.json").write_text(json.dumps(evidence, ensure_ascii=False, indent=2))
print(json.dumps(row, ensure_ascii=False), flush=True)
remote = app.state.runner.client
async def guard_request(request):
if request.method in ("PATCH", "PUT", "DELETE"):
raise RuntimeError("验收禁止平台属性写入")
if request.method == "POST" and request.url.path != "/authentication":
if request.url.path != "/simulations":
raise RuntimeError("验收禁止其他平台写入")
payload = json.loads(request.content)
count = len(payload) if isinstance(payload, list) else 1
if evidence["simulations_sent"] + count > 6:
raise RuntimeError("真实模拟硬上限已达到")
evidence["simulations_sent"] += count
record("simulation_intent", count=count, total=evidence["simulations_sent"])
remote.client.event_hooks["request"].append(guard_request)
async with app.router.lifespan_context(app):
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app), base_url="http://testserver", headers={"X-WQ-Request": "1"}
) as client:
async def api(method, path, body=None):
response = (
await client.request(method, "/api/v1" + path, json=body)
if body is not None
else await client.request(method, "/api/v1" + path)
)
if not response.is_success:
raise RuntimeError(
f"{method} {path}: HTTP {response.status_code}; {response.json().get('detail', '')}"
)
return response.json()
await api("POST", "/auth/login", {"username": "admin", "password": "isolated-live-research-only"})
credentials = json.loads(Path(args.credentials).read_text())
await api(
"PUT",
"/account/credentials",
{"email": credentials["account"], "password": credentials["password"]},
)
del credentials
job = await api("POST", "/account/connect")
await app.state.runner.execute(job["id"])
status = await api("GET", f"/sync-jobs/{job['id']}")
assert status["status"] == "completed", "账户连接未完成"
record("authentication", status="connected")
operators = await api("POST", "/catalog/operators/refresh")
options = await api("POST", "/catalog/setting-options/refresh")
record(
"metadata",
operators=len(operators["content"]["items"]),
setting_rows=len(options["content"]["items"]),
)
async def fixed(scope):
previous = next(
(
row
for row in evidence["stages"]
if row["stage"] == "fixed_input" and row["scope"] == scope
),
None,
)
if previous:
return previous["id"]
for dataset in (None, "pv1"):
job = await api("POST", "/catalog/sync-jobs", {"scope": scope, "dataset_id": dataset})
await app.state.runner.execute(job["id"])
result = await api("GET", f"/sync-jobs/{job['id']}")
if result["status"] != "completed":
record("catalog_error", status=result["status"], error=result.get("error"))
raise RuntimeError("真实目录同步失败")
params = str(httpx.QueryParams(scope))
fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100")
fixed_input = await api(
"POST",
"/catalog/inputs",
{
"scope": scope,
"dataset_id": "pv1",
"collection_version": fields["collection_version"],
"selection": "all",
},
)
availability = await api(
"POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope}
)
assert availability["content"]["status"] == "available", availability["content"]
record(
"fixed_input",
id=fixed_input["id"],
scope=scope,
fields=len(fixed_input["field_ids"]),
availability_rows=len(availability["content"]["items"]),
)
return fixed_input["id"]
input_id = await fixed(SCOPE)
async def save(kind, content):
return await api("POST", "/research/assets", {"kind": kind, "content": content})
async def expand(asset, parents=None):
return await api(
"POST",
"/research/experiments",
{
"asset_id": asset["id"],
"version": asset["version"],
"input_ids": [input_id],
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
"hypothesis": "真实接口验收",
"limit": 1,
"parent_alpha_ids": parents or [],
},
)
async def collect(run_id):
deadline = time.monotonic() + 900
while time.monotonic() < deadline:
async with app.state.sessions() as db:
attempts = list(
await db.scalars(
select(SimulationAttempt).where(SimulationAttempt.run_id == run_id)
)
)
ids = [
a.id
for a in attempts
if a.state in ("queued", "submitting", "submitted", "collecting")
and (
a.next_poll_at is None
or a.next_poll_at.replace(tzinfo=timezone.utc) <= datetime.now(timezone.utc)
)
]
for aid in ids:
await app.state.runner.backtests.step(aid)
run = await api("GET", f"/backtests/runs/{run_id}")
if run["status"] in ("completed", "completed_with_errors", "needs_review", "stopped"):
results = await api("GET", f"/backtests/runs/{run_id}/results")
rows = [
{
"alpha_id": i.get("alpha_id"),
"status": i.get("platform_status"),
"error": i.get("error"),
"complete": (i.get("result") or {}).get("complete"),
}
for i in results["items"]
]
record("backtest", id=run_id, status=run["status"], results=rows)
if run["status"] != "completed":
raise RuntimeError("模拟未正常完成,保留原运行,不自动重提")
return results
await asyncio.sleep(5)
raise RuntimeError("回测等待超时;数据库保留 progress URL,不自动重发")
async def simulate(experiment, label):
previous = next((row for row in evidence["stages"] if row["stage"] == label), None)
if previous:
result = await api("GET", f"/backtests/runs/{previous['run_id']}/results")
return result["items"][0]["alpha_id"]
preview = await api("POST", f"/research/experiments/{experiment['id']}/preview", {})
assert preview["total"] == 1
run = await api(
"POST",
"/backtests/runs",
{
"preview_id": preview["preview_id"],
"version": preview["version"],
"idempotency_key": label,
},
)
result = await collect(run["backtest_run_id"])
record(
label,
experiment_id=experiment["id"],
preview_id=preview["preview_id"],
run_id=run["backtest_run_id"],
)
return result["items"][0]["alpha_id"]
base = await save("template", spec(5))
seed = await simulate(await expand(base), "stage1_template")
variant = await save("template", spec(10))
variant_id = await simulate(await expand(variant, [seed]), "stage1_structure")
target_id = await fixed({**SCOPE, "universe": "TOP1000"})
settings_variant = await api(
"POST",
"/research/variants/settings",
{"alpha_id": seed, "input_ids": [input_id, target_id], "hypothesis": "同表达式不同股票池"},
)
await simulate(settings_variant, "stage1_settings")
feature = await save(
"feature",
{
"name": "真实特征方案",
"hypothesis": "价格短期反转",
"input_ids": [input_id],
"steps": [
{
"name": "变化及排序",
"rationale": "比较截面价格变化",
"expression": "-rank(ts_delta(close, 20))",
}
],
"template": spec(20),
},
)
feature_template = await api(
"POST", f"/research/features/{feature['id']}/template", {"version": 1}
)
feature_alpha = await simulate(await expand(feature_template), "stage2_feature")
report = await api("POST", "/research/evaluations", {"alpha_id": feature_alpha})
lineage = await api("GET", f"/research/lineage?alpha_id={variant_id}")
record(
"stage2_evaluation",
alpha_id=feature_alpha,
evaluation_id=report["id"],
verdict=report["report"]["verdict"],
lineage_nodes=len(lineage.get("items", [])),
)
await configure(app, client)
async def local_model(ai, context, output_type, revision):
return TemplateSpec.model_validate(spec(8)), {
"model": "local deterministic acceptance fixture",
"revision": revision,
}
flow_body = {
"request_id": "live-pipeline",
"name": "真实平台固定流水线",
"input_ids": [input_id],
"hypothesis": "真实平台运行链验收",
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
"budget": {"max_rounds": 1, "max_simulations": 1, "max_model_calls": 1},
"batch_candidates": 1,
"template_id": base["id"],
"template_version": 1,
}
async def drive(body, label):
if any(row["stage"] == label for row in evidence["stages"]):
return
run = await api("POST", "/research/flows/runs", body)
for _ in range(35):
await app.state.research.advance(run["id"])
current = await api("GET", f"/research/flows/runs/{run['id']}")
for step in current["steps"]:
if step["status"] == "waiting":
await collect(step["backtest_run_id"])
if current["status"] not in ("queued", "running"):
record(
label,
run_id=run["id"],
status=current["status"],
simulations=current["simulations_used"],
model_fixture_calls=current["model_calls_used"],
steps=len(current["steps"]),
error=current["error"],
)
assert current["status"] == "completed"
return
raise RuntimeError("研究运行未结束")
with patch("app.research.runtime.request_model", local_model):
await drive(flow_body, "stage3_pipeline")
graph = {
"name": "真实平台原生画布",
"nodes": [
{"id": kind, "type": kind, "label": kind}
for kind in ("input", "expand", "backtest", "evaluate", "condition", "summarize")
],
"edges": [
{"source": a, "target": b}
for a, b in zip(
("input", "expand", "backtest", "evaluate", "condition"),
("expand", "backtest", "evaluate", "condition", "summarize"),
)
],
}
workflow = await save("workflow", graph)
flow_template = await save("template", spec(15))
await drive(
{
**flow_body,
"request_id": "live-quantflow",
"name": "真实平台 QuantFlow",
"template_id": flow_template["id"],
"workflow_id": workflow["id"],
"workflow_version": 1,
},
"stage4_quantflow",
)
record("completed", simulations_sent=evidence["simulations_sent"], official_submissions=0)
secret_file.chmod(0o600)
async def inspect_existing(args):
"""Read platform PnL and refresh snapshots without issuing any simulation."""
out = Path(args.output).resolve()
evidence = json.loads((out / "evidence.json").read_text())
settings = Settings(
_env_file=None,
database_url=f"sqlite+aiosqlite:///{out / 'research.sqlite'}",
admin_password="isolated-live-research-only",
encryption_key=(out / "encryption.key").read_text(),
enable_runner=False,
public_origin="http://testserver",
retry_attempts=4,
)
app = create_app(settings)
async def read_only(request):
if request.method not in ("GET", "OPTIONS") and not (
request.method == "POST" and request.url.path == "/authentication"
):
raise RuntimeError("后续核验禁止平台写入,包括模拟")
app.state.runner.client.client.event_hooks["request"].append(read_only)
async with app.router.lifespan_context(app):
await app.state.runner.ensure_connected()
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app), base_url="http://testserver", headers={"X-WQ-Request": "1"}
) as client:
await client.post(
"/api/v1/auth/login", json={"username": "admin", "password": "isolated-live-research-only"}
)
ids = [row["results"][0]["alpha_id"] for row in evidence["stages"] if row["stage"] == "backtest"][
:2
]
job = (
await client.post("/api/v1/sync-jobs", json={"kind": "pnl_refresh", "alpha_ids": ids})
).json()
await app.state.runner.execute(job["id"])
current = (await client.get(f"/api/v1/sync-jobs/{job['id']}")).json()
assert current["status"] == "completed", current["status"]
comparison = (await client.post("/api/v1/research/compare", json={"alpha_ids": ids})).json()
assert comparison["common_dates"], "真实 PnL 没有共同日期窗口"
evaluation = next(row for row in evidence["stages"] if row["stage"] == "stage2_evaluation")
lineage = (await client.get(f"/api/v1/research/lineage?alpha_id={ids[1]}")).json()
assert lineage["items"] and lineage["edges"] and lineage["sources"], "真实来源链不完整"
evaluation["lineage_nodes"] = len(lineage["items"])
baseline = (await client.get(f"/api/v1/research/lineage?alpha_id={ids[0]}")).json()
kinds = {item["source"]["kind"] for item in baseline["sources"]["items"]}
assert baseline["sources"]["total"] == 2 and {"template", "pipeline"}.issubset(kinds)
before = (await client.get(f"/api/v1/research/evaluations/{evaluation['evaluation_id']}")).json()
job = (
await client.post(
"/api/v1/sync-jobs", json={"kind": "alpha_refresh", "alpha_ids": [evaluation["alpha_id"]]}
)
).json()
await app.state.runner.execute(job["id"])
after = (await client.get(f"/api/v1/research/evaluations/{evaluation['evaluation_id']}")).json()
assert before == after, "同步覆盖了历史评估"
record = {
"stage": "comparison_and_history",
"alpha_ids": ids,
"common_dates": len(comparison["common_dates"]),
"window": comparison["window"],
"different_settings": comparison["different_settings"],
"evaluation_unchanged": True,
"lineage_experiments": len(lineage["items"]),
"lineage_edges": len(lineage["edges"]),
"baseline_sources": baseline["sources"]["total"],
}
evidence["stages"].append(record)
(out / "evidence.json").write_text(json.dumps(evidence, ensure_ascii=False, indent=2))
print(json.dumps(record, ensure_ascii=False), flush=True)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--execute", action="store_true")
parser.add_argument("--inspect-existing", action="store_true")
parser.add_argument("--resume", action="store_true")
parser.add_argument("--credentials", required=True)
parser.add_argument("--output", required=True)
args = parser.parse_args()
os.umask(0o077)
if args.execute == args.inspect_existing:
parser.error("选择 --execute 或 --inspect-existing 之一")
asyncio.run(inspect_existing(args) if args.inspect_existing else main(args))