Files
worldquant-alpha-system/backend/app/main.py
T
yuxuanhui 3cd280d068 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.
2026-09-07 14:54:20 +08:00

499 lines
19 KiB
Python

"""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