merge: integrate alpha management with main and sequence migration 0005
This commit is contained in:
@@ -17,6 +17,23 @@ async def fake_stream(messages, info):
|
||||
str(p.content) for m in messages[latest:] for p in m.parts if isinstance(p, UserPromptPart)
|
||||
)
|
||||
returns = [p for m in messages[latest:] for p in m.parts if isinstance(p, ToolReturnPart)]
|
||||
if returns and returns[-1].tool_name == "prepare_backtest":
|
||||
content = returns[-1].content
|
||||
content = json.loads(content) if isinstance(content, str) else content
|
||||
yield {
|
||||
0: DeltaToolCall(
|
||||
name="start_backtest",
|
||||
json_args=json.dumps(
|
||||
{
|
||||
"preview_id": content["preview_id"],
|
||||
"version": 1,
|
||||
"idempotency_key": content["preview_id"],
|
||||
}
|
||||
),
|
||||
tool_call_id=uuid4().hex,
|
||||
)
|
||||
}
|
||||
return
|
||||
if returns and "LOOP" not in text:
|
||||
if returns[-1].tool_name == "capability_probe":
|
||||
yield str(returns[-1].content)
|
||||
@@ -35,6 +52,23 @@ async def fake_stream(messages, info):
|
||||
await asyncio.sleep(2)
|
||||
yield ",查询完成。"
|
||||
return
|
||||
elif "回测" in text:
|
||||
name, args = (
|
||||
"prepare_backtest",
|
||||
{
|
||||
"inline": {
|
||||
"name": "AI 固定回测",
|
||||
"source": {"kind": "ai"},
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "ai-1",
|
||||
"expression": "rank(close)",
|
||||
"settings": {"region": "USA", "universe": "TOP3000", "delay": 1},
|
||||
}
|
||||
],
|
||||
}
|
||||
},
|
||||
)
|
||||
elif "批量" in text:
|
||||
name, args = "bulk_update_research", {"alpha_ids": ["a0000", "a0001"], "add_tags": ["AI"]}
|
||||
elif "修改" in text or "update" in text:
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Synthetic simulation HTTP used by isolated API and browser acceptance."""
|
||||
|
||||
import json
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class Platform:
|
||||
def __init__(self):
|
||||
self.posts = []
|
||||
self.existing_alpha_ids = None
|
||||
self.simulations = {}
|
||||
self.alphas = {}
|
||||
self.reject = None
|
||||
self.pending = False
|
||||
self.detail_fail = False
|
||||
self.fail_child = None
|
||||
self.missing = False
|
||||
self.secret = "synthetic-platform-secret"
|
||||
|
||||
def __call__(self, request):
|
||||
path = request.url.path
|
||||
if path == "/authentication":
|
||||
return httpx.Response(201, json={})
|
||||
if path == "/simulations" and request.method == "POST":
|
||||
data = json.loads(request.content)
|
||||
data = data if isinstance(data, list) else [data]
|
||||
self.posts.append(data)
|
||||
if self.reject == "unknown":
|
||||
raise httpx.ReadTimeout("synthetic timeout", request=request)
|
||||
if self.reject == "session":
|
||||
self.reject = None
|
||||
return httpx.Response(401)
|
||||
if self.reject == "rate":
|
||||
return httpx.Response(429, headers={"Retry-After": "0.01"})
|
||||
if self.reject == "bad":
|
||||
return httpx.Response(400, json={"error": self.secret})
|
||||
parent = f"p{len(self.posts)}"
|
||||
ids = []
|
||||
for i, item in enumerate(data):
|
||||
child = parent if len(data) == 1 else f"{parent}c{i}"
|
||||
aid = self.existing_alpha_ids[i] if self.existing_alpha_ids else f"alpha{parent}{i}"
|
||||
progress = {
|
||||
"status": "COMPLETE",
|
||||
"alpha": aid,
|
||||
"regular": item["regular"],
|
||||
"settings": item["settings"],
|
||||
}
|
||||
if i == self.fail_child:
|
||||
progress = {
|
||||
"status": "FAILED",
|
||||
"regular": item["regular"],
|
||||
"settings": item["settings"],
|
||||
"message": "invalid expression",
|
||||
}
|
||||
self.simulations[child] = progress
|
||||
self.alphas[aid] = {
|
||||
"id": aid,
|
||||
"regular": {"code": item["regular"]},
|
||||
"type": "REGULAR",
|
||||
"settings": item["settings"],
|
||||
"is": {"sharpe": None, "fitness": 0.8},
|
||||
"status": "UNSUBMITTED",
|
||||
}
|
||||
ids.append(child)
|
||||
if len(data) > 1:
|
||||
self.simulations[parent] = {
|
||||
"status": "COMPLETE",
|
||||
"children": list(reversed(ids[1:] if self.missing else ids)),
|
||||
}
|
||||
if self.reject == "missing_location":
|
||||
return httpx.Response(201)
|
||||
return httpx.Response(
|
||||
201, headers={"Location": f"https://api.worldquantbrain.com/simulations/{parent}"}
|
||||
)
|
||||
if path.startswith("/simulations/"):
|
||||
return httpx.Response(
|
||||
200, json={"status": "PENDING"} if self.pending else self.simulations[path.rsplit("/", 1)[-1]]
|
||||
)
|
||||
if path.startswith("/alphas/"):
|
||||
if self.detail_fail:
|
||||
return httpx.Response(404)
|
||||
return httpx.Response(200, json=self.alphas[path.rsplit("/", 1)[-1]])
|
||||
if path == "/users/self":
|
||||
return httpx.Response(200, json={"id": "TEST_USER"})
|
||||
if path.startswith("/users/self/"):
|
||||
return httpx.Response(200, json={"results": [], "count": 0})
|
||||
raise AssertionError(f"Unexpected HTTP {request.method} {path}")
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Isolated PostgreSQL migration/concurrency acceptance. Never point at a personal database.
|
||||
|
||||
Run with DATABASE_URL ending in /wq_backtest_test, synthetic ADMIN_PASSWORD and
|
||||
ENCRYPTION_KEY. Uses only mock WorldQuant HTTP and a disposable database.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
import httpx
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from app.alphas import upsert_alpha
|
||||
from app.config import Settings
|
||||
from app.db import create_database
|
||||
from app.main import create_app
|
||||
from app.models import BacktestEvent, BacktestResult, BacktestRun, Research, SimulationAttempt
|
||||
from app.worldquant import WqClient
|
||||
from tests.backtest_fake import Platform
|
||||
from tests.test_backtests import candidate, preview, setup, start, tick
|
||||
|
||||
|
||||
async def seed_old(settings):
|
||||
engine, sessions = create_database(settings.database_url)
|
||||
async with sessions.begin() as db:
|
||||
await upsert_alpha(
|
||||
db,
|
||||
{
|
||||
"id": "MIGRATION_ALPHA",
|
||||
"type": "REGULAR",
|
||||
"regular": {"code": "rank(close) + 0"},
|
||||
"settings": candidate()["settings"],
|
||||
},
|
||||
)
|
||||
await db.flush()
|
||||
research = await db.get(Research, "MIGRATION_ALPHA")
|
||||
research.note = "keep old research across upgrade and simulations"
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def acceptance(settings):
|
||||
fake = Platform()
|
||||
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(fake)))
|
||||
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": settings.admin_password.get_secret_value()},
|
||||
)
|
||||
).status_code == 200
|
||||
fake, lane = await setup(app)
|
||||
fake.existing_alpha_ids = ["MIGRATION_ALPHA"]
|
||||
p = await preview(client, [candidate(0), candidate(0) | {"client_item_id": "repeat"}])
|
||||
a, b = await asyncio.gather(
|
||||
start(client, p, "concurrent-confirm"), start(client, p, "concurrent-confirm")
|
||||
)
|
||||
assert a["backtest_run_id"] == b["backtest_run_id"]
|
||||
rid = a["backtest_run_id"]
|
||||
for _ in range(5):
|
||||
await tick(lane)
|
||||
result = (await client.get(f"/api/v1/backtests/runs/{rid}")).json()
|
||||
assert result["status"] == "completed", result
|
||||
assert len(fake.posts) == 2
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestResult)) == 2
|
||||
note = (await db.get(Research, "MIGRATION_ALPHA")).note
|
||||
assert note == "keep old research across upgrade and simulations"
|
||||
events = list(
|
||||
await db.scalars(
|
||||
select(BacktestEvent.seq)
|
||||
.where(BacktestEvent.run_id == rid)
|
||||
.order_by(BacktestEvent.seq)
|
||||
)
|
||||
)
|
||||
assert events == list(range(1, len(events) + 1))
|
||||
# Leave an accepted run for a new application instance to recover.
|
||||
next_run = await start(client, await preview(client, [candidate(2)]), "restart")
|
||||
async with app.state.sessions() as db:
|
||||
aid = await db.scalar(
|
||||
select(SimulationAttempt.id).where(
|
||||
SimulationAttempt.run_id == next_run["backtest_run_id"]
|
||||
)
|
||||
)
|
||||
await lane.step(aid)
|
||||
await lane.interrupt()
|
||||
replacement = create_app(settings, WqClient(settings, transport=httpx.MockTransport(fake)))
|
||||
async with replacement.router.lifespan_context(replacement):
|
||||
lane = replacement.state.runner.backtests
|
||||
await lane.start()
|
||||
await lane.stop()
|
||||
await lane.step(aid)
|
||||
async with replacement.state.sessions() as db:
|
||||
assert (await db.get(BacktestRun, next_run["backtest_run_id"])).status == "completed"
|
||||
assert len(fake.posts) == 3
|
||||
print(
|
||||
"PASS PostgreSQL: concurrent confirmation creates one run; two attempts share one Alpha safely; contiguous transactional events; research preserved; replacement application resumes accepted simulation without POST"
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
if not os.environ.get("DATABASE_URL", "").endswith("/wq_backtest_test"):
|
||||
raise SystemExit("Only an isolated wq_backtest_test database is allowed")
|
||||
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
|
||||
config = Config("alembic.ini")
|
||||
command.upgrade(config, "0002")
|
||||
asyncio.run(seed_old(settings))
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
asyncio.run(acceptance(settings))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -12,6 +12,8 @@ from app.main import create_app
|
||||
from app.models import Base
|
||||
from app.worldquant import WqClient
|
||||
from tests.ai_fake import fake_model
|
||||
from tests.backtest_fake import Platform
|
||||
from tests.catalog_fake import catalog_response
|
||||
|
||||
TEST_PASSWORD = "browser-test-password"
|
||||
|
||||
@@ -81,6 +83,8 @@ def create_test_app():
|
||||
public_origin="http://127.0.0.1:5179",
|
||||
)
|
||||
records = [sample(i) for i in range(620)]
|
||||
simulations = Platform()
|
||||
simulations.existing_alpha_ids = [f"TEST{i:04}" for i in range(1, 100)]
|
||||
|
||||
def upstream(request):
|
||||
path = request.url.path
|
||||
@@ -105,8 +109,15 @@ def create_test_app():
|
||||
},
|
||||
headers={"Set-Cookie": "mock=only; Path=/"},
|
||||
)
|
||||
if path.startswith("/simulations") or (
|
||||
path.startswith("/alphas/") and path.rsplit("/", 1)[-1] in simulations.alphas
|
||||
):
|
||||
return simulations(request)
|
||||
if request.method != "GET":
|
||||
raise AssertionError("Browser acceptance attempted an upstream mutation")
|
||||
catalog = catalog_response(request)
|
||||
if catalog is not None:
|
||||
return catalog
|
||||
if path == "/users/self":
|
||||
return httpx.Response(
|
||||
200,
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Synthetic HTTP catalog, including page overlap and unknown metrics."""
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
def field_records(dataset="TEST_FIN", count=123):
|
||||
return [
|
||||
dict(
|
||||
id=f"{dataset}_{i:03}",
|
||||
name=f"TEST 字段 {i:03}",
|
||||
dataset={"id": dataset},
|
||||
type="FUTURE_TYPE" if i == 122 else "VECTOR" if i % 3 == 0 else "MATRIX",
|
||||
coverage=None if i == 122 else 0.95 if i % 2 else 0.6,
|
||||
userCount=None if i == 122 else i,
|
||||
alphaCount=i * 2,
|
||||
description=None if i == 122 else f"合成字段说明 {i}",
|
||||
)
|
||||
for i in range(count)
|
||||
]
|
||||
|
||||
|
||||
def catalog_response(request, fields=None):
|
||||
path, params = request.url.path, request.url.params
|
||||
if path not in ("/data-sets", "/data-fields"):
|
||||
return None
|
||||
assert request.method == "GET"
|
||||
assert params["instrumentType"] == "EQUITY"
|
||||
assert params["region"] and params["universe"] and params["delay"] in ("0", "1")
|
||||
dataset = params.get("dataset.id", "TEST_FIN")
|
||||
rows = (
|
||||
[
|
||||
{
|
||||
"id": "TEST_FIN",
|
||||
"name": "TEST 财务报表",
|
||||
"category": {"name": "基本面"},
|
||||
"subcategory": {"name": "财务报表"},
|
||||
"fieldCount": 123,
|
||||
"description": "合成数据,仅用于验收",
|
||||
},
|
||||
{
|
||||
"id": "TEST_NEWS",
|
||||
"name": "TEST 新闻",
|
||||
"category": {"name": "新闻"},
|
||||
"subcategory": {"name": "情绪"},
|
||||
"fieldCount": 3,
|
||||
},
|
||||
{"id": "TEST_UNKNOWN", "name": "TEST 未分类", "fieldCount": 0},
|
||||
]
|
||||
if path == "/data-sets"
|
||||
else (fields if fields is not None else field_records(dataset, 123 if dataset == "TEST_FIN" else 3))
|
||||
)
|
||||
if path == "/data-fields" and len(rows) > 50:
|
||||
rows = rows[:50] + [rows[49]] + rows[50:]
|
||||
offset, limit = int(params.get("offset", 0)), int(params.get("limit", 50))
|
||||
return httpx.Response(200, json={"results": rows[offset : offset + limit]})
|
||||
@@ -0,0 +1,124 @@
|
||||
"""One-off acceptance against the dedicated local PostgreSQL catalog_test database."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
|
||||
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
|
||||
|
||||
database_name = os.environ.get("WQ_CATALOG_ACCEPTANCE_DATABASE", "catalog_flow_test")
|
||||
if not re.fullmatch(r"catalog_[a-z0-9_]{1,40}", database_name):
|
||||
raise ValueError("Acceptance requires a dedicated catalog_* database")
|
||||
URL = f"postgresql+asyncpg://postgres:catalog-test-only@127.0.0.1:18436/{database_name}"
|
||||
os.environ.update(
|
||||
DATABASE_URL=URL, ADMIN_PASSWORD="migration-test-only", ENCRYPTION_KEY=Fernet.generate_key().decode()
|
||||
)
|
||||
|
||||
|
||||
async def sql(statement):
|
||||
engine = create_async_engine(URL)
|
||||
async with engine.begin() as connection:
|
||||
result = await connection.execute(text(statement))
|
||||
value = result.fetchall() if result.returns_rows else None
|
||||
await engine.dispose()
|
||||
return value
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
config = Config("alembic.ini")
|
||||
if asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")):
|
||||
raise RuntimeError("Acceptance database must be empty; existing data will not be overwritten")
|
||||
command.upgrade(config, "0002")
|
||||
asyncio.run(
|
||||
sql(
|
||||
"INSERT INTO alphas (id, hidden, settings, is_metrics, os_metrics, checks, synced_at, raw) VALUES ('MIGRATION_TEST', false, '{}', '{}', '{}', '[]', now(), '{}');"
|
||||
)
|
||||
)
|
||||
asyncio.run(
|
||||
sql(
|
||||
"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)
|
||||
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
|
||||
("preserve research", 7)
|
||||
]
|
||||
assert asyncio.run(sql("SELECT count(*) FROM catalog_batches")) == [(0,)]
|
||||
command.downgrade(config, "0002")
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
|
||||
("preserve research", 7)
|
||||
]
|
||||
print(
|
||||
"PostgreSQL 17: 0002 → 0003, downgrade/re-upgrade, metadata check, Alpha/research preservation passed"
|
||||
)
|
||||
|
||||
async def flow():
|
||||
import httpx
|
||||
|
||||
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"})
|
||||
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": "pg@example.com", "password": "synthetic-only"},
|
||||
)
|
||||
job = (await client.post("/api/v1/account/connect")).json()
|
||||
await app.state.runner.execute(job["id"])
|
||||
catalog = (client, app.state.runner, {})
|
||||
assert (await sync(catalog))["status"] == "completed"
|
||||
version = (await sync(catalog, "TEST_FIN"))["id"]
|
||||
result = await search(client, "/datasets/TEST_FIN/fields")
|
||||
assert result["complete_count"] == 123
|
||||
draft = (await prepare(client, version)).json()
|
||||
assert len(draft["field_ids"]) == 123
|
||||
responses = await asyncio.gather(
|
||||
*[
|
||||
client.patch(
|
||||
"/api/v1/catalog/datasets/TEST_FIN/research",
|
||||
params=SCOPE,
|
||||
json={"version": 1, "note": value},
|
||||
)
|
||||
for value in ["one", "two"]
|
||||
]
|
||||
)
|
||||
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()
|
||||
assert persisted == draft
|
||||
print(
|
||||
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
|
||||
)
|
||||
|
||||
asyncio.run(flow())
|
||||
@@ -0,0 +1,413 @@
|
||||
"""End-to-end business tests: real persistence/runtime, only the platform HTTP is replaced."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from app.backtests.contracts import SimulationSettings
|
||||
from app.models import Account, Alpha, BacktestResult, BacktestRun, Research, SimulationAttempt
|
||||
from app.security import cipher
|
||||
from app.worldquant import WqClient
|
||||
from tests.backtest_fake import Platform
|
||||
|
||||
PREFIX = "/api/v1/backtests"
|
||||
PARAMS = SimulationSettings(region="USA", universe="TOP3000", delay=1).model_dump()
|
||||
|
||||
|
||||
def candidate(index=0, **settings):
|
||||
return {
|
||||
"client_item_id": f"item-{index}",
|
||||
"expression": f"rank(close) + {index}",
|
||||
"settings": PARAMS | settings,
|
||||
}
|
||||
|
||||
|
||||
async def setup(app):
|
||||
platform = Platform()
|
||||
runner = app.state.runner
|
||||
await runner.client.close()
|
||||
runner.client = WqClient(app.state.settings, transport=httpx.MockTransport(platform))
|
||||
runner.backtests.client = runner.client
|
||||
runner.backtests.poll_interval = 0
|
||||
async with app.state.sessions.begin() as db:
|
||||
account = await db.get(Account, 1)
|
||||
account.email, account.wq_user_id, account.connection_status = (
|
||||
"synthetic@example.com",
|
||||
"TEST_USER",
|
||||
"connected",
|
||||
)
|
||||
account.password_encrypted = cipher(app.state.settings).encrypt(platform.secret.encode()).decode()
|
||||
return platform, runner.backtests
|
||||
|
||||
|
||||
async def preview(client, candidates=None):
|
||||
response = await client.post(
|
||||
f"{PREFIX}/previews",
|
||||
json={
|
||||
"inline": {
|
||||
"name": "测试研究",
|
||||
"source": {"kind": "test"},
|
||||
"candidates": candidates or [candidate()],
|
||||
}
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()
|
||||
|
||||
|
||||
async def start(client, p, key="request-1"):
|
||||
response = await client.post(
|
||||
f"{PREFIX}/runs",
|
||||
json={"preview_id": p["preview_id"], "version": p["version"], "idempotency_key": key},
|
||||
)
|
||||
assert response.status_code == 202, response.text
|
||||
return response.json()
|
||||
|
||||
|
||||
async def execute(app, lane, run_id):
|
||||
async with app.state.sessions() as db:
|
||||
ids = list(
|
||||
await db.scalars(
|
||||
select(SimulationAttempt.id)
|
||||
.where(SimulationAttempt.run_id == run_id)
|
||||
.order_by(SimulationAttempt.ordinal)
|
||||
)
|
||||
)
|
||||
for aid in ids:
|
||||
await lane.step(aid)
|
||||
await lane.step(aid)
|
||||
return ids
|
||||
|
||||
|
||||
async def test_fixed_preview_grouping_mapping_and_history(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
p = await preview(logged_in, [candidate(0), candidate(1, universe="TOP1000"), candidate(2, delay=0)])
|
||||
assert p["batch_count"] == 2 and p["total"] == 3
|
||||
run = await start(logged_in, p)
|
||||
again = await start(logged_in, p)
|
||||
assert run["backtest_run_id"] == again["backtest_run_id"]
|
||||
rid = run["backtest_run_id"]
|
||||
await execute(app, lane, rid)
|
||||
data = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()
|
||||
assert len(platform.posts) == 2
|
||||
assert all(i["persistence_status"] == "saved" for i in data["items"]), data
|
||||
for item in data["items"]:
|
||||
assert item["result"]["snapshot"]["regular"]["code"] == item["expression"]
|
||||
assert item["result"]["snapshot"]["settings"] == item["settings"]
|
||||
assert item["result"]["snapshot"]["is"]["sharpe"] is None
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
item = data["items"][0]
|
||||
async with app.state.sessions.begin() as db:
|
||||
alpha = await db.get(Alpha, item["alpha_id"])
|
||||
alpha.is_metrics = {"sharpe": 999}
|
||||
assert await db.get(Research, alpha.id)
|
||||
historical = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()
|
||||
assert historical["items"][0]["result"]["snapshot"]["is"]["sharpe"] is None
|
||||
events = (await logged_in.get(f"{PREFIX}/runs/{rid}/events?limit=2")).json()
|
||||
later = (await logged_in.get(f"{PREFIX}/runs/{rid}/events?after={events['next_cursor']}")).json()
|
||||
assert events["has_more"] and later["items"][0]["seq"] > events["next_cursor"]
|
||||
assert (await preview(logged_in))["duplicate_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rejection", ["unknown", "missing_location"])
|
||||
async def test_unknown_submission_never_reposted(app, logged_in, rejection):
|
||||
platform, lane = await setup(app)
|
||||
platform.reject = rejection
|
||||
run = await start(logged_in, await preview(logged_in))
|
||||
rid = run["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
# A process crash/recovery must not turn an unknown POST into queued work.
|
||||
await lane.start()
|
||||
await lane.stop()
|
||||
response = await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "recover", "version": 1})
|
||||
assert response.json()["status"] == "needs_review"
|
||||
async with app.state.sessions() as db:
|
||||
assert (await db.get(SimulationAttempt, ids[0])).state == "needs_review"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_partial_failure_and_rerun_only_selected(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.fail_child = 0
|
||||
run = await start(logged_in, await preview(logged_in, [candidate(0), candidate(1)]))
|
||||
rid = run["backtest_run_id"]
|
||||
await execute(app, lane, rid)
|
||||
result = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()["items"]
|
||||
assert result[0]["platform_status"] == "failed" and result[1]["persistence_status"] == "saved"
|
||||
rerun = await logged_in.post(f"{PREFIX}/runs/{rid}/rerun-preview", json={"item_ids": [result[0]["id"]]})
|
||||
assert rerun.status_code == 201
|
||||
assert rerun.json()["total"] == 1 and rerun.json()["source"]["parent_run_id"] == rid
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_detail_failure_recovers_without_resubmit(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.detail_fail = True
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "needs_review"
|
||||
platform.detail_fail = False
|
||||
await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "recover", "version": 1})
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_draft_version_snapshot_and_pause_stop(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
body = {"name": "草稿", "candidates": [candidate(0), candidate(1, delay=0)]}
|
||||
d = (await logged_in.post(f"{PREFIX}/drafts", json=body)).json()
|
||||
p = (await logged_in.post(f"{PREFIX}/previews", json={"draft_id": d["id"], "draft_version": 1})).json()
|
||||
changed = await logged_in.put(
|
||||
f"{PREFIX}/drafts/{d['id']}", json=body | {"version": 1, "candidates": [candidate(9)]}
|
||||
)
|
||||
assert changed.json()["version"] == 2
|
||||
assert (
|
||||
await logged_in.post(f"{PREFIX}/previews", json={"draft_id": d["id"], "draft_version": 1})
|
||||
).status_code == 409
|
||||
run = await start(logged_in, p)
|
||||
rid = run["backtest_run_id"]
|
||||
async with app.state.sessions() as db:
|
||||
ids = list(
|
||||
await db.scalars(
|
||||
select(SimulationAttempt.id)
|
||||
.where(SimulationAttempt.run_id == rid)
|
||||
.order_by(SimulationAttempt.ordinal)
|
||||
)
|
||||
)
|
||||
await lane.step(ids[0])
|
||||
await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "pause", "version": 1})
|
||||
await lane.step(ids[1])
|
||||
await lane.step(ids[0])
|
||||
assert len(platform.posts) == 1
|
||||
await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "stop", "version": 2})
|
||||
r = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()["items"]
|
||||
assert r[0]["persistence_status"] == "saved" and r[1]["platform_status"] == "skipped"
|
||||
assert r[0]["expression"] == candidate(0)["expression"]
|
||||
|
||||
|
||||
async def test_batch_missing_child_does_not_misattribute(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.missing = True
|
||||
rid = (await start(logged_in, await preview(logged_in, [candidate(0), candidate(1)])))["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
items = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()["items"]
|
||||
assert items[0]["platform_status"] == "unknown"
|
||||
assert items[1]["persistence_status"] == "saved"
|
||||
platform.simulations["p1"]["children"] = ["p1c1", "p1c0"]
|
||||
await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "recover", "version": 1})
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_validation_auth_and_idempotency_conflict(app, logged_in, client):
|
||||
await setup(app)
|
||||
assert (
|
||||
await logged_in.post(
|
||||
f"{PREFIX}/previews",
|
||||
json={"inline": {"name": "x", "candidates": [candidate() | {"alpha_type": "SUPER"}]}},
|
||||
)
|
||||
).status_code == 422
|
||||
p1, p2 = await preview(logged_in), await preview(logged_in, [candidate(2)])
|
||||
await start(logged_in, p1)
|
||||
assert (
|
||||
await logged_in.post(
|
||||
f"{PREFIX}/runs", json={"preview_id": p2["preview_id"], "idempotency_key": "request-1"}
|
||||
)
|
||||
).status_code == 409
|
||||
assert (await logged_in.get(f"{PREFIX}/runs?limit=101")).status_code == 422
|
||||
await client.post("/api/v1/auth/logout")
|
||||
assert (await client.get(f"{PREFIX}/runs")).status_code == 401
|
||||
|
||||
|
||||
async def tick(lane):
|
||||
await lane.tick()
|
||||
await asyncio.gather(*lane.tasks.values(), return_exceptions=False)
|
||||
|
||||
|
||||
async def test_account_budget_round_robin_and_sync_independence(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
assert (
|
||||
await logged_in.put(f"{PREFIX}/config", json={"concurrency": 1, "batch_size": 1, "version": 1})
|
||||
).status_code == 200
|
||||
r1 = await start(logged_in, await preview(logged_in, [candidate(0), candidate(1)]), "first")
|
||||
r2 = await start(logged_in, await preview(logged_in, [candidate(2), candidate(3)]), "second")
|
||||
await tick(lane) # one submission, occupied until remote terminal
|
||||
assert len(platform.posts) == 1
|
||||
await tick(lane) # poll first result
|
||||
await tick(lane) # other run gets next slot
|
||||
assert len(platform.posts) == 2
|
||||
assert platform.posts[0][0]["regular"] == candidate(0)["expression"]
|
||||
assert platform.posts[1][0]["regular"] == candidate(2)["expression"]
|
||||
platform.pending = True
|
||||
sync = await logged_in.post("/api/v1/sync-jobs", json={"kind": "full_sync"})
|
||||
await app.state.runner.run_next()
|
||||
assert (await logged_in.get(f"/api/v1/sync-jobs/{sync.json()['id']}")).json()["status"] == "completed"
|
||||
await logged_in.put(f"{PREFIX}/config", json={"concurrency": 2, "batch_size": 8, "version": 2})
|
||||
await tick(lane)
|
||||
assert len(platform.posts) == 3
|
||||
# Batch sizing of both existing runs remains 1 despite config update.
|
||||
assert all(len(p) == 1 for p in platform.posts)
|
||||
await logged_in.put(f"{PREFIX}/config", json={"concurrency": 1, "batch_size": 8, "version": 3})
|
||||
await tick(lane)
|
||||
assert len(platform.posts) == 3
|
||||
await lane.interrupt()
|
||||
assert r1["batch_size"] == r2["batch_size"] == 1
|
||||
|
||||
|
||||
async def test_rate_limit_and_failed_submit_are_bounded(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.reject = "rate"
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
async with app.state.sessions() as db:
|
||||
aid = await db.scalar(select(SimulationAttempt.id).where(SimulationAttempt.run_id == rid))
|
||||
for _ in range(app.state.settings.retry_attempts):
|
||||
await lane.step(aid)
|
||||
await asyncio.sleep(0.02)
|
||||
assert len(platform.posts) == app.state.settings.retry_attempts
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed_with_errors"
|
||||
assert platform.secret not in (await logged_in.get(f"{PREFIX}/runs/{rid}/attempts")).text
|
||||
|
||||
|
||||
async def test_poll_timeout_and_crash_after_acceptance(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
lane.poll_limit = 1
|
||||
platform.pending = True
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "needs_review"
|
||||
platform.pending = False
|
||||
await logged_in.post(f"{PREFIX}/runs/{rid}/control", json={"action": "recover", "version": 1})
|
||||
# Simulate a crash checkpoint with the Location already persisted.
|
||||
async with app.state.sessions.begin() as db:
|
||||
a = await db.get(SimulationAttempt, ids[0])
|
||||
a.state = "submitting"
|
||||
await lane.start()
|
||||
await lane.stop()
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_result_transaction_failure_recovers_from_saved_receipt(app, logged_in):
|
||||
from sqlalchemy import event
|
||||
from sqlalchemy.exc import OperationalError
|
||||
|
||||
platform, lane = await setup(app)
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
failed = False
|
||||
|
||||
def fail_once(conn, cursor, statement, parameters, context, executemany):
|
||||
nonlocal failed
|
||||
if "INSERT INTO backtest_results" in statement and not failed:
|
||||
failed = True
|
||||
raise OperationalError("synthetic persistence outage", {}, Exception("synthetic"))
|
||||
|
||||
event.listen(app.state.engine.sync_engine, "before_cursor_execute", fail_once)
|
||||
try:
|
||||
ids = await execute(app, lane, rid)
|
||||
finally:
|
||||
event.remove(app.state.engine.sync_engine, "before_cursor_execute", fail_once)
|
||||
assert failed
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestResult)) == 0
|
||||
assert await db.scalar(select(func.count()).select_from(Alpha)) == 0
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_ai_fixed_set_confirmation_and_duplicate_decision(app, logged_in):
|
||||
from tests.test_ai import configure
|
||||
from tests.test_ai import start as start_ai
|
||||
|
||||
platform, lane = await setup(app)
|
||||
await configure(app, logged_in)
|
||||
_, run, _ = await start_ai(app, logged_in, "回测固定候选")
|
||||
assert run["status"] == "waiting_approval", run
|
||||
approval = next(c for c in run["tools"] if c["name"] == "start_backtest")
|
||||
assert approval["preview"]["backtest"]["total"] == 1
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
||||
for _ in range(2):
|
||||
response = await logged_in.post(
|
||||
f"/api/v1/ai/approvals/{approval['id']}/decision", json={"approved": True}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
async with app.state.sessions() as db:
|
||||
rows = list(await db.scalars(select(BacktestRun)))
|
||||
assert len(rows) == 1
|
||||
assert rows[0].ai_context["ai_run_id"] == run["id"]
|
||||
await logged_in.post(f"/api/v1/ai/runs/{run['id']}/cancel")
|
||||
await execute(app, lane, rows[0].id)
|
||||
assert len(platform.posts) == 1
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rows[0].id}")).json()["status"] == "completed"
|
||||
|
||||
|
||||
async def test_duplicate_inputs_are_separate_attempts_and_share_alpha_safely(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.existing_alpha_ids = ["shared_alpha"]
|
||||
p = await preview(logged_in, [candidate(0), candidate(0) | {"client_item_id": "other-experiment"}])
|
||||
assert p["batch_count"] == 2 and p["duplicate_count"] == 1
|
||||
rid = (await start(logged_in, p))["backtest_run_id"]
|
||||
await execute(app, lane, rid)
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(Alpha)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestResult)) == 2
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
|
||||
|
||||
async def test_original_reference_recovery_without_new_post(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.reject = "missing_location"
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
path = f"{PREFIX}/attempts/{ids[0]}/reference"
|
||||
assert (
|
||||
await logged_in.post(
|
||||
path, json={"progress_url": "https://foreign.example/simulations/p1", "version": 1}
|
||||
)
|
||||
).status_code == 422
|
||||
linked = await logged_in.post(path, json={"progress_url": "/simulations/p1", "version": 1})
|
||||
assert linked.status_code == 200, linked.text
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 1
|
||||
|
||||
|
||||
async def test_preview_subset_uses_whole_snapshot_and_does_not_change_original(app, logged_in):
|
||||
await setup(app)
|
||||
p = await preview(logged_in, [candidate(i) for i in range(40)])
|
||||
subset = await logged_in.post(
|
||||
f"{PREFIX}/previews/{p['preview_id']}/subset", json={"exclude_ids": ["item-30"]}
|
||||
)
|
||||
assert subset.json()["total"] == 39 and subset.json()["preview_id"] != p["preview_id"]
|
||||
assert (await logged_in.get(f"{PREFIX}/previews/{p['preview_id']}")).json()["total"] == 40
|
||||
|
||||
|
||||
async def test_session_reauthentication_does_not_retry_accepted_submission(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
platform.reject = "session"
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
ids = await execute(app, lane, rid)
|
||||
await lane.step(ids[0])
|
||||
assert (await logged_in.get(f"{PREFIX}/runs/{rid}")).json()["status"] == "completed"
|
||||
assert len(platform.posts) == 2 # first explicitly rejected with 401, second accepted
|
||||
|
||||
|
||||
async def test_terminal_detail_failure_releases_slot_but_keeps_platform_success(app, logged_in):
|
||||
platform, lane = await setup(app)
|
||||
await logged_in.put(f"{PREFIX}/config", json={"concurrency": 1, "batch_size": 1, "version": 1})
|
||||
platform.detail_fail = True
|
||||
rid = (await start(logged_in, await preview(logged_in)))["backtest_run_id"]
|
||||
await execute(app, lane, rid)
|
||||
item = (await logged_in.get(f"{PREFIX}/runs/{rid}/results")).json()["items"][0]
|
||||
assert item["platform_status"] == "completed" and item["collection_status"] == "failed"
|
||||
await start(logged_in, await preview(logged_in, [candidate(2)]), "next")
|
||||
await tick(lane)
|
||||
assert len(platform.posts) == 2
|
||||
await lane.interrupt()
|
||||
@@ -0,0 +1,257 @@
|
||||
"""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}})
|
||||
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 == 422
|
||||
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
|
||||
state["persona"] = True
|
||||
runner.client.authenticated = False
|
||||
waiting = await sync(catalog, "TEST_FIN")
|
||||
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"
|
||||
Reference in New Issue
Block a user