720 lines
26 KiB
Python
720 lines
26 KiB
Python
"""
|
|
FastAPI routes tests using TestClient
|
|
"""
|
|
from unittest.mock import Mock, patch, AsyncMock
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
|
|
def test_submit_cta_backtest():
|
|
"""Test POST /api/v1/backtest/cta returns task_id"""
|
|
# Create app with temporary DB
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=None)
|
|
token = create_token("admin")
|
|
|
|
# Mock get_orchestrator to return mock orchestrator
|
|
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch, \
|
|
patch("sanguo_api.routes.get_strategy_class", return_value=Mock()):
|
|
mock_orch = Mock()
|
|
mock_orch.submit_cta = AsyncMock(return_value="cta_test_123")
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/api/v1/backtest/cta",
|
|
json={
|
|
"symbol": "600000SH",
|
|
"strategy": "DoubleSMA",
|
|
"params": {"fast": 5, "slow": 20},
|
|
"start": "2024-01-01",
|
|
"end": "2024-12-31"
|
|
},
|
|
headers={"Authorization": f"Bearer {token}"}
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
assert data["task_id"] == "cta_test_123"
|
|
|
|
|
|
def test_get_task_status():
|
|
"""Test GET /api/v1/task/{task_id} returns status"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
from sanguo_orchestrator.task import TaskState
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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:
|
|
mock_orch = Mock()
|
|
mock_orch.get_status.return_value = TaskState.DONE
|
|
mock_orch.pool.get_stage.return_value = "完成"
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
assert data["task_id"] == "cta_test_123"
|
|
assert data["status"] == "done"
|
|
|
|
|
|
def test_get_task_status_not_found():
|
|
"""Test GET /api/v1/task/{task_id} returns 404 for unknown task"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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:
|
|
mock_orch = Mock()
|
|
mock_orch.get_status.return_value = None
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 404
|
|
|
|
|
|
def test_invalid_params_returns_422():
|
|
"""Test POST /api/v1/backtest/cta with missing fields returns 422"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=None)
|
|
token = create_token("admin")
|
|
|
|
client = TestClient(app)
|
|
# Missing required field: strategy
|
|
response = client.post(
|
|
"/api/v1/backtest/cta",
|
|
json={
|
|
"symbol": "600000SH",
|
|
"params": {"fast": 5, "slow": 20},
|
|
"start": "2024-01-01",
|
|
"end": "2024-12-31"
|
|
},
|
|
headers={"Authorization": f"Bearer {token}"}
|
|
)
|
|
|
|
assert response.status_code == 422
|
|
|
|
|
|
def test_get_task_result():
|
|
"""Test GET /api/v1/task/{task_id}/result returns statistics"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
from sanguo_orchestrator.task import TaskState
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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:
|
|
from sanguo_backtest.result_store import BacktestResult
|
|
mock_orch = Mock()
|
|
mock_result = BacktestResult(
|
|
task_id="cta_test_123", type="cta", status="done", strategy="S", symbol="600000",
|
|
params={}, start="2024-01-01", end="2024-12-31",
|
|
statistics={"total_trades": 10, "total_return": 0.15},
|
|
)
|
|
mock_orch.get_result.return_value = mock_result
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/result", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
assert data["task_id"] == "cta_test_123"
|
|
assert "statistics" in data
|
|
assert data["statistics"]["total_trades"] == 10
|
|
|
|
|
|
def test_get_task_result_not_found():
|
|
"""Test GET /api/v1/task/{task_id}/result returns 404 when result not ready"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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:
|
|
mock_orch = Mock()
|
|
mock_orch.get_result.return_value = None
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task/result", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 404
|
|
|
|
|
|
def test_submit_optimize_returns_pending():
|
|
"""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
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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, \
|
|
patch("sanguo_api.routes.get_strategy_class", return_value=Mock()):
|
|
mock_orch = Mock()
|
|
mock_orch.submit_optimize = AsyncMock(return_value="opt_test_123")
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/api/v1/backtest/optimize",
|
|
json={
|
|
"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
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
assert data["task_id"] == "opt_test_123"
|
|
|
|
|
|
def test_submit_factor_returns_pending():
|
|
"""Test POST /api/v1/factor/analyze calls orchestrator submit_factor"""
|
|
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 os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
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:
|
|
mock_orch = Mock()
|
|
mock_orch.submit_factor = AsyncMock(return_value="factor_test_123")
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/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
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
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, \
|
|
patch("sanguo_api.routes.get_strategy_class", return_value=Mock()):
|
|
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"
|
|
|
|
|
|
# ============================================
|
|
# NEW TESTS FOR TASK 4 (API 扩展端点)
|
|
# ============================================
|
|
|
|
def test_backtest_cta_accepts_benchmark():
|
|
"""Test POST /api/v1/backtest/cta accepts benchmark parameter"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch, \
|
|
patch("sanguo_api.routes.get_strategy_class", return_value=Mock()):
|
|
mock_orch = Mock()
|
|
mock_orch.submit_cta = AsyncMock(return_value="cta_test_123")
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/api/v1/backtest/cta",
|
|
json={
|
|
"symbol": "600000SH",
|
|
"strategy": "DoubleSMA",
|
|
"params": {"fast": 5, "slow": 20},
|
|
"start": "2024-01-01",
|
|
"end": "2024-12-31",
|
|
"benchmark": "zz500"
|
|
},
|
|
headers={"Authorization": f"Bearer {token}"}
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "task_id" in data
|
|
assert data["task_id"] == "cta_test_123"
|
|
|
|
|
|
def test_backtest_cta_rejects_invalid_benchmark():
|
|
"""Test POST /api/v1/backtest/cta rejects invalid benchmark"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/api/v1/backtest/cta",
|
|
json={
|
|
"symbol": "600000SH",
|
|
"strategy": "DoubleSMA",
|
|
"params": {"fast": 5, "slow": 20},
|
|
"start": "2024-01-01",
|
|
"end": "2024-12-31",
|
|
"benchmark": "xxx" # Invalid benchmark
|
|
},
|
|
headers={"Authorization": f"Bearer {token}"}
|
|
)
|
|
|
|
assert response.status_code == 422
|
|
|
|
|
|
def test_task_result_includes_relative_metrics():
|
|
"""Test GET /api/v1/task/:id/result includes relative_metrics"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
from sanguo_backtest.result_store import BacktestResult
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
|
|
mock_orch = Mock()
|
|
mock_result = BacktestResult(
|
|
task_id="cta_test_123", type="cta", status="done", strategy="S", symbol="600000",
|
|
params={}, start="2024-01-01", end="2024-12-31",
|
|
statistics={
|
|
"total_trades": 10,
|
|
"total_return": 0.15,
|
|
"alpha": 0.05,
|
|
"beta": 1.2,
|
|
"sortino_ratio": 1.5,
|
|
"information_ratio": 0.8,
|
|
"annual_volatility": 0.2,
|
|
"benchmark_return": 0.1,
|
|
"benchmark_volatility": 0.18
|
|
},
|
|
)
|
|
mock_orch.get_result.return_value = mock_result
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/result", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "relative_metrics" in data
|
|
assert data["relative_metrics"]["alpha"] == 0.05
|
|
assert data["relative_metrics"]["beta"] == 1.2
|
|
|
|
|
|
def test_benchmark_curve_endpoint():
|
|
"""Test GET /api/v1/task/:id/benchmark-curve returns benchmark curve data"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
import json
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
file_dir = tmp
|
|
app = create_app(db_path=db_path, file_dir=file_dir)
|
|
token = create_token("admin")
|
|
|
|
# Create mock metrics file
|
|
metrics_data = {
|
|
"series": {
|
|
"equity_curve": {
|
|
"dates": ["2024-01-02", "2024-01-03"],
|
|
"values": [1000000.0, 1001000.0]
|
|
},
|
|
"benchmark_curve": {
|
|
"dates": ["2024-01-02", "2024-01-03"],
|
|
"values": [1.0, 1.001]
|
|
}
|
|
}
|
|
}
|
|
metrics_file = os.path.join(file_dir, "cta_test_123_metrics.json")
|
|
with open(metrics_file, "w") as f:
|
|
json.dump(metrics_data, f)
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/benchmark-curve", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "dates" in data
|
|
assert "strategy" in data
|
|
assert "benchmark" in data
|
|
assert len(data["dates"]) == 2
|
|
assert len(data["strategy"]) == 2
|
|
assert len(data["benchmark"]) == 2
|
|
|
|
|
|
def test_benchmark_curve_endpoint_not_found():
|
|
"""Test GET /api/v1/task/:id/benchmark-curve returns 404 when file not found"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task/benchmark-curve", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 404
|
|
|
|
|
|
def test_risk_series_endpoint():
|
|
"""Test GET /api/v1/task/:id/risk-series returns risk series data"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
import json
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
file_dir = tmp
|
|
app = create_app(db_path=db_path, file_dir=file_dir)
|
|
token = create_token("admin")
|
|
|
|
# Create mock metrics file
|
|
metrics_data = {
|
|
"series": {
|
|
"alpha": {
|
|
"dates": ["2024-01-02", "2024-01-03"],
|
|
"values": [0.05, 0.06]
|
|
},
|
|
"beta": {
|
|
"dates": ["2024-01-02", "2024-01-03"],
|
|
"values": [1.2, 1.1]
|
|
},
|
|
"drawdown": {
|
|
"dates": ["2024-01-02", "2024-01-03"],
|
|
"values": [-0.01, -0.005]
|
|
}
|
|
}
|
|
}
|
|
metrics_file = os.path.join(file_dir, "cta_test_123_metrics.json")
|
|
with open(metrics_file, "w") as f:
|
|
json.dump(metrics_data, f)
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/risk-series", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "dates" in data
|
|
assert "alpha" in data
|
|
assert "beta" in data
|
|
assert "drawdown" in data
|
|
assert len(data["dates"]) == 2
|
|
|
|
|
|
def test_risk_series_endpoint_not_found():
|
|
"""Test GET /api/v1/task/:id/risk-series returns 404 when file not found"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task/risk-series", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 404
|
|
|
|
|
|
def test_daily_holdings_endpoint():
|
|
"""Test GET /api/v1/task/:id/daily-holdings returns daily holdings data"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
from sanguo_backtest.result_store import BacktestResult
|
|
import tempfile
|
|
import os
|
|
import pandas as pd
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
|
|
# Create mock equity curve data
|
|
equity_df = pd.DataFrame({
|
|
"date": ["2024-01-02", "2024-01-03"],
|
|
"balance": [1000000.0, 1001000.0],
|
|
"return": [0.0, 0.001]
|
|
})
|
|
|
|
mock_result = BacktestResult(
|
|
task_id="cta_test_123", type="cta", status="done", strategy="S", symbol="600000",
|
|
params={}, start="2024-01-01", end="2024-12-31",
|
|
statistics={},
|
|
equity_curve=equity_df
|
|
)
|
|
mock_orch = Mock()
|
|
mock_orch.get_result.return_value = mock_result
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/daily-holdings", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "daily_holdings" in data
|
|
assert len(data["daily_holdings"]) == 2
|
|
|
|
|
|
def test_daily_holdings_endpoint_not_found():
|
|
"""Test GET /api/v1/task/:id/daily-holdings returns 404 when result not ready"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
with patch("sanguo_api.routes.get_orchestrator") as mock_get_orch:
|
|
mock_orch = Mock()
|
|
mock_orch.get_result.return_value = None
|
|
mock_get_orch.return_value = mock_orch
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task/daily-holdings", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 404
|
|
|
|
|
|
def test_log_endpoint():
|
|
"""Test GET /api/v1/task/:id/log returns log data"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
file_dir = tmp
|
|
app = create_app(db_path=db_path, file_dir=file_dir)
|
|
token = create_token("admin")
|
|
|
|
# Create mock log file
|
|
log_file = os.path.join(file_dir, "cta_test_123.log")
|
|
with open(log_file, "w") as f:
|
|
f.write("Backtest started\nBacktest completed\n")
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/cta_test_123/log", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "log" in data
|
|
assert "Backtest started" in data["log"]
|
|
|
|
|
|
def test_log_endpoint_not_found():
|
|
"""Test GET /api/v1/task/:id/log returns empty log when file not found"""
|
|
from sanguo_api.app import create_app
|
|
from sanguo_api.auth import set_jwt_config, create_token
|
|
import tempfile
|
|
import os
|
|
|
|
set_jwt_config(secret="test", expire_minutes=60)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = os.path.join(tmp, "test.db")
|
|
app = create_app(db_path=db_path, file_dir=tmp)
|
|
token = create_token("admin")
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/api/v1/task/unknown_task/log", headers={"Authorization": f"Bearer {token}"})
|
|
|
|
# Log endpoint returns empty log instead of 404
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "log" in data
|
|
assert data["log"] == ""
|