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 FastAPI application factory for Sanguo Quant API
""" """
from fastapi import FastAPI 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 from sanguo_orchestrator.runner import Orchestrator
def create_app(db_path: str, file_dir=None) -> FastAPI: def create_app(db_path: str, file_dir=None, auth_config=None, max_workers: int = 2) -> FastAPI:
"""Create FastAPI application with orchestrator""" """Create FastAPI application with orchestrator and optional authentication"""
app = FastAPI(title="Sanguo Quant API") app = FastAPI(title="Sanguo Quant API")
# Initialize orchestrator # 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) 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 # Include routes
app.include_router(router, prefix="/api/v1") 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 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 .schemas import CtaBacktestRequest, OptimizeRequest, FactorAnalysisRequest
from .auth import verify_token as verify_token_impl, verify_password, create_token
from .ws import manager
router = APIRouter() router = APIRouter()
_orchestrator = None _orchestrator = None
_auth_config = {"username": "admin", "password_hash": "", "jwt_secret": "x", "expire_minutes": 60}
def set_orchestrator(orch): def set_orchestrator(orch):
@@ -20,10 +24,41 @@ def get_orchestrator():
return _orchestrator return _orchestrator
@router.post("/backtest/cta") def set_auth_config(cfg):
def submit_cta(req: CtaBacktestRequest): """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""" """Submit CTA backtest task"""
tid = get_orchestrator().submit_cta( tid = await get_orchestrator().submit_cta(
strategy_class=req.strategy, strategy_class=req.strategy,
symbol=req.symbol, symbol=req.symbol,
params=req.params, params=req.params,
@@ -34,39 +69,73 @@ def submit_cta(req: CtaBacktestRequest):
return {"task_id": tid} return {"task_id": tid}
@router.post("/backtest/optimize") @router.post("/backtest/optimize", dependencies=[Depends(verify_token)])
def submit_optimize(req: OptimizeRequest): async def submit_optimize(req: OptimizeRequest):
"""Submit optimization task (placeholder)""" """Submit optimization task"""
# TODO: Implement optimize submission in Phase 3 tid = await get_orchestrator().submit_optimize(
return {"task_id": "pending_impl"} 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") @router.post("/factor/analyze", dependencies=[Depends(verify_token)])
def submit_factor(req: FactorAnalysisRequest): async def submit_factor(req: FactorAnalysisRequest):
"""Submit factor analysis task (placeholder)""" """Submit factor analysis task"""
# TODO: Implement factor analysis submission in Phase 3 tid = await get_orchestrator().submit_factor(
return {"task_id": "pending_impl"} 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): def get_status(task_id: str):
"""Get task status""" """Get task status"""
status = get_orchestrator().get_status(task_id) s = get_orchestrator().get_status(task_id)
if status is None: if s is None:
raise HTTPException(status_code=404, detail="task not found") raise HTTPException(status_code=404, detail="task not found")
stage = get_orchestrator().pool.get_stage(task_id)
return { return {
"task_id": task_id, "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): def get_result(task_id: str):
"""Get task result""" """Get task result"""
result = get_orchestrator().get_result(task_id) r = get_orchestrator().get_result(task_id)
if result is None: if r is None:
raise HTTPException(status_code=404, detail="result not ready") raise HTTPException(status_code=404, detail="result not ready")
return { return {"task_id": task_id, "statistics": r.statistics}
"task_id": task_id,
"statistics": result.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)
+151 -34
View File
@@ -1,7 +1,7 @@
""" """
FastAPI routes tests using TestClient FastAPI routes tests using TestClient
""" """
from unittest.mock import Mock, patch from unittest.mock import Mock, patch, AsyncMock
import pytest import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
@@ -10,17 +10,21 @@ def test_submit_cta_backtest():
"""Test POST /api/v1/backtest/cta returns task_id""" """Test POST /api/v1/backtest/cta returns task_id"""
# Create app with temporary DB # Create app with temporary DB
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
# Mock get_orchestrator to return mock orchestrator # Mock get_orchestrator to return mock orchestrator
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch: with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
mock_orch = Mock() mock_orch = Mock()
mock_orch.submit_cta.return_value = "cta_test_123" mock_orch.submit_cta = AsyncMock(return_value="cta_test_123")
mock_get_orch.return_value = mock_orch mock_get_orch.return_value = mock_orch
client = TestClient(app) client = TestClient(app)
@@ -32,7 +36,8 @@ def test_submit_cta_backtest():
"params": {"fast": 5, "slow": 20}, "params": {"fast": 5, "slow": 20},
"start": "2024-01-01", "start": "2024-01-01",
"end": "2024-12-31" "end": "2024-12-31"
} },
headers={"Authorization": f"Bearer {token}"}
) )
assert response.status_code == 200 assert response.status_code == 200
@@ -44,21 +49,26 @@ def test_submit_cta_backtest():
def test_get_task_status(): def test_get_task_status():
"""Test GET /api/v1/task/{task_id} returns status""" """Test GET /api/v1/task/{task_id} returns status"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
from sanguo_orchestrator.task import TaskState from sanguo_orchestrator.task import TaskState
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch: with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
mock_orch = Mock() mock_orch = Mock()
mock_orch.get_status.return_value = TaskState.DONE mock_orch.get_status.return_value = TaskState.DONE
mock_orch.pool.get_stage.return_value = "完成"
mock_get_orch.return_value = mock_orch mock_get_orch.return_value = mock_orch
client = TestClient(app) client = TestClient(app)
response = client.get("/api/v1/task/cta_test_123") response = client.get("/api/v1/task/cta_test_123", headers={"Authorization": f"Bearer {token}"})
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
@@ -70,12 +80,16 @@ def test_get_task_status():
def test_get_task_status_not_found(): def test_get_task_status_not_found():
"""Test GET /api/v1/task/{task_id} returns 404 for unknown task""" """Test GET /api/v1/task/{task_id} returns 404 for unknown task"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch: with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
mock_orch = Mock() mock_orch = Mock()
@@ -83,7 +97,7 @@ def test_get_task_status_not_found():
mock_get_orch.return_value = mock_orch mock_get_orch.return_value = mock_orch
client = TestClient(app) client = TestClient(app)
response = client.get("/api/v1/task/unknown_task") response = client.get("/api/v1/task/unknown_task", headers={"Authorization": f"Bearer {token}"})
assert response.status_code == 404 assert response.status_code == 404
@@ -91,12 +105,16 @@ def test_get_task_status_not_found():
def test_invalid_params_returns_422(): def test_invalid_params_returns_422():
"""Test POST /api/v1/backtest/cta with missing fields returns 422""" """Test POST /api/v1/backtest/cta with missing fields returns 422"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
client = TestClient(app) client = TestClient(app)
# Missing required field: strategy # Missing required field: strategy
@@ -107,7 +125,8 @@ def test_invalid_params_returns_422():
"params": {"fast": 5, "slow": 20}, "params": {"fast": 5, "slow": 20},
"start": "2024-01-01", "start": "2024-01-01",
"end": "2024-12-31" "end": "2024-12-31"
} },
headers={"Authorization": f"Bearer {token}"}
) )
assert response.status_code == 422 assert response.status_code == 422
@@ -116,13 +135,17 @@ def test_invalid_params_returns_422():
def test_get_task_result(): def test_get_task_result():
"""Test GET /api/v1/task/{task_id}/result returns statistics""" """Test GET /api/v1/task/{task_id}/result returns statistics"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
from sanguo_orchestrator.task import TaskState from sanguo_orchestrator.task import TaskState
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch: with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
mock_orch = Mock() mock_orch = Mock()
@@ -132,7 +155,7 @@ def test_get_task_result():
mock_get_orch.return_value = mock_orch mock_get_orch.return_value = mock_orch
client = TestClient(app) client = TestClient(app)
response = client.get("/api/v1/task/cta_test_123/result") response = client.get("/api/v1/task/cta_test_123/result", headers={"Authorization": f"Bearer {token}"})
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
@@ -145,12 +168,16 @@ def test_get_task_result():
def test_get_task_result_not_found(): def test_get_task_result_not_found():
"""Test GET /api/v1/task/{task_id}/result returns 404 when result not ready""" """Test GET /api/v1/task/{task_id}/result returns 404 when result not ready"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch: with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
mock_orch = Mock() mock_orch = Mock()
@@ -158,62 +185,152 @@ def test_get_task_result_not_found():
mock_get_orch.return_value = mock_orch mock_get_orch.return_value = mock_orch
client = TestClient(app) client = TestClient(app)
response = client.get("/api/v1/task/unknown_task/result") response = client.get("/api/v1/task/unknown_task/result", headers={"Authorization": f"Bearer {token}"})
assert response.status_code == 404 assert response.status_code == 404
def test_submit_optimize_returns_pending(): def test_submit_optimize_returns_pending():
"""Test POST /api/v1/backtest/optimize returns pending placeholder""" """Test POST /api/v1/backtest/optimize calls orchestrator submit_optimize"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
from unittest.mock import AsyncMock
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
client = TestClient(app) with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
response = client.post( mock_orch = Mock()
"/api/v1/backtest/optimize", mock_orch.submit_optimize = AsyncMock(return_value="opt_test_123")
json={ mock_get_orch.return_value = mock_orch
"symbol": "600000SH",
"strategy": "DoubleSMA", client = TestClient(app)
"grid": {"fast": [5, 10], "slow": [20, 30]}, response = client.post(
"start": "2024-01-01", "/api/v1/backtest/optimize",
"end": "2024-12-31", json={
"max_workers": 2 "symbol": "600000SH",
} "strategy": "DoubleSMA",
) "grid": {"fast": [5, 10], "slow": [20, 30]},
"start": "2024-01-01",
"end": "2024-12-31",
"max_workers": 2
},
headers={"Authorization": f"Bearer {token}"}
)
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
assert "task_id" in data assert "task_id" in data
assert data["task_id"] == "pending_impl" assert data["task_id"] == "opt_test_123"
def test_submit_factor_returns_pending(): def test_submit_factor_returns_pending():
"""Test POST /api/v1/factor/analyze returns pending placeholder""" """Test POST /api/v1/factor/analyze calls orchestrator submit_factor"""
from sanguo_api.app import create_app from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
from unittest.mock import AsyncMock
import tempfile import tempfile
import os import os
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = os.path.join(tmp, "test.db") db_path = os.path.join(tmp, "test.db")
app = create_app(db_path=db_path, file_dir=None) app = create_app(db_path=db_path, file_dir=None)
token = create_token("admin")
client = TestClient(app) with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
response = client.post( mock_orch = Mock()
"/api/v1/factor/analyze", mock_orch.submit_factor = AsyncMock(return_value="factor_test_123")
json={ mock_get_orch.return_value = mock_orch
"symbols": ["600000SH", "000001SZ"],
"factor_names": ["ts_mean_5", "ts_mean_20"], client = TestClient(app)
"start": "2024-01-01", response = client.post(
"end": "2024-12-31" "/api/v1/factor/analyze",
} json={
) "symbols": ["600000SH", "000001SZ"],
"factor_names": ["ts_mean_5", "ts_mean_20"],
"start": "2024-01-01",
"end": "2024-12-31"
},
headers={"Authorization": f"Bearer {token}"}
)
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
assert "task_id" in data assert "task_id" in data
assert data["task_id"] == "pending_impl" assert data["task_id"] == "factor_test_123"
# ============================================
# NEW TESTS FOR TASK 5 (JWT + LOGIN + OPTIMIZE/FCTOR + WS)
# ============================================
def test_login_returns_token(tmp_path):
"""Test POST /api/v1/auth/login returns JWT token on successful login"""
from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, hash_password
import tempfile
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp:
db_path = f"{tmp}/t.db"
app = create_app(db_path=db_path, auth_config={
"username": "admin",
"password_hash": hash_password("pass123"),
"jwt_secret": "test",
"expire_minutes": 60
})
client = TestClient(app)
resp = client.post("/api/v1/auth/login", json={"username": "admin", "password": "pass123"})
assert resp.status_code == 200
assert "token" in resp.json()
def test_protected_route_without_token_401(tmp_path):
"""Test that protected routes return 401 without JWT token"""
from sanguo_api.app import create_app
import tempfile
with tempfile.TemporaryDirectory() as tmp:
db_path = f"{tmp}/t.db"
app = create_app(db_path=db_path)
client = TestClient(app)
resp = client.get("/api/v1/task/t1")
assert resp.status_code == 401
def test_optimize_route_calls_submit(tmp_path):
"""Test POST /api/v1/backtest/optimize calls orchestrator submit_optimize"""
from sanguo_api.app import create_app
from sanguo_api.auth import set_jwt_config, create_token
from unittest.mock import AsyncMock
import tempfile
set_jwt_config(secret="test", expire_minutes=60)
with tempfile.TemporaryDirectory() as tmp:
db_path = f"{tmp}/t.db"
app = create_app(db_path=db_path)
client = TestClient(app)
token = create_token("admin")
with patch("sanguo_api.routes.get_orchestrator") as m:
orch = Mock()
orch.submit_optimize = AsyncMock(return_value="opt_1")
m.return_value = orch
resp = client.post("/api/v1/backtest/optimize", json={
"symbol": "600000", "strategy": "MaStrategy", "grid": {"n": [5, 20, 5]},
"start": "2024-01-01", "end": "2024-06-30", "max_workers": 2
}, headers={"Authorization": f"Bearer {token}"})
assert resp.status_code == 200
assert resp.json()["task_id"] == "opt_1"