Files
worldquant-alpha-system/backend/app/main.py
T

492 lines
20 KiB
Python

"""FastAPI application and authenticated public API."""
import asyncio
import csv
import io
import time
from collections import defaultdict
from contextlib import AsyncExitStack, 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 pydantic import ValidationError
from sqlalchemy import delete, select, text
from .ai.routes import router as ai_router
from .ai.runtime import AIRuntime
from .alphas import failed_checks, list_statement, sorted_statement
from .backtests.routes import router as backtest_router
from .business import Business, notify_job
from .catalog.research_routes import router as research_catalog_router
from .catalog.routes import router as catalog_router
from .config import Settings
from .dashboard import router as dashboard_router
from .db import create_database
from .jobs import AUTH_KINDS, Runner, create_job
from .mcp_api.token_routes import router as mcp_token_router
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession
from .research.routes import router as research_router
from .research.runtime import ResearchRuntime
from .schemas import (
AccountOutput,
AlphaDetail,
AlphaFilters,
AlphaPage,
AlphaSourcePage,
BulkOutput,
BulkUpdate,
CredentialsInput,
ErrorOutput,
FacetsOutput,
HealthOutput,
JobErrorOutput,
JobInput,
JobOutput,
LoginInput,
OkOutput,
PnlOutput,
PreferencesInput,
ResearchUpdate,
SelfCorrelationOutput,
SessionOutput,
)
from .security import bootstrap, cipher, issue_session, require_auth, token_hash, valid_password
from .submission import router as submission_router
def account_output(account, client, settings):
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),
"credentials_source": "environment" if settings.wq_email else "database",
"session": client.session_info(),
}
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, ai_model_factory=None):
settings = settings or Settings()
engine, sessions = create_database(settings.database_url)
runner = Runner(sessions, settings, client=wq_client)
ai_runtime = AIRuntime(sessions, settings, runner, ai_model_factory)
research_runtime = ResearchRuntime(sessions, ai_runtime, runner)
mcp_runtime = None
if settings.mcp_enabled:
from .mcp_api.server import MCPResearchServer
mcp_runtime = MCPResearchServer(sessions, runner, settings)
@asynccontextmanager
async def lifespan(app):
async with sessions() as db:
await bootstrap(db, settings)
async with sessions.begin() as db:
if not await db.get(BacktestConfig, 1):
db.add(BacktestConfig(id=1))
await ai_runtime.start()
if settings.enable_runner:
await runner.start()
await research_runtime.start()
async with AsyncExitStack() as stack:
if mcp_runtime:
await stack.enter_async_context(mcp_runtime.server.session_manager.run())
yield
if settings.enable_runner:
await research_runtime.stop()
await ai_runtime.stop()
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
app.state.ai = ai_runtime
app.state.research = research_runtime
app.state.mcp = mcp_runtime
login_failures = defaultdict(list)
@app.exception_handler(RequestValidationError)
@app.exception_handler(ValidationError)
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):
is_mcp = request.url.path in ("/api/v1/mcp", "/api/v1/mcp/")
if is_mcp:
if not mcp_runtime:
return JSONResponse({"detail": "MCP 未启用"}, status_code=404)
from urllib.parse import urlsplit
from .mcp_api.auth import authenticate
from .mcp_api.server import TOOLS
if request.headers.get("host", "").lower() != urlsplit(settings.public_origin).netloc.lower():
return JSONResponse({"detail": "MCP Host 不被允许"}, status_code=403)
origin = request.headers.get("origin")
if origin and origin.rstrip("/") != settings.public_origin.rstrip("/"):
return JSONResponse({"detail": "MCP Origin 不被允许"}, status_code=403)
scheme, _, secret = request.headers.get("authorization", "").partition(" ")
if scheme.lower() != "bearer" or not secret or len(secret) > 256:
return JSONResponse({"detail": "需要 MCP Bearer 令牌"}, status_code=401,
headers={"WWW-Authenticate": "Bearer"})
try:
async with sessions() as db:
principal = await authenticate(db, secret)
except HTTPException as exc:
return JSONResponse({"detail": exc.detail}, status_code=exc.status_code,
headers={"WWW-Authenticate": "Bearer"})
request.state.mcp_principal = principal
if "research:read" not in principal.scopes:
return JSONResponse({"detail": "缺少读取权限"}, status_code=403)
if request.method == "POST":
body = bytearray()
async for chunk in request.stream():
body.extend(chunk)
if len(body) > 4 * 1024 * 1024:
return JSONResponse({"detail": "MCP 请求过大"}, status_code=413)
# BaseHTTPMiddleware replays cached bytes to the SDK; never log this payload.
request._body = bytes(body)
try:
import json
message = json.loads(body)
except (ValueError, UnicodeDecodeError):
return JSONResponse({"detail": "无效 JSON"}, status_code=400)
if isinstance(message, dict) and message.get("method") == "tools/call":
params = message.get("params")
tool = params.get("name") if isinstance(params, dict) else None
definition = TOOLS.get(tool) if isinstance(tool, str) else None
if definition and definition[2] not in principal.scopes:
return JSONResponse({"detail": "MCP 令牌缺少所需权限"}, status_code=403)
elif 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), runner.client, settings)
@api.put("/account/credentials", response_model=AccountOutput, tags=["account"])
async def credentials(body: CredentialsInput):
if settings.wq_email:
raise HTTPException(409, "WorldQuant 凭据由环境变量管理,请修改部署配置并重启服务")
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, runner.client, settings)
@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, runner.client, settings)
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()]):
async with sessions() as db:
return await Business(db).search_alphas(filters)
@api.get("/alphas/facets", response_model=FacetsOutput, tags=["alphas"])
async def facets():
async with sessions() as db:
return await Business(db).get_alpha_facets()
@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",
"sub_universe_sharpe",
"robust_universe_sharpe",
"two_year_sharpe",
"prod_correlation",
"pnl",
"neutralization",
"check_type",
"date_created",
"date_submitted",
"synced_at",
]
writer.writerow(columns + ["failed_checks", "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]
+ [csv_cell(";".join(failed_checks(a.checks))), 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: BulkUpdate):
async with sessions.begin() as db:
return await Business(db).bulk_update_research(body)
@api.get("/alphas/{alpha_id}", response_model=AlphaDetail, tags=["alphas"])
async def detail(alpha_id: str):
async with sessions() as db:
return await Business(db).get_alpha(alpha_id)
@api.patch("/alphas/{alpha_id}/research", response_model=OkOutput, tags=["alphas"])
async def research(alpha_id: str, body: ResearchUpdate):
async with sessions.begin() as db:
return await Business(db).update_research(alpha_id, body)
@api.get("/alphas/{alpha_id}/sources", response_model=AlphaSourcePage, tags=["alphas"])
async def sources(alpha_id: str, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)):
async with sessions() as db:
return await Business(db).get_alpha_sources(alpha_id, limit, offset)
@api.get("/alphas/{alpha_id}/pnl", response_model=PnlOutput, tags=["alphas"])
async def pnl(alpha_id: str):
async with sessions() as db:
return await Business(db).get_alpha_pnl(alpha_id)
@api.get("/alphas/{alpha_id}/self-correlation", response_model=SelfCorrelationOutput, tags=["alphas"])
async def self_correlation(alpha_id: str):
async with sessions() as db:
return await Business(db).get_self_correlation(alpha_id)
@api.post("/sync-jobs", status_code=202, response_model=JobOutput, tags=["sync-jobs"])
async def new_job(body: JobInput):
async with sessions.begin() as db:
result = await Business(db).create_sync_job(body)
await notify_job(runner, "create_sync_job", result)
return result
@api.get("/sync-jobs", response_model=list[JobOutput], tags=["sync-jobs"])
async def get_jobs():
async with sessions() as db:
return await Business(db).list_jobs()
@api.get("/sync-jobs/{job_id}", response_model=JobOutput, tags=["sync-jobs"])
async def get_job(job_id: str):
async with sessions() as db:
return await Business(db).get_job_status(job_id)
@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.begin() as db:
result = await Business(db).cancel_job(job_id)
await notify_job(runner, "cancel_job", result)
return result
@api.post("/sync-jobs/{job_id}/retry", response_model=JobOutput, tags=["sync-jobs"])
async def retry_job(job_id: str):
async with sessions.begin() as db:
result = await Business(db).retry_job(job_id)
await notify_job(runner, "retry_job", result)
return result
if mcp_runtime:
app.mount("/api/v1/mcp", mcp_runtime.app)
app.include_router(mcp_token_router)
app.include_router(dashboard_router)
app.include_router(backtest_router)
app.include_router(api)
app.include_router(catalog_router)
app.include_router(research_catalog_router)
app.include_router(research_router)
app.include_router(ai_router(ai_runtime))
app.include_router(submission_router(runner, ai_runtime))
return app