feat(api): routes 完整(login + JWT 依赖 + optimize/factor + WS route)

This commit is contained in:
2026-07-06 19:04:26 +08:00
parent f5a69b312a
commit efaac41c06
3 changed files with 267 additions and 64 deletions
+22 -5
View File
@@ -2,19 +2,36 @@
FastAPI application factory for Sanguo Quant API
"""
from fastapi import FastAPI
from .routes import router, set_orchestrator
from .routes import router, set_orchestrator, set_auth_config
from .auth import set_jwt_config
from .ws import manager
from sanguo_orchestrator.runner import Orchestrator
def create_app(db_path: str, file_dir=None) -> FastAPI:
"""Create FastAPI application with orchestrator"""
def create_app(db_path: str, file_dir=None, auth_config=None, max_workers: int = 2) -> FastAPI:
"""Create FastAPI application with orchestrator and optional authentication"""
app = FastAPI(title="Sanguo Quant API")
# Initialize orchestrator
orch = Orchestrator(db_path=db_path, file_dir=file_dir)
orch = Orchestrator(db_path=db_path, file_dir=file_dir, max_workers=max_workers)
# Set up WebSocket stage callback
async def _on_stage(task_id, stage):
"""Broadcast stage updates to WebSocket subscribers"""
await manager.broadcast(task_id, {"task_id": task_id, "stage": stage})
orch.set_on_stage(_on_stage)
set_orchestrator(orch)
# Configure authentication if provided
if auth_config:
set_auth_config(auth_config)
set_jwt_config(
auth_config.get("jwt_secret", "x"),
auth_config.get("expire_minutes", 60)
)
# Include routes
app.include_router(router, prefix="/api/v1")
return app
return app
+94 -25
View File
@@ -1,12 +1,16 @@
"""
FastAPI routes for Sanguo Quant API
"""
from fastapi import APIRouter, HTTPException
from fastapi import APIRouter, HTTPException, Depends, WebSocket, Query, Header
from pydantic import BaseModel
from .schemas import CtaBacktestRequest, OptimizeRequest, FactorAnalysisRequest
from .auth import verify_token as verify_token_impl, verify_password, create_token
from .ws import manager
router = APIRouter()
_orchestrator = None
_auth_config = {"username": "admin", "password_hash": "", "jwt_secret": "x", "expire_minutes": 60}
def set_orchestrator(orch):
@@ -20,10 +24,41 @@ def get_orchestrator():
return _orchestrator
@router.post("/backtest/cta")
def submit_cta(req: CtaBacktestRequest):
def set_auth_config(cfg):
"""Set authentication configuration"""
_auth_config.update(cfg)
class LoginRequest(BaseModel):
"""Login request schema"""
username: str
password: str
async def verify_token(authorization: str | None = Header(None)):
"""Dependency to verify JWT token from Authorization header"""
if authorization is None:
raise HTTPException(status_code=401, detail="Missing authorization header")
if not authorization.startswith("Bearer "):
raise HTTPException(status_code=401, detail="Invalid authorization header format")
token = authorization.split(" ")[1]
return verify_token_impl(token)
@router.post("/auth/login")
def login(req: LoginRequest):
"""Authenticate user and return JWT token"""
if req.username != _auth_config["username"] or not verify_password(req.password, _auth_config["password_hash"]):
raise HTTPException(status_code=401, detail="用户名或密码错误")
return {"token": create_token(req.username)}
@router.post("/backtest/cta", dependencies=[Depends(verify_token)])
async def submit_cta(req: CtaBacktestRequest):
"""Submit CTA backtest task"""
tid = get_orchestrator().submit_cta(
tid = await get_orchestrator().submit_cta(
strategy_class=req.strategy,
symbol=req.symbol,
params=req.params,
@@ -34,39 +69,73 @@ def submit_cta(req: CtaBacktestRequest):
return {"task_id": tid}
@router.post("/backtest/optimize")
def submit_optimize(req: OptimizeRequest):
"""Submit optimization task (placeholder)"""
# TODO: Implement optimize submission in Phase 3
return {"task_id": "pending_impl"}
@router.post("/backtest/optimize", dependencies=[Depends(verify_token)])
async def submit_optimize(req: OptimizeRequest):
"""Submit optimization task"""
tid = await get_orchestrator().submit_optimize(
strategy_class=req.strategy,
symbol=req.symbol,
grid=req.grid,
start=req.start,
end=req.end,
cfg=None
)
return {"task_id": tid}
@router.post("/factor/analyze")
def submit_factor(req: FactorAnalysisRequest):
"""Submit factor analysis task (placeholder)"""
# TODO: Implement factor analysis submission in Phase 3
return {"task_id": "pending_impl"}
@router.post("/factor/analyze", dependencies=[Depends(verify_token)])
async def submit_factor(req: FactorAnalysisRequest):
"""Submit factor analysis task"""
tid = await get_orchestrator().submit_factor(
symbols=req.symbols,
factor_names=req.factor_names,
start=req.start,
end=req.end,
cfg=None,
output_dir="/tmp/factor"
)
return {"task_id": tid}
@router.get("/task/{task_id}")
@router.get("/task/{task_id}", dependencies=[Depends(verify_token)])
def get_status(task_id: str):
"""Get task status"""
status = get_orchestrator().get_status(task_id)
if status is None:
s = get_orchestrator().get_status(task_id)
if s is None:
raise HTTPException(status_code=404, detail="task not found")
stage = get_orchestrator().pool.get_stage(task_id)
return {
"task_id": task_id,
"status": status.value if hasattr(status, "value") else str(status)
"status": s.value if hasattr(s, "value") else str(s),
"stage": stage or ""
}
@router.get("/task/{task_id}/result")
@router.get("/task/{task_id}/result", dependencies=[Depends(verify_token)])
def get_result(task_id: str):
"""Get task result"""
result = get_orchestrator().get_result(task_id)
if result is None:
r = get_orchestrator().get_result(task_id)
if r is None:
raise HTTPException(status_code=404, detail="result not ready")
return {
"task_id": task_id,
"statistics": result.statistics
}
return {"task_id": task_id, "statistics": r.statistics}
@router.websocket("/ws/task/{task_id}")
async def task_ws(websocket: WebSocket, task_id: str, token: str = Query(...)):
"""WebSocket endpoint for task status updates"""
try:
verify_token(token)
except Exception:
await websocket.close(code=4401)
return
await websocket.accept()
manager.connect(task_id, websocket)
try:
while True:
await websocket.receive_text() # Keep connection alive
except Exception:
pass
finally:
manager.disconnect(task_id, websocket)