添加None确认
This commit is contained in:
+53
-2
@@ -1,3 +1,8 @@
|
||||
"""
|
||||
Chan 数据服务:用 ccxt 从交易所拉取 K 线,内存缓存 + CSV 落盘;
|
||||
后台线程定期增量刷新,断线时记录 resume_since 以免漏 K;
|
||||
配置中的基础周期(如 1m/1h)可合成 DERIVED_TIMEFRAME_PLAN 中的衍生周期。
|
||||
"""
|
||||
import asyncio
|
||||
import csv
|
||||
import json
|
||||
@@ -20,14 +25,17 @@ from technical.util import resample_to_interval
|
||||
# docker compose logs --tail=200
|
||||
# docker compose down && docker compose build --no-cache && docker compose up -d
|
||||
|
||||
# 基础周期枚举顺序(用于衍生周期展示顺序);仅允许集合内周期作为交易所直接拉取的 tf
|
||||
TIMEFRAME_ORDER = ["1m", "1h", "1d", "1w"]
|
||||
ALLOWED_TIMEFRAMES = set(TIMEFRAME_ORDER)
|
||||
# 各基础周期一根 K 线的毫秒长度(用于历史分页与断线回退)
|
||||
TIMEFRAME_TO_MS: Dict[str, int] = {
|
||||
"1m": 60_000,
|
||||
"1h": 3_600_000,
|
||||
"1d": 86_400_000,
|
||||
"1w": 604_800_000,
|
||||
}
|
||||
# 每个基础周期可派生出的合成周期列表(由该基础周期 K 线 resample 得到)
|
||||
DERIVED_TIMEFRAME_PLAN: Dict[str, List[str]] = {
|
||||
"1m": ["2m", "3m", "4m", "5m", "10m", "15m", "20m", "25m", "30m", "45m"],
|
||||
"1h": ["2h", "3h", "4h", "5h", "6h", "7h", "8h", "9h", "10", "11h", "12h", "16h", "20h"],
|
||||
@@ -37,8 +45,8 @@ DERIVED_TIMEFRAME_PLAN: Dict[str, List[str]] = {
|
||||
CSV_FIELDNAMES = ["timestamp", "datetime", "open", "high", "low", "close", "volume"]
|
||||
DEFAULT_LIMIT = 500
|
||||
RECENT_CANDLE_LIMIT = 10
|
||||
RECENT_FETCH_INTERVAL = 5
|
||||
PERSIST_INTERVAL = 600
|
||||
RECENT_FETCH_INTERVAL = 5 # 后台刷新循环休眠秒数
|
||||
PERSIST_INTERVAL = 600 # 全量落盘周期(秒)
|
||||
|
||||
|
||||
logger = logging.getLogger("data_provider")
|
||||
@@ -49,11 +57,13 @@ logging.basicConfig(
|
||||
|
||||
|
||||
def to_utc_iso(timestamp_ms: int) -> str:
|
||||
"""将毫秒时间戳格式化为 UTC ISO 字符串(末尾 Z)。"""
|
||||
dt = datetime.fromtimestamp(timestamp_ms / 1000, tz=timezone.utc)
|
||||
return dt.isoformat().replace("+00:00", "Z")
|
||||
|
||||
|
||||
def parse_timestamp(value: Optional[object]) -> Optional[int]:
|
||||
"""解析查询参数中的时间为 UTC 毫秒时间戳;支持数字或 ISO 字符串。"""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, (int, float)):
|
||||
@@ -79,6 +89,7 @@ def parse_timestamp(value: Optional[object]) -> Optional[int]:
|
||||
|
||||
|
||||
def candle_to_dict(candle: Iterable[float]) -> Dict[str, float]:
|
||||
"""ccxt OHLCV 单根 [ts, o, h, l, c, v] 转为内部字典结构。"""
|
||||
ts = int(candle[0])
|
||||
return {
|
||||
"timestamp": ts,
|
||||
@@ -92,6 +103,7 @@ def candle_to_dict(candle: Iterable[float]) -> Dict[str, float]:
|
||||
|
||||
|
||||
def timeframe_to_minutes(tf: str) -> Optional[int]:
|
||||
"""将如 15m、2h 转为「分钟数」,供 resample 与衍生周期计算。"""
|
||||
if not tf:
|
||||
return None
|
||||
unit = tf[-1]
|
||||
@@ -111,6 +123,8 @@ def timeframe_to_minutes(tf: str) -> Optional[int]:
|
||||
|
||||
|
||||
class DataProvider:
|
||||
"""封装交易所连接、本地 CSV、内存缓存、断线恢复与衍生周期聚合。"""
|
||||
|
||||
def __init__(self, config_path: Path) -> None:
|
||||
self.config_path = config_path
|
||||
self.config = self._load_config()
|
||||
@@ -126,10 +140,12 @@ class DataProvider:
|
||||
self.data: Dict[str, Dict[str, List[Dict[str, float]]]] = {
|
||||
symbol: {tf: [] for tf in self.timeframes} for symbol in self.symbols
|
||||
}
|
||||
# 衍生周期 -> 用于合成的交易所基础周期(每个衍生只对应一个 base)
|
||||
self.derived_map: Dict[str, str] = {}
|
||||
for base_tf in self.timeframes:
|
||||
for derived_tf in DERIVED_TIMEFRAME_PLAN.get(base_tf, []):
|
||||
self.derived_map.setdefault(derived_tf, base_tf)
|
||||
# 衍生周期展示顺序:按 TIMEFRAME_ORDER 中的基础周期依次展开
|
||||
derived_order: List[str] = []
|
||||
for base_tf in TIMEFRAME_ORDER:
|
||||
if base_tf not in self.timeframes:
|
||||
@@ -151,12 +167,14 @@ class DataProvider:
|
||||
self._load_resume_since()
|
||||
|
||||
def _load_config(self) -> Dict[str, object]:
|
||||
"""读取 JSON 配置文件。"""
|
||||
if not self.config_path.exists():
|
||||
raise FileNotFoundError(f"未找到配置文件: {self.config_path}")
|
||||
with self.config_path.open("r", encoding="utf-8") as fp:
|
||||
return json.load(fp)
|
||||
|
||||
def _load_symbols(self, config: Dict[str, object]) -> List[str]:
|
||||
"""从 symbols 列表、逗号分隔字符串或单字段 symbol 解析交易对,去重保序。"""
|
||||
raw_symbols: List[str] = []
|
||||
symbols_value = config.get("symbols")
|
||||
if isinstance(symbols_value, list):
|
||||
@@ -175,6 +193,7 @@ class DataProvider:
|
||||
return unique
|
||||
|
||||
def _validate_timeframes(self, configured: Optional[Iterable[str]]) -> List[str]:
|
||||
"""校验周期在允许集合内;未配置则默认 TIMEFRAME_ORDER 全部;顺序优先按 TIMEFRAME_ORDER。"""
|
||||
if not configured:
|
||||
return list(TIMEFRAME_ORDER)
|
||||
invalid = [tf for tf in configured if tf not in ALLOWED_TIMEFRAMES]
|
||||
@@ -193,6 +212,7 @@ class DataProvider:
|
||||
return unique
|
||||
|
||||
def _init_exchange(self):
|
||||
"""实例化 ccxt 交易所,币安期货默认 defaultType=future,并 load_markets。"""
|
||||
if not hasattr(ccxt, self.exchange_name):
|
||||
raise ValueError(f"不支持的交易所: {self.exchange_name}")
|
||||
exchange_class = getattr(ccxt, self.exchange_name)
|
||||
@@ -204,10 +224,12 @@ class DataProvider:
|
||||
return exchange
|
||||
|
||||
def _data_file_path(self, symbol: str, timeframe: str) -> Path:
|
||||
"""单交易对单周期的 CSV 路径:data_dir/tf/exchange_symbol_tf.csv。"""
|
||||
symbol_safe = symbol.replace("/", "_").replace(":", "_")
|
||||
return self.data_dir / timeframe / f"{self.exchange.id}_{symbol_safe}_{timeframe}.csv"
|
||||
|
||||
def _load_local(self, symbol: str, timeframe: str) -> List[Dict[str, float]]:
|
||||
"""启动时从磁盘加载已有 K 线,损坏行跳过,按时间排序。"""
|
||||
path = self._data_file_path(symbol, timeframe)
|
||||
if not path.exists():
|
||||
return []
|
||||
@@ -239,6 +261,7 @@ class DataProvider:
|
||||
base: List[Dict[str, float]],
|
||||
new_candles: Iterable[Iterable[float]],
|
||||
) -> List[Dict[str, float]]:
|
||||
"""按 timestamp 去重合并,新数据覆盖同时间戳旧数据。"""
|
||||
merged = {entry["timestamp"]: entry for entry in base}
|
||||
for candle in new_candles:
|
||||
entry = candle_to_dict(candle)
|
||||
@@ -248,6 +271,7 @@ class DataProvider:
|
||||
return ordered
|
||||
|
||||
def _write_to_disk(self, symbol: str, timeframe: str, data: List[Dict[str, float]]) -> None:
|
||||
"""先写临时文件再 replace,避免写入中断导致 CSV 损坏。"""
|
||||
path = self._data_file_path(symbol, timeframe)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp_path = path.with_suffix(path.suffix + ".tmp")
|
||||
@@ -266,6 +290,7 @@ class DataProvider:
|
||||
logger.info("交易对 %s 时间周期 %s 已写入磁盘 (%s 根K线)", symbol, timeframe, len(data))
|
||||
|
||||
def _fetch_history(self, symbol: str, timeframe: str, since_ms: int) -> List[List[float]]:
|
||||
"""从 since_ms 分页拉取直到接近当前时间;遇限频则 sleep 重试。"""
|
||||
results: List[List[float]] = []
|
||||
limit = 1500
|
||||
now_ms = self.exchange.milliseconds()
|
||||
@@ -302,6 +327,7 @@ class DataProvider:
|
||||
return results
|
||||
|
||||
def initialize(self) -> None:
|
||||
"""阻塞式启动:加载本地、从倒数第二根或配置起点补历史、写盘并 set _ready。"""
|
||||
logger.info("开始初始化数据提供商")
|
||||
for symbol in self.symbols:
|
||||
for timeframe in self.timeframes:
|
||||
@@ -310,6 +336,7 @@ class DataProvider:
|
||||
last_ts = existing[-1]["timestamp"] if existing else None
|
||||
if last_ts is not None:
|
||||
if len(existing) >= 2:
|
||||
# 从倒数第二根起拉,避免最后一根未收盘重复/缺口
|
||||
fetch_since = existing[-2]["timestamp"]
|
||||
else:
|
||||
fetch_since = max(0, last_ts - tf_ms)
|
||||
@@ -334,9 +361,11 @@ class DataProvider:
|
||||
logger.info("数据初始化完成")
|
||||
|
||||
def resample_df(self, df: pd.DataFrame, interval: int) -> pd.DataFrame:
|
||||
"""将基础周期 DataFrame 聚合为 interval 分钟周期(freqtrade technical.util)。"""
|
||||
return resample_to_interval(df, interval)
|
||||
|
||||
def _save_resume_since(self) -> None:
|
||||
"""将断线恢复点持久化到 resume_since.json(原子替换)。"""
|
||||
path = self._resume_file
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp_path = path.with_suffix(path.suffix + ".tmp")
|
||||
@@ -358,6 +387,7 @@ class DataProvider:
|
||||
logger.debug("恢复点已保存到磁盘: %s", path)
|
||||
|
||||
def _load_resume_since(self) -> None:
|
||||
"""启动时加载恢复点;与内存合并时取更早的 since,避免漏拉。"""
|
||||
path = self._resume_file
|
||||
if not path.exists():
|
||||
return
|
||||
@@ -395,10 +425,12 @@ class DataProvider:
|
||||
logger.info("已加载恢复点: %s", path)
|
||||
|
||||
def _get_resume_since(self, symbol: str, timeframe: str) -> Optional[int]:
|
||||
"""若曾断线,返回应从哪一毫秒起补拉该 symbol/tf。"""
|
||||
with self._lock:
|
||||
return self._resume_since.get(symbol, {}).get(timeframe)
|
||||
|
||||
def _set_resume_since(self, symbol: str, timeframe: str, since_ms: int) -> None:
|
||||
"""断线时写入恢复点(取更早的 since 以免漏数据),并持久化到磁盘。"""
|
||||
with self._lock:
|
||||
per_symbol = self._resume_since.setdefault(symbol, {})
|
||||
prev = per_symbol.get(timeframe)
|
||||
@@ -416,6 +448,7 @@ class DataProvider:
|
||||
self._save_resume_since()
|
||||
|
||||
def _clear_resume_since(self, symbol: str, timeframe: str) -> None:
|
||||
"""补数成功后清除该 symbol/tf 的恢复点。"""
|
||||
with self._lock:
|
||||
if symbol in self._resume_since and timeframe in self._resume_since[symbol]:
|
||||
del self._resume_since[symbol][timeframe]
|
||||
@@ -426,6 +459,7 @@ class DataProvider:
|
||||
self._save_resume_since()
|
||||
|
||||
def start_background_workers(self) -> None:
|
||||
"""启动增量刷新线程与周期性落盘线程。"""
|
||||
if self._fetch_thread and self._fetch_thread.is_alive():
|
||||
return
|
||||
self._stop_event.clear()
|
||||
@@ -436,6 +470,7 @@ class DataProvider:
|
||||
logger.info("后台线程已启动")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""停止后台线程(应用关闭时 lifespan finally 调用)。"""
|
||||
self._stop_event.set()
|
||||
if self._fetch_thread:
|
||||
self._fetch_thread.join(timeout=5)
|
||||
@@ -444,6 +479,7 @@ class DataProvider:
|
||||
logger.info("数据提供商已停止")
|
||||
|
||||
def _refresh_loop(self) -> None:
|
||||
"""轮询各 symbol/tf:有恢复点则先补历史,否则 fetch 最近 RECENT_CANDLE_LIMIT 根。"""
|
||||
while not self._stop_event.is_set():
|
||||
for symbol in self.symbols:
|
||||
for timeframe in self.timeframes:
|
||||
@@ -496,10 +532,12 @@ class DataProvider:
|
||||
break
|
||||
|
||||
def _persist_loop(self) -> None:
|
||||
"""每隔 PERSIST_INTERVAL 秒把内存快照写 CSV 并保存恢复点。"""
|
||||
while not self._stop_event.wait(PERSIST_INTERVAL):
|
||||
self._persist_all()
|
||||
|
||||
def _persist_all(self) -> None:
|
||||
"""在锁内复制 data 后落盘,避免长时间持锁。"""
|
||||
if not self._ready.is_set():
|
||||
return
|
||||
with self._lock:
|
||||
@@ -533,6 +571,7 @@ class DataProvider:
|
||||
end_ms: Optional[int],
|
||||
limit: Optional[int],
|
||||
) -> List[Dict[str, float]]:
|
||||
"""从内存读取已缓存的基础周期 K 线并按时间/limit 裁剪。"""
|
||||
with self._lock:
|
||||
candles = list(self.data.get(symbol, {}).get(timeframe, []))
|
||||
if start_ms is not None:
|
||||
@@ -551,6 +590,7 @@ class DataProvider:
|
||||
end_time: Optional[object] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> List[Dict[str, float]]:
|
||||
"""对外查询:基础周期直接返回;衍生周期从 derived_map 取 base,resample 后对齐时间戳再裁剪。"""
|
||||
if symbol not in self.symbols:
|
||||
raise HTTPException(status_code=404, detail=f"symbol {symbol} 不可用")
|
||||
self.wait_ready()
|
||||
@@ -565,6 +605,7 @@ class DataProvider:
|
||||
if target_minutes is None:
|
||||
raise HTTPException(status_code=400, detail=f"不支持的时间周期: {timeframe}")
|
||||
target_ms = target_minutes * 60_000
|
||||
# 起点前移一根目标周期长度,保证首根合成 K 边界完整
|
||||
adjusted_start = None if start_ms is None else max(0, start_ms - target_ms)
|
||||
base_candles = self._get_base_klines(symbol, base_tf, adjusted_start, end_ms, None)
|
||||
if not base_candles:
|
||||
@@ -574,9 +615,11 @@ class DataProvider:
|
||||
return []
|
||||
df = df.drop_duplicates(subset=["timestamp"], keep="last").sort_values("timestamp")
|
||||
df["date"] = pd.to_datetime(df["timestamp"], unit="ms", utc=True)
|
||||
# resample_to_interval 按「分钟」目标周期聚合 OHLCV
|
||||
resampled = self.resample_df(df, target_minutes)
|
||||
if resampled is None or resampled.empty:
|
||||
return []
|
||||
# 统一得到毫秒 timestamp 列(resample 可能返回 date 或 DatetimeIndex)
|
||||
if "timestamp" in resampled.columns:
|
||||
resampled_df = resampled.copy()
|
||||
else:
|
||||
@@ -621,6 +664,8 @@ class DataProvider:
|
||||
|
||||
|
||||
def create_app(provider: DataProvider) -> FastAPI:
|
||||
"""构造 FastAPI 应用:lifespan 内同步 initialize 并启动后台拉数。"""
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
loop = asyncio.get_running_loop()
|
||||
@@ -643,6 +688,7 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
|
||||
@app.get("/health")
|
||||
async def health() -> Dict[str, object]:
|
||||
"""存活检查:交易所、交易对、基础/衍生周期、是否已完成冷启动。"""
|
||||
return {
|
||||
"status": "ok",
|
||||
"exchange": provider.exchange_name,
|
||||
@@ -655,6 +701,7 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
|
||||
@app.get("/timeframes")
|
||||
async def list_timeframes() -> Dict[str, List[str]]:
|
||||
"""返回配置的基础周期与可合成的衍生周期列表。"""
|
||||
provider.wait_ready()
|
||||
return {
|
||||
"base_timeframes": provider.timeframes,
|
||||
@@ -670,11 +717,13 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
end: Optional[int] = Query(None, description="结束时间戳(ms)"),
|
||||
limit: Optional[int] = Query(None, description="可选,限制返回数量"),
|
||||
):
|
||||
"""按交易对与时间周期返回 OHLCV;tf 支持配置的基础周期及衍生合成周期。"""
|
||||
data = provider.get_klines(symbol=symbol, timeframe=tf, start_time=start, end_time=end, limit=limit)
|
||||
return data
|
||||
|
||||
@app.get("/")
|
||||
async def root() -> Dict[str, object]:
|
||||
"""根路径:服务名、交易所、交易对与可用周期(含 ready 标志)。"""
|
||||
return {
|
||||
"service": "Data Provider",
|
||||
"exchange": provider.exchange_name,
|
||||
@@ -689,6 +738,7 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
|
||||
|
||||
def build_app() -> FastAPI:
|
||||
"""默认入口:从环境变量 CONFIG_PATH(或 config.json)加载配置并创建 FastAPI app。"""
|
||||
config_path = Path(os.getenv("CONFIG_PATH", "config.json"))
|
||||
provider = DataProvider(config_path)
|
||||
return create_app(provider)
|
||||
@@ -698,6 +748,7 @@ app = build_app()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""直接运行本模块时启动 uvicorn(监听 UVICORN_HOST / UVICORN_PORT)。"""
|
||||
host = os.getenv("UVICORN_HOST", "0.0.0.0")
|
||||
port = int(os.getenv("UVICORN_PORT", "9009"))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user