Files
stock/backend/app/data/aggregation.py
2026-08-15 15:36:11 +08:00

69 lines
2.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""K 线周期聚合:日线 -> 周/月/年。
生产环境用 TimescaleDB Continuous Aggregates 在库里预物化(性能);
MVP 在应用层用 pandas resample 即可,逻辑等价、便于切换。
OHLCV 聚合规则:开=周期内首根开、高=最高、低=最低、收=末根收、量=求和。
成交额/换手率为名义量:求和(全缺则保持 None不伪造 0
"""
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 _sum_or_none(s: pd.Series):
"""求和;全为 NaN 返回 None部分缺失则忽略缺失项求和"""
if s.isna().all():
return None
return float(s.sum())
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,
"amount": b.amount if b.amount is not None else float("nan"),
"turnover": b.turnover if b.turnover is not None else float("nan")}
for b in bars]
).set_index("ts").sort_index()
agg = (
df.resample(rule)
.agg({"open": "first", "high": "max", "low": "min", "close": "last",
"volume": "sum", "amount": _sum_or_none, "turnover": _sum_or_none})
.dropna(subset=["open"])
)
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"]),
amount=row["amount"] if row["amount"] == row["amount"] else None, # NaN -> None
turnover=row["turnover"] if row["turnover"] == row["turnover"] else None,
)
for ts, row in agg.iterrows()
]