fix: align research metadata with live WorldQuant responses
This commit is contained in:
@@ -0,0 +1,479 @@
|
||||
"""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))
|
||||
@@ -255,3 +255,21 @@ async def test_restart_reuses_preview_and_known_backtest(app, logged_in, flow_se
|
||||
after = await get(logged_in, run["id"])
|
||||
assert after["simulations_used"] == 2 and len(platform.posts) == 1
|
||||
assert next(s for s in after["steps"] if s["node_id"] == "simulate")["backtest_run_id"] == backtest_id
|
||||
|
||||
|
||||
async def test_legacy_authorization_keeps_default_execution_settings(app, logged_in, flow_setup):
|
||||
from app.models import ResearchFlowRun
|
||||
|
||||
body, _, lane, _ = flow_setup
|
||||
run = await begin(logged_in, body)
|
||||
async with app.state.sessions.begin() as db:
|
||||
saved = await db.get(ResearchFlowRun, run["id"])
|
||||
authorization = dict(saved.authorization)
|
||||
authorization["settings"] = {k: v for k, v in authorization["settings"].items() if k != "maxPosition"}
|
||||
authorization["allowed_settings"] = [
|
||||
{k: v for k, v in item.items() if k != "maxPosition"}
|
||||
for item in authorization["allowed_settings"]
|
||||
]
|
||||
saved.authorization = authorization
|
||||
final = await drive(app, logged_in, run["id"], lane)
|
||||
assert final["status"] == "completed" and final["simulations_used"] == 4
|
||||
|
||||
@@ -444,3 +444,73 @@ def test_partial_availability_and_deep_expression_fail_closed():
|
||||
)
|
||||
assert result["status"] == "needs_review"
|
||||
assert analyze("+".join(["close"] * 2000), {"close": "MATRIX"}, set())["status"] == "invalid"
|
||||
|
||||
|
||||
def test_real_field_detail_data_requires_explicit_instrument_context():
|
||||
from app.catalog.research_metadata import normalize_availability
|
||||
|
||||
response = {
|
||||
"id": "close",
|
||||
"type": "MATRIX",
|
||||
"data": [
|
||||
{"region": "USA", "delay": 1, "universe": "TOP3000", "coverage": 1.0},
|
||||
{"region": "EUR", "delay": 1, "universe": "TOP2500", "coverage": 1.0},
|
||||
],
|
||||
}
|
||||
assert normalize_availability(response)["status"] == "needs_review"
|
||||
result = normalize_availability(response, instrument_type="EQUITY")
|
||||
assert result["status"] == "available" and result["items"] == [
|
||||
SCOPE,
|
||||
{**SCOPE, "region": "EUR", "universe": "TOP2500"},
|
||||
]
|
||||
response["data"].append({"region": "USA", "delay": 1})
|
||||
assert normalize_availability(response, instrument_type="EQUITY")["status"] == "needs_review"
|
||||
assert (
|
||||
normalize_availability(
|
||||
{"availability": [{"region": "USA", "delay": 1, "universe": "TOP3000"}]}, instrument_type="EQUITY"
|
||||
)["status"]
|
||||
== "needs_review"
|
||||
)
|
||||
|
||||
|
||||
async def test_field_detail_identity_mismatch_preserves_snapshot(app):
|
||||
from fastapi import HTTPException
|
||||
|
||||
from app.catalog.contracts import Scope
|
||||
from app.catalog.research_metadata import availability_key
|
||||
from app.research.workspace_contracts import FieldAvailabilityInput
|
||||
|
||||
class WrongField:
|
||||
async def field_availability(self, field_id, scope):
|
||||
return {"id": "open", "data": [{"region": "USA", "delay": 1, "universe": "TOP3000"}]}
|
||||
|
||||
body = FieldAvailabilityInput(field_id="close", scope=Scope(**SCOPE))
|
||||
key = availability_key("close", body.scope)
|
||||
async with app.state.sessions.begin() as db:
|
||||
service = ResearchMetadata(db, WrongField())
|
||||
original = {"field_id": "close", "status": "needs_review", "items": []}
|
||||
await service.publish(key, "availability", original)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await service.refresh_availability(body)
|
||||
assert exc.value.status_code == 502
|
||||
assert (await service.get(key))["content"] == original
|
||||
|
||||
|
||||
def test_real_seed_settings_preserve_execution_options_and_reject_unknowns():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.research.experiments import seed_settings
|
||||
|
||||
snapshot = {
|
||||
"region": "USA",
|
||||
"universe": "TOP3000",
|
||||
"delay": 1,
|
||||
"maxPosition": "ON",
|
||||
"startDate": "2014-01-01",
|
||||
"endDate": "2023-12-31",
|
||||
}
|
||||
settings = seed_settings(snapshot)
|
||||
assert settings.maxPosition == "ON" and "startDate" not in settings.model_dump()
|
||||
assert snapshot["startDate"] == "2014-01-01"
|
||||
with pytest.raises(ValidationError):
|
||||
seed_settings({**snapshot, "unknownOption": True})
|
||||
|
||||
Reference in New Issue
Block a user