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:
@@ -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
|
||||
Reference in New Issue
Block a user