"""公共配置 + 回测组合 API 测试(TestClient + 内存 SQLite)。 覆盖: - /api/config GET 返回默认值、PUT 持久化; - /api/combos CRUD(name 唯一、原地更新保留 created_at、删除); - /api/combos/{id}/run 在策略缺失时提前 400(不等到后台才失败)。 """ from __future__ import annotations import pytest from app.api import deps from app.infrastructure.persistence.sqlalchemy.base import Base from app.main import app from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker @pytest.fixture() def client(tmp_path) -> TestClient: engine = create_engine(f"sqlite:///{tmp_path / 'combo_api.db'}", future=True) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine, expire_on_commit=False) def _session_override(): with Session() as s: yield s app.dependency_overrides[deps.get_session] = _session_override with TestClient(app) as c: yield c app.dependency_overrides.clear() class TestConfigApi: def test_get_returns_defaults(self, client: TestClient) -> None: cfg = client.get("/api/config").json() assert cfg["id"] == "default" assert cfg["price_adjustment"] == "hfq" assert cfg["min_commission"] == 5.0 def test_put_persists(self, client: TestClient) -> None: body = { "commission_rate": 0.00025, "stamp_tax_rate": 0.0005, "slippage_rate": 0.0008, "min_commission": 3.0, "price_adjustment": "qfq", "benchmark": "000905.SH", } resp = client.put("/api/config", json=body) assert resp.status_code == 200 got = client.get("/api/config").json() assert got["commission_rate"] == pytest.approx(0.00025) assert got["price_adjustment"] == "qfq" assert got["benchmark"] == "000905.SH" def test_put_rejects_unknown_field(self, client: TestClient) -> None: """拼错键名必须报错:否则「以为改了滑点、其实没生效」是静默降级。""" resp = client.put("/api/config", json={"slippage": 0.001}) assert resp.status_code == 422 assert "slippage" in resp.text # 值未被改动(仍是默认滑点) assert client.get("/api/config").json()["slippage_rate"] == pytest.approx(0.001) class TestCombosApi: def _make_strategy(self, client: TestClient, name: str) -> str: resp = client.post( "/api/strategies", json={"name": name, "factors": [{"name": "dividend_yield", "weight": 1}]}, ) assert resp.status_code == 200 return resp.json()["id"] def test_crud(self, client: TestClient) -> None: sid = self._make_strategy(client, "高股息") body = { "name": "组合A", "strategy_ids": [sid], "initial_capital": 500000, "hold_count": 10, "hold_min_days": 5, "hold_max_days": 30, "rebalance_freq": "weekly", "period": ["2024-01-01", "2024-06-01"], } created = client.post("/api/combos", json=body) assert created.status_code == 200 cid = created.json()["id"] assert cid.startswith("CMB-") assert created.json()["hold_max_days"] == 30 assert len(client.get("/api/combos").json()) == 1 detail = client.get(f"/api/combos/{cid}").json() assert detail["strategy_ids"] == [sid] assert detail["rebalance_freq"] == "weekly" # 原地更新保留 id 与 created_at created_at = detail["created_at"] upd = client.put(f"/api/combos/{cid}", json={**body, "name": "组合A", "hold_count": 15}) assert upd.status_code == 200 assert upd.json()["id"] == cid assert upd.json()["created_at"] == created_at assert upd.json()["hold_count"] == 15 assert client.delete(f"/api/combos/{cid}").status_code == 200 assert client.get(f"/api/combos/{cid}").status_code == 404 def test_duplicate_name_400(self, client: TestClient) -> None: sid = self._make_strategy(client, "S") body = {"name": "重名", "strategy_ids": [sid], "hold_count": 5, "rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]} assert client.post("/api/combos", json=body).status_code == 200 assert client.post("/api/combos", json=body).status_code == 400 def test_run_missing_strategy_400(self, client: TestClient) -> None: """引用不存在的策略 → 提交时即 400,而非等后台 Job 才失败。""" body = {"name": "缺策略组合", "strategy_ids": ["STG-NOT-EXIST"], "hold_count": 5, "rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]} resp = client.post("/api/combos/run", json=body) assert resp.status_code == 400 assert "不存在" in resp.json()["detail"] def test_unknown_param_rejected(self, client: TestClient) -> None: """回测参数写错键名(capital/hold_days)→ 422,而不是静默用默认值。""" sid = self._make_strategy(client, "S2") body = { "name": "拼错参数", "strategy_ids": [sid], "hold_count": 5, "rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"], "capital": 500000, "hold_days": 30, } resp = client.post("/api/combos", json=body) assert resp.status_code == 422 assert "capital" in resp.text and "hold_days" in resp.text assert client.get("/api/combos").json() == []