first commit

This commit is contained in:
2026-08-07 16:08:34 +08:00
commit e0b5228008
51 changed files with 5175 additions and 0 deletions

17
backend/.env.example Normal file
View File

@@ -0,0 +1,17 @@
# ---- Database ----
# MVP 默认 SQLite零配置即可跑
DATABASE_URL=sqlite+aiosqlite:///./stock.db
# 切到你自己的 PostgreSQLTimescaleDB 是 Postgres 扩展,目标库执行 CREATE EXTENSION timescaledb; 即可)
# DATABASE_URL=postgresql+asyncpg://user:password@localhost:5432/stock
# ---- 真实数据源Tushare Pro免费版即可。留空则仅 DEMO 合成数据可用)----
TUSHARE_TOKEN=你的token
DATA_ADJUST=qfq # 复权qfq 前复权 / hfq 后复权 / 留空不复权
DATA_DEFAULT_START=20200101
# ---- A股交易成本基准日 2026-08可覆盖详见 app/commission.py----
# STAMP_DUTY_RATE=0.0005 # 印花税 0.05%,单边卖出
# TRANSFER_FEE_RATE=0.00001 # 过户费 0.001%,沪深双边
# COMMISSION_RATE=0.0001 # 佣金 万1
# COMMISSION_MIN=5.0 # 最低 5 元

0
backend/app/__init__.py Normal file
View File

167
backend/app/api.py Normal file
View 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 -> Nonelightweight-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,
)

View File

View 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

View 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,
}

View 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

View File

@@ -0,0 +1,24 @@
"""策略注册表与工厂。
新增策略:实现 Strategybase.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)

View 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:
...

View 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

View 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
View 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
View 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()

View File

View 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()
]

View 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

View 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}

View 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())

View 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"

View 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()

View 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
View 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
View 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+1locked_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
View 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:
"""RSIWilder 平滑)。"""
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:
"""KDJA股常用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
View 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
View 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
View 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

21
backend/pyproject.toml Normal file
View File

@@ -0,0 +1,21 @@
[project]
name = "stock-backend"
version = "0.1.0"
description = "Stock backtest platform backend (FastAPI + self-built backtest engine)"
requires-python = ">=3.12"
dependencies = [
"fastapi>=0.115",
"uvicorn[standard]>=0.30",
"pydantic>=2.7",
"pydantic-settings>=2.3",
"sqlalchemy>=2.0",
"aiosqlite>=0.20",
"asyncpg>=0.29", # PostgreSQL 异步驱动(连你已有的 Postgres / TimescaleDB
"numpy>=1.26",
"pandas>=2.2",
"tushare>=1.4",
]
[tool.uv]
# 应用型项目(非库):不把自身打包安装,只管理依赖到 .venv
package = false

54
backend/smoke_test.py Normal file
View File

@@ -0,0 +1,54 @@
"""开发自检脚本:跑一遍 /health 与 /backtest打印结果。
用 FastAPI TestClient无需起服务器同进程验证全链路
用法: uv run --with httpx --directory backend python smoke_test.py
"""
import json
from fastapi.testclient import TestClient
from app.main import app
# 必须用 withlifespan建表只在进入上下文时执行
with TestClient(app) as c:
r = c.get("/api/health")
print("== /api/health ==", r.status_code, r.json())
r = c.post(
"/api/backtest",
json={
"symbol": "DEMO",
"strategy": "macd_cross",
"params": {"fast": 12, "slow": 26, "signal": 9},
"initial_cash": 100000.0,
"fast_mode": False,
},
)
print("== /api/backtest ==", r.status_code)
if r.status_code != 200:
print("ERROR:", r.text)
raise SystemExit(1)
d = r.json()
print("candles :", len(d["candles"]))
print("signals :", len(d["signals"]), "(买卖点)")
print("equity pts :", len(d["equity"]))
print("final_cash :", round(d["final_cash"], 2))
print("final_pos :", d["final_position"])
print("metrics :", json.dumps(d["metrics"], ensure_ascii=False, indent=2))
print("first signal :", d["signals"][0] if d["signals"] else None)
assert len(d["candles"]) > 100
# 周期聚合:周线 K 线数应明显少于日线
rw = c.post(
"/api/backtest",
json={"symbol": "DEMO", "timeframe": "1w", "strategy": "macd_cross",
"params": {"fast": 12, "slow": 26, "signal": 9}, "initial_cash": 100000.0},
)
wd = rw.json()
print("== weekly ==", rw.status_code, "candles:", len(wd["candles"]), "vs daily", len(d["candles"]))
assert rw.status_code == 200
assert len(wd["candles"]) < len(d["candles"])
print("\n✅ 后端全链路自检通过(含周期聚合)")

1159
backend/uv.lock generated Normal file

File diff suppressed because it is too large Load Diff