refactor: unify data preparations and research input snapshots
Deploy production / deploy (push) Successful in 53s
Deploy production / deploy (push) Successful in 53s
This commit is contained in:
@@ -43,8 +43,7 @@ if __name__ == "__main__":
|
||||
"INSERT INTO research (alpha_id, note, tags, favorite, state, updated_at, version) VALUES ('MIGRATION_TEST', 'preserve research', '[]', false, 'inbox', now(), 7);"
|
||||
)
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
command.upgrade(config, "0014")
|
||||
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
|
||||
("preserve research", 7)
|
||||
]
|
||||
@@ -56,7 +55,7 @@ if __name__ == "__main__":
|
||||
("preserve research", 7)
|
||||
]
|
||||
print(
|
||||
"PostgreSQL 17: 0002 → 0003, downgrade/re-upgrade, metadata check, Alpha/research preservation passed"
|
||||
"PostgreSQL 17: catalog migration, 0014 downgrade/re-upgrade then head, metadata and Alpha preservation passed"
|
||||
)
|
||||
|
||||
async def flow():
|
||||
@@ -73,7 +72,6 @@ if __name__ == "__main__":
|
||||
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
||||
if request.url.path == "/users/self":
|
||||
return httpx.Response(200, json={"id": "PG_TEST_USER"})
|
||||
assert request.method == "GET"
|
||||
return catalog_response(request) or httpx.Response(404)
|
||||
|
||||
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
|
||||
@@ -115,7 +113,7 @@ if __name__ == "__main__":
|
||||
assert sorted(r.status_code for r in responses) == [200, 409]
|
||||
await sync(catalog, "TEST_FIN")
|
||||
assert (await prepare(client, version)).status_code == 409
|
||||
persisted = (await client.get("/api/v1/catalog/inputs/" + draft["id"])).json()
|
||||
persisted = (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json()
|
||||
assert persisted == draft
|
||||
print(
|
||||
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Isolated PostgreSQL migration and concurrency acceptance, using synthetic upstream only."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from cryptography.fernet import Fernet
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
URL = "postgresql+asyncpg://postgres:preparations-test-only@127.0.0.1:18437/preparations_test"
|
||||
os.environ.update(
|
||||
DATABASE_URL=URL,
|
||||
ADMIN_PASSWORD="migration-test-only",
|
||||
ENCRYPTION_KEY=Fernet.generate_key().decode(),
|
||||
WQ_EMAIL="",
|
||||
WQ_PASSWORD="",
|
||||
)
|
||||
|
||||
|
||||
async def sql(query):
|
||||
engine = create_async_engine(URL)
|
||||
try:
|
||||
async with engine.begin() as db:
|
||||
result = await db.execute(text(query))
|
||||
return result.fetchall() if result.returns_rows else None
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def acceptance():
|
||||
import httpx
|
||||
|
||||
from app.catalog.contracts import CatalogJobInput, Scope
|
||||
from app.catalog.service import Catalog
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.worldquant import WqClient
|
||||
from tests.catalog_fake import catalog_response
|
||||
from tests.test_catalog import SCOPE, prepare, search, sync
|
||||
|
||||
def upstream(request):
|
||||
if request.url.path == "/authentication":
|
||||
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
||||
if request.url.path == "/users/self":
|
||||
return httpx.Response(200, json={"id": "PG_TEST_USER"})
|
||||
from tests.research_metadata_fake import response
|
||||
|
||||
metadata = response(request)
|
||||
if metadata is not None:
|
||||
return metadata
|
||||
assert request.method == "GET"
|
||||
return catalog_response(request) or httpx.Response(404)
|
||||
|
||||
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
|
||||
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app),
|
||||
base_url="http://testserver",
|
||||
headers={"X-WQ-Request": "1"},
|
||||
) as client:
|
||||
assert (
|
||||
await client.post(
|
||||
"/api/v1/auth/login", json={"username": "admin", "password": "migration-test-only"}
|
||||
)
|
||||
).status_code == 200
|
||||
await client.put(
|
||||
"/api/v1/account/credentials",
|
||||
json={"email": "synthetic@example.com", "password": "synthetic-only"},
|
||||
)
|
||||
job = (await client.post("/api/v1/account/connect")).json()
|
||||
await app.state.runner.execute(job["id"])
|
||||
fixture = (client, app.state.runner, {})
|
||||
await sync(fixture)
|
||||
await sync(fixture, "TEST_FIN")
|
||||
snapshot = (
|
||||
await prepare(
|
||||
client, (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
|
||||
)
|
||||
).json()
|
||||
assert len(snapshot["field_ids"]) == 123
|
||||
fields = await client.get("/api/v1/catalog/fields", params={**SCOPE, "category": "基本面", "limit": 2, "offset": 2})
|
||||
assert fields.status_code == 200, fields.text
|
||||
assert fields.json()["total"] == 123 and len(fields.json()["items"]) == 2
|
||||
assert fields.json()["items"][0]["category"] == "基本面"
|
||||
|
||||
|
||||
async def enqueue():
|
||||
async with app.state.sessions.begin() as db:
|
||||
return (await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)).id
|
||||
|
||||
jobs = await asyncio.gather(*(enqueue() for _ in range(5)))
|
||||
assert len(set(jobs)) == 1, jobs
|
||||
collection = (await client.get("/api/v1/data-preparations")).json()["items"][0]
|
||||
ref = {"id": collection["id"], "version": collection["version"]}
|
||||
|
||||
async def freeze():
|
||||
response = await client.post("/api/v1/data-preparations/freeze", json={"items": [ref]})
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()["items"][0]["id"]
|
||||
|
||||
assert len(set(await asyncio.gather(*(freeze() for _ in range(5))))) == 1
|
||||
await client.delete(f"/api/v1/data-preparations/{ref['id']}?version={ref['version']}")
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
assert not asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")), (
|
||||
"Requires an empty isolated database"
|
||||
)
|
||||
config = Config("alembic.ini")
|
||||
command.upgrade(config, "0014")
|
||||
# Existing catalog checkpoints survive the structural change; no research data migration.
|
||||
asyncio.run(
|
||||
sql(
|
||||
"INSERT INTO sync_jobs(id,kind,status,payload,checkpoint,total,processed,failed,cancel_requested,created_at,updated_at) VALUES ('checkpoint-test','catalog_sync','failed','{}','{\"offset\": 100}',0,100,0,false,now(),now())"
|
||||
)
|
||||
)
|
||||
asyncio.run(sql("INSERT INTO catalog_scopes(key,scope) VALUES ('checkpoint-scope','{}')"))
|
||||
asyncio.run(
|
||||
sql(
|
||||
"INSERT INTO catalog_batches(id,scope_key,dataset_id,complete,count) VALUES ('checkpoint-test','checkpoint-scope',NULL,false,100)"
|
||||
)
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
assert asyncio.run(
|
||||
sql("SELECT job_id, catalog_batches.\"offset\" FROM catalog_batches WHERE id='checkpoint-test'")
|
||||
) == [("checkpoint-test", 100)]
|
||||
assert asyncio.run(sql("SELECT to_regclass('template_inputs')")) == [(None,)]
|
||||
asyncio.run(sql("DELETE FROM catalog_batches WHERE id='checkpoint-test'"))
|
||||
asyncio.run(sql("DELETE FROM catalog_scopes WHERE key='checkpoint-scope'"))
|
||||
asyncio.run(sql("DELETE FROM sync_jobs WHERE id='checkpoint-test'"))
|
||||
asyncio.run(acceptance())
|
||||
print(
|
||||
"PostgreSQL 17: 0014 → 0015 metadata, retained catalog checkpoint, concurrent job deduplication/freeze and independent snapshot passed"
|
||||
)
|
||||
@@ -22,7 +22,7 @@ def research_step(text, returns, history):
|
||||
return "get_backtest_results", {"run_id": run_id}
|
||||
data = content(returns[-1])
|
||||
return f"已读取本次研究的 {data['total']} 条真实保存结果;缺失 Sharpe 仍为未知。"
|
||||
if context.get("unsaved_field_selection") and not context.get("template_input_id"):
|
||||
if context.get("unsaved_field_selection") and not context.get("input_snapshot_id"):
|
||||
return "请先保存字段选择,再点击用此输入研究。"
|
||||
if not returns:
|
||||
return "get_backtest_capabilities", {}
|
||||
@@ -30,39 +30,23 @@ def research_step(text, returns, history):
|
||||
data = content(last)
|
||||
if "error" in data:
|
||||
return f"研究尚未完成:{data['error']}"
|
||||
scope = context.get("catalog_scope") or {
|
||||
"instrument_type": "EQUITY",
|
||||
"region": "USA",
|
||||
"universe": "TOP3000",
|
||||
"delay": 1,
|
||||
}
|
||||
if last.tool_name == "get_backtest_capabilities":
|
||||
if context.get("template_input_id"):
|
||||
if context.get("input_snapshot_id"):
|
||||
return "get_research_input", {
|
||||
"input_id": context["template_input_id"],
|
||||
"input_id": context["input_snapshot_id"],
|
||||
"field_type": "MATRIX",
|
||||
"limit": 1,
|
||||
}
|
||||
return "search_catalog", {"filters": {**scope, "q": "TEST_FIN", "limit": 1}}
|
||||
if last.tool_name == "search_catalog":
|
||||
if data["dataset_id"] is None:
|
||||
return "search_catalog", {
|
||||
"dataset_id": data["items"][0]["id"],
|
||||
"filters": {**scope, "field_type": "MATRIX", "limit": 1},
|
||||
}
|
||||
return "prepare_research_input", {
|
||||
"scope": scope,
|
||||
"dataset_id": data["dataset_id"],
|
||||
"collection_version": data["collection_version"],
|
||||
"field_ids": [data["items"][0]["id"]],
|
||||
}
|
||||
return "search_data_preparations", {"limit": 1}
|
||||
if last.tool_name == "search_data_preparations":
|
||||
return "prepare_research_input", {"items": [{"id": data["items"][0]["id"], "version": data["items"][0]["version"]}]}
|
||||
if last.tool_name in ("get_research_input", "prepare_research_input"):
|
||||
field = data["items"][0]
|
||||
field = next(f for f in data["items"] if f["field_type"] == "MATRIX")
|
||||
saved_scope = data["scope"]
|
||||
return "prepare_research_backtest", {
|
||||
"name": "Chatbox 数据集研究",
|
||||
"hypothesis": "验证所选合成字段的横截面排序信号",
|
||||
"template_input_id": data["id"],
|
||||
"input_snapshot_id": data["id"],
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "research-1",
|
||||
|
||||
@@ -30,7 +30,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import TemplateInput
|
||||
from app.models import ResearchInputSnapshot
|
||||
from app.research.workspace_contracts import TemplateSpec
|
||||
from tests.test_ai import configure
|
||||
from tests.test_backtests import setup
|
||||
@@ -66,7 +66,7 @@ async def acceptance():
|
||||
}
|
||||
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
body = {
|
||||
"request_id": "finite-run",
|
||||
"name": "PG 有限研究",
|
||||
|
||||
@@ -170,14 +170,15 @@ async def main(args):
|
||||
fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100")
|
||||
fixed_input = await api(
|
||||
"POST",
|
||||
"/catalog/inputs",
|
||||
"/data-preparations/from-dataset",
|
||||
{
|
||||
"scope": scope,
|
||||
"dataset_id": "pv1",
|
||||
"collection_version": fields["collection_version"],
|
||||
"selection": "all",
|
||||
},
|
||||
)
|
||||
fixed_input = (await api("POST", "/data-preparations/freeze", {
|
||||
"items": [{"id": fixed_input["id"], "version": fixed_input["version"]}]}))["items"][0]
|
||||
availability = await api(
|
||||
"POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope}
|
||||
)
|
||||
|
||||
@@ -28,7 +28,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import ResearchExperiment, ResearchParent, TemplateInput
|
||||
from app.models import ResearchExperiment, ResearchInputSnapshot, ResearchParent
|
||||
from tests.test_research_outcomes import (
|
||||
test_feature_conversion_keeps_original_version_through_experiment,
|
||||
test_lineage_retains_multiple_parents_and_descendants,
|
||||
@@ -46,7 +46,7 @@ async def acceptance():
|
||||
)
|
||||
assert response.status_code == 200
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
for experiment in await db.scalars(select(ResearchExperiment)):
|
||||
for parent in experiment.parents:
|
||||
assert await db.get(ResearchParent, (experiment.id, parent["kind"], parent["id"]))
|
||||
|
||||
@@ -30,7 +30,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import TemplateInput
|
||||
from app.models import ResearchInputSnapshot
|
||||
from app.research.workspace_contracts import TemplateSpec
|
||||
from tests.test_ai import configure
|
||||
from tests.test_backtests import setup
|
||||
@@ -67,7 +67,7 @@ async def acceptance():
|
||||
}
|
||||
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
body = {
|
||||
"request_id": "finite-run",
|
||||
"name": "PG 有限研究",
|
||||
|
||||
@@ -10,11 +10,10 @@ from sqlalchemy import func, select
|
||||
from app.ai.capabilities import ToolContext, assemble
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.business import Business
|
||||
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, TemplateInput
|
||||
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, ResearchInputSnapshot
|
||||
from app.research.service import ResearchBuilder
|
||||
from tests.test_ai import configure, single_tool_factory, start
|
||||
from tests.test_api import seed
|
||||
from tests.test_catalog import SCOPE
|
||||
from tests.test_catalog import catalog as catalog_fixture
|
||||
from tests.test_research_integration import fixed_input as fixed_input_fixture
|
||||
|
||||
@@ -62,15 +61,13 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
|
||||
# select_input has already persisted the new input before requesting its result page.
|
||||
raise HTTPException(422, "准备输入后的校验失败")
|
||||
|
||||
updated = await logged_in.patch(f"/api/v1/data-preparations/{fixed_input['preparation_id']}",
|
||||
json={"version": fixed_input["preparation_version"], "name": "new version"})
|
||||
assert updated.status_code == 200
|
||||
monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
|
||||
app.state.ai.model_factory = single_tool_factory(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
{"items": [{"id": fixed_input["preparation_id"], "version": updated.json()["version"]}]},
|
||||
)
|
||||
_, run, _ = await start(app, logged_in, "保存研究输入")
|
||||
call = run["tools"][0]
|
||||
@@ -78,7 +75,7 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
|
||||
assert call["presentation"]["effect"] == "prepare"
|
||||
assert call["result"]["error"] == "准备输入后的校验失败"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 1
|
||||
assert (await db.get(AIToolCall, call["id"])).status == "failed"
|
||||
|
||||
|
||||
|
||||
@@ -212,8 +212,7 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
|
||||
},
|
||||
)
|
||||
for _ in range(2):
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
command.upgrade(config, "0014")
|
||||
alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine)
|
||||
with engine.connect() as db:
|
||||
rows = db.execute(sa.select(alphas).order_by(alphas.c.id)).mappings().all()
|
||||
@@ -224,4 +223,6 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
|
||||
record = db.execute(sa.select(research)).mappings().one()
|
||||
assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7
|
||||
command.downgrade(config, "0010")
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
engine.dispose()
|
||||
|
||||
@@ -203,6 +203,7 @@ async def test_export_formula_injection_and_detail_variants(app, logged_in):
|
||||
assert row["name"] == "'=DANGEROUS()" and row["note"] == "' @formula()"
|
||||
assert (await logged_in.get(f"{PREFIX}/alphas/super1/pnl")).json() == {
|
||||
"cached": False,
|
||||
"series": [],
|
||||
"points": [],
|
||||
"fetched_at": None,
|
||||
}
|
||||
|
||||
@@ -86,16 +86,24 @@ async def search(client, suffix="/datasets", **params):
|
||||
|
||||
|
||||
async def prepare(client, version, **changes):
|
||||
return await client.post(
|
||||
BASE + "/inputs",
|
||||
json={
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": version,
|
||||
"selection": "all",
|
||||
**changes,
|
||||
},
|
||||
)
|
||||
response = await client.post("/api/v1/data-preparations/from-dataset", json={
|
||||
"scope": changes.get("scope", SCOPE), "dataset_id": changes.get("dataset_id", "TEST_FIN"),
|
||||
"collection_version": version,
|
||||
})
|
||||
if response.status_code != 201:
|
||||
return response
|
||||
collection = response.json()
|
||||
if changes.get("excluded_ids"):
|
||||
response = await client.patch(f"/api/v1/data-preparations/{collection['id']}/fields", json={
|
||||
"version": collection["version"], "remove_ids": changes["excluded_ids"]})
|
||||
if response.status_code != 200:
|
||||
return response
|
||||
collection = response.json()
|
||||
response = await client.post("/api/v1/data-preparations/freeze", json={
|
||||
"items": [{"id": collection["id"], "version": collection["version"]}]})
|
||||
if response.status_code != 201:
|
||||
return response
|
||||
return httpx.Response(201, json=response.json()["items"][0])
|
||||
|
||||
|
||||
async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
@@ -112,7 +120,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
response = await prepare(client, version)
|
||||
assert response.status_code == 201, response.text
|
||||
draft = response.json()
|
||||
assert len(draft["field_ids"]) == 123 and draft["status"] == "draft"
|
||||
assert len(draft["field_ids"]) == 123 and draft["preparation_version"] == 1
|
||||
for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
|
||||
detail = await search(client, suffix)
|
||||
assert detail["research"]["version"] == 1
|
||||
@@ -135,7 +143,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
assert newer["collection_version"] != version
|
||||
assert (await prepare(client, version)).status_code == 409
|
||||
assert len((await prepare(client, newer["collection_version"])).json()["field_ids"]) == 125
|
||||
assert (await client.get(BASE + "/inputs/" + draft["id"])).json() == draft
|
||||
assert (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json() == draft
|
||||
assert (await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
|
||||
"note"
|
||||
] == "保留研究备注"
|
||||
@@ -219,7 +227,7 @@ async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
|
||||
await sync(catalog, scope=other)
|
||||
await sync(catalog, "TEST_FIN", scope=other)
|
||||
assert (await prepare(client, version, scope=other)).status_code == 409
|
||||
assert len((await client.get(BASE + "/inputs", params=SCOPE)).json()) == 1
|
||||
assert (await client.get(BASE + "/inputs", params=SCOPE)).status_code == 404
|
||||
|
||||
|
||||
async def test_catalog_authentication_and_origin(app, client):
|
||||
|
||||
@@ -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) == 18
|
||||
assert len(listed.tools) == 20
|
||||
assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools}
|
||||
caps = await client.call_tool("get_research_capabilities", {})
|
||||
assert caps.structured_content["max_candidates"] == 100
|
||||
result = await client.call_tool("submit_backtests", submission())
|
||||
@@ -459,3 +460,27 @@ async def test_worldquant_authentication_permissions_and_challenge(mcp_app):
|
||||
assert missing.structured_content["error"]["code"] == "CREDENTIALS_NOT_CONFIGURED"
|
||||
wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"})
|
||||
assert wrong.structured_content["error"]["code"] == "NOT_FOUND"
|
||||
|
||||
async def test_mcp_preparations_freeze_at_submit_and_survive_deletion(mcp_app):
|
||||
from app.catalog.contracts import Scope
|
||||
from app.preparations.contracts import PreparationReference
|
||||
from app.preparations.service import Preparations
|
||||
principal, _ = await credentials(mcp_app)
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
collection = await Preparations(db).create("MCP prepared", "", Scope(region="USA", universe="TOP3000", delay=1),
|
||||
[{"id": "close", "field_id": "close", "name": "Close", "dataset_id": "pv1", "dataset_name": "Price",
|
||||
"description": "Synthetic close", "field_type": "MATRIX", "source": "local", "fetched_at": now().isoformat()}])
|
||||
refs = [{"id": collection["id"], "version": 1}]
|
||||
found = await invoke(mcp_app, principal, "search_data_preparations", {"q": "MCP prepared"})
|
||||
assert found["items"][0]["id"] == collection["id"]
|
||||
detail = await invoke(mcp_app, principal, "get_data_preparation", {**refs[0], "limit": 1})
|
||||
assert detail["fields"]["items"][0]["dataset_id"] == "pv1"
|
||||
result = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
run = await db.get(BacktestRun, result["backtest_run_id"])
|
||||
snapshot_id = run.source["input_snapshot_ids"][0]
|
||||
await Preparations(db).remove([PreparationReference(**refs[0])])
|
||||
assert (await Preparations(db).snapshot(snapshot_id))["fields"][0]["description"] == "Synthetic close"
|
||||
# Idempotent replay uses the already fixed run even after the collection is gone.
|
||||
replay = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
||||
assert replay["backtest_run_id"] == result["backtest_run_id"]
|
||||
|
||||
@@ -0,0 +1,352 @@
|
||||
"""Preparation contracts across real HTTP, transactions and immutable research inputs."""
|
||||
|
||||
from argparse import Namespace
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.catalog.contracts import CatalogJobInput, Scope
|
||||
from app.catalog.service import Catalog
|
||||
from app.models import CatalogBatch, Job, ResearchInputSnapshot
|
||||
from tests.catalog_fake import field_records
|
||||
from tests.test_catalog import SCOPE, sync
|
||||
from tests.test_catalog import catalog as catalog_fixture
|
||||
|
||||
catalog = catalog_fixture
|
||||
BASE = "/api/v1/data-preparations"
|
||||
|
||||
|
||||
def reference(field):
|
||||
return {k: field[k] for k in ("scope", "dataset_id", "field_id", "source", "collection_version")}
|
||||
|
||||
|
||||
async def copied(catalog):
|
||||
client = catalog[0]
|
||||
await sync(catalog)
|
||||
job = await sync(catalog, "TEST_FIN")
|
||||
response = await client.post(
|
||||
BASE + "/from-dataset",
|
||||
json={"scope": SCOPE, "dataset_id": "TEST_FIN", "collection_version": job["id"]},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()
|
||||
|
||||
|
||||
async def test_multi_dataset_collection_atomic_changes_and_snapshot_independence(catalog):
|
||||
client, runner, state = catalog
|
||||
collection = await copied(catalog)
|
||||
state["fields"] = field_records("TEST_NEWS", 3)
|
||||
await sync(catalog, "TEST_NEWS")
|
||||
query = (await client.get("/api/v1/catalog/fields", params={**SCOPE, "dataset_id": "TEST_NEWS"})).json()
|
||||
assert query["total"] == 3
|
||||
ref = reference(query["items"][0])
|
||||
response = await client.patch(
|
||||
f"{BASE}/{collection['id']}/fields", json={"version": 1, "fields": [ref, ref]}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
collection = response.json()
|
||||
assert collection["field_count"] == 124 and collection["dataset_count"] == 2
|
||||
bad = {**ref, "scope": {**SCOPE, "delay": 0}}
|
||||
response = await client.patch(
|
||||
f"{BASE}/{collection['id']}/fields",
|
||||
json={"version": 2, "fields": [ref, bad], "remove_ids": ["TEST_FIN_001"]},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
assert (await client.get(f"{BASE}/{collection['id']}")).json()["version"] == 2
|
||||
snapshot = (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
|
||||
).json()["items"][0]
|
||||
assert snapshot["dataset_ids"] == ["TEST_FIN", "TEST_NEWS"]
|
||||
assert all("description" in f and f["dataset_id"] for f in snapshot["fields"])
|
||||
assert (
|
||||
await client.patch(f"{BASE}/{collection['id']}", json={"version": 2, "name": "edited"})
|
||||
).status_code == 200
|
||||
assert (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
|
||||
).status_code == 409
|
||||
assert (await client.delete(f"{BASE}/{collection['id']}?version=3")).status_code == 200
|
||||
assert (await client.get(f"{BASE}/{collection['id']}")).status_code == 404
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
|
||||
async with runner.sessions() as db:
|
||||
assert await db.get(ResearchInputSnapshot, snapshot["id"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key,value", [("region", "CHN"), ("universe", "TOP1000"), ("delay", 0), ("instrument_type", "FUTURE")]
|
||||
)
|
||||
async def test_scope_dimensions_rejected_without_creating_collection(catalog, key, value):
|
||||
client = catalog[0]
|
||||
await copied(catalog)
|
||||
field = (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["items"][0]
|
||||
before = (await client.get(BASE)).json()["total"]
|
||||
response = await client.post(
|
||||
BASE, json={"name": "bad", "scope": {**SCOPE, key: value}, "fields": [reference(field)]}
|
||||
)
|
||||
assert response.status_code == 422
|
||||
assert (await client.get(BASE)).json()["total"] == before
|
||||
|
||||
|
||||
async def test_online_fields_without_sync_and_local_search_before_pagination(catalog):
|
||||
client, _, _ = catalog
|
||||
response = await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100})
|
||||
assert response.status_code == 200, response.text
|
||||
first = response.json()
|
||||
assert len(first["items"]) == 100 and first["has_more"]
|
||||
second = (
|
||||
await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100, "offset": 100})
|
||||
).json()
|
||||
assert len(second["items"]) > 0
|
||||
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
|
||||
field = reference(second["items"][-1])
|
||||
response = await client.post(BASE, json={"name": "online only", "scope": SCOPE, "fields": [field]})
|
||||
assert response.status_code == 201, response.text
|
||||
assert response.json()["field_count"] == 1
|
||||
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
|
||||
await copied(catalog)
|
||||
query = (
|
||||
await client.get(
|
||||
"/api/v1/catalog/fields", params={**SCOPE, "q": "合成字段说明 11", "limit": 2, "offset": 2}
|
||||
)
|
||||
).json()
|
||||
assert query["total"] == 11 and len(query["items"]) == 2
|
||||
response = await client.get(
|
||||
"/api/v1/catalog/fields", params={**SCOPE, "coverage_min": 0.9, "coverage_max": 0.5}
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
async def test_empty_crud_copy_and_batch_delete_are_version_checked(catalog):
|
||||
client = catalog[0]
|
||||
first = (await client.post(BASE, json={"name": "empty", "scope": SCOPE})).json()
|
||||
assert first["field_count"] == 0
|
||||
assert (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": first["id"], "version": 1}]})
|
||||
).status_code == 422
|
||||
second = (await client.post(f"{BASE}/{first['id']}/copy", json={"version": 1})).json()
|
||||
response = await client.post(
|
||||
BASE + "/batch-delete",
|
||||
json={"items": [{"id": first["id"], "version": 1}, {"id": second["id"], "version": 99}]},
|
||||
)
|
||||
assert response.status_code == 409
|
||||
assert (await client.get(BASE)).json()["total"] == 2
|
||||
assert (
|
||||
await client.patch(
|
||||
f"{BASE}/{first['id']}", json={"name": "x", "version": 1, "scope": {**SCOPE, "delay": 0}}
|
||||
)
|
||||
).status_code == 422
|
||||
|
||||
|
||||
async def test_full_sync_partial_failure_and_resume_publish_per_dataset(catalog):
|
||||
client, runner, state = catalog
|
||||
await copied(catalog)
|
||||
async with runner.sessions.begin() as db:
|
||||
service = Catalog(db)
|
||||
job = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
duplicate = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
assert duplicate.id == job.id
|
||||
# The fixture deliberately returns TEST_FIN ownership for the other datasets.
|
||||
await runner.execute(job.id)
|
||||
failed = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert failed["status"] == "completed_with_errors" and failed["processed"] == 1
|
||||
version = (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
|
||||
"collection_version"
|
||||
]
|
||||
assert version != job.id
|
||||
async with runner.sessions() as db:
|
||||
assert (await db.get(CatalogBatch, version)).job_id == job.id
|
||||
state["fields"] = field_records("TEST_NEWS", 3)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
|
||||
await runner.execute(job.id)
|
||||
retried = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert retried["processed"] == 2 and retried["failed"] == 1
|
||||
assert (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
|
||||
"collection_version"
|
||||
] == version
|
||||
state["fields"] = field_records("TEST_UNKNOWN", 0)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
|
||||
await runner.execute(job.id)
|
||||
done = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert done["status"] == "completed" and done["processed"] == 3 and done["failed"] == 0
|
||||
|
||||
|
||||
async def test_cli_timeout_reuses_active_job_and_does_not_cancel(catalog, monkeypatch, capsys):
|
||||
from app import cli
|
||||
|
||||
_, runner, _ = catalog
|
||||
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
|
||||
args = Namespace(
|
||||
region="USA",
|
||||
universe="TOP3000",
|
||||
delay=1,
|
||||
instrument_type="EQUITY",
|
||||
resume_job=None,
|
||||
wait_timeout=0.01,
|
||||
)
|
||||
assert await cli.catalog_sync_command(args) == 4
|
||||
assert await cli.catalog_sync_command(args) == 4
|
||||
async with runner.sessions() as db:
|
||||
rows = list(await db.scalars(select(Job).where(Job.kind == "catalog_full_sync")))
|
||||
assert len(rows) == 1 and rows[0].status == "queued" and not rows[0].cancel_requested
|
||||
assert "等待超时" in capsys.readouterr().out
|
||||
|
||||
|
||||
async def test_collection_version_checked_at_research_submission(catalog):
|
||||
client = catalog[0]
|
||||
collection = await copied(catalog)
|
||||
response = await client.patch(f"{BASE}/{collection['id']}", json={"version": 1, "name": "changed"})
|
||||
assert response.status_code == 200
|
||||
body = {
|
||||
"inline": {
|
||||
"name": "prepared backtest",
|
||||
"preparation_refs": [{"id": collection["id"], "version": 1}],
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "one",
|
||||
"expression": "rank(TEST_FIN_001)",
|
||||
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
|
||||
}
|
||||
],
|
||||
}
|
||||
}
|
||||
assert (await client.post("/api/v1/backtests/previews", json=body)).status_code == 409
|
||||
body["inline"]["preparation_refs"][0]["version"] = 2
|
||||
response = await client.post("/api/v1/backtests/previews", json=body)
|
||||
assert response.status_code == 201, response.text
|
||||
snapshot_id = response.json()["source"]["input_snapshot_ids"][0]
|
||||
await client.delete(f"{BASE}/{collection['id']}?version=2")
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot_id}")).json()["fields"][1][
|
||||
"dataset_id"
|
||||
] == "TEST_FIN"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status,code",
|
||||
[
|
||||
("completed", 0),
|
||||
("completed_with_errors", 1),
|
||||
("failed", 1),
|
||||
("cancelled", 1),
|
||||
("waiting_connection", 3),
|
||||
("waiting_auth", 3),
|
||||
],
|
||||
)
|
||||
async def test_cli_terminal_exit_codes(catalog, monkeypatch, status, code):
|
||||
from app import cli
|
||||
|
||||
_, runner, _ = catalog
|
||||
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
|
||||
original = Catalog.create_job
|
||||
|
||||
async def terminal(self, body, **kwargs):
|
||||
result = await original(self, body, **kwargs)
|
||||
row = await self.db.get(Job, result.id)
|
||||
row.status = status
|
||||
return result
|
||||
|
||||
monkeypatch.setattr(Catalog, "create_job", terminal)
|
||||
args = Namespace(
|
||||
region="USA", universe="TOP3000", delay=1, instrument_type="EQUITY", resume_job=None, wait_timeout=1
|
||||
)
|
||||
assert await cli.catalog_sync_command(args) == code
|
||||
|
||||
|
||||
async def test_full_sync_restart_keeps_page_and_auth_pauses_all_datasets(catalog, monkeypatch):
|
||||
import asyncio
|
||||
|
||||
from app.worldquant import WqError
|
||||
|
||||
client, runner, _ = catalog
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
original = runner.client.catalog_page
|
||||
interrupted = False
|
||||
|
||||
async def pages(scope, dataset_id, offset):
|
||||
nonlocal interrupted
|
||||
if dataset_id == "TEST_FIN" and offset == 100 and not interrupted:
|
||||
interrupted = True
|
||||
runner.stopping = True
|
||||
raise asyncio.CancelledError()
|
||||
if dataset_id == "TEST_NEWS":
|
||||
raise WqError("synthetic network interruption", "network_error")
|
||||
return await original(scope, dataset_id, offset)
|
||||
|
||||
monkeypatch.setattr(runner.client, "catalog_page", pages)
|
||||
await runner.execute(job.id)
|
||||
async with runner.sessions() as db:
|
||||
batch = await db.scalar(
|
||||
select(CatalogBatch).where(CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_FIN")
|
||||
)
|
||||
assert batch.offset == 100 and not batch.complete
|
||||
assert (await db.get(Job, job.id)).status == "queued"
|
||||
runner.stopping = False
|
||||
await runner.execute(job.id)
|
||||
row = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert row["status"] == "waiting_connection" and row["processed"] == 1
|
||||
assert row["checkpoint"]["dataset_id"] == "TEST_NEWS"
|
||||
assert row["checkpoint"]["datasets_completed"] == 1
|
||||
async with runner.sessions() as db:
|
||||
assert not await db.scalar(
|
||||
select(CatalogBatch).where(
|
||||
CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_UNKNOWN"
|
||||
)
|
||||
)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/cancel")
|
||||
assert (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()["status"] == "cancelled"
|
||||
|
||||
|
||||
async def test_online_instrument_type_is_verified(catalog):
|
||||
client, _, state = catalog
|
||||
state["fields"][0]["instrumentType"] = "FUTURE"
|
||||
assert (await client.get("/api/v1/catalog/worldquant/fields", params=SCOPE)).status_code == 502
|
||||
|
||||
|
||||
async def test_mcp_collection_reads_and_versioned_submit_contract(catalog):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from app.research_access.contracts import PreparationRead, PreparationSearch, Submit
|
||||
from app.research_access.service import ResearchAccess
|
||||
|
||||
client, runner, _ = catalog
|
||||
collection = await copied(catalog)
|
||||
async with runner.sessions.begin() as db:
|
||||
access = ResearchAccess(
|
||||
db, SimpleNamespace(token_id="test", admin_id=1), runner.client, "http://testserver"
|
||||
)
|
||||
result = await access.preparations(PreparationSearch(q=collection["name"]))
|
||||
assert result["items"][0]["version"] == 1
|
||||
detail = await access.preparation(PreparationRead(id=collection["id"], version=1, limit=1))
|
||||
assert detail["fields"]["has_more"] and detail["fields"]["items"][0]["dataset_id"] == "TEST_FIN"
|
||||
# Schema accepts versioned references and rejects accidental old dataset input payloads.
|
||||
from tests.test_mcp import submission
|
||||
|
||||
payload = submission()
|
||||
payload["preparation_refs"] = [{"id": collection["id"], "version": 1}]
|
||||
assert Submit.model_validate(payload).preparation_refs[0].id == collection["id"]
|
||||
|
||||
async def test_retry_full_job_reuses_another_active_job(catalog):
|
||||
from app.business import Business
|
||||
_, runner, _ = catalog
|
||||
body = CatalogJobInput(scope=Scope(**SCOPE))
|
||||
async with runner.sessions.begin() as db:
|
||||
old = await Catalog(db).create_job(body, full=True)
|
||||
(await db.get(Job, old.id)).status = "failed"
|
||||
async with runner.sessions.begin() as db:
|
||||
new = await Catalog(db).create_job(body, full=True)
|
||||
assert new.id != old.id
|
||||
async with runner.sessions.begin() as db:
|
||||
result = await Business(db).retry_job(old.id)
|
||||
assert result["id"] == new.id
|
||||
assert (await db.get(Job, old.id)).status == "failed"
|
||||
|
||||
|
||||
async def test_retry_waiting_full_job_requeues_its_checkpoint(catalog):
|
||||
from app.business import Business
|
||||
_, runner, _ = catalog
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
row = await db.get(Job, job.id)
|
||||
row.status, row.checkpoint = "waiting_connection", {"offset": 100}
|
||||
async with runner.sessions.begin() as db:
|
||||
result = await Business(db).retry_job(job.id)
|
||||
assert result["status"] == "queued" and result["checkpoint"]["offset"] == 100
|
||||
@@ -11,7 +11,7 @@ from app.ai.tools import CAPABILITIES
|
||||
from app.alphas import upsert_alpha
|
||||
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
|
||||
from app.business import Business
|
||||
from app.models import BacktestPreview, BacktestRun, Research, TemplateInput
|
||||
from app.models import BacktestPreview, BacktestRun, Research, ResearchInputSnapshot
|
||||
from tests.test_ai import configure, single_tool_factory
|
||||
from tests.test_backtests import execute, setup, start
|
||||
from tests.test_catalog import SCOPE, prepare, sync
|
||||
@@ -47,7 +47,7 @@ def construction(input_id):
|
||||
return {
|
||||
"name": "字段研究",
|
||||
"hypothesis": "显式字段排序",
|
||||
"template_input_id": input_id,
|
||||
"input_snapshot_id": input_id,
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "one",
|
||||
@@ -66,7 +66,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
|
||||
conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
|
||||
context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"}
|
||||
if use_saved_input:
|
||||
context["template_input_id"] = fixed_input["id"]
|
||||
context["input_snapshot_id"] = fixed_input["id"]
|
||||
run = await ask(logged_in, conversation, "研究此输入" if use_saved_input else "自行选字段研究", context)
|
||||
assert run["status"] == "waiting_approval", run
|
||||
assert not platform.posts
|
||||
@@ -75,7 +75,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
|
||||
assert source["kind"] == "chatbox"
|
||||
assert source["reference"] == conversation
|
||||
assert source["research_id"] == run["id"]
|
||||
assert source["template_input_id"]
|
||||
assert source["input_snapshot_id"]
|
||||
assert approval["preview"]["backtest"]["items"][0]["expression"] == "rank(TEST_FIN_001)"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
||||
@@ -128,7 +128,7 @@ async def test_invalid_construction_has_no_partial_preview(logged_in, app, fixed
|
||||
elif invalid == "unknown_type":
|
||||
item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"}
|
||||
else:
|
||||
body["template_input_id"] = "missing"
|
||||
body["input_snapshot_id"] = "missing"
|
||||
response = await logged_in.post("/api/v1/backtests/research-previews", json=body)
|
||||
assert response.status_code == (404 if invalid == "missing_input" else 422), response.text
|
||||
async with app.state.sessions() as db:
|
||||
@@ -144,15 +144,10 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
|
||||
assert page["field_count"] == page["total"] == 123 and len(page["items"]) == 23
|
||||
assert page["items"][-1]["field_type"] == "FUTURE_TYPE" and not page["has_more"]
|
||||
assert page["_meta"]["source"] == "local_database"
|
||||
selected = await tool(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
)
|
||||
subset = (await logged_in.post("/api/v1/data-preparations", json={"name": "one field", "scope": SCOPE,
|
||||
"fields": [{"scope": SCOPE, "dataset_id": "TEST_FIN", "field_id": "TEST_FIN_001", "source": "local",
|
||||
"collection_version": fixed_input["fields"][0]["collection_version"]}]})).json()
|
||||
selected = await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
||||
assert selected["field_count"] == 1
|
||||
bad = construction(selected["id"])
|
||||
bad["candidates"][0]["bindings"]["signal"]["field_id"] = "TEST_FIN_002"
|
||||
@@ -166,19 +161,12 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
|
||||
"/api/v1/backtests/research-previews", json=construction(fixed_input["id"])
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
await logged_in.patch(f"/api/v1/data-preparations/{subset['id']}", json={"name": "edited", "version": subset["version"]})
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await tool(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
)
|
||||
await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
||||
assert exc.value.status_code == 409
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 2
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 2
|
||||
|
||||
|
||||
async def test_multiple_origins_preserve_research_and_do_not_duplicate_alphas(app, logged_in):
|
||||
|
||||
@@ -213,7 +213,7 @@ async def test_operator_annotation_and_refresh_preserve_local(app, logged_in, re
|
||||
|
||||
|
||||
async def test_settings_variant_requires_all_fields_in_target(app, logged_in, research_input, catalog):
|
||||
from app.models import CatalogScope, TemplateInput
|
||||
from app.models import CatalogScope, ResearchInputSnapshot
|
||||
from app.research.experiments import Experiments
|
||||
from app.research.workspace_contracts import SettingVariants
|
||||
|
||||
@@ -232,14 +232,13 @@ async def test_settings_variant_requires_all_fields_in_target(app, logged_in, re
|
||||
db.add(CatalogScope(key=target_key, scope=target_scope))
|
||||
await db.flush()
|
||||
db.add(
|
||||
TemplateInput(
|
||||
ResearchInputSnapshot(
|
||||
id="target",
|
||||
scope_key=target_key,
|
||||
dataset_id="TEST_FIN",
|
||||
collection_version=research_input["collection_version"],
|
||||
selection="explicit",
|
||||
field_ids=["TEST_FIN_001"],
|
||||
field_types={"TEST_FIN_001": "MATRIX"},
|
||||
preparation_id="target-preparation",
|
||||
preparation_version=1,
|
||||
content={**{k: v for k, v in research_input.items() if k not in ("id", "preparation_id", "preparation_version", "created_at")}, "scope": target_scope,
|
||||
"field_ids": ["TEST_FIN_001"], "field_types": {"TEST_FIN_001": "MATRIX"},
|
||||
"fields": [f for f in research_input["fields"] if f["id"] == "TEST_FIN_001"]},
|
||||
)
|
||||
)
|
||||
await db.flush()
|
||||
@@ -514,3 +513,16 @@ def test_real_seed_settings_preserve_execution_options_and_reject_unknowns():
|
||||
assert snapshot["startDate"] == "2014-01-01"
|
||||
with pytest.raises(ValidationError):
|
||||
seed_settings({**snapshot, "unknownOption": True})
|
||||
|
||||
async def test_model_cannot_append_preparation_references_to_fixed_inputs(app, logged_in, research_input, monkeypatch):
|
||||
from app.models import ResearchAsset
|
||||
from app.research import routes
|
||||
from app.research.workspace_contracts import FeatureSpec
|
||||
async def model(*args):
|
||||
return FeatureSpec(name="untrusted", hypothesis="test", input_ids=[research_input["id"]],
|
||||
preparation_refs=[{"id": research_input["preparation_id"], "version": research_input["preparation_version"]}]), {}
|
||||
monkeypatch.setattr(routes, "request_model", model)
|
||||
response = await logged_in.post("/api/v1/research/generate", json={"name": "test", "hypothesis": "test", "method": "feature", "input_ids": [research_input["id"]]})
|
||||
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"))
|
||||
|
||||
Reference in New Issue
Block a user