3cd280d068
- 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.
499 lines
19 KiB
Python
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
|