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:
yuxuanhui
2026-09-07 14:54:20 +08:00
commit 3cd280d068
53 changed files with 10888 additions and 0 deletions
View File
+199
View File
@@ -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"])
+33
View File
@@ -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())
+28
View File
@@ -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
+8
View File
@@ -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)
+390
View File
@@ -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()
+498
View File
@@ -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
+123
View File
@@ -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)
+281
View File
@@ -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
+65
View File
@@ -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, "系统登录已过期,请重新登录")
+211
View File
@@ -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")