282 lines
13 KiB
Python
282 lines
13 KiB
Python
"""Public API through real business/runner/database; only upstream HTTP is replaced."""
|
|
|
|
import asyncio
|
|
|
|
import httpx
|
|
import pytest
|
|
|
|
from app.jobs import Runner
|
|
from app.models import Job
|
|
from app.worldquant import WqClient
|
|
from tests.catalog_fake import catalog_response, field_records
|
|
|
|
SCOPE = dict(instrument_type="EQUITY", region="USA", universe="TOP3000", delay=1)
|
|
BASE = "/api/v1/catalog"
|
|
|
|
|
|
@pytest.fixture
|
|
async def catalog(logged_in, app):
|
|
state = {"fail": False, "fields": field_records(), "calls": [], "mode": "", "block": None}
|
|
|
|
async def upstream(request):
|
|
state["calls"].append((request.url.path, int(request.url.params.get("offset", 0))))
|
|
if request.url.path == "/authentication":
|
|
if state.get("persona"):
|
|
return httpx.Response(
|
|
401, headers={"WWW-Authenticate": "persona", "Location": "/authentication/persona/test"}
|
|
)
|
|
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
|
if request.method == "OPTIONS":
|
|
return catalog_response(request)
|
|
assert request.method == "GET"
|
|
if request.url.path == "/users/self":
|
|
return httpx.Response(200, json={"id": "TEST_USER"})
|
|
if request.url.path == "/data-fields":
|
|
if state.get("throttle"):
|
|
state["throttle"] = False
|
|
return httpx.Response(429, headers={"Retry-After": "2"})
|
|
if state["mode"] == "invalid-next":
|
|
return httpx.Response(200, json={"results": state["fields"][:50], "next": []})
|
|
if state["mode"] == "missing-owner":
|
|
return httpx.Response(200, json={"results": [{"id": "UNOWNED"}], "next": None})
|
|
if state["mode"] == "coverage-unit":
|
|
return httpx.Response(
|
|
200, json={"results": [{**state["fields"][0], "coverage": 95}], "next": None}
|
|
)
|
|
if int(request.url.params["offset"]) >= 50:
|
|
if state["block"]:
|
|
state["block"].set()
|
|
await asyncio.Future()
|
|
if state["fail"]:
|
|
return httpx.Response(403)
|
|
if state["mode"] == "early":
|
|
return httpx.Response(200, json={"results": [], "next": "/next", "count": 123})
|
|
if state["mode"] == "repeat":
|
|
return httpx.Response(200, json={"results": state["fields"][:50], "next": "/next"})
|
|
if state["mode"] == "wrong-owner":
|
|
return httpx.Response(200, json={"results": field_records("OTHER", 1)})
|
|
return catalog_response(request, state["fields"]) or httpx.Response(404)
|
|
|
|
await app.state.runner.client.close()
|
|
app.state.runner.client = WqClient(app.state.settings, transport=httpx.MockTransport(upstream))
|
|
client = logged_in
|
|
assert (
|
|
await client.put(
|
|
"/api/v1/account/credentials", json={"email": "test@example.com", "password": "test-only"}
|
|
)
|
|
).status_code == 200
|
|
connect = (await client.post("/api/v1/account/connect")).json()
|
|
await app.state.runner.execute(connect["id"])
|
|
return client, app.state.runner, state
|
|
|
|
|
|
async def sync(catalog, dataset=None, scope=SCOPE):
|
|
client, runner, _ = catalog
|
|
response = await client.post(BASE + "/sync-jobs", json={"scope": scope, "dataset_id": dataset})
|
|
assert response.status_code == 202, response.text
|
|
job = response.json()
|
|
await runner.execute(job["id"])
|
|
return (await client.get("/api/v1/sync-jobs/" + job["id"])).json()
|
|
|
|
|
|
async def search(client, suffix="/datasets", **params):
|
|
response = await client.get(BASE + suffix, params={**SCOPE, **params})
|
|
assert response.status_code == 200, response.text
|
|
return response.json()
|
|
|
|
|
|
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,
|
|
},
|
|
)
|
|
|
|
|
|
async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
|
client, _, state = catalog
|
|
assert (await search(client))["total"] == 0
|
|
assert (await sync(catalog))["status"] == "completed"
|
|
datasets = await search(client, category="基本面", subcategory="财务报表")
|
|
assert [r["id"] for r in datasets["items"]] == ["TEST_FIN"]
|
|
assert datasets["items"][0]["complete_count"] is None
|
|
assert (await sync(catalog, "TEST_FIN"))["processed"] == 123
|
|
fields = await search(client, "/datasets/TEST_FIN/fields", q="字段 12", limit=1)
|
|
assert fields["total"] == 3 and fields["complete_count"] == 123 and len(fields["items"]) == 1
|
|
version = fields["collection_version"]
|
|
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"
|
|
for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
|
|
detail = await search(client, suffix)
|
|
assert detail["research"]["version"] == 1
|
|
response = await client.patch(
|
|
BASE + suffix + "/research", params=SCOPE, json={"version": 1, "note": "保留研究备注"}
|
|
)
|
|
assert response.status_code == 200
|
|
assert (
|
|
await client.patch(
|
|
BASE + suffix + "/research", params=SCOPE, json={"version": 1, "note": "不能覆盖"}
|
|
)
|
|
).status_code == 409
|
|
detail = await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122")
|
|
assert detail["coverage"] is None and detail["unit"] is None and detail["field_type"] == "FUTURE_TYPE"
|
|
assert (await search(client, "/datasets/TEST_FIN/fields", coverage_min=0))["total"] == 122
|
|
state["fields"] = field_records(count=125)
|
|
assert (await sync(catalog, "TEST_FIN"))["processed"] == 125
|
|
assert (await sync(catalog))["status"] == "completed"
|
|
newer = await search(client, "/datasets/TEST_FIN/fields")
|
|
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 search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
|
|
"note"
|
|
] == "保留研究备注"
|
|
assert (await search(client, "/datasets/TEST_FIN"))["research"]["note"] == "保留研究备注"
|
|
|
|
|
|
async def test_partial_refresh_resume_cancel_restart_keeps_old_version(catalog):
|
|
client, runner, state = catalog
|
|
await sync(catalog)
|
|
state["fail"] = True
|
|
job = await sync(catalog, "TEST_FIN")
|
|
assert job["status"] == "failed" and job["processed"] == 50
|
|
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] is None
|
|
assert (await prepare(client, job["id"])).status_code == 409
|
|
state["fail"] = False
|
|
state["calls"].clear()
|
|
assert (await client.post("/api/v1/sync-jobs/" + job["id"] + "/retry")).status_code == 200
|
|
await runner.execute(job["id"])
|
|
assert state["calls"][0] == ("/data-fields", 50)
|
|
old_version = (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
|
|
state["fail"] = True
|
|
refresh = await sync(catalog, "TEST_FIN")
|
|
assert refresh["status"] == "failed"
|
|
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] == old_version
|
|
state["fail"] = False
|
|
state["calls"].clear()
|
|
async with runner.sessions() as db:
|
|
row = await db.get(Job, refresh["id"])
|
|
row.status = "running"
|
|
await db.commit()
|
|
restarted = Runner(runner.sessions, runner.settings, runner.client)
|
|
await restarted.start()
|
|
async with asyncio.timeout(5):
|
|
while True:
|
|
response = (await client.get("/api/v1/sync-jobs/" + refresh["id"])).json()
|
|
if response["status"] in ("completed", "failed"):
|
|
break
|
|
await asyncio.sleep(0.02)
|
|
assert response["status"] == "completed"
|
|
assert state["calls"][0] == ("/data-fields", 50)
|
|
state["block"] = asyncio.Event()
|
|
response = await client.post(BASE + "/sync-jobs", json={"scope": SCOPE, "dataset_id": "TEST_FIN"})
|
|
cancel_id = response.json()["id"]
|
|
restarted.wake.set()
|
|
await asyncio.wait_for(state["block"].wait(), 5)
|
|
await client.post("/api/v1/sync-jobs/" + cancel_id + "/cancel")
|
|
await restarted.cancel(cancel_id)
|
|
assert (await client.get("/api/v1/sync-jobs/" + cancel_id)).json()["status"] == "cancelled"
|
|
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] == refresh["id"]
|
|
await restarted.stop()
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"mode", ["early", "repeat", "wrong-owner", "invalid-next", "missing-owner", "coverage-unit"]
|
|
)
|
|
async def test_anomalous_pagination_is_never_complete(catalog, mode):
|
|
client, _, state = catalog
|
|
await sync(catalog)
|
|
state["mode"] = mode
|
|
assert (await sync(catalog, "TEST_FIN"))["status"] == "failed"
|
|
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] is None
|
|
|
|
|
|
async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
|
|
client, _, _ = catalog
|
|
await sync(catalog)
|
|
version = (await sync(catalog, "TEST_FIN"))["id"]
|
|
assert (await prepare(client, version, selection="explicit", excluded_ids=["OTHER"])).status_code == 422
|
|
assert (
|
|
await prepare(
|
|
client, version, selection="explicit", excluded_ids=[f"TEST_FIN_{i:03}" for i in range(123)]
|
|
)
|
|
).status_code == 422
|
|
assert (await prepare(client, version, dataset_id="TEST_NEWS")).status_code == 409
|
|
assert (await prepare(client, version, scope={**SCOPE, "delay": 0})).status_code == 404
|
|
assert (await prepare(client, version, scope={**SCOPE, "region": "CHN"})).status_code == 404
|
|
subset = await prepare(client, version, selection="explicit", excluded_ids=["TEST_FIN_110"])
|
|
assert subset.status_code == 201 and len(subset.json()["field_ids"]) == 122
|
|
assert "TEST_FIN_110" not in subset.json()["field_ids"]
|
|
other = {**SCOPE, "delay": 0}
|
|
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
|
|
|
|
|
|
async def test_catalog_authentication_and_origin(app, client):
|
|
assert (await client.get(BASE + "/datasets", params=SCOPE)).status_code == 401
|
|
assert (
|
|
await client.post(BASE + "/sync-jobs", headers={"Origin": "http://evil.test"}, json={"scope": SCOPE})
|
|
).status_code == 403
|
|
|
|
|
|
async def test_retry_after_auth_wait_disconnect_and_collection_manifest(catalog):
|
|
client, runner, state = catalog
|
|
await sync(catalog)
|
|
delays = []
|
|
|
|
async def sleep(delay):
|
|
delays.append(delay)
|
|
|
|
runner.client.sleep = sleep
|
|
state["throttle"] = True
|
|
job = await sync(catalog, "TEST_FIN")
|
|
assert job["status"] == "completed" and delays == [2]
|
|
manifest = await search(client, "/datasets/TEST_FIN/collection")
|
|
assert manifest["collection_version"] == job["id"] and len(manifest["field_ids"]) == 123
|
|
queued = await client.post(BASE + "/sync-jobs", json={"scope": SCOPE, "dataset_id": "TEST_FIN"})
|
|
state["persona"] = True
|
|
runner.client.authenticated = False
|
|
await runner.execute(queued.json()["id"])
|
|
waiting = (await client.get("/api/v1/sync-jobs/" + queued.json()["id"])).json()
|
|
assert waiting["status"] == "waiting_auth"
|
|
assert (await search(client, "/datasets/TEST_FIN/collection")) == manifest
|
|
await runner.disconnect()
|
|
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "waiting_connection"
|
|
assert (await client.post(BASE + "/sync-jobs", json={"scope": SCOPE})).status_code == 409
|
|
# Explicit reconnect verifies the original account and resumes the same task.
|
|
state["persona"] = False
|
|
connect = (await client.post("/api/v1/account/connect")).json()
|
|
await runner.execute(connect["id"])
|
|
await runner.execute(waiting["id"])
|
|
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "completed"
|
|
|
|
|
|
async def test_dynamic_platform_scopes_and_validation(catalog):
|
|
client, runner, _ = catalog
|
|
options = (await client.get(BASE + "/scopes")).json()
|
|
assert "IND" in options["regions_by_type"]["EQUITY"]
|
|
assert (await sync(catalog, scope={**SCOPE, "region": "IND", "universe": "TOP500"}))["status"] == "completed"
|
|
invalid = await client.post(BASE + "/sync-jobs", json={"scope": {**SCOPE, "region": "IND"}})
|
|
assert invalid.status_code == 422
|
|
from app.ai.capabilities import ToolContext
|
|
from app.ai.tools import CAPABILITIES
|
|
from app.business import Business
|
|
|
|
async with runner.sessions() as db:
|
|
ai = await CAPABILITIES["get_catalog_scopes"].invoke(ToolContext(Business(db), runner.client), {})
|
|
assert ai["instrument_options"] == options["instrument_options"]
|
|
assert ai["_meta"]["source"] == "worldquant_platform"
|
|
runner.client.disconnect()
|
|
assert (await client.get(BASE + "/scopes")).status_code == 409
|
|
assert (await search(client))["total"] == 0
|