256 lines
10 KiB
Python
256 lines
10 KiB
Python
import asyncio
|
|
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.exc import OperationalError
|
|
|
|
from app.jobs import Runner, create_job
|
|
from app.models import Account, Alpha, Job, JobItem, Pnl, Research
|
|
from app.security import cipher
|
|
from app.worldquant import VerificationRequired, WqError
|
|
from tests.conftest import alpha
|
|
|
|
|
|
class FakePlatform:
|
|
def __init__(self):
|
|
self.verification_url = None
|
|
self.on_retry = None
|
|
self.page_calls = []
|
|
self.fail_page = True
|
|
self.fail_id = True
|
|
self.require_verification = False
|
|
self.block = None
|
|
self.permissions = ["CONSULTANT", "SUPER_ALPHA"]
|
|
|
|
async def authenticate(self, *args, **kwargs):
|
|
if self.require_verification:
|
|
self.verification_url = "https://api.worldquantbrain.com/authentication/persona/test"
|
|
raise VerificationRequired(self.verification_url)
|
|
|
|
async def verify(self):
|
|
self.require_verification = False
|
|
self.verification_url = None
|
|
|
|
async def profile(self):
|
|
return {"id": "user1", "email": "test@example.com", "password": "never-save", "unknown": 1}
|
|
|
|
async def account_usage(self):
|
|
return {"date": "2026-09-07", "alphas": {"active": 99}, "errors": {}}
|
|
|
|
async def alphas(self, submission, hidden, offset, before):
|
|
self.page_calls.append((submission, hidden, offset, before))
|
|
if self.block:
|
|
self.block.set()
|
|
await asyncio.Future()
|
|
if submission == "UNSUBMITTED" and not hidden:
|
|
if offset == 0:
|
|
return {"results": [alpha("a1"), alpha("a2")], "count": 4}
|
|
if self.fail_page:
|
|
raise WqError("模拟第二页网络故障", "network_error")
|
|
return {"results": [alpha("a2", **{"is": {"sharpe": 7.0}}), alpha("a3")], "count": 4}
|
|
if submission == "SUBMITTED" and hidden:
|
|
return {"results": [alpha("hidden1", hidden=True, stage="OS", status="ACTIVE")], "next": None}
|
|
return {"results": [], "next": None}
|
|
|
|
async def alpha(self, alpha_id):
|
|
if alpha_id == "bad" and self.fail_id:
|
|
raise WqError("模拟无权访问", "access_denied")
|
|
return alpha(alpha_id, **{"is": {"sharpe": 8.0}})
|
|
|
|
async def pnl(self, alpha_id):
|
|
return {
|
|
"schema": {"properties": [{"name": "date"}, {"name": "pnl"}]},
|
|
"records": [["2025-01-01", 100]],
|
|
}
|
|
|
|
def disconnect(self):
|
|
self.verification_url = None
|
|
|
|
async def close(self):
|
|
pass
|
|
|
|
|
|
async def ready_runner(app):
|
|
async with app.state.sessions() as db:
|
|
account = await db.get(Account, 1)
|
|
account.email = "test@example.com"
|
|
account.password_encrypted = cipher(app.state.settings).encrypt(b"test-password").decode()
|
|
account.connection_status = "connected"
|
|
await db.commit()
|
|
return Runner(app.state.sessions, app.state.settings, FakePlatform())
|
|
|
|
|
|
async def job_for(runner, kind="full_sync", ids=None):
|
|
async with runner.sessions() as db:
|
|
return (await create_job(db, kind, {"alpha_ids": ids or []})).id
|
|
|
|
|
|
async def result(runner, job_id):
|
|
async with runner.sessions() as db:
|
|
return await db.get(Job, job_id)
|
|
|
|
|
|
async def test_pages_commit_resume_dedupe_and_preserve_research(app):
|
|
runner = await ready_runner(app)
|
|
job_id = await job_for(runner)
|
|
await runner.execute(job_id)
|
|
failed = await result(runner, job_id)
|
|
assert failed.status == "failed" and failed.processed == 2
|
|
assert failed.checkpoint == {"partition": 0, "offset": 2}
|
|
async with runner.sessions() as db:
|
|
research = await db.get(Research, "a2")
|
|
research.note, research.tags, research.state = "keep me", ["keep"], "candidate"
|
|
await db.commit()
|
|
runner.client.fail_page = False
|
|
runner.client.page_calls.clear()
|
|
await runner.execute(job_id)
|
|
finished = await result(runner, job_id)
|
|
assert finished.status == "completed" and finished.processed == 4 and finished.total == 4
|
|
assert runner.client.page_calls[0][:3] == ("UNSUBMITTED", False, 2)
|
|
assert len({call[3] for call in runner.client.page_calls}) == 1
|
|
async with runner.sessions() as db:
|
|
assert await db.scalar(select(func.count()).select_from(Alpha)) == 4
|
|
assert (await db.get(Alpha, "a2")).sharpe == 7.0
|
|
assert (await db.get(Alpha, "hidden1")).hidden is True
|
|
research = await db.get(Research, "a2")
|
|
assert (research.note, research.tags, research.state) == ("keep me", ["keep"], "candidate")
|
|
assert await db.scalar(select(func.count()).select_from(JobItem)) == 4
|
|
# A later full scan updates existing IDs, without deleting unseen local records.
|
|
second = await job_for(runner)
|
|
await runner.execute(second)
|
|
assert (await result(runner, second)).processed == 4
|
|
|
|
|
|
async def test_partial_ids_retry_only_errors_and_pnl_cache(app):
|
|
runner = await ready_runner(app)
|
|
job_id = await job_for(runner, "alpha_refresh", ["good", "bad"])
|
|
await runner.execute(job_id)
|
|
partial = await result(runner, job_id)
|
|
assert (partial.status, partial.processed, partial.failed, partial.total) == (
|
|
"completed_with_errors",
|
|
1,
|
|
1,
|
|
2,
|
|
)
|
|
runner.client.fail_id = False
|
|
await runner.execute(job_id)
|
|
complete = await result(runner, job_id)
|
|
assert (complete.status, complete.processed, complete.failed) == ("completed", 2, 0)
|
|
pnl_job = await job_for(runner, "pnl_refresh", ["good", "missing"])
|
|
await runner.execute(pnl_job)
|
|
assert (await result(runner, pnl_job)).failed == 1
|
|
async with runner.sessions() as db:
|
|
pnl = await db.get(Pnl, "good")
|
|
assert pnl.points == [{"date": "2025-01-01", "value": 100.0}]
|
|
|
|
|
|
async def test_human_verification_waits_then_resumes_pending_jobs(app):
|
|
runner = await ready_runner(app)
|
|
runner.client.require_verification = True
|
|
job_id = await job_for(runner, "alpha_refresh", ["good"])
|
|
await runner.execute(job_id)
|
|
assert (await result(runner, job_id)).status == "waiting_auth"
|
|
async with runner.sessions() as db:
|
|
account = await db.get(Account, 1)
|
|
assert account.connection_status == "verification_required" and account.verification_url
|
|
verification = await job_for(runner, "verify")
|
|
await runner.execute(verification)
|
|
assert (await result(runner, verification)).status == "completed"
|
|
assert (await result(runner, job_id)).status == "queued"
|
|
async with runner.sessions() as db:
|
|
account = await db.get(Account, 1)
|
|
assert (
|
|
account.wq_user_id == "user1"
|
|
and "password" not in account.profile
|
|
and "unknown" not in account.profile
|
|
)
|
|
await runner.execute(job_id)
|
|
assert (await result(runner, job_id)).status == "completed"
|
|
|
|
|
|
async def wait_status(runner, job_id, expected):
|
|
async with asyncio.timeout(5):
|
|
while (await result(runner, job_id)).status not in expected:
|
|
await asyncio.sleep(0.02)
|
|
|
|
|
|
async def test_user_cancel_interrupts_request_and_disconnect_pauses(app):
|
|
runner = await ready_runner(app)
|
|
runner.client.block = asyncio.Event()
|
|
job_id = await job_for(runner)
|
|
await runner.start()
|
|
await asyncio.wait_for(runner.client.block.wait(), 3)
|
|
async with runner.sessions() as db:
|
|
job = await db.get(Job, job_id)
|
|
job.cancel_requested = True
|
|
await db.commit()
|
|
await runner.cancel(job_id)
|
|
await wait_status(runner, job_id, {"cancelled"})
|
|
runner.client.block.clear()
|
|
pending = await job_for(runner)
|
|
runner.wake.set()
|
|
await asyncio.wait_for(runner.client.block.wait(), 3)
|
|
await runner.disconnect()
|
|
assert (await result(runner, pending)).status == "waiting_connection"
|
|
async with runner.sessions() as db:
|
|
assert (await db.get(Account, 1)).connection_status == "disconnected"
|
|
await runner.stop()
|
|
|
|
|
|
async def test_restart_recovers_running_job_and_keeps_checkpoint(app):
|
|
runner = await ready_runner(app)
|
|
runner.client.block = asyncio.Event()
|
|
job_id = await job_for(runner)
|
|
await runner.start()
|
|
await asyncio.wait_for(runner.client.block.wait(), 3)
|
|
await runner.stop()
|
|
assert (await result(runner, job_id)).status == "queued"
|
|
# Simulate abrupt process loss after a committed page, before the running flag was cleared.
|
|
async with runner.sessions() as db:
|
|
job = await db.get(Job, job_id)
|
|
job.status, job.checkpoint = "running", {"partition": 0, "offset": 2}
|
|
await db.commit()
|
|
resumed = Runner(runner.sessions, runner.settings, FakePlatform())
|
|
resumed.client.fail_page = False
|
|
await resumed.start()
|
|
await wait_status(resumed, job_id, {"completed", "failed"})
|
|
await resumed.stop()
|
|
assert (await result(resumed, job_id)).status == "completed"
|
|
assert resumed.client.page_calls[0][:3] == ("UNSUBMITTED", False, 2)
|
|
|
|
|
|
async def test_identity_change_fails_without_overwriting_profile(app):
|
|
runner = await ready_runner(app)
|
|
async with runner.sessions() as db:
|
|
account = await db.get(Account, 1)
|
|
account.wq_user_id = "original-user"
|
|
await db.commit()
|
|
job_id = await job_for(runner, "connect")
|
|
await runner.execute(job_id)
|
|
assert (await result(runner, job_id)).status == "waiting_connection"
|
|
async with runner.sessions() as db:
|
|
assert (await db.get(Account, 1)).wq_user_id == "original-user"
|
|
|
|
|
|
async def test_scheduler_recovers_after_database_outage(app):
|
|
runner = await ready_runner(app)
|
|
job_id = await job_for(runner, "alpha_refresh", ["good"])
|
|
original = runner.run_next
|
|
attempts = 0
|
|
|
|
async def failing_once():
|
|
nonlocal attempts
|
|
attempts += 1
|
|
if attempts == 1:
|
|
async with runner.sessions() as db:
|
|
job = await db.get(Job, job_id)
|
|
job.status = "running"
|
|
await db.commit()
|
|
raise OperationalError("", {}, Exception("simulated database restart"))
|
|
await original()
|
|
|
|
runner.run_next = failing_once
|
|
await runner.start()
|
|
await wait_status(runner, job_id, {"completed", "failed"})
|
|
await runner.stop()
|
|
assert attempts >= 2 and (await result(runner, job_id)).status == "completed"
|