first commit
This commit is contained in:
0
backend/app/__init__.py
Normal file
0
backend/app/__init__.py
Normal file
167
backend/app/api.py
Normal file
167
backend/app/api.py
Normal file
@@ -0,0 +1,167 @@
|
||||
"""HTTP 路由(OpenAPI 契约的载体)。
|
||||
|
||||
GET /api/health 健康检查
|
||||
GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合)
|
||||
POST /api/backtest 跑回测,返回 K线+指标+买卖点+净值+绩效
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pandas as pd
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from .backtest.engine import BacktestConfig, run_backtest
|
||||
from .backtest.strategies import build_strategy
|
||||
from .data import fetcher, repository
|
||||
from .data.aggregation import bars_per_year, resample_bars
|
||||
from .data.synthetic import seed_if_empty
|
||||
from .db import get_session
|
||||
from .domain import Bar
|
||||
from .models import BacktestRun
|
||||
from .schemas import (
|
||||
BacktestRequest,
|
||||
BacktestResponse,
|
||||
CandleOut,
|
||||
EquityPoint,
|
||||
IndicatorOut,
|
||||
MetricsOut,
|
||||
SignalOut,
|
||||
SyncRequest,
|
||||
SyncResponse,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api")
|
||||
|
||||
|
||||
def _series_to_jsonable(s: pd.Series) -> list[float | None]:
|
||||
"""NaN -> None(lightweight-charts 的 whitespace data,跳过指标预热期)。"""
|
||||
out: list[float | None] = []
|
||||
for v in s.tolist():
|
||||
if v is None or (isinstance(v, float) and v != v):
|
||||
out.append(None)
|
||||
else:
|
||||
out.append(float(v))
|
||||
return out
|
||||
|
||||
|
||||
def _rows_to_bars(rows) -> list[Bar]:
|
||||
return [Bar(ts=r.ts, open=r.open, high=r.high, low=r.low, close=r.close, volume=r.volume) for r in rows]
|
||||
|
||||
|
||||
@router.get("/health")
|
||||
async def health() -> dict:
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@router.get("/candles/{symbol}", response_model=list[CandleOut])
|
||||
async def get_candles(
|
||||
symbol: str,
|
||||
timeframe: str = "1d",
|
||||
limit: int = 5000,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> list[CandleOut]:
|
||||
await seed_if_empty(session, symbol="DEMO")
|
||||
# 始终以日线为基底,再聚合到目标周期
|
||||
rows = await repository.get_candles(session, symbol, "1d", limit=limit)
|
||||
bars = resample_bars(_rows_to_bars(rows), timeframe)
|
||||
return [CandleOut(ts=b.ts, open=b.open, high=b.high, low=b.low, close=b.close, volume=b.volume) for b in bars]
|
||||
|
||||
|
||||
@router.post("/data/sync", response_model=SyncResponse)
|
||||
async def sync_data(req: SyncRequest, session: AsyncSession = Depends(get_session)) -> SyncResponse:
|
||||
"""主动拉取并缓存某标的的日线(Tushare 主 -> AKShare 兜底)。"""
|
||||
try:
|
||||
res = await fetcher.sync_symbol(
|
||||
session, req.symbol, start=req.start, end=req.end, source=req.source, force=req.force
|
||||
)
|
||||
return SyncResponse(**res)
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise HTTPException(status_code=502, detail=str(e))
|
||||
|
||||
|
||||
@router.post("/backtest", response_model=BacktestResponse)
|
||||
async def backtest(
|
||||
req: BacktestRequest,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> BacktestResponse:
|
||||
await seed_if_empty(session, symbol="DEMO")
|
||||
|
||||
# 非演示标的:首次自动拉取真实数据并缓存
|
||||
if req.symbol != "DEMO" and not await fetcher.is_cached(session, req.symbol):
|
||||
try:
|
||||
await fetcher.sync_symbol(session, req.symbol, source="auto")
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise HTTPException(status_code=502, detail=f"数据拉取失败: {e}")
|
||||
|
||||
# 日线为基底,聚合到请求周期
|
||||
rows = await repository.get_candles(
|
||||
session, req.symbol, "1d", start=req.start, end=req.end, limit=100000
|
||||
)
|
||||
if not rows:
|
||||
raise HTTPException(status_code=404, detail=f"无数据: symbol={req.symbol}")
|
||||
|
||||
bars = resample_bars(_rows_to_bars(rows), req.timeframe)
|
||||
if len(bars) < 2:
|
||||
raise HTTPException(status_code=400, detail=f"周期 {req.timeframe} 下数据不足,无法回测")
|
||||
|
||||
try:
|
||||
strategy = build_strategy(req.strategy, req.params)
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise HTTPException(status_code=400, detail=f"策略构建失败: {e}")
|
||||
cfg = BacktestConfig(
|
||||
initial_cash=req.initial_cash,
|
||||
fast_mode=req.fast_mode,
|
||||
bars_per_year=bars_per_year(req.timeframe),
|
||||
)
|
||||
result = run_backtest(bars, strategy, cfg)
|
||||
|
||||
df: pd.DataFrame = result["df"]
|
||||
m = result["metrics"]
|
||||
|
||||
# 记录到回测运行注册表(可复现/可审计的基础)
|
||||
session.add(
|
||||
BacktestRun(
|
||||
symbol=req.symbol,
|
||||
strategy=req.strategy,
|
||||
timeframe=req.timeframe,
|
||||
params_json=json.dumps(req.params, ensure_ascii=False),
|
||||
initial_cash=req.initial_cash,
|
||||
total_return=m["total_return"],
|
||||
max_drawdown=m["max_drawdown"],
|
||||
sharpe=m["sharpe"],
|
||||
num_trades=m["num_trades"],
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
candles = [
|
||||
CandleOut(ts=r["ts"], open=r["open"], high=r["high"], low=r["low"],
|
||||
close=r["close"], volume=r["volume"])
|
||||
for _, r in df.iterrows()
|
||||
]
|
||||
signals = [
|
||||
SignalOut(ts=f.ts, side=f.side.value, price=f.price, qty=f.qty)
|
||||
for f in result["fills"]
|
||||
]
|
||||
indicators = IndicatorOut(
|
||||
strategy=req.strategy,
|
||||
data={col: _series_to_jsonable(df[col]) for col in result["indicator_cols"]},
|
||||
)
|
||||
equity = [EquityPoint(ts=t.to_pydatetime(), value=float(v))
|
||||
for t, v in result["equity"].items()]
|
||||
|
||||
return BacktestResponse(
|
||||
symbol=req.symbol,
|
||||
timeframe=req.timeframe,
|
||||
strategy=req.strategy,
|
||||
candles=candles,
|
||||
indicators=indicators,
|
||||
signals=signals,
|
||||
equity=equity,
|
||||
metrics=MetricsOut(**m),
|
||||
final_cash=result["final_cash"],
|
||||
final_position=result["final_position"],
|
||||
initial_cash=req.initial_cash,
|
||||
)
|
||||
0
backend/app/backtest/__init__.py
Normal file
0
backend/app/backtest/__init__.py
Normal file
104
backend/app/backtest/broker.py
Normal file
104
backend/app/backtest/broker.py
Normal file
@@ -0,0 +1,104 @@
|
||||
"""PaperBroker —— 回测中的虚拟撮合 / 账户。
|
||||
|
||||
建模 A 股规则:
|
||||
- 100 股整数手(1手=100股)
|
||||
- T+1:当日买入次日才可卖(fast_mode 关闭此约束)
|
||||
- 印花税(卖出)、过户费(双边)、佣金(万1 最低5元)、滑点
|
||||
- 撮合价:以当根 bar 收盘价近似(阶段1 接 VWAP / 限价单)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from ..commission import CostSchedule, DEFAULT, buy_cost, sell_cost
|
||||
from ..domain import Fill, Side
|
||||
|
||||
LOT = 100 # A股 1 手 = 100 股
|
||||
|
||||
|
||||
@dataclass
|
||||
class PaperBroker:
|
||||
initial_cash: float = 100000.0
|
||||
schedule: CostSchedule = field(default_factory=lambda: DEFAULT)
|
||||
enable_costs: bool = True
|
||||
enable_t_plus_1: bool = True
|
||||
|
||||
cash: float = field(init=False)
|
||||
holdings: float = 0.0 # 可卖数量
|
||||
locked: float = 0.0 # 当日买入(T+1 锁定)
|
||||
avg_price: float = 0.0
|
||||
fills: list[Fill] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.cash = self.initial_cash
|
||||
|
||||
@property
|
||||
def position(self) -> float:
|
||||
return self.holdings + self.locked
|
||||
|
||||
def equity(self, price: float) -> float:
|
||||
return self.cash + self.position * price
|
||||
|
||||
@staticmethod
|
||||
def _to_lots(qty: float) -> int:
|
||||
return int(qty // LOT) * LOT
|
||||
|
||||
def buy_max(self, ts, price: float) -> Fill | None:
|
||||
"""用当前现金买尽可能多的整手。"""
|
||||
if price <= 0:
|
||||
return None
|
||||
rate = (
|
||||
self.schedule.commission_rate
|
||||
+ self.schedule.transfer_fee_rate
|
||||
+ self.schedule.slippage_rate
|
||||
) if self.enable_costs else 0.0
|
||||
affordable_qty = self.cash / (price * (1 + rate))
|
||||
qty = self._to_lots(affordable_qty)
|
||||
if qty <= 0:
|
||||
return None
|
||||
return self._execute_buy(ts, price, qty)
|
||||
|
||||
def sell_all(self, ts, price: float) -> Fill | None:
|
||||
"""卖出全部可卖持仓(整手)。"""
|
||||
qty = self._to_lots(self.holdings)
|
||||
if qty <= 0:
|
||||
return None
|
||||
return self._execute_sell(ts, price, qty)
|
||||
|
||||
def _execute_buy(self, ts, price: float, qty: int) -> Fill:
|
||||
if self.enable_costs:
|
||||
fp, comm, tf = buy_cost(price, qty, self.schedule)
|
||||
else:
|
||||
fp, comm, tf = price, 0.0, 0.0
|
||||
cost = fp * qty + comm + tf
|
||||
prev_pos = self.position
|
||||
new_pos = prev_pos + qty
|
||||
self.avg_price = (self.avg_price * prev_pos + fp * qty) / new_pos if new_pos else 0.0
|
||||
if self.enable_t_plus_1:
|
||||
self.locked += qty
|
||||
else:
|
||||
self.holdings += qty
|
||||
self.cash -= cost
|
||||
f = Fill(ts=ts, side=Side.BUY, price=fp, qty=qty, commission=comm, transfer_fee=tf)
|
||||
self.fills.append(f)
|
||||
return f
|
||||
|
||||
def _execute_sell(self, ts, price: float, qty: int) -> Fill:
|
||||
if self.enable_costs:
|
||||
fp, comm, tf, sd = sell_cost(price, qty, self.schedule)
|
||||
else:
|
||||
fp, comm, tf, sd = price, 0.0, 0.0, 0.0
|
||||
proceeds = fp * qty - comm - tf - sd
|
||||
self.holdings -= qty
|
||||
self.cash += proceeds
|
||||
if self.position == 0:
|
||||
self.avg_price = 0.0
|
||||
f = Fill(ts=ts, side=Side.SELL, price=fp, qty=qty, commission=comm, stamp_duty=sd, transfer_fee=tf)
|
||||
self.fills.append(f)
|
||||
return f
|
||||
|
||||
def release_t_plus_1(self) -> None:
|
||||
"""每根 bar 结束时调用:当日锁定转为次日可卖。"""
|
||||
if self.enable_t_plus_1:
|
||||
self.holdings += self.locked
|
||||
self.locked = 0.0
|
||||
114
backend/app/backtest/engine.py
Normal file
114
backend/app/backtest/engine.py
Normal file
@@ -0,0 +1,114 @@
|
||||
"""单一回测引擎(fast/strict 两档,同一代码路径、同一撮合逻辑)。
|
||||
|
||||
不做"快层 vectorbt + 真层自研"双引擎——用户必然信任更快的那层,
|
||||
一旦快层省略 T+1/费用,两套结论背离即反复扯皮"以谁为准"。
|
||||
这里用开关控制规则是否注入,引擎只有一份。
|
||||
|
||||
数据流:
|
||||
bars -> DataFrame -> strategy.compute(指标) -> 逐 bar 喂 strategy.on_bar(broker)
|
||||
-> broker 产生 fills -> 每根 bar 记录 equity -> metrics
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from ..domain import Bar
|
||||
from .broker import PaperBroker
|
||||
from .metrics import compute_metrics
|
||||
|
||||
|
||||
@dataclass
|
||||
class BacktestConfig:
|
||||
initial_cash: float = 1000000.0
|
||||
fast_mode: bool = False # True: 关 T+1/费用,交互试探
|
||||
bars_per_year: int = 252 # 日线 252;分钟级另算
|
||||
|
||||
|
||||
def run_backtest(bars: list[Bar], strategy, cfg: BacktestConfig | None = None) -> dict:
|
||||
cfg = cfg or BacktestConfig()
|
||||
if not bars:
|
||||
return _empty_result()
|
||||
|
||||
df = pd.DataFrame(
|
||||
[{"ts": b.ts, "open": b.open, "high": b.high, "low": b.low,
|
||||
"close": b.close, "volume": b.volume} for b in bars]
|
||||
).sort_values("ts").reset_index(drop=True)
|
||||
|
||||
# 指标预计算(由策略持有,单一事实源)
|
||||
ind = strategy.compute(df["close"], df["high"], df["low"])
|
||||
indicator_cols = list(ind.columns)
|
||||
df = pd.concat([df, ind], axis=1)
|
||||
|
||||
broker = PaperBroker(
|
||||
initial_cash=cfg.initial_cash,
|
||||
enable_costs=not cfg.fast_mode,
|
||||
enable_t_plus_1=not cfg.fast_mode,
|
||||
)
|
||||
|
||||
equity_values = []
|
||||
for i in range(len(df)):
|
||||
row = df.iloc[i]
|
||||
strategy.on_bar(i, row, broker)
|
||||
equity_values.append(broker.equity(row["close"]))
|
||||
broker.release_t_plus_1()
|
||||
|
||||
equity = pd.Series(equity_values, index=df["ts"], name="equity")
|
||||
metrics = compute_metrics(equity, cfg.bars_per_year)
|
||||
metrics["num_trades"] = len(broker.fills)
|
||||
metrics["win_rate"] = _win_rate(broker.fills)
|
||||
|
||||
return {
|
||||
"df": df,
|
||||
"indicator_cols": indicator_cols,
|
||||
"fills": broker.fills,
|
||||
"equity": equity,
|
||||
"metrics": metrics,
|
||||
"final_cash": broker.cash,
|
||||
"final_position": broker.position,
|
||||
}
|
||||
|
||||
|
||||
def _win_rate(fills) -> float:
|
||||
"""FIFO 配对计算卖出胜率。"""
|
||||
from collections import deque
|
||||
buys: deque = deque()
|
||||
wins = 0
|
||||
total_sells = 0
|
||||
for f in fills:
|
||||
if f.side.value == "buy":
|
||||
buys.append((f.price, f.qty))
|
||||
else: # sell
|
||||
remaining = f.qty
|
||||
total_sells += 1
|
||||
profitable = True
|
||||
while remaining > 0 and buys:
|
||||
buy_price, buy_qty = buys[0]
|
||||
if f.price >= buy_price:
|
||||
pass
|
||||
else:
|
||||
profitable = False
|
||||
take = min(remaining, buy_qty)
|
||||
buy_qty -= take
|
||||
remaining -= take
|
||||
if buy_qty <= 0:
|
||||
buys.popleft()
|
||||
else:
|
||||
buys[0] = (buy_price, buy_qty)
|
||||
if profitable and f.qty > 0:
|
||||
wins += 1
|
||||
return wins / total_sells if total_sells else 0.0
|
||||
|
||||
|
||||
def _empty_result() -> dict:
|
||||
return {
|
||||
"df": pd.DataFrame(),
|
||||
"indicator_cols": [],
|
||||
"fills": [],
|
||||
"equity": pd.Series(dtype=float),
|
||||
"metrics": {"total_return": 0.0, "max_drawdown": 0.0, "sharpe": 0.0,
|
||||
"volatility": 0.0, "num_trades": 0, "win_rate": 0.0},
|
||||
"final_cash": 0.0,
|
||||
"final_position": 0.0,
|
||||
}
|
||||
41
backend/app/backtest/metrics.py
Normal file
41
backend/app/backtest/metrics.py
Normal file
@@ -0,0 +1,41 @@
|
||||
"""绩效统计(阶段1 补基准归因:超额/信息比率/beta/alpha)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def compute_metrics(equity: pd.Series, bars_per_year: int = 252) -> dict:
|
||||
equity = equity.dropna()
|
||||
if len(equity) < 2 or equity.iloc[0] == 0:
|
||||
return {"total_return": 0.0, "max_drawdown": 0.0, "sharpe": 0.0,
|
||||
"volatility": 0.0, "win_rate": 0.0}
|
||||
|
||||
total_return = float(equity.iloc[-1] / equity.iloc[0] - 1)
|
||||
|
||||
returns = equity.pct_change().dropna()
|
||||
cummax = equity.cummax()
|
||||
drawdown = (equity - cummax) / cummax
|
||||
max_drawdown = float(abs(drawdown.min()))
|
||||
|
||||
std = float(returns.std())
|
||||
sharpe = float(returns.mean() / std * np.sqrt(bars_per_year)) if std > 0 else 0.0
|
||||
volatility = std * np.sqrt(bars_per_year)
|
||||
|
||||
return {
|
||||
"total_return": total_return,
|
||||
"max_drawdown": max_drawdown,
|
||||
"sharpe": sharpe,
|
||||
"volatility": float(volatility),
|
||||
"win_rate": 0.0, # 由 engine 用成交对计算后注入
|
||||
}
|
||||
|
||||
|
||||
def win_rate_from_fills(fills) -> float:
|
||||
"""按"卖出-对应买入"配对估算胜率(粗略,阶段1 用 FIFO 精确配对)。"""
|
||||
sells = [f for f in fills if f.side.value == "sell"]
|
||||
if not sells:
|
||||
return 0.0
|
||||
wins = sum(1 for f in fills if f.side.value == "sell" and f.price > 0)
|
||||
# 简化:有成交即计;真实胜率需配对,这里先返回 0 占位,由 engine 精算
|
||||
return 0.0
|
||||
24
backend/app/backtest/strategies/__init__.py
Normal file
24
backend/app/backtest/strategies/__init__.py
Normal file
@@ -0,0 +1,24 @@
|
||||
"""策略注册表与工厂。
|
||||
|
||||
新增策略:实现 Strategy(base.py),在此注册 {name: Class},前端下拉即可选。
|
||||
策略构造参数由请求的 params(dict) 以 **kwargs 传入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from .base import Strategy
|
||||
from .macd_cross import MACDCrossStrategy
|
||||
from .ma_strategies import MACrossStrategy, SingleMAStrategy
|
||||
|
||||
STRATEGIES: dict[str, type[Strategy]] = {
|
||||
"macd_cross": MACDCrossStrategy,
|
||||
"ma_cross": MACrossStrategy,
|
||||
"single_ma": SingleMAStrategy,
|
||||
}
|
||||
|
||||
|
||||
def build_strategy(name: str, params: dict | None) -> Strategy:
|
||||
cls = STRATEGIES.get(name)
|
||||
if cls is None:
|
||||
raise ValueError(f"未知策略: {name}(可用: {', '.join(STRATEGIES)})")
|
||||
kwargs = {k: v for k, v in (params or {}).items()}
|
||||
return cls(**kwargs)
|
||||
20
backend/app/backtest/strategies/base.py
Normal file
20
backend/app/backtest/strategies/base.py
Normal file
@@ -0,0 +1,20 @@
|
||||
"""策略抽象基类。策略只产生买卖意图(向 broker 下单),不负责撮合/费用。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from ..broker import PaperBroker
|
||||
|
||||
|
||||
class Strategy(ABC):
|
||||
"""compute 预算指标;on_bar 逐 bar 决策并向 broker 下单。"""
|
||||
|
||||
@abstractmethod
|
||||
def compute(self, close: pd.Series, high: pd.Series, low: pd.Series) -> pd.DataFrame:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def on_bar(self, i: int, row: pd.Series, broker: PaperBroker) -> None:
|
||||
...
|
||||
74
backend/app/backtest/strategies/ma_strategies.py
Normal file
74
backend/app/backtest/strategies/ma_strategies.py
Normal file
@@ -0,0 +1,74 @@
|
||||
"""均线类策略:双均线交叉、单均线(价格上穿/下穿)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from ...indicators import ma
|
||||
from ..broker import PaperBroker
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
class MACrossStrategy(Strategy):
|
||||
"""双均线交叉:快线上穿慢线买入,下穿卖出(金叉/死叉)。"""
|
||||
|
||||
def __init__(self, fast: float = 5, slow: float = 20):
|
||||
self.fast = int(fast)
|
||||
self.slow = int(slow)
|
||||
self._ind: pd.DataFrame | None = None
|
||||
self._prev_fast: float | None = None
|
||||
self._prev_slow: float | None = None
|
||||
|
||||
def compute(self, close: pd.Series, high: pd.Series, low: pd.Series) -> pd.DataFrame:
|
||||
self._ind = pd.DataFrame({"fast": ma(close, self.fast), "slow": ma(close, self.slow)})
|
||||
return self._ind
|
||||
|
||||
def on_bar(self, i: int, row: pd.Series, broker: PaperBroker) -> None:
|
||||
f = float(self._ind["fast"].iloc[i])
|
||||
s = float(self._ind["slow"].iloc[i])
|
||||
if self._prev_fast is None:
|
||||
self._prev_fast, self._prev_slow = f, s
|
||||
return
|
||||
|
||||
golden = self._prev_fast <= self._prev_slow and f > s
|
||||
death = self._prev_fast >= self._prev_slow and f < s
|
||||
|
||||
price = float(row["close"])
|
||||
ts = row["ts"]
|
||||
if golden and pd.notna(f):
|
||||
broker.buy_max(ts, price)
|
||||
elif death and broker.position > 0 and pd.notna(f):
|
||||
broker.sell_all(ts, price)
|
||||
|
||||
self._prev_fast, self._prev_slow = f, s
|
||||
|
||||
|
||||
class SingleMAStrategy(Strategy):
|
||||
"""单均线:收盘价上穿均线买入,下穿均线卖出。"""
|
||||
|
||||
def __init__(self, period: float = 20):
|
||||
self.period = int(period)
|
||||
self._ind: pd.DataFrame | None = None
|
||||
self._prev_above: bool | None = None
|
||||
|
||||
def compute(self, close: pd.Series, high: pd.Series, low: pd.Series) -> pd.DataFrame:
|
||||
self._ind = pd.DataFrame({"ma": ma(close, self.period)})
|
||||
return self._ind
|
||||
|
||||
def on_bar(self, i: int, row: pd.Series, broker: PaperBroker) -> None:
|
||||
m = float(self._ind["ma"].iloc[i])
|
||||
price = float(row["close"])
|
||||
ts = row["ts"]
|
||||
above = price > m
|
||||
if self._prev_above is None or pd.isna(m):
|
||||
self._prev_above = above
|
||||
return
|
||||
|
||||
cross_up = above and not self._prev_above # 上穿
|
||||
cross_down = (not above) and self._prev_above # 下穿
|
||||
|
||||
if cross_up:
|
||||
broker.buy_max(ts, price)
|
||||
elif cross_down and broker.position > 0:
|
||||
broker.sell_all(ts, price)
|
||||
|
||||
self._prev_above = above
|
||||
42
backend/app/backtest/strategies/macd_cross.py
Normal file
42
backend/app/backtest/strategies/macd_cross.py
Normal file
@@ -0,0 +1,42 @@
|
||||
"""MACD 金叉死叉策略。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from ...indicators import macd
|
||||
from ..broker import PaperBroker
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
class MACDCrossStrategy(Strategy):
|
||||
def __init__(self, fast: float = 12, slow: float = 26, signal: float = 9):
|
||||
self.fast = int(fast)
|
||||
self.slow = int(slow)
|
||||
self.signal = int(signal)
|
||||
self._ind: pd.DataFrame | None = None
|
||||
self._prev_dif: float | None = None
|
||||
self._prev_dea: float | None = None
|
||||
|
||||
def compute(self, close: pd.Series, high: pd.Series, low: pd.Series) -> pd.DataFrame:
|
||||
self._ind = macd(close, self.fast, self.slow, self.signal)
|
||||
return self._ind
|
||||
|
||||
def on_bar(self, i: int, row: pd.Series, broker: PaperBroker) -> None:
|
||||
dif = float(self._ind["macd"].iloc[i])
|
||||
dea = float(self._ind["signal"].iloc[i])
|
||||
|
||||
if self._prev_dif is None:
|
||||
self._prev_dif, self._prev_dea = dif, dea
|
||||
return
|
||||
|
||||
golden = self._prev_dif <= self._prev_dea and dif > dea # 金叉
|
||||
death = self._prev_dif >= self._prev_dea and dif < dea # 死叉
|
||||
|
||||
price = float(row["close"])
|
||||
ts = row["ts"]
|
||||
if golden and pd.notna(dif):
|
||||
broker.buy_max(ts, price)
|
||||
elif death and broker.position > 0 and pd.notna(dif):
|
||||
broker.sell_all(ts, price)
|
||||
|
||||
self._prev_dif, self._prev_dea = dif, dea
|
||||
47
backend/app/commission.py
Normal file
47
backend/app/commission.py
Normal file
@@ -0,0 +1,47 @@
|
||||
"""A股交易成本(基准日 2026-08,已修正历史错误)。
|
||||
|
||||
⚠️ 重要:费率必须是"参数表 + 生效日期版本化 + 显式基准日",不能硬编码成常量。
|
||||
费率会再变(如 2023-08-28 印花税减半、2022 过户费下调)。
|
||||
MVP 用单一 CostSchedule;阶段1 扩展为按生效日期区间查找的多版本表。
|
||||
|
||||
参考(已核实):
|
||||
印花税 0.05% 单边卖出 —— 2023-08-28 财政部减半(原 0.1%)
|
||||
过户费 0.001% 沪深双边 —— 2022 年统一下调(原沪市万 0.2 单边)
|
||||
佣金 万1 含规费,最低 5 元 —— 2026 主流
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .config import settings
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CostSchedule:
|
||||
stamp_duty_rate: float = settings.stamp_duty_rate # 印花税,卖出
|
||||
transfer_fee_rate: float = settings.transfer_fee_rate # 过户费,双边
|
||||
commission_rate: float = settings.commission_rate # 佣金
|
||||
commission_min: float = settings.commission_min # 最低佣金
|
||||
slippage_rate: float = settings.slippage_rate # 滑点(价格比例近似)
|
||||
|
||||
|
||||
DEFAULT = CostSchedule()
|
||||
|
||||
|
||||
def buy_cost(price: float, qty: float, sch: CostSchedule = DEFAULT) -> tuple[float, float, float]:
|
||||
"""买入成本。返回 (成交价, 佣金, 过户费)。买入无印花税。"""
|
||||
fill_price = price * (1 + sch.slippage_rate)
|
||||
gross = fill_price * qty
|
||||
commission = max(gross * sch.commission_rate, sch.commission_min)
|
||||
transfer_fee = gross * sch.transfer_fee_rate
|
||||
return fill_price, commission, transfer_fee
|
||||
|
||||
|
||||
def sell_cost(price: float, qty: float, sch: CostSchedule = DEFAULT) -> tuple[float, float, float, float]:
|
||||
"""卖出成本。返回 (成交价, 佣金, 过户费, 印花税)。"""
|
||||
fill_price = price * (1 - sch.slippage_rate)
|
||||
gross = fill_price * qty
|
||||
commission = max(gross * sch.commission_rate, sch.commission_min)
|
||||
transfer_fee = gross * sch.transfer_fee_rate
|
||||
stamp_duty = gross * sch.stamp_duty_rate
|
||||
return fill_price, commission, transfer_fee, stamp_duty
|
||||
26
backend/app/config.py
Normal file
26
backend/app/config.py
Normal file
@@ -0,0 +1,26 @@
|
||||
"""应用配置(pydantic-settings)。可由 .env / 环境变量覆盖。"""
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
|
||||
|
||||
app_name: str = "Stock Backtest"
|
||||
|
||||
# 默认 SQLite 零配置;切 Postgres/TimescaleDB 只改这一行
|
||||
database_url: str = "sqlite+aiosqlite:///./stock.db"
|
||||
|
||||
# 真实数据源
|
||||
tushare_token: str = "" # Tushare Pro token(主数据源)
|
||||
data_adjust: str = "qfq" # 复权:qfq 前复权 / hfq 后复权 / "" 不复权
|
||||
data_default_start: str = "20200101" # 默认拉取起点(约近 5 年)
|
||||
|
||||
# A股交易成本(基准日 2026-08)——做成可配置参数,便于将来按生效日期版本化
|
||||
stamp_duty_rate: float = 0.0005 # 印花税 0.05%,单边卖出(2023-08-28 减半)
|
||||
transfer_fee_rate: float = 0.00001 # 过户费 0.001%,沪深双边(2022 调整)
|
||||
commission_rate: float = 0.0001 # 佣金 万1(含规费)
|
||||
commission_min: float = 5.0 # 最低 5 元
|
||||
slippage_rate: float = 0.0005 # 滑点近似(按价格比例)
|
||||
|
||||
|
||||
settings = Settings()
|
||||
0
backend/app/data/__init__.py
Normal file
0
backend/app/data/__init__.py
Normal file
54
backend/app/data/aggregation.py
Normal file
54
backend/app/data/aggregation.py
Normal file
@@ -0,0 +1,54 @@
|
||||
"""K 线周期聚合:日线 -> 周/月/年。
|
||||
|
||||
生产环境用 TimescaleDB Continuous Aggregates 在库里预物化(性能);
|
||||
MVP 在应用层用 pandas resample 即可,逻辑等价、便于切换。
|
||||
|
||||
OHLCV 聚合规则:开=周期内首根开、高=最高、低=最低、收=末根收、量=求和。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from ..domain import Bar
|
||||
|
||||
# pandas resample 规则(周一为周首;月/年以首日对齐)
|
||||
_RULES = {"1w": "W-MON", "1M": "MS", "1y": "YS"}
|
||||
|
||||
# 各周期的"年交易日数"(用于夏普等指标的年化)
|
||||
_BARS_PER_YEAR = {"1d": 252, "1w": 52, "1M": 12, "1y": 1}
|
||||
|
||||
|
||||
def bars_per_year(timeframe: str) -> int:
|
||||
return _BARS_PER_YEAR.get(timeframe, 252)
|
||||
|
||||
|
||||
def resample_bars(bars: list[Bar], timeframe: str) -> list[Bar]:
|
||||
"""把日线 bars 聚合为目标周期;日线或未知周期原样返回。"""
|
||||
if not bars or timeframe in ("1d", "d", "day", "", None):
|
||||
return bars
|
||||
rule = _RULES.get(timeframe)
|
||||
if rule is None:
|
||||
return bars
|
||||
|
||||
df = pd.DataFrame(
|
||||
[{"ts": b.ts, "open": b.open, "high": b.high, "low": b.low, "close": b.close, "volume": b.volume}
|
||||
for b in bars]
|
||||
).set_index("ts").sort_index()
|
||||
|
||||
agg = (
|
||||
df.resample(rule)
|
||||
.agg({"open": "first", "high": "max", "low": "min", "close": "last", "volume": "sum"})
|
||||
.dropna()
|
||||
)
|
||||
|
||||
return [
|
||||
Bar(
|
||||
ts=ts.to_pydatetime(),
|
||||
open=float(row["open"]),
|
||||
high=float(row["high"]),
|
||||
low=float(row["low"]),
|
||||
close=float(row["close"]),
|
||||
volume=float(row["volume"]),
|
||||
)
|
||||
for ts, row in agg.iterrows()
|
||||
]
|
||||
38
backend/app/data/akshare_provider.py
Normal file
38
backend/app/data/akshare_provider.py
Normal file
@@ -0,0 +1,38 @@
|
||||
"""AKShare 数据源(兜底/校验)。免费、无需 token。
|
||||
|
||||
默认不安装(依赖较重);如需启用:`uv add akshare`。
|
||||
fetcher 在 Tushare 失败时会尝试本模块;未安装则该路径自动跳过。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from ..domain import Bar
|
||||
from .symbols import plain_code
|
||||
|
||||
|
||||
def fetch_daily(code: str, start: str = "20200101", end: str | None = None,
|
||||
adjust: str = "qfq") -> list[Bar]:
|
||||
import akshare as ak # 延迟导入
|
||||
|
||||
end = end or datetime.now().strftime("%Y%m%d")
|
||||
symbol = plain_code(code)
|
||||
adj_map = {"qfq": "qfq", "hfq": "hfq", "": "", None: ""}
|
||||
df = ak.stock_zh_a_hist(
|
||||
symbol=symbol, period="daily",
|
||||
start_date=start, end_date=end, adjust=adj_map.get(adjust, ""),
|
||||
)
|
||||
if df is None or df.empty:
|
||||
raise RuntimeError(f"AKShare 无数据: {symbol}")
|
||||
|
||||
bars: list[Bar] = []
|
||||
for _, r in df.iterrows():
|
||||
bars.append(
|
||||
Bar(
|
||||
ts=datetime.strptime(str(r["日期"]), "%Y-%m-%d"),
|
||||
open=float(r["开盘"]), high=float(r["最高"]),
|
||||
low=float(r["最低"]), close=float(r["收盘"]),
|
||||
volume=float(r["成交量"]) * 100.0, # AKShare 成交量单位为手 -> 股
|
||||
)
|
||||
)
|
||||
return bars
|
||||
81
backend/app/data/fetcher.py
Normal file
81
backend/app/data/fetcher.py
Normal file
@@ -0,0 +1,81 @@
|
||||
"""数据编排:拉取(Tushare 主 -> AKShare 兜底)+ 本地缓存。
|
||||
|
||||
真实行情落库到 candles 表(timeframe='1d'),回测统一从库读,与 DEMO 同路径。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from sqlalchemy import delete, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ..config import settings
|
||||
from ..domain import Bar
|
||||
from ..models import Candle
|
||||
from . import akshare_provider, tushare_provider
|
||||
|
||||
DEFAULT_START = settings.data_default_start or "20200101"
|
||||
|
||||
|
||||
def _providers(source: str):
|
||||
"""按优先级返回 (名称, 同步拉取函数) 列表。"""
|
||||
seq = []
|
||||
if source in ("auto", "tushare") and settings.tushare_token:
|
||||
seq.append(("tushare", tushare_provider.fetch_daily))
|
||||
if source in ("auto", "akshare"):
|
||||
seq.append(("akshare", akshare_provider.fetch_daily))
|
||||
return seq
|
||||
|
||||
|
||||
async def count_cached(session: AsyncSession, symbol: str) -> int:
|
||||
res = await session.execute(
|
||||
select(func.count()).select_from(Candle).where(
|
||||
Candle.symbol == symbol, Candle.timeframe == "1d"
|
||||
)
|
||||
)
|
||||
return int(res.scalar() or 0)
|
||||
|
||||
|
||||
async def is_cached(session: AsyncSession, symbol: str) -> bool:
|
||||
return await count_cached(session, symbol) > 0
|
||||
|
||||
|
||||
async def sync_symbol(
|
||||
session: AsyncSession,
|
||||
code: str,
|
||||
start: str | None = None,
|
||||
end: str | None = None,
|
||||
source: str = "auto",
|
||||
force: bool = False,
|
||||
) -> dict:
|
||||
"""拉取并缓存某标的日线。已缓存且非 force 时直接返回缓存计数。"""
|
||||
if not force and await is_cached(session, code):
|
||||
return {"symbol": code, "bars": await count_cached(session, code), "source": "cache"}
|
||||
|
||||
start = start or DEFAULT_START
|
||||
adjust = settings.data_adjust
|
||||
errors: list[str] = []
|
||||
bars: list[Bar] = []
|
||||
used = None
|
||||
|
||||
for name, fn in _providers(source):
|
||||
try:
|
||||
# tushare/akshare 是同步网络 IO,丢到线程池避免阻塞事件循环
|
||||
bars = await asyncio.to_thread(fn, code, start, end, adjust)
|
||||
used = name
|
||||
break
|
||||
except Exception as e: # noqa: BLE001
|
||||
errors.append(f"{name}: {e}")
|
||||
|
||||
if not bars:
|
||||
raise RuntimeError("所有数据源均失败 -> " + " | ".join(errors) if errors else "无可用数据源")
|
||||
|
||||
# 全量替换该标的日线(避免重复主键)
|
||||
await session.execute(delete(Candle).where(Candle.symbol == code, Candle.timeframe == "1d"))
|
||||
for b in bars:
|
||||
session.add(
|
||||
Candle(symbol=code, timeframe="1d", ts=b.ts, open=b.open, high=b.high,
|
||||
low=b.low, close=b.close, volume=b.volume)
|
||||
)
|
||||
await session.commit()
|
||||
return {"symbol": code, "bars": len(bars), "source": used}
|
||||
34
backend/app/data/repository.py
Normal file
34
backend/app/data/repository.py
Normal file
@@ -0,0 +1,34 @@
|
||||
"""K 线数据访问(从库读)。
|
||||
|
||||
写入由 DataProvider 适配器负责(阶段1 接 Tushare/AKShare)。
|
||||
MVP 的数据由 synthetic.seed_if_empty 灌入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ..models import Candle
|
||||
|
||||
|
||||
async def get_candles(
|
||||
session: AsyncSession,
|
||||
symbol: str,
|
||||
timeframe: str = "1d",
|
||||
start: datetime | None = None,
|
||||
end: datetime | None = None,
|
||||
limit: int = 5000,
|
||||
) -> list[Candle]:
|
||||
stmt = select(Candle).where(
|
||||
Candle.symbol == symbol,
|
||||
Candle.timeframe == timeframe,
|
||||
)
|
||||
if start is not None:
|
||||
stmt = stmt.where(Candle.ts >= start)
|
||||
if end is not None:
|
||||
stmt = stmt.where(Candle.ts <= end)
|
||||
stmt = stmt.order_by(Candle.ts.asc()).limit(limit)
|
||||
result = await session.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
25
backend/app/data/symbols.py
Normal file
25
backend/app/data/symbols.py
Normal file
@@ -0,0 +1,25 @@
|
||||
"""A 股代码归一化。支持 6 位纯数字或带交易所后缀(000001 / 000001.SZ)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def plain_code(code: str) -> str:
|
||||
"""000001.SZ -> 000001"""
|
||||
return code.strip().upper().split(".")[0]
|
||||
|
||||
|
||||
def to_ts_code(code: str) -> str:
|
||||
"""转 Tushare ts_code(带交易所后缀)。"""
|
||||
c = code.strip().upper()
|
||||
if "." in c:
|
||||
return c
|
||||
c = plain_code(c)
|
||||
# 沪市:60xxxx 主板、68xxxx 科创、9xxxxx B 股
|
||||
if c.startswith(("60", "68", "9")):
|
||||
return c + ".SH"
|
||||
# 深市:00xxxx 主板/中小、30xxxx 创业、20xxxx B 股
|
||||
if c.startswith(("00", "30", "20")):
|
||||
return c + ".SZ"
|
||||
# 北交所:8xxxxx / 4xxxxx
|
||||
if c.startswith(("8", "4")):
|
||||
return c + ".BJ"
|
||||
return c + ".SZ"
|
||||
78
backend/app/data/synthetic.py
Normal file
78
backend/app/data/synthetic.py
Normal file
@@ -0,0 +1,78 @@
|
||||
"""合成数据(MVP 零依赖可跑)。
|
||||
|
||||
生成随机游走 OHLCV,灌入 DB。仅用于让回测链路在没有真实数据源时也能跑通演示。
|
||||
阶段1 接 Tushare/AKShare 后,这里仅保留为"离线测试夹具"。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import numpy as np
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ..domain import Bar
|
||||
from ..models import Candle
|
||||
|
||||
|
||||
def _trading_days(n: int) -> list[datetime]:
|
||||
"""粗略生成 n 个工作日(跳过周末;节假日由阶段1 的交易日历服务处理)。"""
|
||||
start = datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(days=int(n * 1.6))
|
||||
days: list[datetime] = []
|
||||
d = start
|
||||
while len(days) < n:
|
||||
if d.weekday() < 5:
|
||||
days.append(d.replace(hour=15, minute=0, second=0, microsecond=0))
|
||||
d += timedelta(days=1)
|
||||
return days
|
||||
|
||||
|
||||
def generate_ohlcv(n: int = 500, seed: int = 42) -> list[Bar]:
|
||||
"""随机游走 + A 股风格的价格区间(5~30 元)。"""
|
||||
rng = np.random.default_rng(seed)
|
||||
rets = rng.normal(loc=0.0003, scale=0.018, size=n)
|
||||
price = 10.0 * np.cumprod(1 + rets)
|
||||
days = _trading_days(n)
|
||||
bars: list[Bar] = []
|
||||
for i in range(n):
|
||||
close = float(price[i])
|
||||
op = close * (1 + rng.normal(0, 0.005))
|
||||
hi = max(op, close) * (1 + abs(rng.normal(0, 0.006)))
|
||||
lo = min(op, close) * (1 - abs(rng.normal(0, 0.006)))
|
||||
vol = float(rng.integers(1_000_000, 10_000_000))
|
||||
bars.append(
|
||||
Bar(
|
||||
ts=days[i],
|
||||
open=round(op, 2),
|
||||
high=round(hi, 2),
|
||||
low=round(lo, 2),
|
||||
close=round(close, 2),
|
||||
volume=vol,
|
||||
)
|
||||
)
|
||||
return bars
|
||||
|
||||
|
||||
async def seed_if_empty(session: AsyncSession, symbol: str = "DEMO", n: int = 500) -> None:
|
||||
"""若库中无该 symbol 数据,则灌入合成数据。"""
|
||||
existing = await session.execute(
|
||||
select(Candle.id).where(Candle.symbol == symbol).limit(1)
|
||||
)
|
||||
if existing.scalars().first() is not None:
|
||||
return
|
||||
|
||||
bars = generate_ohlcv(n=n)
|
||||
for b in bars:
|
||||
session.add(
|
||||
Candle(
|
||||
symbol=symbol,
|
||||
timeframe="1d",
|
||||
ts=b.ts,
|
||||
open=b.open,
|
||||
high=b.high,
|
||||
low=b.low,
|
||||
close=b.close,
|
||||
volume=b.volume,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
51
backend/app/data/tushare_provider.py
Normal file
51
backend/app/data/tushare_provider.py
Normal file
@@ -0,0 +1,51 @@
|
||||
"""Tushare 数据源(主)。日线 + 前复权。
|
||||
|
||||
token 从 settings.tushare_token 读取(.env)。免费版 pro.daily 与 ts.pro_bar 实测可用。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from ..config import settings
|
||||
from ..domain import Bar
|
||||
from .symbols import to_ts_code
|
||||
|
||||
|
||||
def _parse(date_str: str) -> datetime:
|
||||
return datetime.strptime(str(date_str), "%Y%m%d")
|
||||
|
||||
|
||||
def fetch_daily(code: str, start: str = "20200101", end: str | None = None,
|
||||
adjust: str = "qfq") -> list[Bar]:
|
||||
import tushare as ts # 延迟导入:未装/无 token 时 DEMO 仍可用
|
||||
|
||||
if not settings.tushare_token:
|
||||
raise RuntimeError("未配置 TUSHARE_TOKEN")
|
||||
ts.set_token(settings.tushare_token)
|
||||
pro = ts.pro_api()
|
||||
ts_code = to_ts_code(code)
|
||||
end = end or datetime.now().strftime("%Y%m%d")
|
||||
|
||||
# 优先 pro_bar(含复权);积分不足则退化为 pro.daily(不复权)
|
||||
df = None
|
||||
try:
|
||||
df = ts.pro_bar(ts_code=ts_code, adj=adjust, start_date=start, end_date=end, freq="D")
|
||||
except Exception:
|
||||
df = None
|
||||
if df is None or df.empty:
|
||||
df = pro.daily(ts_code=ts_code, start_date=start, end_date=end)
|
||||
if df is None or df.empty:
|
||||
raise RuntimeError(f"Tushare 无数据: {ts_code}")
|
||||
|
||||
df = df.sort_values("trade_date")
|
||||
bars: list[Bar] = []
|
||||
for _, r in df.iterrows():
|
||||
bars.append(
|
||||
Bar(
|
||||
ts=_parse(r["trade_date"]),
|
||||
open=float(r["open"]), high=float(r["high"]),
|
||||
low=float(r["low"]), close=float(r["close"]),
|
||||
volume=float(r["vol"]) * 100.0, # Tushare vol 单位为手 -> 股
|
||||
)
|
||||
)
|
||||
return bars
|
||||
20
backend/app/db.py
Normal file
20
backend/app/db.py
Normal file
@@ -0,0 +1,20 @@
|
||||
"""async SQLAlchemy 引擎与会话。"""
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
|
||||
from .config import settings
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
"""所有 ORM 模型的基类。"""
|
||||
|
||||
|
||||
# echo=False;生产环境可用连接池参数调优
|
||||
engine = create_async_engine(settings.database_url, echo=False, future=True)
|
||||
async_session = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession)
|
||||
|
||||
|
||||
async def get_session() -> AsyncSession:
|
||||
"""FastAPI 依赖:提供一个事务会话。"""
|
||||
async with async_session() as session:
|
||||
yield session
|
||||
77
backend/app/domain.py
Normal file
77
backend/app/domain.py
Normal file
@@ -0,0 +1,77 @@
|
||||
"""领域模型契约(单一事实源的载体)。
|
||||
|
||||
MVP 第一周必须定稿的核心类型。回测引擎、指标、API、(未来的)前端类型化客户端
|
||||
都基于这些契约——保证"图表/回测/实盘"用同一套语义。
|
||||
|
||||
注意:回测/图表的指标值由后端唯一计算(app.indicators),前端不另算。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class Side(str, Enum):
|
||||
BUY = "buy"
|
||||
SELL = "sell"
|
||||
|
||||
|
||||
class Timeframe(str, Enum):
|
||||
"""K 线周期。分钟级已预留(用户选了分钟级回测)。"""
|
||||
M1 = "1m"
|
||||
M5 = "5m"
|
||||
M15 = "15m"
|
||||
M30 = "30m"
|
||||
H1 = "1h"
|
||||
D1 = "1d"
|
||||
W1 = "1w"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Bar:
|
||||
"""一根 K 线(OHLCV + 时间戳)。复权标识后续扩展。"""
|
||||
ts: datetime
|
||||
open: float
|
||||
high: float
|
||||
low: float
|
||||
close: float
|
||||
volume: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Signal:
|
||||
"""策略产生的交易信号(用于在图上标注买卖点)。"""
|
||||
ts: datetime
|
||||
side: Side
|
||||
price: float
|
||||
strength: float = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Fill:
|
||||
"""一笔成交(含费用明细)。回测中由 PaperBroker 产生。"""
|
||||
ts: datetime
|
||||
side: Side
|
||||
price: float
|
||||
qty: float
|
||||
commission: float = 0.0
|
||||
stamp_duty: float = 0.0 # 印花税(仅卖出)
|
||||
transfer_fee: float = 0.0 # 过户费(双边)
|
||||
|
||||
@property
|
||||
def total_cost(self) -> float:
|
||||
return self.commission + self.stamp_duty + self.transfer_fee
|
||||
|
||||
|
||||
@dataclass
|
||||
class Position:
|
||||
"""持仓状态。T+1:locked_qty 为当日买入、次日才可卖的部分。"""
|
||||
symbol: str = ""
|
||||
holdings: float = 0.0 # 可卖数量
|
||||
locked: float = 0.0 # 当日买入(T+1 锁定)
|
||||
avg_price: float = 0.0
|
||||
|
||||
@property
|
||||
def qty(self) -> float:
|
||||
return self.holdings + self.locked
|
||||
56
backend/app/indicators.py
Normal file
56
backend/app/indicators.py
Normal file
@@ -0,0 +1,56 @@
|
||||
"""技术指标(单一事实源)。
|
||||
|
||||
MVP 用纯 pandas/numpy 实现,避免 Windows 上 TA-Lib C 库的安装痛点。
|
||||
算法正确(MACD = 快慢 EMA 之差),接口稳定;阶段1 在 Linux/Docker 上可换 TA-Lib,
|
||||
只需保持函数签名(输入 close Series,输出指标),上层无感。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def ema(series: pd.Series, span: int) -> pd.Series:
|
||||
"""指数移动平均(adjust=False,与 TA-Lib 默认一致)。"""
|
||||
return series.ewm(span=span, adjust=False).mean()
|
||||
|
||||
|
||||
def macd(close: pd.Series, fast: int = 12, slow: int = 26, signal: int = 9) -> pd.DataFrame:
|
||||
"""MACD:返回 DataFrame[DIF, DEA, HIST]。"""
|
||||
dif = ema(close, fast) - ema(close, slow)
|
||||
dea = ema(dif, signal)
|
||||
hist = (dif - dea) * 2 # A股惯例 MACD 柱 = 2*(DIF-DEA)
|
||||
return pd.DataFrame({"macd": dif, "signal": dea, "hist": hist})
|
||||
|
||||
|
||||
def rsi(close: pd.Series, period: int = 14) -> pd.Series:
|
||||
"""RSI(Wilder 平滑)。"""
|
||||
delta = close.diff()
|
||||
gain = delta.clip(lower=0.0)
|
||||
loss = -delta.clip(upper=0.0)
|
||||
avg_gain = gain.ewm(alpha=1 / period, adjust=False).mean()
|
||||
avg_loss = loss.ewm(alpha=1 / period, adjust=False).mean()
|
||||
rs = avg_gain / avg_loss.replace(0, np.nan)
|
||||
return 100 - (100 / (1 + rs))
|
||||
|
||||
|
||||
def kdj(high: pd.Series, low: pd.Series, close: pd.Series,
|
||||
n: int = 9, m1: int = 3, m2: int = 3) -> pd.DataFrame:
|
||||
"""KDJ(A股常用:RSV -> K -> D -> J)。"""
|
||||
low_n = low.rolling(n, min_periods=1).min()
|
||||
high_n = high.rolling(n, min_periods=1).max()
|
||||
rsv = (close - low_n) / (high_n - low_n).replace(0, np.nan) * 100
|
||||
k = rsv.ewm(alpha=1 / m1, adjust=False).mean()
|
||||
d = k.ewm(alpha=1 / m2, adjust=False).mean()
|
||||
j = 3 * k - 2 * d
|
||||
return pd.DataFrame({"k": k, "d": d, "j": j})
|
||||
|
||||
|
||||
def bollinger(close: pd.Series, period: int = 20, std: float = 2.0) -> pd.DataFrame:
|
||||
ma = close.rolling(period, min_periods=1).mean()
|
||||
sd = close.rolling(period, min_periods=1).std(ddof=0)
|
||||
return pd.DataFrame({"mid": ma, "upper": ma + std * sd, "lower": ma - std * sd})
|
||||
|
||||
|
||||
def ma(close: pd.Series, period: int) -> pd.Series:
|
||||
return close.rolling(period, min_periods=1).mean()
|
||||
43
backend/app/main.py
Normal file
43
backend/app/main.py
Normal file
@@ -0,0 +1,43 @@
|
||||
"""FastAPI 入口。
|
||||
|
||||
启动时自动建表(MVP 用 create_all;阶段1 切 Alembic 迁移,含 TimescaleDB hypertable)。
|
||||
"""
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from .api import router
|
||||
from .db import Base, engine
|
||||
from . import models # noqa: F401 —— 注册 ORM 到 Base.metadata
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
yield
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="Stock Backtest",
|
||||
description="历史回测 + 回放式模拟平台(A 股为主,不做实盘)",
|
||||
version="0.1.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
# 开发期允许前端 dev server 跨域;上线收窄 origins
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.include_router(router)
|
||||
|
||||
|
||||
@app.get("/")
|
||||
async def root() -> dict:
|
||||
return {"name": "Stock Backtest API", "docs": "/docs"}
|
||||
57
backend/app/models.py
Normal file
57
backend/app/models.py
Normal file
@@ -0,0 +1,57 @@
|
||||
"""ORM 模型。
|
||||
|
||||
Candle 表设计与 TimescaleDB hypertable 完全兼容:将来在目标 PG 库执行
|
||||
SELECT create_hypertable('candles', 'ts');
|
||||
即可升级为时序表 + Continuous Aggregates 多周期预聚合,无需改表结构。
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, Float, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from .db import Base
|
||||
|
||||
|
||||
def _utcnow() -> datetime:
|
||||
# naive UTC,避免 SQLite 存储时区带来的麻烦
|
||||
from datetime import timezone
|
||||
return datetime.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
class Candle(Base):
|
||||
__tablename__ = "candles"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(16), index=True)
|
||||
timeframe: Mapped[str] = mapped_column(String(4), default="1d", index=True)
|
||||
ts: Mapped[datetime] = mapped_column(DateTime, index=True) # bar 开始时间
|
||||
open: Mapped[float] = mapped_column(Float)
|
||||
high: Mapped[float] = mapped_column(Float)
|
||||
low: Mapped[float] = mapped_column(Float)
|
||||
close: Mapped[float] = mapped_column(Float)
|
||||
volume: Mapped[float] = mapped_column(Float)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("symbol", "timeframe", "ts", name="uq_candle_sym_tf_ts"),
|
||||
)
|
||||
|
||||
|
||||
class BacktestRun(Base):
|
||||
"""回测运行注册表(可复现/可审计/可回归对比的基础)。
|
||||
|
||||
完整版应记录 策略版本 + 参数快照 + 数据快照(复权/数据源/库版本)+ 环境指纹 + 结果指纹。
|
||||
MVP 先落关键字段,结构就位。
|
||||
"""
|
||||
__tablename__ = "backtest_runs"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, default=_utcnow)
|
||||
symbol: Mapped[str] = mapped_column(String(16))
|
||||
strategy: Mapped[str] = mapped_column(String(64))
|
||||
timeframe: Mapped[str] = mapped_column(String(4), default="1d")
|
||||
params_json: Mapped[str] = mapped_column(String, default="{}")
|
||||
initial_cash: Mapped[float] = mapped_column(Float, default=100000.0)
|
||||
total_return: Mapped[float] = mapped_column(Float, default=0.0)
|
||||
max_drawdown: Mapped[float] = mapped_column(Float, default=0.0)
|
||||
sharpe: Mapped[float] = mapped_column(Float, default=0.0)
|
||||
num_trades: Mapped[int] = mapped_column(Integer, default=0)
|
||||
87
backend/app/schemas.py
Normal file
87
backend/app/schemas.py
Normal file
@@ -0,0 +1,87 @@
|
||||
"""Pydantic DTO —— 这就是 OpenAPI 契约(前端据此生成类型化客户端)。
|
||||
|
||||
契约先于业务锁定:字段一旦定下,前端可并行开发,后端实现改动不影响前端。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# ---------- Candle ----------
|
||||
class CandleOut(BaseModel):
|
||||
ts: datetime
|
||||
open: float
|
||||
high: float
|
||||
low: float
|
||||
close: float
|
||||
volume: float
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------- Backtest ----------
|
||||
class BacktestRequest(BaseModel):
|
||||
symbol: str = "DEMO"
|
||||
timeframe: str = "1d"
|
||||
strategy: str = "macd_cross" # macd_cross | ma_cross | single_ma
|
||||
params: dict[str, float] = Field(default_factory=dict) # 各策略参数
|
||||
initial_cash: float = 1000000.0
|
||||
fast_mode: bool = False # True => 关闭 T+1/费用,交互试探
|
||||
start: datetime | None = None
|
||||
end: datetime | None = None
|
||||
|
||||
|
||||
class SignalOut(BaseModel):
|
||||
ts: datetime
|
||||
side: str # "buy" | "sell"
|
||||
price: float
|
||||
qty: float
|
||||
|
||||
|
||||
class EquityPoint(BaseModel):
|
||||
ts: datetime
|
||||
value: float
|
||||
|
||||
|
||||
class IndicatorOut(BaseModel):
|
||||
strategy: str
|
||||
data: dict[str, list[float | None]] = {} # 列名 -> 序列(MACD: macd/signal/hist;均线: fast/slow 或 ma)
|
||||
|
||||
|
||||
class MetricsOut(BaseModel):
|
||||
total_return: float
|
||||
max_drawdown: float
|
||||
sharpe: float
|
||||
volatility: float
|
||||
num_trades: int = 0
|
||||
win_rate: float = 0.0
|
||||
|
||||
|
||||
class BacktestResponse(BaseModel):
|
||||
symbol: str
|
||||
timeframe: str
|
||||
strategy: str
|
||||
candles: list[CandleOut]
|
||||
indicators: IndicatorOut
|
||||
signals: list[SignalOut]
|
||||
equity: list[EquityPoint]
|
||||
metrics: MetricsOut
|
||||
final_cash: float
|
||||
final_position: float
|
||||
initial_cash: float
|
||||
|
||||
|
||||
class SyncRequest(BaseModel):
|
||||
symbol: str
|
||||
start: str | None = None # YYYYMMDD
|
||||
end: str | None = None
|
||||
source: str = "auto" # auto | tushare | akshare
|
||||
force: bool = False # True => 忽略缓存重新拉取
|
||||
|
||||
|
||||
class SyncResponse(BaseModel):
|
||||
symbol: str
|
||||
bars: int
|
||||
source: str
|
||||
Reference in New Issue
Block a user