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,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
|
||||
Reference in New Issue
Block a user