feat: add initial frontend setup with styles, types, and testing framework
- Created a global CSS file for styling the frontend with responsive design. - Introduced TypeScript types for various entities including Research, Alpha, and Account. - Implemented Playwright tests for account management and data synchronization workflows. - Configured TypeScript with strict settings and included necessary libraries. - Set up Vite as the build tool with React plugin and API proxy configuration. - Added a Python script to initialize environment variables securely.
This commit is contained in:
@@ -0,0 +1,199 @@
|
||||
"""Normalize upstream business data, without making platform or research decisions."""
|
||||
|
||||
import math
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
|
||||
from .models import Alpha, Research, ResearchTag, now
|
||||
|
||||
SENSITIVE_KEYS = {
|
||||
"password",
|
||||
"token",
|
||||
"cookies",
|
||||
"cookie",
|
||||
"authorization",
|
||||
"credentials",
|
||||
"secret",
|
||||
"accesstoken",
|
||||
"refreshtoken",
|
||||
"sessiontoken",
|
||||
"clientsecret",
|
||||
"authorizationheader",
|
||||
"setcookie",
|
||||
"apikey",
|
||||
"csrftoken",
|
||||
"xsrftoken",
|
||||
"authentication",
|
||||
}
|
||||
|
||||
|
||||
def sanitize(value):
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
k: sanitize(v) for k, v in value.items() if re.sub(r"[^a-z]", "", k.lower()) not in SENSITIVE_KEYS
|
||||
}
|
||||
if isinstance(value, list):
|
||||
return [sanitize(v) for v in value]
|
||||
if isinstance(value, float) and not math.isfinite(value):
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
def number(value):
|
||||
if value is None or isinstance(value, bool):
|
||||
return None
|
||||
try:
|
||||
result = float(value)
|
||||
return result if math.isfinite(result) else None
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def date(value):
|
||||
try:
|
||||
return datetime.fromisoformat(value.replace("Z", "+00:00")) if value else None
|
||||
except (ValueError, TypeError, AttributeError):
|
||||
return None
|
||||
|
||||
|
||||
def code(value):
|
||||
return value.get("code") if isinstance(value, dict) else value if isinstance(value, str) else None
|
||||
|
||||
|
||||
async def upsert_alpha(db, raw: dict):
|
||||
alpha_id = raw.get("id")
|
||||
if not isinstance(alpha_id, str) or not alpha_id:
|
||||
raise ValueError("Alpha 数据缺少 ID")
|
||||
item = await db.get(Alpha, alpha_id)
|
||||
if item is None:
|
||||
item = Alpha(id=alpha_id)
|
||||
db.add(item)
|
||||
settings = raw.get("settings") or {}
|
||||
metrics = raw.get("is") if isinstance(raw.get("is"), dict) else {}
|
||||
item.name = raw.get("name")
|
||||
item.expression = code(raw.get("regular"))
|
||||
item.selection, item.combo = code(raw.get("selection")), code(raw.get("combo"))
|
||||
item.alpha_type, item.language = raw.get("type"), settings.get("language")
|
||||
item.stage, item.status, item.hidden = raw.get("stage"), raw.get("status"), raw.get("hidden") is True
|
||||
item.region, item.universe = settings.get("region"), settings.get("universe")
|
||||
item.settings, item.is_metrics = sanitize(settings), sanitize(metrics)
|
||||
item.os_metrics = sanitize(raw.get("os")) if isinstance(raw.get("os"), dict) else {}
|
||||
item.checks = sanitize(metrics.get("checks") or raw.get("checks") or [])
|
||||
for key in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
|
||||
setattr(item, key, number(metrics.get(key)))
|
||||
item.date_created, item.date_submitted = date(raw.get("dateCreated")), date(raw.get("dateSubmitted"))
|
||||
item.synced_at, item.raw = now(), sanitize(raw)
|
||||
await db.flush()
|
||||
if await db.get(Research, alpha_id) is None:
|
||||
db.add(Research(alpha_id=alpha_id))
|
||||
return item
|
||||
|
||||
|
||||
def list_statement(filters):
|
||||
query = select(Alpha, Research).join(Research, Research.alpha_id == Alpha.id)
|
||||
q = filters.q
|
||||
if q:
|
||||
pattern = "%" + q.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
|
||||
query = query.where(
|
||||
or_(
|
||||
*(
|
||||
getattr(Alpha, f).ilike(pattern, escape="\\")
|
||||
for f in ("id", "name", "expression", "selection", "combo")
|
||||
)
|
||||
)
|
||||
)
|
||||
for name in ("region", "universe", "alpha_type", "language", "status", "stage", "hidden"):
|
||||
value = getattr(filters, name)
|
||||
if value is not None:
|
||||
query = query.where(getattr(Alpha, name) == value)
|
||||
if filters.research_state:
|
||||
query = query.where(Research.state == filters.research_state)
|
||||
if filters.favorite is not None:
|
||||
query = query.where(Research.favorite == filters.favorite)
|
||||
if filters.tag:
|
||||
query = query.where(Alpha.id.in_(select(ResearchTag.alpha_id).where(ResearchTag.tag == filters.tag)))
|
||||
if filters.created_from:
|
||||
query = query.where(Alpha.date_created >= filters.created_from)
|
||||
if filters.created_to:
|
||||
query = query.where(Alpha.date_created <= filters.created_to)
|
||||
for name in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
|
||||
for suffix, compare in (("min", "ge"), ("max", "le")):
|
||||
value = getattr(filters, f"{name}_{suffix}")
|
||||
if value is not None:
|
||||
column = getattr(Alpha, name)
|
||||
query = query.where(column >= value if compare == "ge" else column <= value)
|
||||
return query
|
||||
|
||||
|
||||
def sorted_statement(query, sort, direction):
|
||||
column = getattr(Alpha, sort)
|
||||
order = column.desc() if direction == "desc" else column.asc()
|
||||
return query.order_by(order.nullslast(), Alpha.id.asc())
|
||||
|
||||
|
||||
def summary(item: Alpha, research: Research):
|
||||
keys = (
|
||||
"id",
|
||||
"name",
|
||||
"alpha_type",
|
||||
"language",
|
||||
"stage",
|
||||
"status",
|
||||
"hidden",
|
||||
"region",
|
||||
"universe",
|
||||
"sharpe",
|
||||
"fitness",
|
||||
"returns",
|
||||
"turnover",
|
||||
"margin",
|
||||
"drawdown",
|
||||
"date_created",
|
||||
"date_submitted",
|
||||
"synced_at",
|
||||
)
|
||||
result = {k: getattr(item, k) for k in keys}
|
||||
result["expression_preview"] = (item.expression or item.selection or "")[:240]
|
||||
result["research"] = {
|
||||
k: getattr(research, k) for k in ("note", "tags", "favorite", "state", "updated_at")
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def pnl_points(raw):
|
||||
"""Use schema column names, preserving missing values rather than creating zero PnL."""
|
||||
records = raw.get("records")
|
||||
schema = raw.get("schema") or {}
|
||||
properties = schema.get("properties", []) if isinstance(schema, dict) else schema
|
||||
if isinstance(properties, dict):
|
||||
names = list(properties)
|
||||
else:
|
||||
names = [p.get("name", "") if isinstance(p, dict) else str(p) for p in properties]
|
||||
normalized = [name.lower() for name in names]
|
||||
if not isinstance(records, list):
|
||||
raise ValueError("PnL 缺少 records")
|
||||
points = []
|
||||
for row in records:
|
||||
if isinstance(row, dict):
|
||||
row = {str(k).lower(): v for k, v in row.items()}
|
||||
timestamp = next((row[k] for k in ("date", "datetime", "timestamp") if k in row), None)
|
||||
value = next((row[k] for k in ("pnl", "value") if k in row), None)
|
||||
else:
|
||||
date_i = next(
|
||||
(i for i, n in enumerate(normalized) if n in ("date", "datetime", "timestamp")), None
|
||||
)
|
||||
pnl_i = next((i for i, n in enumerate(normalized) if n in ("pnl", "value")), None)
|
||||
if date_i is None or pnl_i is None or not isinstance(row, list) or len(row) <= max(date_i, pnl_i):
|
||||
raise ValueError("PnL schema 无法识别日期或数值列")
|
||||
timestamp, value = row[date_i], row[pnl_i]
|
||||
if timestamp is not None:
|
||||
if isinstance(timestamp, (int, float)):
|
||||
from datetime import timezone
|
||||
|
||||
timestamp = datetime.fromtimestamp(
|
||||
timestamp / 1000 if timestamp > 1e11 else timestamp, tz=timezone.utc
|
||||
).isoformat()
|
||||
points.append({"date": str(timestamp), "value": number(value)})
|
||||
return sorted(points, key=lambda p: p["date"])
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Administrative commands run inside the backend container."""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import getpass
|
||||
|
||||
from sqlalchemy import delete
|
||||
|
||||
from .config import Settings
|
||||
from .db import create_database
|
||||
from .models import Admin, LoginSession
|
||||
from .security import password_hasher
|
||||
|
||||
|
||||
async def reset_password():
|
||||
password = getpass.getpass("New admin password (12+ characters): ")
|
||||
if len(password) < 12 or password != getpass.getpass("Confirm password: "):
|
||||
raise SystemExit("Password too short or confirmation does not match")
|
||||
engine, sessions = create_database(Settings().database_url)
|
||||
async with sessions() as db:
|
||||
admin = await db.get(Admin, 1)
|
||||
admin.password_hash = password_hasher.hash(password)
|
||||
await db.execute(delete(LoginSession))
|
||||
await db.commit()
|
||||
await engine.dispose()
|
||||
print("Admin password updated; all system sessions revoked.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("command", choices=["reset-password"])
|
||||
parser.parse_args()
|
||||
asyncio.run(reset_password())
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Deployment configuration; secrets are required and never included in API responses."""
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
from pydantic import Field, SecretStr, model_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(env_file="../.env", extra="ignore")
|
||||
|
||||
database_url: str = "postgresql+asyncpg://wq:wq@localhost:5432/wq"
|
||||
admin_username: str = "admin"
|
||||
admin_password: SecretStr = Field(min_length=12)
|
||||
encryption_key: SecretStr
|
||||
public_origin: str = "http://localhost:8080"
|
||||
cookie_secure: bool = False
|
||||
session_hours: int = Field(default=24, ge=1, le=168)
|
||||
wq_base_url: str = "https://api.worldquantbrain.com"
|
||||
request_timeout: float = 30
|
||||
retry_attempts: int = Field(default=4, ge=1, le=8)
|
||||
enable_runner: bool = True
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_secrets(self):
|
||||
Fernet(self.encryption_key.get_secret_value().encode())
|
||||
if self.cookie_secure and not self.public_origin.startswith("https://"):
|
||||
raise ValueError("COOKIE_SECURE requires an HTTPS PUBLIC_ORIGIN")
|
||||
return self
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Async database sessions; each request and background operation owns its transaction."""
|
||||
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
|
||||
|
||||
def create_database(url: str):
|
||||
engine = create_async_engine(url, pool_pre_ping=True, hide_parameters=True)
|
||||
return engine, async_sessionmaker(engine, expire_on_commit=False)
|
||||
@@ -0,0 +1,390 @@
|
||||
"""Durable, single-process task runner.
|
||||
|
||||
Each page and its checkpoint commit in one transaction. A single Uvicorn process
|
||||
owns the runner and shared upstream session; do not scale replicas before adding
|
||||
a distributed lease and session coordinator.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from contextlib import suppress
|
||||
from datetime import timedelta
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from .alphas import pnl_points, sanitize, upsert_alpha
|
||||
from .models import Account, Alpha, Job, JobItem, Pnl, now
|
||||
from .security import cipher
|
||||
from .worldquant import VerificationRequired, WqClient, WqError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
ACTIVE = ("queued", "running", "waiting_auth", "waiting_connection")
|
||||
AUTH_KINDS = ("connect", "verify", "profile")
|
||||
|
||||
|
||||
async def create_job(db, kind, payload=None):
|
||||
job = Job(id=str(uuid4()), kind=kind, payload=payload or {}, checkpoint={})
|
||||
db.add(job)
|
||||
await db.commit()
|
||||
return job
|
||||
|
||||
|
||||
class Runner:
|
||||
def __init__(self, sessions, settings, client=None):
|
||||
self.sessions, self.settings = sessions, settings
|
||||
self.client = client or WqClient(settings)
|
||||
self.loop_task = None
|
||||
self.active_task = None
|
||||
self.active_id = None
|
||||
self.stopping = False
|
||||
self.disconnecting = False
|
||||
self.control_lock = asyncio.Lock()
|
||||
self.recover_database = False
|
||||
self.wake = asyncio.Event()
|
||||
|
||||
async def start(self):
|
||||
async with self.sessions() as db:
|
||||
await db.execute(update(Job).where(Job.status == "running").values(status="queued"))
|
||||
account = await db.get(Account, 1)
|
||||
if account.connection_status in ("connected", "connecting", "verification_required"):
|
||||
account.connection_status = "expired"
|
||||
account.verification_url = None
|
||||
await db.commit()
|
||||
self.loop_task = asyncio.create_task(self.run_loop())
|
||||
|
||||
async def stop(self):
|
||||
self.stopping = True
|
||||
self.wake.set()
|
||||
if self.active_task:
|
||||
self.active_task.cancel()
|
||||
if self.loop_task:
|
||||
await self.loop_task
|
||||
await self.client.close()
|
||||
|
||||
async def cancel(self, job_id):
|
||||
if self.active_id == job_id and self.active_task:
|
||||
self.active_task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await self.active_task
|
||||
|
||||
async def disconnect(self):
|
||||
# Prevent the scheduler from starting another request while credentials are cleared.
|
||||
async with self.control_lock:
|
||||
self.disconnecting = True
|
||||
try:
|
||||
if self.active_task:
|
||||
await self.cancel(self.active_id)
|
||||
self.client.disconnect()
|
||||
async with self.sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
account.connection_status = "disconnected"
|
||||
account.verification_url = None
|
||||
account.connection_error = None
|
||||
await db.execute(
|
||||
update(Job).where(Job.status.in_(ACTIVE)).values(status="waiting_connection")
|
||||
)
|
||||
await db.commit()
|
||||
finally:
|
||||
self.disconnecting = False
|
||||
|
||||
async def run_loop(self):
|
||||
while not self.stopping:
|
||||
try:
|
||||
await self.run_next()
|
||||
except (SQLAlchemyError, OSError):
|
||||
# A temporary DB outage must not silently kill the scheduler.
|
||||
self.recover_database = True
|
||||
logger.warning("Task runner is waiting for database recovery")
|
||||
await self.wait_for_work()
|
||||
|
||||
async def wait_for_work(self):
|
||||
self.wake.clear()
|
||||
if not self.stopping:
|
||||
try:
|
||||
await asyncio.wait_for(self.wake.wait(), timeout=1)
|
||||
except TimeoutError:
|
||||
pass
|
||||
|
||||
async def run_next(self):
|
||||
async with self.control_lock:
|
||||
async with self.sessions() as db:
|
||||
if self.recover_database:
|
||||
await db.execute(update(Job).where(Job.status == "running").values(status="queued"))
|
||||
await db.commit()
|
||||
self.recover_database = False
|
||||
job = (
|
||||
await db.scalars(
|
||||
select(Job)
|
||||
.where(Job.status == "queued")
|
||||
.order_by(Job.kind.in_(AUTH_KINDS).desc(), Job.created_at)
|
||||
.limit(1)
|
||||
)
|
||||
).first()
|
||||
job_id = job.id if job else None
|
||||
if job_id and not self.stopping:
|
||||
self.active_id = job_id
|
||||
self.active_task = asyncio.create_task(self.execute(job_id))
|
||||
if job_id and self.active_task:
|
||||
try:
|
||||
with suppress(asyncio.CancelledError):
|
||||
await self.active_task
|
||||
finally:
|
||||
self.active_id, self.active_task = None, None
|
||||
else:
|
||||
await self.wait_for_work()
|
||||
|
||||
async def set_account(self, status, error=None, url=None):
|
||||
async with self.sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
account.connection_status, account.connection_error, account.verification_url = status, error, url
|
||||
await db.commit()
|
||||
|
||||
async def ensure_connected(self, force=False):
|
||||
async with self.sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
if (
|
||||
not account.email
|
||||
or not account.password_encrypted
|
||||
or (account.connection_status == "disconnected" and not force)
|
||||
):
|
||||
raise WqError("请先配置并连接 WorldQuant", "disconnected")
|
||||
password = cipher(self.settings).decrypt(account.password_encrypted.encode()).decode()
|
||||
email = account.email
|
||||
if self.client.verification_url and not force:
|
||||
raise VerificationRequired(self.client.verification_url)
|
||||
await self.client.authenticate(email, password, force=force)
|
||||
await self.set_account("connected")
|
||||
|
||||
async def refresh_profile(self):
|
||||
raw = await self.client.profile()
|
||||
async with self.sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
user_id = raw.get("id")
|
||||
if not user_id or (account.wq_user_id and account.wq_user_id != str(user_id)):
|
||||
self.client.disconnect()
|
||||
raise WqError("平台账户身份不匹配,请核对凭据", "identity_mismatch")
|
||||
account.wq_user_id = str(user_id)
|
||||
# An allowlist avoids persisting unknown personal/security fields.
|
||||
account.profile = sanitize(
|
||||
{
|
||||
k: raw[k]
|
||||
for k in (
|
||||
"id",
|
||||
"username",
|
||||
"name",
|
||||
"email",
|
||||
"firstName",
|
||||
"lastName",
|
||||
"role",
|
||||
"roles",
|
||||
"permissions",
|
||||
"type",
|
||||
)
|
||||
if k in raw
|
||||
}
|
||||
)
|
||||
account.connection_status, account.connection_error = "connected", None
|
||||
account.verification_url, account.last_synced_at = None, now()
|
||||
await db.execute(
|
||||
update(Job)
|
||||
.where(Job.status.in_(("waiting_auth", "waiting_connection")), ~Job.kind.in_(AUTH_KINDS))
|
||||
.values(status="queued", error=None)
|
||||
)
|
||||
# Superseded connect/verify attempts are complete after identity confirmation.
|
||||
await db.execute(
|
||||
update(Job)
|
||||
.where(Job.status.in_(("waiting_auth", "waiting_connection")), Job.kind.in_(AUTH_KINDS))
|
||||
.values(status="completed", error=None)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
async def checkpoint(self, job_id, values):
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job.cancel_requested:
|
||||
raise asyncio.CancelledError()
|
||||
for key, value in values.items():
|
||||
setattr(job, key, value)
|
||||
job.updated_at = now()
|
||||
await db.commit()
|
||||
|
||||
async def execute(self, job_id):
|
||||
async def retrying(delay):
|
||||
await self.checkpoint(job_id, {"next_retry_at": now() + timedelta(seconds=delay)})
|
||||
|
||||
self.client.on_retry = retrying
|
||||
try:
|
||||
await self.checkpoint(job_id, {"status": "running", "error": None, "next_retry_at": None})
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
kind, payload = job.kind, job.payload
|
||||
if kind == "verify":
|
||||
if not self.client.verification_url:
|
||||
await self.ensure_connected(force=True)
|
||||
else:
|
||||
await self.client.verify()
|
||||
await self.refresh_profile()
|
||||
else:
|
||||
await self.ensure_connected(force=kind == "connect")
|
||||
if kind in ("connect", "profile"):
|
||||
await self.refresh_profile()
|
||||
elif kind == "full_sync":
|
||||
await self.sync_all(job_id)
|
||||
else:
|
||||
await self.sync_ids(job_id, kind, payload["alpha_ids"])
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.status = "completed_with_errors" if job.failed else "completed"
|
||||
job.next_retry_at, job.updated_at = None, now()
|
||||
await db.commit()
|
||||
except asyncio.CancelledError:
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.status = (
|
||||
"cancelled"
|
||||
if job.cancel_requested
|
||||
else "waiting_connection"
|
||||
if self.disconnecting
|
||||
else "queued"
|
||||
if self.stopping
|
||||
else "cancelled"
|
||||
)
|
||||
job.next_retry_at, job.updated_at = None, now()
|
||||
await db.commit()
|
||||
except VerificationRequired as exc:
|
||||
await self.set_account("verification_required", str(exc), exc.url)
|
||||
await self.checkpoint(
|
||||
job_id, {"status": "waiting_auth", "error": str(exc), "next_retry_at": None}
|
||||
)
|
||||
except WqError as exc:
|
||||
waiting = exc.code in ("disconnected", "authentication_failed", "identity_mismatch")
|
||||
if waiting:
|
||||
await self.set_account("disconnected" if exc.code == "disconnected" else "error", str(exc))
|
||||
await self.checkpoint(
|
||||
job_id,
|
||||
{
|
||||
"status": "waiting_connection" if waiting else "failed",
|
||||
"error": str(exc),
|
||||
"next_retry_at": None,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
# Do not expose upstream bodies, decrypted credentials, or SQL bind parameters.
|
||||
logger.error("Job %s failed with an internal error", job_id)
|
||||
await self.checkpoint(
|
||||
job_id,
|
||||
{"status": "failed", "error": "任务内部错误,请检查服务日志后重试", "next_retry_at": None},
|
||||
)
|
||||
finally:
|
||||
self.client.on_retry = None
|
||||
|
||||
async def sync_all(self, job_id):
|
||||
partitions = [
|
||||
("UNSUBMITTED", False),
|
||||
("UNSUBMITTED", True),
|
||||
("SUBMITTED", False),
|
||||
("SUBMITTED", True),
|
||||
]
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
checkpoint = job.checkpoint
|
||||
before = job.created_at.isoformat()
|
||||
start_partition, offset = checkpoint.get("partition", 0), checkpoint.get("offset", 0)
|
||||
for partition in range(start_partition, len(partitions)):
|
||||
submission, hidden = partitions[partition]
|
||||
while True:
|
||||
await self.checkpoint(job_id, {"next_retry_at": None})
|
||||
raw = await self.client.alphas(submission, hidden, offset, before)
|
||||
rows = raw.get("results")
|
||||
if not isinstance(rows, list):
|
||||
raise WqError("Alpha 列表缺少 results,已保留当前进度", "invalid_response")
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job.cancel_requested:
|
||||
raise asyncio.CancelledError()
|
||||
for raw_alpha in rows:
|
||||
await upsert_alpha(db, raw_alpha)
|
||||
if not await db.get(JobItem, (job_id, raw_alpha["id"])):
|
||||
db.add(JobItem(job_id=job_id, alpha_id=raw_alpha["id"]))
|
||||
job.processed += 1
|
||||
# Trust explicit next/count when available; otherwise an empty page ends iteration.
|
||||
count = raw.get("count")
|
||||
more = bool(rows)
|
||||
if "next" in raw:
|
||||
more = bool(raw["next"])
|
||||
elif isinstance(count, int):
|
||||
more = offset + len(rows) < count and bool(rows)
|
||||
if more and not rows:
|
||||
raise WqError("平台分页未前进", "invalid_response")
|
||||
offset += len(rows)
|
||||
job.checkpoint = {
|
||||
"partition": partition if more else partition + 1,
|
||||
"offset": offset if more else 0,
|
||||
}
|
||||
job.updated_at = now()
|
||||
await db.commit()
|
||||
if not more:
|
||||
break
|
||||
offset = 0
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.total = job.processed
|
||||
account = await db.get(Account, 1)
|
||||
account.last_synced_at = now()
|
||||
await db.commit()
|
||||
|
||||
async def sync_ids(self, job_id, kind, alpha_ids):
|
||||
await self.checkpoint(job_id, {"total": len(alpha_ids)})
|
||||
for alpha_id in alpha_ids:
|
||||
async with self.sessions() as db:
|
||||
previous = await db.get(JobItem, (job_id, alpha_id))
|
||||
if previous and not previous.error:
|
||||
continue
|
||||
await self.checkpoint(job_id, {"next_retry_at": None})
|
||||
error = None
|
||||
try:
|
||||
raw = await (
|
||||
self.client.pnl(alpha_id) if kind == "pnl_refresh" else self.client.alpha(alpha_id)
|
||||
)
|
||||
if kind != "pnl_refresh" and raw.get("id") != alpha_id:
|
||||
raise ValueError("平台返回的 Alpha ID 与请求不一致")
|
||||
points = pnl_points(raw) if kind == "pnl_refresh" else None
|
||||
except VerificationRequired:
|
||||
raise
|
||||
except WqError as exc:
|
||||
if exc.code not in ("not_found", "access_denied"):
|
||||
raise
|
||||
error = str(exc)
|
||||
except ValueError as exc:
|
||||
error = str(exc)
|
||||
async with self.sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job.cancel_requested:
|
||||
raise asyncio.CancelledError()
|
||||
previous = await db.get(JobItem, (job_id, alpha_id))
|
||||
if previous and previous.error:
|
||||
job.failed -= 1
|
||||
if not previous:
|
||||
previous = JobItem(job_id=job_id, alpha_id=alpha_id)
|
||||
db.add(previous)
|
||||
if not error:
|
||||
if kind == "pnl_refresh":
|
||||
if not await db.get(Alpha, alpha_id):
|
||||
error = "请先导入此 Alpha"
|
||||
else:
|
||||
pnl = await db.get(Pnl, alpha_id)
|
||||
if pnl is None:
|
||||
pnl = Pnl(alpha_id=alpha_id)
|
||||
db.add(pnl)
|
||||
pnl.raw, pnl.points, pnl.fetched_at = sanitize(raw), points, now()
|
||||
else:
|
||||
await upsert_alpha(db, raw)
|
||||
previous.error = error
|
||||
if error:
|
||||
job.failed += 1
|
||||
else:
|
||||
job.processed += 1
|
||||
job.updated_at = now()
|
||||
await db.commit()
|
||||
@@ -0,0 +1,498 @@
|
||||
"""FastAPI application and authenticated public API."""
|
||||
|
||||
import asyncio
|
||||
import csv
|
||||
import io
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, FastAPI, HTTPException, Query, Request, Response
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from sqlalchemy import delete, func, select, text
|
||||
|
||||
from .alphas import list_statement, sorted_statement, summary
|
||||
from .config import Settings
|
||||
from .db import create_database
|
||||
from .jobs import ACTIVE, AUTH_KINDS, Runner, create_job
|
||||
from .models import Account, Admin, Alpha, Job, JobItem, LoginSession, Pnl, Research, ResearchTag, now
|
||||
from .schemas import (
|
||||
AccountOutput,
|
||||
AlphaDetail,
|
||||
AlphaFilters,
|
||||
AlphaPage,
|
||||
BulkInput,
|
||||
BulkOutput,
|
||||
CredentialsInput,
|
||||
ErrorOutput,
|
||||
FacetsOutput,
|
||||
HealthOutput,
|
||||
JobErrorOutput,
|
||||
JobInput,
|
||||
JobOutput,
|
||||
LoginInput,
|
||||
OkOutput,
|
||||
PnlOutput,
|
||||
PreferencesInput,
|
||||
ResearchInput,
|
||||
SessionOutput,
|
||||
normalize_tags,
|
||||
)
|
||||
from .security import bootstrap, cipher, issue_session, require_auth, token_hash, valid_password
|
||||
|
||||
|
||||
def account_output(account):
|
||||
keys = (
|
||||
"email",
|
||||
"wq_user_id",
|
||||
"profile",
|
||||
"connection_status",
|
||||
"connection_error",
|
||||
"verification_url",
|
||||
"last_synced_at",
|
||||
"display_name",
|
||||
"theme",
|
||||
"timezone",
|
||||
"page_size",
|
||||
)
|
||||
return {**{k: getattr(account, k) for k in keys}, "configured": bool(account.password_encrypted)}
|
||||
|
||||
|
||||
async def set_tags(db, record, tags):
|
||||
record.tags = normalize_tags(tags)
|
||||
await db.execute(delete(ResearchTag).where(ResearchTag.alpha_id == record.alpha_id))
|
||||
db.add_all(ResearchTag(alpha_id=record.alpha_id, tag=t) for t in record.tags)
|
||||
|
||||
|
||||
def csv_cell(value):
|
||||
"""Neutralize spreadsheet formulas in untrusted names, expressions, notes and tags."""
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, str) and (
|
||||
value.lstrip().startswith(("=", "+", "-", "@")) or value.startswith(("\t", "\r", "\n"))
|
||||
):
|
||||
return "'" + value
|
||||
return value
|
||||
|
||||
|
||||
def create_app(settings=None, wq_client=None):
|
||||
settings = settings or Settings()
|
||||
engine, sessions = create_database(settings.database_url)
|
||||
runner = Runner(sessions, settings, client=wq_client)
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app):
|
||||
async with sessions() as db:
|
||||
await bootstrap(db, settings)
|
||||
if settings.enable_runner:
|
||||
await runner.start()
|
||||
yield
|
||||
if settings.enable_runner:
|
||||
await runner.stop()
|
||||
else:
|
||||
await runner.client.close()
|
||||
await engine.dispose()
|
||||
|
||||
app = FastAPI(
|
||||
title="WorldQuant Alpha Research API",
|
||||
version="0.1.0",
|
||||
lifespan=lifespan,
|
||||
responses={code: {"model": ErrorOutput} for code in (401, 403, 404, 409, 422, 429)},
|
||||
)
|
||||
app.state.engine, app.state.sessions, app.state.runner = engine, sessions, runner
|
||||
app.state.settings = settings
|
||||
login_failures = defaultdict(list)
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_error(request, exc):
|
||||
# Pydantic's default error includes the submitted value, possibly a password.
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
content={
|
||||
"detail": "; ".join(f"{'.'.join(str(v) for v in e['loc'])}: {e['msg']}" for e in exc.errors())
|
||||
},
|
||||
)
|
||||
|
||||
@app.middleware("http")
|
||||
async def browser_security(request, call_next):
|
||||
if request.method not in ("GET", "HEAD", "OPTIONS"):
|
||||
if request.headers.get("X-WQ-Request") != "1":
|
||||
return JSONResponse({"detail": "缺少请求校验头"}, status_code=403)
|
||||
origin = request.headers.get("Origin")
|
||||
if origin and origin.rstrip("/") != settings.public_origin.rstrip("/"):
|
||||
return JSONResponse({"detail": "请求来源不被允许"}, status_code=403)
|
||||
response = await call_next(request)
|
||||
response.headers["X-Content-Type-Options"] = "nosniff"
|
||||
response.headers["Referrer-Policy"] = "no-referrer"
|
||||
response.headers["Cache-Control"] = "no-store"
|
||||
return response
|
||||
|
||||
@app.get("/api/v1/health", response_model=HealthOutput, tags=["health"])
|
||||
async def health():
|
||||
async with sessions() as db:
|
||||
await db.execute(text("SELECT 1"))
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.post("/api/v1/auth/login", response_model=SessionOutput, tags=["auth"])
|
||||
async def login(body: LoginInput, request: Request, response: Response):
|
||||
key = request.client.host if request.client else "unknown"
|
||||
current = time.monotonic()
|
||||
# Bound the rate-limit map and remove expired entries.
|
||||
for peer in list(login_failures):
|
||||
login_failures[peer] = [t for t in login_failures[peer] if current - t < 60]
|
||||
if not login_failures[peer]:
|
||||
del login_failures[peer]
|
||||
if len(login_failures[key]) >= 5:
|
||||
raise HTTPException(429, "登录尝试过多,请一分钟后重试")
|
||||
async with sessions() as db:
|
||||
admin = await db.get(Admin, 1)
|
||||
matches = await asyncio.to_thread(valid_password, admin.password_hash, body.password)
|
||||
if body.username != admin.username or not matches:
|
||||
login_failures[key].append(current)
|
||||
raise HTTPException(401, "用户名或密码不正确")
|
||||
login_failures.pop(key, None)
|
||||
token = await issue_session(db, settings.session_hours)
|
||||
response.set_cookie(
|
||||
"wq_session",
|
||||
token,
|
||||
httponly=True,
|
||||
secure=settings.cookie_secure,
|
||||
samesite="strict",
|
||||
max_age=settings.session_hours * 3600,
|
||||
path="/",
|
||||
)
|
||||
return {"username": admin.username}
|
||||
|
||||
api = APIRouter(prefix="/api/v1", dependencies=[Depends(require_auth)])
|
||||
|
||||
@api.get("/auth/me", response_model=SessionOutput, tags=["auth"])
|
||||
async def me():
|
||||
async with sessions() as db:
|
||||
return {"username": (await db.get(Admin, 1)).username}
|
||||
|
||||
@api.post("/auth/logout", response_model=OkOutput, tags=["auth"])
|
||||
async def logout(request: Request, response: Response):
|
||||
async with sessions() as db:
|
||||
await db.execute(
|
||||
delete(LoginSession).where(
|
||||
LoginSession.token_hash == token_hash(request.cookies["wq_session"])
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
response.delete_cookie(
|
||||
"wq_session", path="/", httponly=True, secure=settings.cookie_secure, samesite="strict"
|
||||
)
|
||||
return {"ok": True}
|
||||
|
||||
@api.get("/account", response_model=AccountOutput, tags=["account"])
|
||||
async def get_account():
|
||||
async with sessions() as db:
|
||||
return account_output(await db.get(Account, 1))
|
||||
|
||||
@api.put("/account/credentials", response_model=AccountOutput, tags=["account"])
|
||||
async def credentials(body: CredentialsInput):
|
||||
async with sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
if account.wq_user_id and account.email.casefold() != body.email.casefold():
|
||||
raise HTTPException(409, "本系统已绑定一个账户;不允许混入其他账户的 Alpha")
|
||||
await runner.disconnect()
|
||||
async with sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
account.email = body.email
|
||||
account.password_encrypted = cipher(settings).encrypt(body.password.encode()).decode()
|
||||
await db.commit()
|
||||
return account_output(account)
|
||||
|
||||
@api.patch("/account/preferences", response_model=AccountOutput, tags=["account"])
|
||||
async def preferences(body: PreferencesInput):
|
||||
async with sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
for key, value in body.model_dump().items():
|
||||
setattr(account, key, value)
|
||||
await db.commit()
|
||||
return account_output(account)
|
||||
|
||||
async def account_job(kind):
|
||||
async with sessions() as db:
|
||||
account = await db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||||
if not account.password_encrypted:
|
||||
raise HTTPException(409, "请先保存 WorldQuant 账户配置")
|
||||
existing = (
|
||||
await db.scalars(
|
||||
select(Job).where(Job.kind.in_(AUTH_KINDS), Job.status.in_(("running", "queued")))
|
||||
)
|
||||
).first()
|
||||
if existing:
|
||||
return existing
|
||||
if kind == "connect":
|
||||
account.connection_status = "connecting"
|
||||
account.connection_error = None
|
||||
job = await create_job(db, kind)
|
||||
runner.wake.set()
|
||||
return job
|
||||
|
||||
@api.post("/account/connect", status_code=202, response_model=JobOutput, tags=["account"])
|
||||
async def connect():
|
||||
return await account_job("connect")
|
||||
|
||||
@api.post("/account/verify", status_code=202, response_model=JobOutput, tags=["account"])
|
||||
async def verify():
|
||||
return await account_job("verify")
|
||||
|
||||
@api.post("/account/refresh", status_code=202, response_model=JobOutput, tags=["account"])
|
||||
async def refresh():
|
||||
return await account_job("profile")
|
||||
|
||||
@api.post("/account/disconnect", response_model=OkOutput, tags=["account"])
|
||||
async def disconnect():
|
||||
await runner.disconnect()
|
||||
return {"ok": True}
|
||||
|
||||
@api.get("/alphas", response_model=AlphaPage, tags=["alphas"])
|
||||
async def get_alphas(filters: Annotated[AlphaFilters, Query()]):
|
||||
query = list_statement(filters)
|
||||
async with sessions() as db:
|
||||
total = await db.scalar(select(func.count()).select_from(query.subquery()))
|
||||
rows = (
|
||||
await db.execute(
|
||||
sorted_statement(query, filters.sort, filters.direction)
|
||||
.limit(filters.limit)
|
||||
.offset(filters.offset)
|
||||
)
|
||||
).all()
|
||||
return {
|
||||
"items": [summary(a, r) for a, r in rows],
|
||||
"total": total,
|
||||
"limit": filters.limit,
|
||||
"offset": filters.offset,
|
||||
}
|
||||
|
||||
@api.get("/alphas/facets", response_model=FacetsOutput, tags=["alphas"])
|
||||
async def facets():
|
||||
async with sessions() as db:
|
||||
result = {}
|
||||
for key in ("region", "universe", "alpha_type", "language", "status", "stage"):
|
||||
column = getattr(Alpha, key)
|
||||
result[key] = list(
|
||||
(
|
||||
await db.scalars(
|
||||
select(column).where(column.is_not(None)).distinct().order_by(column)
|
||||
)
|
||||
).all()
|
||||
)
|
||||
result["tags"] = list(
|
||||
(await db.scalars(select(ResearchTag.tag).distinct().order_by(ResearchTag.tag))).all()
|
||||
)
|
||||
result["total"] = await db.scalar(select(func.count()).select_from(Alpha))
|
||||
result["favorites"] = await db.scalar(
|
||||
select(func.count()).select_from(Research).where(Research.favorite.is_(True))
|
||||
)
|
||||
result["last_sync"] = await db.scalar(select(func.max(Alpha.synced_at)))
|
||||
return result
|
||||
|
||||
@api.get(
|
||||
"/alphas/export",
|
||||
response_class=StreamingResponse,
|
||||
responses={200: {"content": {"text/csv": {"schema": {"type": "string"}}}}},
|
||||
tags=["alphas"],
|
||||
)
|
||||
async def export(filters: Annotated[AlphaFilters, Query()]):
|
||||
query = sorted_statement(list_statement(filters), filters.sort, filters.direction)
|
||||
|
||||
async def content():
|
||||
yield "\ufeff"
|
||||
buffer = io.StringIO()
|
||||
writer = csv.writer(buffer)
|
||||
columns = [
|
||||
"id",
|
||||
"name",
|
||||
"expression",
|
||||
"selection",
|
||||
"combo",
|
||||
"alpha_type",
|
||||
"language",
|
||||
"stage",
|
||||
"status",
|
||||
"hidden",
|
||||
"region",
|
||||
"universe",
|
||||
"sharpe",
|
||||
"fitness",
|
||||
"returns",
|
||||
"turnover",
|
||||
"margin",
|
||||
"drawdown",
|
||||
"date_created",
|
||||
"date_submitted",
|
||||
"synced_at",
|
||||
]
|
||||
writer.writerow(columns + ["research_state", "favorite", "tags", "note"])
|
||||
yield buffer.getvalue()
|
||||
buffer.seek(0)
|
||||
buffer.truncate(0)
|
||||
async with sessions() as db:
|
||||
rows = await db.stream(query.execution_options(yield_per=200))
|
||||
async for a, r in rows:
|
||||
writer.writerow(
|
||||
[csv_cell(getattr(a, key)) for key in columns]
|
||||
+ [r.state, r.favorite, csv_cell(";".join(r.tags)), csv_cell(r.note)]
|
||||
)
|
||||
yield buffer.getvalue()
|
||||
buffer.seek(0)
|
||||
buffer.truncate(0)
|
||||
|
||||
return StreamingResponse(
|
||||
content(),
|
||||
media_type="text/csv; charset=utf-8",
|
||||
headers={"Content-Disposition": 'attachment; filename="alphas.csv"'},
|
||||
)
|
||||
|
||||
@api.patch("/alphas/research/bulk", response_model=BulkOutput, tags=["alphas"])
|
||||
async def bulk(body: BulkInput):
|
||||
async with sessions() as db:
|
||||
rows = (
|
||||
await db.scalars(
|
||||
select(Research).where(Research.alpha_id.in_(body.alpha_ids)).with_for_update()
|
||||
)
|
||||
).all()
|
||||
if len(rows) != len(body.alpha_ids):
|
||||
raise HTTPException(404, "部分 Alpha 尚未同步,本次未修改任何记录")
|
||||
for row in rows:
|
||||
tags = (set(row.tags) | set(body.add_tags)) - set(body.remove_tags)
|
||||
try:
|
||||
await set_tags(db, row, list(tags))
|
||||
except ValueError as exc:
|
||||
raise HTTPException(422, str(exc)) from None
|
||||
if body.state:
|
||||
row.state = body.state
|
||||
row.updated_at = now()
|
||||
await db.commit()
|
||||
return {"updated": len(rows)}
|
||||
|
||||
@api.get("/alphas/{alpha_id}", response_model=AlphaDetail, tags=["alphas"])
|
||||
async def detail(alpha_id: str):
|
||||
async with sessions() as db:
|
||||
a, r = await db.get(Alpha, alpha_id), await db.get(Research, alpha_id)
|
||||
if a is None:
|
||||
raise HTTPException(404, "Alpha 尚未同步")
|
||||
return {
|
||||
**summary(a, r),
|
||||
**{
|
||||
key: getattr(a, key)
|
||||
for key in (
|
||||
"expression",
|
||||
"selection",
|
||||
"combo",
|
||||
"settings",
|
||||
"is_metrics",
|
||||
"os_metrics",
|
||||
"checks",
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
@api.patch("/alphas/{alpha_id}/research", response_model=OkOutput, tags=["alphas"])
|
||||
async def research(alpha_id: str, body: ResearchInput):
|
||||
async with sessions() as db:
|
||||
row = await db.scalar(select(Research).where(Research.alpha_id == alpha_id).with_for_update())
|
||||
if row is None:
|
||||
raise HTTPException(404, "Alpha 尚未同步")
|
||||
for key, value in body.model_dump(exclude_unset=True).items():
|
||||
if key == "tags":
|
||||
await set_tags(db, row, value)
|
||||
else:
|
||||
setattr(row, key, value)
|
||||
row.updated_at = now()
|
||||
await db.commit()
|
||||
return {"ok": True}
|
||||
|
||||
@api.get("/alphas/{alpha_id}/pnl", response_model=PnlOutput, tags=["alphas"])
|
||||
async def pnl(alpha_id: str):
|
||||
async with sessions() as db:
|
||||
if not await db.get(Alpha, alpha_id):
|
||||
raise HTTPException(404, "Alpha 尚未同步")
|
||||
record = await db.get(Pnl, alpha_id)
|
||||
return {
|
||||
"cached": record is not None,
|
||||
"points": record.points if record else [],
|
||||
"fetched_at": record.fetched_at if record else None,
|
||||
}
|
||||
|
||||
@api.post("/sync-jobs", status_code=202, response_model=JobOutput, tags=["sync-jobs"])
|
||||
async def new_job(body: JobInput):
|
||||
async with sessions() as db:
|
||||
# Serialize creation against the singleton account, avoiding duplicate full scans.
|
||||
account = await db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||||
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
|
||||
raise HTTPException(409, "请先连接 WorldQuant")
|
||||
payload = {"alpha_ids": body.alpha_ids}
|
||||
existing = (
|
||||
await db.scalars(select(Job).where(Job.kind == body.kind, Job.status.in_(ACTIVE)))
|
||||
).all()
|
||||
for job in existing:
|
||||
if job.payload == payload:
|
||||
return job
|
||||
job = await create_job(db, body.kind, payload)
|
||||
runner.wake.set()
|
||||
return job
|
||||
|
||||
@api.get("/sync-jobs", response_model=list[JobOutput], tags=["sync-jobs"])
|
||||
async def get_jobs():
|
||||
async with sessions() as db:
|
||||
return (await db.scalars(select(Job).order_by(Job.created_at.desc()).limit(100))).all()
|
||||
|
||||
@api.get("/sync-jobs/{job_id}", response_model=JobOutput, tags=["sync-jobs"])
|
||||
async def get_job(job_id: str):
|
||||
async with sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job is None:
|
||||
raise HTTPException(404, "任务不存在")
|
||||
return job
|
||||
|
||||
@api.get("/sync-jobs/{job_id}/errors", response_model=list[JobErrorOutput], tags=["sync-jobs"])
|
||||
async def job_errors(job_id: str):
|
||||
async with sessions() as db:
|
||||
rows = (
|
||||
await db.scalars(select(JobItem).where(JobItem.job_id == job_id, JobItem.error.is_not(None)))
|
||||
).all()
|
||||
return [{"alpha_id": row.alpha_id, "error": row.error} for row in rows]
|
||||
|
||||
@api.post("/sync-jobs/{job_id}/cancel", response_model=OkOutput, tags=["sync-jobs"])
|
||||
async def cancel_job(job_id: str):
|
||||
async with sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job is None:
|
||||
raise HTTPException(404, "任务不存在")
|
||||
if job.status not in ACTIVE:
|
||||
return {"ok": True}
|
||||
job.cancel_requested = True
|
||||
if job.status != "running":
|
||||
job.status = "cancelled"
|
||||
await db.commit()
|
||||
await runner.cancel(job_id)
|
||||
return {"ok": True}
|
||||
|
||||
@api.post("/sync-jobs/{job_id}/retry", response_model=JobOutput, tags=["sync-jobs"])
|
||||
async def retry_job(job_id: str):
|
||||
async with sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
if job is None:
|
||||
raise HTTPException(404, "任务不存在")
|
||||
if job.status not in (
|
||||
"failed",
|
||||
"cancelled",
|
||||
"completed_with_errors",
|
||||
"waiting_connection",
|
||||
"waiting_auth",
|
||||
):
|
||||
raise HTTPException(409, "该任务当前不需要重试")
|
||||
job.status, job.error, job.cancel_requested, job.next_retry_at = "queued", None, False, None
|
||||
job.updated_at = now()
|
||||
await db.commit()
|
||||
runner.wake.set()
|
||||
return job
|
||||
|
||||
app.include_router(api)
|
||||
return app
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Platform snapshots and local research are intentionally stored separately."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import JSON, Boolean, DateTime, Float, ForeignKey, Index, Integer, String, Text
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
|
||||
def now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
class Admin(Base):
|
||||
__tablename__ = "admins"
|
||||
id: Mapped[int] = mapped_column(primary_key=True, default=1)
|
||||
username: Mapped[str] = mapped_column(String(100), unique=True)
|
||||
password_hash: Mapped[str] = mapped_column(Text)
|
||||
|
||||
|
||||
class LoginSession(Base):
|
||||
__tablename__ = "login_sessions"
|
||||
token_hash: Mapped[str] = mapped_column(String(64), primary_key=True)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True)
|
||||
|
||||
|
||||
class Account(Base):
|
||||
__tablename__ = "accounts"
|
||||
id: Mapped[int] = mapped_column(primary_key=True, default=1)
|
||||
email: Mapped[str | None] = mapped_column(String(254))
|
||||
password_encrypted: Mapped[str | None] = mapped_column(Text)
|
||||
wq_user_id: Mapped[str | None] = mapped_column(String(100))
|
||||
profile: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
connection_status: Mapped[str] = mapped_column(String(30), default="disconnected")
|
||||
connection_error: Mapped[str | None] = mapped_column(Text)
|
||||
verification_url: Mapped[str | None] = mapped_column(Text)
|
||||
last_synced_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
display_name: Mapped[str] = mapped_column(String(100), default="研究员")
|
||||
theme: Mapped[str] = mapped_column(String(10), default="light")
|
||||
timezone: Mapped[str] = mapped_column(String(64), default="Asia/Shanghai")
|
||||
page_size: Mapped[int] = mapped_column(Integer, default=25)
|
||||
|
||||
|
||||
class Alpha(Base):
|
||||
__tablename__ = "alphas"
|
||||
id: Mapped[str] = mapped_column(String(100), primary_key=True)
|
||||
name: Mapped[str | None] = mapped_column(Text)
|
||||
expression: Mapped[str | None] = mapped_column(Text)
|
||||
selection: Mapped[str | None] = mapped_column(Text)
|
||||
combo: Mapped[str | None] = mapped_column(Text)
|
||||
alpha_type: Mapped[str | None] = mapped_column(String(50), index=True)
|
||||
language: Mapped[str | None] = mapped_column(String(50))
|
||||
stage: Mapped[str | None] = mapped_column(String(50), index=True)
|
||||
status: Mapped[str | None] = mapped_column(String(80), index=True)
|
||||
hidden: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
region: Mapped[str | None] = mapped_column(String(50), index=True)
|
||||
universe: Mapped[str | None] = mapped_column(String(100))
|
||||
settings: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
is_metrics: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
os_metrics: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
checks: Mapped[list] = mapped_column(JSON, default=list)
|
||||
sharpe: Mapped[float | None] = mapped_column(Float, index=True)
|
||||
fitness: Mapped[float | None] = mapped_column(Float, index=True)
|
||||
returns: Mapped[float | None] = mapped_column(Float)
|
||||
turnover: Mapped[float | None] = mapped_column(Float)
|
||||
margin: Mapped[float | None] = mapped_column(Float)
|
||||
drawdown: Mapped[float | None] = mapped_column(Float)
|
||||
date_created: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), index=True)
|
||||
date_submitted: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
synced_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
raw: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
__table_args__ = (Index("ix_alphas_region_status", "region", "status"),)
|
||||
|
||||
|
||||
class Research(Base):
|
||||
__tablename__ = "research"
|
||||
alpha_id: Mapped[str] = mapped_column(ForeignKey("alphas.id"), primary_key=True)
|
||||
note: Mapped[str] = mapped_column(Text, default="")
|
||||
tags: Mapped[list] = mapped_column(JSON, default=list)
|
||||
favorite: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
state: Mapped[str] = mapped_column(String(30), default="inbox", index=True)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class Pnl(Base):
|
||||
__tablename__ = "pnl_cache"
|
||||
alpha_id: Mapped[str] = mapped_column(ForeignKey("alphas.id"), primary_key=True)
|
||||
raw: Mapped[dict] = mapped_column(JSON)
|
||||
points: Mapped[list] = mapped_column(JSON)
|
||||
fetched_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class ResearchTag(Base):
|
||||
__tablename__ = "research_tags"
|
||||
alpha_id: Mapped[str] = mapped_column(ForeignKey("alphas.id"), primary_key=True)
|
||||
tag: Mapped[str] = mapped_column(String(60), primary_key=True, index=True)
|
||||
|
||||
|
||||
class Job(Base):
|
||||
__tablename__ = "sync_jobs"
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(30))
|
||||
status: Mapped[str] = mapped_column(String(30), default="queued", index=True)
|
||||
payload: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
checkpoint: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
processed: Mapped[int] = mapped_column(Integer, default=0)
|
||||
failed: Mapped[int] = mapped_column(Integer, default=0)
|
||||
total: Mapped[int | None] = mapped_column(Integer)
|
||||
error: Mapped[str | None] = mapped_column(Text)
|
||||
next_retry_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
cancel_requested: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class JobItem(Base):
|
||||
__tablename__ = "job_items"
|
||||
job_id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), primary_key=True)
|
||||
alpha_id: Mapped[str] = mapped_column(String(100), primary_key=True)
|
||||
error: Mapped[str | None] = mapped_column(Text)
|
||||
@@ -0,0 +1,281 @@
|
||||
"""Validated public API contracts. Platform state is intentionally not a closed enum."""
|
||||
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from typing import Literal
|
||||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
ResearchState = Literal["inbox", "candidate", "optimizing", "archived"]
|
||||
SortField = Literal[
|
||||
"id",
|
||||
"name",
|
||||
"sharpe",
|
||||
"fitness",
|
||||
"returns",
|
||||
"turnover",
|
||||
"margin",
|
||||
"drawdown",
|
||||
"date_created",
|
||||
"date_submitted",
|
||||
"synced_at",
|
||||
]
|
||||
|
||||
|
||||
class Contract(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", allow_inf_nan=False)
|
||||
|
||||
|
||||
class LoginInput(Contract):
|
||||
username: str = Field(min_length=1, max_length=100)
|
||||
password: str = Field(min_length=1, max_length=1024)
|
||||
|
||||
|
||||
class CredentialsInput(Contract):
|
||||
email: str = Field(max_length=254)
|
||||
password: str = Field(min_length=1, max_length=1024)
|
||||
|
||||
@field_validator("email")
|
||||
@classmethod
|
||||
def valid_email(cls, value):
|
||||
value = value.strip()
|
||||
if not re.fullmatch(r"[^\s@]+@[^\s@]+\.[^\s@]+", value):
|
||||
raise ValueError("邮箱格式不正确")
|
||||
return value
|
||||
|
||||
|
||||
class PreferencesInput(Contract):
|
||||
display_name: str = Field(min_length=1, max_length=100)
|
||||
theme: Literal["light", "dark"]
|
||||
timezone: str
|
||||
page_size: Literal[25, 50, 100]
|
||||
|
||||
@field_validator("timezone")
|
||||
@classmethod
|
||||
def valid_timezone(cls, value):
|
||||
try:
|
||||
ZoneInfo(value)
|
||||
except (ZoneInfoNotFoundError, ValueError):
|
||||
raise ValueError("未知时区") from None
|
||||
return value
|
||||
|
||||
|
||||
class AlphaFilters(Contract):
|
||||
q: str | None = Field(default=None, max_length=300)
|
||||
region: str | None = None
|
||||
universe: str | None = None
|
||||
alpha_type: str | None = None
|
||||
language: str | None = None
|
||||
status: str | None = None
|
||||
stage: str | None = None
|
||||
hidden: bool | None = None
|
||||
research_state: ResearchState | None = None
|
||||
favorite: bool | None = None
|
||||
tag: str | None = Field(default=None, max_length=60)
|
||||
created_from: datetime | None = None
|
||||
created_to: datetime | None = None
|
||||
sharpe_min: float | None = None
|
||||
sharpe_max: float | None = None
|
||||
fitness_min: float | None = None
|
||||
fitness_max: float | None = None
|
||||
returns_min: float | None = None
|
||||
returns_max: float | None = None
|
||||
turnover_min: float | None = None
|
||||
turnover_max: float | None = None
|
||||
margin_min: float | None = None
|
||||
margin_max: float | None = None
|
||||
drawdown_min: float | None = None
|
||||
drawdown_max: float | None = None
|
||||
sort: SortField = "date_created"
|
||||
direction: Literal["asc", "desc"] = "desc"
|
||||
limit: int = Field(default=25, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
@field_validator("created_from", "created_to")
|
||||
@classmethod
|
||||
def utc_dates(cls, value):
|
||||
return value.replace(tzinfo=timezone.utc) if value and value.tzinfo is None else value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def range_order(self):
|
||||
for key in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
|
||||
lo, hi = getattr(self, f"{key}_min"), getattr(self, f"{key}_max")
|
||||
if lo is not None and hi is not None and lo > hi:
|
||||
raise ValueError(f"{key} 最小值不能大于最大值")
|
||||
if self.created_from and self.created_to and self.created_from > self.created_to:
|
||||
raise ValueError("开始时间不能晚于结束时间")
|
||||
return self
|
||||
|
||||
|
||||
def normalize_tags(values):
|
||||
result = sorted(set(v.strip() for v in values if v.strip()))
|
||||
if len(result) > 30 or any(len(v) > 60 for v in result):
|
||||
raise ValueError("最多 30 个标签,每个标签最多 60 字符")
|
||||
return result
|
||||
|
||||
|
||||
class ResearchInput(Contract):
|
||||
note: str = Field(default="", max_length=20000)
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
favorite: bool = False
|
||||
state: ResearchState = "inbox"
|
||||
_tags = field_validator("tags")(normalize_tags)
|
||||
|
||||
|
||||
def valid_ids(values):
|
||||
result = list(dict.fromkeys(values))
|
||||
if not result or len(result) > 100 or any(not re.fullmatch(r"[A-Za-z0-9_-]{1,100}", v) for v in result):
|
||||
raise ValueError("请输入 1–100 个合法 Alpha ID(字母、数字、下划线或短横线)")
|
||||
return result
|
||||
|
||||
|
||||
class BulkInput(Contract):
|
||||
alpha_ids: list[str]
|
||||
add_tags: list[str] = Field(default_factory=list)
|
||||
remove_tags: list[str] = Field(default_factory=list)
|
||||
state: ResearchState | None = None
|
||||
_ids = field_validator("alpha_ids")(valid_ids)
|
||||
_tags = field_validator("add_tags", "remove_tags")(normalize_tags)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def disjoint(self):
|
||||
if set(self.add_tags) & set(self.remove_tags):
|
||||
raise ValueError("同一标签不能同时添加和移除")
|
||||
return self
|
||||
|
||||
|
||||
class JobInput(Contract):
|
||||
kind: Literal["full_sync", "alpha_refresh", "pnl_refresh"]
|
||||
alpha_ids: list[str] = Field(default_factory=list)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_ids(self):
|
||||
if self.kind == "full_sync":
|
||||
if self.alpha_ids:
|
||||
raise ValueError("全量同步不接受 Alpha ID")
|
||||
else:
|
||||
self.alpha_ids = valid_ids(self.alpha_ids)
|
||||
return self
|
||||
|
||||
|
||||
class ResearchOutput(ResearchInput):
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class AlphaSummary(BaseModel):
|
||||
id: str
|
||||
name: str | None
|
||||
expression_preview: str
|
||||
alpha_type: str | None
|
||||
language: str | None
|
||||
stage: str | None
|
||||
status: str | None
|
||||
hidden: bool
|
||||
region: str | None
|
||||
universe: str | None
|
||||
sharpe: float | None
|
||||
fitness: float | None
|
||||
returns: float | None
|
||||
turnover: float | None
|
||||
margin: float | None
|
||||
drawdown: float | None
|
||||
date_created: datetime | None
|
||||
date_submitted: datetime | None
|
||||
synced_at: datetime
|
||||
research: ResearchOutput
|
||||
|
||||
|
||||
class AlphaDetail(AlphaSummary):
|
||||
expression: str | None
|
||||
selection: str | None
|
||||
combo: str | None
|
||||
settings: dict
|
||||
is_metrics: dict
|
||||
os_metrics: dict
|
||||
checks: list
|
||||
|
||||
|
||||
class AlphaPage(BaseModel):
|
||||
items: list[AlphaSummary]
|
||||
total: int
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
|
||||
class JobOutput(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
id: str
|
||||
kind: str
|
||||
status: str
|
||||
processed: int
|
||||
failed: int
|
||||
total: int | None
|
||||
error: str | None
|
||||
next_retry_at: datetime | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class AccountOutput(BaseModel):
|
||||
email: str | None
|
||||
configured: bool
|
||||
wq_user_id: str | None
|
||||
profile: dict
|
||||
connection_status: str
|
||||
connection_error: str | None
|
||||
verification_url: str | None
|
||||
last_synced_at: datetime | None
|
||||
display_name: str
|
||||
theme: Literal["light", "dark"]
|
||||
timezone: str
|
||||
page_size: int
|
||||
|
||||
|
||||
class SessionOutput(BaseModel):
|
||||
username: str
|
||||
|
||||
|
||||
class OkOutput(BaseModel):
|
||||
ok: bool
|
||||
|
||||
|
||||
class ErrorOutput(BaseModel):
|
||||
detail: str
|
||||
|
||||
|
||||
class HealthOutput(BaseModel):
|
||||
status: Literal["ok"]
|
||||
|
||||
|
||||
class BulkOutput(BaseModel):
|
||||
updated: int
|
||||
|
||||
|
||||
class PnlPoint(BaseModel):
|
||||
date: str
|
||||
value: float | None
|
||||
|
||||
|
||||
class PnlOutput(BaseModel):
|
||||
cached: bool
|
||||
points: list[PnlPoint]
|
||||
fetched_at: datetime | None
|
||||
|
||||
|
||||
class FacetsOutput(BaseModel):
|
||||
region: list[str]
|
||||
universe: list[str]
|
||||
alpha_type: list[str]
|
||||
language: list[str]
|
||||
status: list[str]
|
||||
stage: list[str]
|
||||
tags: list[str]
|
||||
total: int
|
||||
favorites: int
|
||||
last_sync: datetime | None
|
||||
|
||||
|
||||
class JobErrorOutput(BaseModel):
|
||||
alpha_id: str
|
||||
error: str
|
||||
@@ -0,0 +1,65 @@
|
||||
"""Server-side authentication and encryption utilities."""
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from datetime import timedelta
|
||||
|
||||
from argon2 import PasswordHasher
|
||||
from argon2.exceptions import VerificationError
|
||||
from cryptography.fernet import Fernet
|
||||
from fastapi import Depends, HTTPException, Request
|
||||
from fastapi.security import APIKeyCookie
|
||||
from sqlalchemy import delete
|
||||
|
||||
from .models import Account, Admin, LoginSession, now
|
||||
|
||||
password_hasher = PasswordHasher()
|
||||
session_cookie = APIKeyCookie(name="wq_session", auto_error=False)
|
||||
|
||||
|
||||
def token_hash(token: str) -> str:
|
||||
return hashlib.sha256(token.encode()).hexdigest()
|
||||
|
||||
|
||||
def cipher(settings) -> Fernet:
|
||||
return Fernet(settings.encryption_key.get_secret_value().encode())
|
||||
|
||||
|
||||
async def bootstrap(db, settings):
|
||||
"""Only initialize missing singleton records; deployments never reset existing passwords."""
|
||||
if not await db.get(Admin, 1):
|
||||
db.add(
|
||||
Admin(
|
||||
id=1,
|
||||
username=settings.admin_username,
|
||||
password_hash=password_hasher.hash(settings.admin_password.get_secret_value()),
|
||||
)
|
||||
)
|
||||
if not await db.get(Account, 1):
|
||||
db.add(Account(id=1))
|
||||
await db.execute(delete(LoginSession).where(LoginSession.expires_at < now()))
|
||||
await db.commit()
|
||||
|
||||
|
||||
def valid_password(encoded: str, value: str) -> bool:
|
||||
try:
|
||||
return password_hasher.verify(encoded, value)
|
||||
except VerificationError:
|
||||
return False
|
||||
|
||||
|
||||
async def issue_session(db, hours: int) -> str:
|
||||
token = secrets.token_urlsafe(32)
|
||||
db.add(LoginSession(token_hash=token_hash(token), expires_at=now() + timedelta(hours=hours)))
|
||||
await db.commit()
|
||||
return token
|
||||
|
||||
|
||||
async def require_auth(request: Request, token: str | None = Depends(session_cookie)):
|
||||
if not token:
|
||||
raise HTTPException(401, "请先登录系统")
|
||||
async with request.app.state.sessions() as db:
|
||||
row = await db.get(LoginSession, token_hash(token))
|
||||
# SQLite fixtures lose tzinfo; production PostgreSQL preserves it.
|
||||
if row is None or row.expires_at.replace(tzinfo=now().tzinfo) <= now():
|
||||
raise HTTPException(401, "系统登录已过期,请重新登录")
|
||||
@@ -0,0 +1,211 @@
|
||||
"""Read-only WorldQuant adapter. Authentication is the only allowed upstream POST.
|
||||
|
||||
No upstream response body or request headers are included in exceptions: they may
|
||||
contain credentials, cookies, or temporary authentication links.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import random
|
||||
from datetime import datetime, timezone
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Awaitable, Callable
|
||||
from urllib.parse import urljoin, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class WqError(Exception):
|
||||
def __init__(self, message: str, code: str = "upstream_error"):
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
|
||||
|
||||
class VerificationRequired(WqError):
|
||||
def __init__(self, url: str):
|
||||
super().__init__("请在 WorldQuant 完成人工验证,再点击继续验证", "verification_required")
|
||||
self.url = url
|
||||
|
||||
|
||||
class WqClient:
|
||||
def __init__(self, settings, transport=None, sleep=asyncio.sleep):
|
||||
self.settings = settings
|
||||
self.client = httpx.AsyncClient(
|
||||
base_url=settings.wq_base_url,
|
||||
timeout=settings.request_timeout,
|
||||
transport=transport,
|
||||
follow_redirects=False,
|
||||
)
|
||||
self.lock = asyncio.Lock()
|
||||
self.authenticated = False
|
||||
self.auth_generation = 0
|
||||
self.credentials: tuple[str, str] | None = None
|
||||
self.verification_url: str | None = None
|
||||
self.sleep = sleep
|
||||
self.on_retry: Callable[[float], Awaitable[None]] | None = None
|
||||
|
||||
async def close(self):
|
||||
await self.client.aclose()
|
||||
|
||||
def disconnect(self):
|
||||
self.authenticated = False
|
||||
self.credentials = None
|
||||
self.verification_url = None
|
||||
self.client.cookies.clear()
|
||||
|
||||
def safe_verification_url(self, response: httpx.Response) -> str:
|
||||
url = urljoin(str(response.url), response.headers.get("Location", ""))
|
||||
base = urlparse(self.settings.wq_base_url)
|
||||
parsed = urlparse(url)
|
||||
if (
|
||||
not response.headers.get("Location")
|
||||
or parsed.scheme != base.scheme
|
||||
or parsed.netloc != base.netloc
|
||||
):
|
||||
raise WqError("平台返回了无法确认的验证地址,请到 BRAIN 官网检查账户", "invalid_verification")
|
||||
return url
|
||||
|
||||
@staticmethod
|
||||
def retry_delay(value: str | None, attempt: int) -> float:
|
||||
if value:
|
||||
try:
|
||||
seconds = float(value)
|
||||
except ValueError:
|
||||
try:
|
||||
seconds = (parsedate_to_datetime(value) - datetime.now(timezone.utc)).total_seconds()
|
||||
except (TypeError, ValueError):
|
||||
seconds = -1
|
||||
if math.isfinite(seconds) and seconds >= 0:
|
||||
return seconds
|
||||
return min(2**attempt + random.uniform(0, 0.5), 30)
|
||||
|
||||
async def _request(self, method: str, path: str, **kwargs) -> httpx.Response:
|
||||
for attempt in range(self.settings.retry_attempts):
|
||||
try:
|
||||
response = await self.client.request(method, path, **kwargs)
|
||||
except (httpx.TimeoutException, httpx.NetworkError):
|
||||
if attempt + 1 == self.settings.retry_attempts:
|
||||
raise WqError("WorldQuant 网络请求失败,请稍后重试", "network_error") from None
|
||||
response = None
|
||||
if response is not None and response.status_code not in (429, 500, 502, 503, 504):
|
||||
return response
|
||||
if attempt + 1 == self.settings.retry_attempts:
|
||||
raise WqError("WorldQuant 暂时限流或不可用,重试预算已用完", "retry_exhausted")
|
||||
delay = self.retry_delay(
|
||||
response.headers.get("Retry-After") if response is not None else None, attempt
|
||||
)
|
||||
if self.on_retry:
|
||||
await self.on_retry(delay)
|
||||
await self.sleep(delay)
|
||||
raise WqError("重试失败")
|
||||
|
||||
async def authenticate(self, email: str, password: str, force=False, stale_generation=None):
|
||||
async with self.lock:
|
||||
if (
|
||||
self.authenticated
|
||||
and not force
|
||||
and (stale_generation is None or stale_generation != self.auth_generation)
|
||||
):
|
||||
return
|
||||
if self.verification_url and not force:
|
||||
raise VerificationRequired(self.verification_url)
|
||||
self.credentials = (email, password)
|
||||
self.authenticated = False
|
||||
if force:
|
||||
self.client.cookies.clear()
|
||||
self.verification_url = None
|
||||
response = await self._request("POST", "/authentication", auth=httpx.BasicAuth(email, password))
|
||||
if (
|
||||
response.status_code == 401
|
||||
and response.headers.get("WWW-Authenticate", "").lower() == "persona"
|
||||
):
|
||||
self.verification_url = self.safe_verification_url(response)
|
||||
raise VerificationRequired(self.verification_url)
|
||||
if response.status_code in (401, 403):
|
||||
raise WqError("平台认证未通过,请检查凭据或账户访问权限", "authentication_failed")
|
||||
if response.status_code not in (200, 201):
|
||||
raise WqError(f"平台认证返回异常状态 {response.status_code}", "authentication_failed")
|
||||
self.authenticated = True
|
||||
self.auth_generation += 1
|
||||
self.verification_url = None
|
||||
|
||||
async def verify(self):
|
||||
"""Continue the same server-side challenge session after the human visits the URL."""
|
||||
if not self.verification_url:
|
||||
if not self.credentials:
|
||||
raise WqError("验证会话已结束,请重新连接", "authentication_failed")
|
||||
await self.authenticate(*self.credentials, force=True)
|
||||
return
|
||||
response = await self._request("POST", self.verification_url)
|
||||
if response.status_code in (200, 201):
|
||||
self.authenticated = True
|
||||
self.auth_generation += 1
|
||||
self.verification_url = None
|
||||
elif response.status_code in (401, 403, 202):
|
||||
raise VerificationRequired(self.verification_url)
|
||||
else:
|
||||
raise WqError("验证会话已失效,请重新连接", "authentication_failed")
|
||||
|
||||
async def get(self, path: str, params=None):
|
||||
if not self.credentials:
|
||||
raise WqError("请先连接 WorldQuant", "disconnected")
|
||||
if not self.authenticated:
|
||||
await self.authenticate(*self.credentials)
|
||||
refreshed = False
|
||||
for attempt in range(self.settings.retry_attempts):
|
||||
generation = self.auth_generation
|
||||
response = await self._request("GET", path, params=params)
|
||||
if response.status_code == 401 and not refreshed:
|
||||
await self.authenticate(*self.credentials, stale_generation=generation)
|
||||
refreshed = True
|
||||
continue
|
||||
if response.status_code in (401, 403):
|
||||
raise WqError("无权访问该平台资源", "access_denied")
|
||||
if response.status_code == 404:
|
||||
raise WqError("平台资源不存在或不可访问", "not_found")
|
||||
if response.status_code not in (200, 201, 202):
|
||||
raise WqError(f"WorldQuant 请求返回状态 {response.status_code}")
|
||||
# Recordsets may return 200/202 with Retry-After before results exist.
|
||||
if (
|
||||
response.headers.get("Retry-After")
|
||||
and self.retry_delay(response.headers["Retry-After"], 0) > 0
|
||||
):
|
||||
if attempt + 1 == self.settings.retry_attempts:
|
||||
raise WqError("平台数据仍在准备,请稍后重试", "pending")
|
||||
delay = self.retry_delay(response.headers["Retry-After"], attempt)
|
||||
if self.on_retry:
|
||||
await self.on_retry(delay)
|
||||
await self.sleep(delay)
|
||||
continue
|
||||
try:
|
||||
result = response.json()
|
||||
if not isinstance(result, dict):
|
||||
raise ValueError()
|
||||
return result
|
||||
except ValueError:
|
||||
raise WqError("平台返回的数据格式无法识别", "invalid_response") from None
|
||||
raise WqError("平台数据尚未就绪", "pending")
|
||||
|
||||
async def profile(self):
|
||||
return await self.get("/users/self")
|
||||
|
||||
async def alphas(self, submission, hidden, offset, before):
|
||||
# Cover every platform stage; submitted records are not assumed to be OS only.
|
||||
status_key = "status" if submission == "UNSUBMITTED" else "status!"
|
||||
return await self.get(
|
||||
"/users/self/alphas",
|
||||
{
|
||||
status_key: "UNSUBMITTED",
|
||||
"hidden": str(hidden).lower(),
|
||||
"limit": 100,
|
||||
"offset": offset,
|
||||
"order": "dateCreated",
|
||||
"dateCreated<": before,
|
||||
},
|
||||
)
|
||||
|
||||
async def alpha(self, alpha_id):
|
||||
return await self.get(f"/alphas/{alpha_id}")
|
||||
|
||||
async def pnl(self, alpha_id):
|
||||
return await self.get(f"/alphas/{alpha_id}/recordsets/pnl")
|
||||
Reference in New Issue
Block a user