Files
Chan/datasvc/app/main.py
T

1458 lines
52 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.
import os
import asyncio
import json
import logging
from contextlib import suppress
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from time import time
from pathlib import Path
from collections import defaultdict
from threading import Event, RLock
from typing import Dict, List, Optional, Set, Tuple, Union
import ccxt
import ccxt.async_support as ccxt_async
import pandas as pd
import websockets
from fastapi import FastAPI, WebSocket, WebSocketDisconnect, Query, HTTPException, status
from fastapi.responses import JSONResponse
from fastapi.middleware.cors import CORSMiddleware
# docker compose down && docker compose build --no-cache && docker compose up -d
# docker compose down && docker compose build && docker compose up -d
try:
from technical.util import resample_to_interval # type: ignore
except ImportError: # pragma: no cover - 环境缺失依赖时自动降级
resample_to_interval = None # type: ignore
from .storage import candle_path, ensure_storage, read_candles, write_candles_snapshot
LOG_LEVEL = os.environ.get("LOG_LEVEL", "INFO").upper()
logging.basicConfig(
level=LOG_LEVEL,
format="%(asctime)s %(levelname)s [%(name)s] %(message)s",
)
logger = logging.getLogger("datasvc")
BASE_DIR = Path(__file__).resolve().parent.parent
DEFAULT_PAIRS_FILE = BASE_DIR / "pairs.json"
RESAMPLE_AVAILABLE = resample_to_interval is not None
RESAMPLE_WARNING_EMITTED = False
WS_ENABLED = os.environ.get("WS_ENABLED", "false").lower() in {"1", "true", "yes"}
REST_POLL_INTERVAL = max(1.0, float(os.environ.get("REST_POLL_INTERVAL", "5")))
REST_POLL_WINDOW = max(1, int(os.environ.get("REST_POLL_WINDOW", "10")))
REST_MAX_CONCURRENCY = int(os.environ.get("REST_MAX_CONCURRENCY", "1"))
REST_FETCH_SEMAPHORE = asyncio.Semaphore(max(1, REST_MAX_CONCURRENCY))
AGGREGATION_PLAN: Dict[str, List[str]] = {
"1m": ["2m", "3m", "4m", "5m", "10m", "15m", "20m", "25m", "30m"],
"1h": ["2h", "3h", "4h", "6h", "8h", "12h", "16h"],
"1d": ["2d", "3d", "4d", "5d", "6d"],
"1w": ["2w"],
"1M": ["2M", "3M", "6M"],
}
CandleRow = List[Union[int, float]]
CANDLE_COLUMNS = ["timestamp", "open", "high", "low", "close", "volume"]
CANDLES_CACHE: Dict[Tuple[str, str], pd.DataFrame] = {}
CACHE_LOCK = RLock()
BASE_TIMEFRAMES: Set[str] = set()
BASE_DIRTY_VERSION: Dict[Tuple[str, str], int] = {}
BASE_FLUSH_INTERVAL_SECONDS = 600
CACHE_FILE_MTIME: Dict[Tuple[str, str], float] = {}
IS_ENGINE_PROCESS = os.environ.get("DATASVC_ENGINE") == "1"
def _empty_frame() -> pd.DataFrame:
return pd.DataFrame(columns=CANDLE_COLUMNS)
def _normalize_dataframe(df: pd.DataFrame) -> pd.DataFrame:
if df.empty:
return _empty_frame()
normalized = df.copy()
missing_columns = [col for col in CANDLE_COLUMNS if col not in normalized.columns]
for column in missing_columns:
normalized[column] = 0.0 if column != "timestamp" else 0
normalized = normalized[CANDLE_COLUMNS]
normalized["timestamp"] = normalized["timestamp"].astype("int64")
for column in CANDLE_COLUMNS[1:]:
normalized[column] = normalized[column].astype("float64")
normalized = normalized.drop_duplicates(subset=["timestamp"], keep="last").sort_values("timestamp").reset_index(drop=True)
return normalized
def preload_candles_cache(symbols: List[str], timeframes: List[str]) -> None:
new_cache: Dict[Tuple[str, str], pd.DataFrame] = {}
new_mtime: Dict[Tuple[str, str], float] = {}
for symbol in symbols:
for timeframe in timeframes:
df = read_candles(DATA_DIR, symbol, timeframe, None, None)
normalized = _normalize_dataframe(df)
new_cache[(symbol, timeframe)] = normalized
try:
mtime = os.path.getmtime(candle_path(DATA_DIR, symbol, timeframe))
except OSError:
mtime = 0.0
new_mtime[(symbol, timeframe)] = mtime
with CACHE_LOCK:
CANDLES_CACHE.clear()
CANDLES_CACHE.update(new_cache)
BASE_DIRTY_VERSION.clear()
CACHE_FILE_MTIME.clear()
for key in new_cache:
if key[1] in BASE_TIMEFRAMES:
BASE_DIRTY_VERSION[key] = 0
CACHE_FILE_MTIME[key] = new_mtime.get(key, 0.0)
def refresh_cache_from_disk(symbol: str, timeframe: str) -> None:
if timeframe not in BASE_TIMEFRAMES:
return
if IS_ENGINE_PROCESS:
return
key = (symbol, timeframe)
path = candle_path(DATA_DIR, symbol, timeframe)
try:
mtime = os.path.getmtime(path)
except FileNotFoundError:
with CACHE_LOCK:
if key not in CANDLES_CACHE:
CANDLES_CACHE[key] = _empty_frame()
CACHE_FILE_MTIME[key] = 0.0
return
except OSError:
return
with CACHE_LOCK:
cached_mtime = CACHE_FILE_MTIME.get(key, 0.0)
if mtime <= cached_mtime:
return
df = read_candles(DATA_DIR, symbol, timeframe, None, None)
normalized = _normalize_dataframe(df)
with CACHE_LOCK:
CANDLES_CACHE[key] = normalized
CACHE_FILE_MTIME[key] = mtime
if timeframe in BASE_TIMEFRAMES:
BASE_DIRTY_VERSION[key] = 0
def update_cache_mtime(symbol: str, timeframe: str) -> None:
path = candle_path(DATA_DIR, symbol, timeframe)
try:
mtime = os.path.getmtime(path)
except OSError:
mtime = time()
with CACHE_LOCK:
CACHE_FILE_MTIME[(symbol, timeframe)] = mtime
def rebuild_all_derived_timeframes(symbols: List[str]) -> None:
if not RESAMPLE_AVAILABLE:
return
for symbol in symbols:
for base_tf, targets in AGGREGATION_TARGETS.items():
if not targets:
continue
resample_and_store(symbol, base_tf, targets)
def cache_get(symbol: str, timeframe: str, start: Optional[int] = None, end: Optional[int] = None) -> pd.DataFrame:
refresh_cache_from_disk(symbol, timeframe)
key = (symbol, timeframe)
with CACHE_LOCK:
df = CANDLES_CACHE.get(key)
if df is None:
df = _empty_frame()
result = df
if start is not None:
result = result[result["timestamp"] >= int(start)]
if end is not None:
result = result[result["timestamp"] <= int(end)]
return result.copy()
def cache_get_last_timestamp(symbol: str, timeframe: str) -> Optional[int]:
if timeframe in BASE_TIMEFRAMES:
refresh_cache_from_disk(symbol, timeframe)
key = (symbol, timeframe)
with CACHE_LOCK:
df = CANDLES_CACHE.get(key)
if df is None:
CANDLES_CACHE[key] = _empty_frame()
return None
if df.empty:
return None
return int(df["timestamp"].iloc[-1])
def cache_update(symbol: str, timeframe: str, candles: List[CandleRow]) -> Optional[pd.DataFrame]:
if not candles:
return None
new_df = _normalize_dataframe(pd.DataFrame(candles, columns=CANDLE_COLUMNS))
if new_df.empty:
return None
key = (symbol, timeframe)
with CACHE_LOCK:
existing = CANDLES_CACHE.get(key)
if existing is None or existing.empty:
merged = new_df
else:
merged = pd.concat([existing, new_df], ignore_index=True)
merged = _normalize_dataframe(merged)
with CACHE_LOCK:
CANDLES_CACHE[key] = merged
if timeframe in BASE_TIMEFRAMES:
BASE_DIRTY_VERSION[key] = BASE_DIRTY_VERSION.get(key, 0) + 1
snapshot = merged.copy()
return snapshot
def collect_engine_status() -> dict:
updated_at = datetime.utcnow().replace(microsecond=0).isoformat() + "Z"
tasks = [state.to_payload() for state in fetch_states.values()]
return {
"updated_at": updated_at,
"tasks": tasks,
"queues": {},
}
def write_engine_status_snapshot() -> None:
try:
ENGINE_STATUS_PATH.parent.mkdir(parents=True, exist_ok=True)
snapshot = collect_engine_status()
ENGINE_STATUS_PATH.write_text(json.dumps(snapshot, ensure_ascii=False), encoding="utf-8")
except Exception:
logger.warning("写入引擎状态快照失败", exc_info=True)
async def status_flush_worker(stop_event: Event, interval: float = 5.0) -> None:
await asyncio.to_thread(write_engine_status_snapshot)
try:
while not stop_event.is_set():
await asyncio.sleep(interval)
await asyncio.to_thread(write_engine_status_snapshot)
except asyncio.CancelledError:
raise
def load_engine_status_snapshot() -> Optional[dict]:
try:
content = ENGINE_STATUS_PATH.read_text(encoding="utf-8")
except FileNotFoundError:
return None
except Exception:
logger.warning("读取引擎状态快照失败", exc_info=True)
return None
try:
return json.loads(content)
except json.JSONDecodeError:
logger.warning("解析引擎状态快照失败")
return None
async def flush_dirty_base_snapshots(force_all: bool = False) -> None:
with CACHE_LOCK:
if force_all:
target_entries = []
for key in CANDLES_CACHE.keys():
symbol, timeframe = key
if timeframe in BASE_TIMEFRAMES:
version = BASE_DIRTY_VERSION.get(key, 0)
target_entries.append((key, version))
else:
target_entries = [(key, version) for key, version in BASE_DIRTY_VERSION.items() if version > 0]
snapshots = {key: CANDLES_CACHE.get(key, _empty_frame()).copy() for key, _ in target_entries}
if not snapshots:
return
failed: Set[Tuple[str, str]] = set()
for key, snapshot in snapshots.items():
symbol, timeframe = key
try:
await asyncio.to_thread(write_candles_snapshot, DATA_DIR, symbol, timeframe, snapshot)
except Exception:
failed.add(key)
logger.exception(
"基础周期快照写入失败",
extra={"symbol": symbol, "timeframe": timeframe},
)
else:
update_cache_mtime(symbol, timeframe)
if not failed:
logger.debug(
"基础周期快照写入完成",
extra={"count": len(snapshots), "force_all": force_all},
)
with CACHE_LOCK:
for key, version in target_entries:
if key in failed:
continue
current_version = BASE_DIRTY_VERSION.get(key, 0)
if current_version == version:
BASE_DIRTY_VERSION[key] = 0
async def base_flush_worker():
try:
logger.info(
"基础周期定时写盘任务已启动",
extra={"interval_seconds": BASE_FLUSH_INTERVAL_SECONDS},
)
while True:
await asyncio.sleep(BASE_FLUSH_INTERVAL_SECONDS)
await flush_dirty_base_snapshots()
except asyncio.CancelledError:
raise
finally:
with suppress(Exception):
await flush_dirty_base_snapshots(force_all=True)
async def run_engine(stop_event: Optional[Event] = None):
if stop_event is None:
stop_event = Event()
logger.info("数据引擎启动")
fetch_tasks.clear()
try:
await asyncio.to_thread(rebuild_all_derived_timeframes, SYMBOLS)
except Exception:
logger.exception("初始化衍生周期失败,继续启动引擎")
try:
status_task = asyncio.create_task(status_flush_worker(stop_event), name="status::flush")
fetch_tasks.append(status_task)
flush_task = asyncio.create_task(base_flush_worker(), name="flush::base")
fetch_tasks.append(flush_task)
for s in SYMBOLS:
for tf in FETCH_TIMEFRAMES:
fetch_task = asyncio.create_task(fetch_loop(s, tf), name=f"fetch::{s}::{tf}")
fetch_tasks.append(fetch_task)
while not stop_event.is_set():
await asyncio.sleep(1.0)
finally:
stop_event.set()
if fetch_tasks:
logger.info("数据引擎正在停止")
tasks = list(fetch_tasks)
for task in tasks:
task.cancel()
results = await asyncio.gather(*tasks, return_exceptions=True)
for result in results:
if isinstance(result, Exception) and not isinstance(result, asyncio.CancelledError):
logger.warning("任务停止时出现异常:%s", result)
fetch_tasks.clear()
with suppress(Exception):
await flush_dirty_base_snapshots(force_all=True)
with suppress(Exception):
await asyncio.to_thread(write_engine_status_snapshot)
logger.info("数据引擎已停止")
def _split_env_list(value: str) -> List[str]:
return [item.strip() for item in value.split(",") if item.strip()]
def _unique_preserve(values: List[str]) -> List[str]:
seen = set()
ordered: List[str] = []
for item in values:
if item not in seen:
ordered.append(item)
seen.add(item)
return ordered
def _load_symbols() -> List[str]:
path = DEFAULT_PAIRS_FILE
if not path.is_file():
default_symbols = ["BTC/USDT:USDT"]
logger.warning("交易对配置文件不存在,使用默认值", extra={"file": str(path), "symbols": default_symbols})
return default_symbols
try:
content = path.read_text(encoding="utf-8")
data = json.loads(content)
except Exception:
default_symbols = ["BTC/USDT:USDT"]
logger.exception("读取交易对配置文件失败,使用默认值", extra={"file": str(path), "symbols": default_symbols})
return default_symbols
raw_symbols: List[str] = []
if isinstance(data, list):
raw_symbols = [str(item).strip() for item in data if isinstance(item, str) and item.strip()]
elif isinstance(data, dict):
candidates = data.get("symbols") or data.get("pairs")
if isinstance(candidates, list):
raw_symbols = [str(item).strip() for item in candidates if isinstance(item, str) and item.strip()]
if not raw_symbols:
default_symbols = ["BTC/USDT:USDT"]
logger.warning("交易对配置文件未提供有效列表,使用默认值", extra={"file": str(path), "symbols": default_symbols})
return default_symbols
symbols = _unique_preserve(raw_symbols)
logger.info("已从配置文件载入交易对", extra={"file": str(path), "symbols": symbols})
return symbols
def timeframe_to_minutes(tf: str) -> Optional[int]:
if not tf:
return None
unit = tf[-1]
try:
value = int(tf[:-1])
except ValueError:
return None
multiplier = {
"m": 1,
"h": 60,
"d": 1440,
"w": 10080,
"M": 43200, # 30 天近似
}.get(unit)
if multiplier is None:
return None
return value * multiplier
def binance_stream_symbol(symbol: str) -> str:
try:
base, rest = symbol.split("/", 1)
except ValueError:
cleaned = symbol.replace("/", "").split(":")[0]
return cleaned.lower()
quote = rest.split(":")[0]
return f"{base}{quote}".lower()
def build_stream_url(symbol: str, timeframe: str) -> str:
stream_symbol = binance_stream_symbol(symbol)
return f"{BINANCE_WS_BASE}/{stream_symbol}@kline_{timeframe}"
DATA_DIR = os.environ.get("DATA_DIR", "/data")
EXCHANGE = os.environ.get("EXCHANGE", "binance")
SYMBOLS = _load_symbols()
_default_timeframes = ["1m", "1h", "1d", "1w", "1M"]
requested_timeframes = _split_env_list(os.environ.get("TIMEFRAMES", ",".join(_default_timeframes)))
if not requested_timeframes:
requested_timeframes = _default_timeframes
FETCH_TIMEFRAMES = _unique_preserve(requested_timeframes)
AVAILABLE_TIMEFRAMES = list(FETCH_TIMEFRAMES)
for base_tf in FETCH_TIMEFRAMES:
for derived_tf in AGGREGATION_PLAN.get(base_tf, []):
if derived_tf not in AVAILABLE_TIMEFRAMES:
AVAILABLE_TIMEFRAMES.append(derived_tf)
DERIVED_TIMEFRAMES = [tf for tf in AVAILABLE_TIMEFRAMES if tf not in FETCH_TIMEFRAMES]
AGGREGATION_TARGETS = {tf: AGGREGATION_PLAN.get(tf, []) for tf in FETCH_TIMEFRAMES}
BASE_TIMEFRAMES = set(FETCH_TIMEFRAMES)
START_FROM = os.environ.get("START_FROM", "2025-09-01") # 首次启动拉取起始日期(UTC
POLL_FACTOR = float(os.environ.get("POLL_FACTOR", "0.5")) # 轮询间隔 = tf_ms * factor
BACKOFF_BASE = float(os.environ.get("BACKOFF_BASE", "2.0"))
BACKOFF_MAX = float(os.environ.get("BACKOFF_MAX", "30.0"))
BINANCE_WS_BASE = os.environ.get("BINANCE_WS_BASE", "wss://fstream.binance.com/ws").rstrip("/")
BASE_FLUSH_INTERVAL_MINUTES = max(1, int(os.environ.get("BASE_FLUSH_INTERVAL_MINUTES", "10")))
BASE_FLUSH_INTERVAL_SECONDS = BASE_FLUSH_INTERVAL_MINUTES * 60
ENGINE_STATUS_PATH = Path(DATA_DIR) / "engine_status.json"
VALID_SYMBOLS = set(SYMBOLS)
VALID_TIMEFRAMES = set(AVAILABLE_TIMEFRAMES)
ensure_storage(DATA_DIR)
preload_candles_cache(SYMBOLS, AVAILABLE_TIMEFRAMES)
app = FastAPI(title="Local Data Service", version="0.1.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def tf_to_ms(tf: str) -> int:
minutes = timeframe_to_minutes(tf)
if minutes is None:
logger.warning("无法解析时间周期,默认使用 60 秒", extra={"timeframe": tf})
return 60_000
return minutes * 60_000
def ensure_symbol_timeframe(symbol: str, timeframe: str) -> None:
if symbol not in VALID_SYMBOLS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"symbol 必须为 {sorted(VALID_SYMBOLS)} 之一。",
)
if timeframe not in VALID_TIMEFRAMES:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"tf 必须为 {sorted(VALID_TIMEFRAMES)} 之一。",
)
def parse_start_from_ms(val: str) -> int:
"""将 START_FROM 解析成毫秒级时间戳。
支持两种格式:
- YYYY-MM-DDUTC 00:00:00
- 整型毫秒时间戳字符串
"""
try:
return int(val)
except Exception:
pass
try:
dt = datetime.fromisoformat(val) # 允许 '2022-01-01' 或 '2022-01-01T00:00:00'
except Exception:
# 回退到固定日期
dt = datetime(2025, 9, 1)
return int(dt.timestamp() * 1000)
class Hub:
def __init__(self) -> None:
self.subscribers: Dict[str, List[WebSocket]] = {}
def topic(self, symbol: str, timeframe: str) -> str:
return f"candles::{symbol}::{timeframe}"
async def subscribe(self, ws: WebSocket, symbol: str, timeframe: str):
topic = self.topic(symbol, timeframe)
await ws.accept()
self.subscribers.setdefault(topic, []).append(ws)
def _clean(self, topic: str):
conns = self.subscribers.get(topic, [])
self.subscribers[topic] = [w for w in conns if not w.client_state.name == "DISCONNECTED"]
async def publish(self, symbol: str, timeframe: str, payload: dict):
topic = self.topic(symbol, timeframe)
conns = self.subscribers.get(topic, [])
if not conns:
return
message = json.dumps(payload, ensure_ascii=False)
dead: List[WebSocket] = []
for ws in conns:
try:
await ws.send_text(message)
except Exception:
dead.append(ws)
if dead:
self.subscribers[topic] = [w for w in conns if w not in dead]
hub = Hub()
fetch_tasks: List[asyncio.Task] = []
engine_runner_stop: Optional[Event] = None
engine_runner_task: Optional[asyncio.Task] = None
def resample_and_store(symbol: str, base_timeframe: str, derived_timeframes: List[str]) -> List[Tuple[str, List[CandleRow]]]:
if not derived_timeframes:
return []
base_tf_ms = tf_to_ms(base_timeframe)
updates: List[Tuple[str, List[List[float]]]] = []
for target_tf in derived_timeframes:
minutes = timeframe_to_minutes(target_tf)
if minutes is None:
logger.warning("无法解析聚合周期", extra={"target_timeframe": target_tf})
continue
last_ts = cache_get_last_timestamp(symbol, target_tf)
start_ts: Optional[int] = None
if last_ts is not None and base_tf_ms is not None and base_tf_ms > 0:
buffer_ms = minutes * 60_000 + base_tf_ms
start_ts = max(0, int(last_ts) - buffer_ms)
base_slice = cache_get(symbol, base_timeframe, start_ts, None)
if base_slice.empty:
continue
base_slice = base_slice.copy()
if "date" not in base_slice.columns:
base_slice["date"] = pd.to_datetime(base_slice["timestamp"], unit="ms", utc=True)
base_slice = (
base_slice.drop_duplicates(subset=["timestamp"], keep="last")
.sort_values("timestamp")
.reset_index(drop=True)
)
base_slice["timestamp"] = base_slice["timestamp"].astype("int64")
try:
if RESAMPLE_AVAILABLE:
derived_df = resample_to_interval(base_slice, minutes) # type: ignore[misc]
else:
derived_df = _fallback_resample_to_interval(base_slice, minutes)
except Exception:
logger.exception(
"聚合周期计算失败",
extra={"symbol": symbol, "base_timeframe": base_timeframe, "target_timeframe": target_tf},
)
continue
if derived_df is None or derived_df.empty:
continue
derived_df = derived_df.copy()
if "timestamp" not in derived_df.columns:
if "date" in derived_df.columns:
dates = pd.to_datetime(derived_df["date"], utc=True, errors="coerce")
derived_df["timestamp"] = (dates.astype("int64") // 1_000_000)
elif isinstance(derived_df.index, pd.DatetimeIndex):
idx = derived_df.index
if idx.tz is None:
idx = idx.tz_localize("UTC")
else:
idx = idx.tz_convert("UTC")
derived_df["timestamp"] = (idx.astype("int64") // 1_000_000)
if "timestamp" not in derived_df.columns:
logger.warning(
"聚合结果缺少 timestamp 列,已跳过",
extra={"target_timeframe": target_tf},
)
continue
derived_df = derived_df.dropna(subset=["timestamp", "open", "high", "low", "close", "volume"])
if derived_df.empty:
continue
derived_df["timestamp"] = derived_df["timestamp"].astype("int64")
derived_df = derived_df.sort_values("timestamp")
if last_ts is not None:
derived_df = derived_df[derived_df["timestamp"] > last_ts]
if derived_df.empty:
continue
numpy_rows = derived_df[["timestamp", "open", "high", "low", "close", "volume"]].to_numpy()
records: List[CandleRow] = []
for ts, o, h, l, c, v in numpy_rows:
records.append(
[
int(ts),
float(o),
float(h),
float(l),
float(c),
float(v),
]
)
if not records:
continue
cache_update(symbol, target_tf, records)
updates.append((target_tf, records[-3:] if len(records) > 3 else records))
return updates
def normalize_candles_for_timeframe(candles: List[CandleRow], tf_ms: int) -> Tuple[List[CandleRow], List[int]]:
if not candles:
return [], []
normalized_map: Dict[int, CandleRow] = {}
for row in candles:
if not row:
continue
try:
ts = int(row[0])
o = float(row[1])
h = float(row[2])
l = float(row[3])
c = float(row[4])
v = float(row[5])
except (TypeError, ValueError, IndexError):
continue
normalized_map[ts] = [ts, o, h, l, c, v]
ordered_ts = sorted(normalized_map.keys())
normalized: List[CandleRow] = []
missing: List[int] = []
last_ts: Optional[int] = None
for ts in ordered_ts:
normalized.append(normalized_map[ts])
if last_ts is not None and tf_ms > 0:
delta = ts - last_ts
if delta > tf_ms:
gap_ts = last_ts + tf_ms
while gap_ts < ts:
missing.append(gap_ts)
gap_ts += tf_ms
last_ts = ts
return normalized, missing
def compute_live_derived_updates(
symbol: str,
base_timeframe: str,
derived_timeframes: List[str],
base_tf_ms: int,
candles: List[CandleRow],
last_closed_ts: Optional[int],
) -> Dict[str, List[Tuple[CandleRow, bool]]]:
updates: Dict[str, List[Tuple[CandleRow, bool]]] = {}
if not candles or not derived_timeframes or base_tf_ms <= 0:
return updates
pending_updates: Dict[str, List[CandleRow]] = defaultdict(list)
derived_ms_map: Dict[str, int] = {}
max_multiplier = 1
for target_tf in derived_timeframes:
derived_ms = tf_to_ms(target_tf)
if derived_ms is None or derived_ms <= 0 or derived_ms % base_tf_ms != 0:
continue
multiplier = derived_ms // base_tf_ms
derived_ms_map[target_tf] = derived_ms
if multiplier > max_multiplier:
max_multiplier = multiplier
if not derived_ms_map:
return updates
window_ms = max_multiplier * base_tf_ms
newest_ts = max(int(row[0]) for row in candles if row)
base_start = newest_ts - window_ms + base_tf_ms
if base_start < 0:
base_start = 0
base_df = cache_get(symbol, base_timeframe, base_start, newest_ts)
if base_df.empty:
return updates
base_df = base_df.sort_values("timestamp")
base_rows: List[Tuple[int, float, float, float, float, float]] = []
for record in candles:
try:
ts = int(record[0])
if ts < base_start:
continue
base_rows.append(
(
ts,
float(record[1]),
float(record[2]),
float(record[3]),
float(record[4]),
float(record[5]),
)
)
except (TypeError, ValueError, IndexError):
continue
if base_rows:
temp_df = pd.DataFrame(
base_rows,
columns=["timestamp", "open", "high", "low", "close", "volume"],
)
base_df = pd.concat([base_df, temp_df], ignore_index=True)
if base_df.empty:
return updates
base_df = (
base_df.drop_duplicates(subset=["timestamp"], keep="last")
.sort_values("timestamp")
.reset_index(drop=True)
)
base_df_indexed = base_df.set_index("timestamp", drop=False)
if base_df_indexed.empty:
return updates
for target_tf, derived_ms in derived_ms_map.items():
multiplier = derived_ms // base_tf_ms
rows_with_status: List[Tuple[CandleRow, bool]] = []
latest_available_ts = int(base_df_indexed.index.max())
candidate_start = max(base_start, int(base_df_indexed.index.min()))
first_bucket = (candidate_start // derived_ms) * derived_ms
if first_bucket < candidate_start:
first_bucket += derived_ms
last_possible_start = latest_available_ts - (multiplier - 1) * base_tf_ms
current_start = first_bucket
while current_start <= last_possible_start:
expected_ts = [current_start + i * base_tf_ms for i in range(multiplier)]
subset = base_df_indexed.reindex(expected_ts)
if subset.isna().any().any():
current_start += derived_ms
continue
start_ts = current_start
end_ts = start_ts + derived_ms - base_tf_ms
row: CandleRow = [
start_ts,
float(subset.iloc[0]["open"]),
float(subset["high"].max()),
float(subset["low"].min()),
float(subset.iloc[-1]["close"]),
float(subset["volume"].sum()),
]
closed = last_closed_ts is not None and last_closed_ts >= end_ts
pending_updates[target_tf].append(row)
rows_with_status.append((row, closed))
current_start += derived_ms
if rows_with_status:
updates[target_tf] = rows_with_status
for target_tf, rows in pending_updates.items():
cache_update(symbol, target_tf, rows)
return updates
@dataclass
class FetchState:
symbol: str
timeframe: str
started_at: datetime = field(default_factory=datetime.utcnow)
last_fetch_at: Optional[datetime] = None
last_candle_ts: Optional[int] = None
consecutive_errors: int = 0
last_error: Optional[str] = None
def to_payload(self) -> dict:
def serialize_dt(dt: Optional[datetime]) -> Optional[str]:
if not dt:
return None
return dt.replace(microsecond=0).isoformat() + "Z"
return {
"symbol": self.symbol,
"timeframe": self.timeframe,
"started_at": serialize_dt(self.started_at),
"last_fetch_at": serialize_dt(self.last_fetch_at),
"last_candle_ts": self.last_candle_ts,
"consecutive_errors": self.consecutive_errors,
"last_error": self.last_error,
}
fetch_states: Dict[Tuple[str, str], FetchState] = {}
async def process_candles(
symbol: str,
timeframe: str,
candles: List[CandleRow],
derived_timeframes: List[str],
tf_ms: int,
finalized: bool,
closed_flags: Optional[List[bool]] = None,
) -> None:
if not candles:
return
if closed_flags is None or len(closed_flags) != len(candles):
closed_flags = [finalized] * len(candles)
state_key = (symbol, timeframe)
cache_update(symbol, timeframe, candles)
base_records = list(zip(candles, closed_flags))
last_closed_ts: Optional[int] = None
for row, is_closed in base_records:
if is_closed:
if last_closed_ts is None or row[0] > last_closed_ts:
last_closed_ts = row[0]
if last_closed_ts is None:
last_closed_ts = candles[-1][0] - tf_ms
logger.info(
"基础周期 K 线更新完成",
extra={
"symbol": symbol,
"timeframe": timeframe,
"count": len(candles),
"finalized": finalized,
"last_closed_ts": last_closed_ts,
},
)
derived_updates: List[Tuple[str, List[CandleRow]]] = []
live_derived_updates: Dict[str, List[Tuple[CandleRow, bool]]] = {}
if derived_timeframes:
needs_resample = finalized or len(candles) > 1
if needs_resample:
derived_updates = await asyncio.to_thread(
resample_and_store,
symbol,
timeframe,
derived_timeframes,
)
if derived_updates:
logger.info(
"衍生周期批量聚合完成",
extra={
"symbol": symbol,
"base_timeframe": timeframe,
"targets": [item[0] for item in derived_updates],
"origin": "resample" if RESAMPLE_AVAILABLE else "fallback",
},
)
live_derived_updates = await asyncio.to_thread(
compute_live_derived_updates,
symbol,
timeframe,
derived_timeframes,
tf_ms,
candles,
last_closed_ts,
)
if live_derived_updates:
logger.info(
"衍生周期实时聚合完成",
extra={
"symbol": symbol,
"base_timeframe": timeframe,
"targets": list(live_derived_updates.keys()),
"origin": "live",
},
)
for row, is_closed in base_records[-3:]:
payload = {
"topic": f"candles.{symbol}.{timeframe}",
"type": "upsert",
"data": {
"t": row[0],
"o": row[1],
"h": row[2],
"l": row[3],
"c": row[4],
"v": row[5],
"closed": bool(is_closed),
},
}
await hub.publish(symbol, timeframe, payload)
last_closed_ts_for_derived = last_closed_ts
for target_tf, rows in derived_updates:
if not rows:
continue
target_tf_ms = tf_to_ms(target_tf)
for row in rows:
ts = int(row[0])
o, h, l, c, v = map(float, row[1:])
if target_tf_ms and target_tf_ms > 0:
derived_closed = last_closed_ts_for_derived is not None and last_closed_ts_for_derived >= ts + target_tf_ms - tf_ms
else:
derived_closed = last_closed_ts_for_derived is not None and last_closed_ts_for_derived >= ts
derived_state_key = (symbol, target_tf)
derived_state = fetch_states.get(derived_state_key)
if derived_state is None:
derived_state = FetchState(symbol=symbol, timeframe=target_tf)
fetch_states[derived_state_key] = derived_state
derived_state.last_fetch_at = datetime.utcnow()
derived_state.last_candle_ts = ts
derived_state.consecutive_errors = 0
derived_state.last_error = None
payload = {
"topic": f"candles.{symbol}.{target_tf}",
"type": "upsert",
"data": {
"t": ts,
"o": o,
"h": h,
"l": l,
"c": c,
"v": v,
"closed": bool(derived_closed),
},
}
await hub.publish(symbol, target_tf, payload)
if live_derived_updates:
for target_tf, items in live_derived_updates.items():
if not items:
continue
for row, derived_closed in items:
ts = int(row[0])
derived_state_key = (symbol, target_tf)
derived_state = fetch_states.get(derived_state_key)
if derived_state is None:
derived_state = FetchState(symbol=symbol, timeframe=target_tf)
fetch_states[derived_state_key] = derived_state
derived_state.last_fetch_at = datetime.utcnow()
derived_state.last_candle_ts = ts
derived_state.consecutive_errors = 0
derived_state.last_error = None
payload = {
"topic": f"candles.{symbol}.{target_tf}",
"type": "upsert",
"data": {
"t": ts,
"o": float(row[1]),
"h": float(row[2]),
"l": float(row[3]),
"c": float(row[4]),
"v": float(row[5]),
"closed": bool(derived_closed),
},
}
await hub.publish(symbol, target_tf, payload)
state = fetch_states.get(state_key)
if state:
state.last_fetch_at = datetime.utcnow()
state.last_candle_ts = candles[-1][0]
state.consecutive_errors = 0
state.last_error = None
# 验证逻辑已移除,衍生周期的缺口依赖轮询补齐
async def rest_catchup(
symbol: str,
timeframe: str,
derived_timeframes: List[str],
tf_ms: int,
start_since: int,
) -> None:
state_key = (symbol, timeframe)
state = fetch_states[state_key]
exchange = build_exchange()
since = start_since
backoff = 1.0
gap_retry: Dict[int, int] = {}
logger.info("开始 REST 补齐历史", extra={"symbol": symbol, "timeframe": timeframe, "since": since})
try:
while True:
now_ms = int(datetime.utcnow().timestamp() * 1000)
if since >= now_ms - tf_ms:
break
try:
async with REST_FETCH_SEMAPHORE:
candles = await exchange.fetch_ohlcv(symbol, timeframe, since=since, limit=1000)
except asyncio.CancelledError:
raise
except (ccxt.NetworkError, ccxt.ExchangeNotAvailable, ccxt.RequestTimeout) as exc:
logger.warning(
"历史补齐网络异常,准备重试",
extra={"symbol": symbol, "timeframe": timeframe, "error": str(exc)},
)
state.last_error = str(exc)
state.consecutive_errors += 1
backoff = min(backoff * BACKOFF_BASE, BACKOFF_MAX)
await asyncio.sleep(backoff)
continue
except Exception as exc:
logger.exception(
"历史补齐发生异常,准备重试",
extra={"symbol": symbol, "timeframe": timeframe},
)
state.last_error = str(exc)
state.consecutive_errors += 1
backoff = min(backoff * BACKOFF_BASE, BACKOFF_MAX)
await asyncio.sleep(backoff)
continue
if not candles:
break
candles, missing_ts = normalize_candles_for_timeframe(candles, tf_ms)
if not candles:
since += tf_ms
backoff = 1.0
await asyncio.sleep(0.2)
continue
await process_candles(
symbol,
timeframe,
candles,
derived_timeframes,
tf_ms,
finalized=True,
closed_flags=[True] * len(candles),
)
state.consecutive_errors = 0
state.last_error = None
backoff = 1.0
if missing_ts:
gap_start = missing_ts[0]
attempts = gap_retry.get(gap_start, 0) + 1
gap_retry[gap_start] = attempts
if attempts <= 3:
logger.warning(
"检测到缺失 K 线,准备回补",
extra={
"symbol": symbol,
"timeframe": timeframe,
"missing_from": gap_start,
"missing_to": missing_ts[-1],
"attempt": attempts,
},
)
since = gap_start
await asyncio.sleep(0.2)
continue
logger.error(
"缺失 K 线多次回补失败,已跳过",
extra={
"symbol": symbol,
"timeframe": timeframe,
"missing_from": gap_start,
"missing_to": missing_ts[-1],
},
)
gap_retry.pop(gap_start, None)
else:
gap_retry.clear()
since = candles[-1][0] + tf_ms
now_ms = int(datetime.utcnow().timestamp() * 1000)
lag = now_ms - since
if lag > tf_ms * 10:
await asyncio.sleep(0.2)
else:
await asyncio.sleep(max(1.0, tf_ms * POLL_FACTOR / 1000.0))
finally:
with suppress(Exception):
await exchange.close()
logger.info("REST 补齐完成", extra={"symbol": symbol, "timeframe": timeframe, "latest": state.last_candle_ts})
async def rest_poll_loop(
symbol: str,
timeframe: str,
derived_timeframes: List[str],
tf_ms: int,
) -> None:
state_key = (symbol, timeframe)
window = max(REST_POLL_WINDOW, 1)
interval = max(REST_POLL_INTERVAL, 1.0)
exchange = build_exchange()
try:
while True:
state = fetch_states.get(state_key)
latest_ts = state.last_candle_ts if state else None
if latest_ts is None or latest_ts <= 0:
since = parse_start_from_ms(START_FROM)
else:
since = max(0, latest_ts - (window - 1) * tf_ms)
try:
async with REST_FETCH_SEMAPHORE:
candles = await exchange.fetch_ohlcv(
symbol,
timeframe,
since=since,
limit=max(window + 2, window),
)
except asyncio.CancelledError:
raise
except (ccxt.NetworkError, ccxt.ExchangeNotAvailable, ccxt.RequestTimeout) as exc:
logger.warning(
"实时轮询网络异常,准备重试",
extra={"symbol": symbol, "timeframe": timeframe, "error": str(exc)},
)
await asyncio.sleep(interval)
continue
except Exception as exc:
logger.exception(
"实时轮询发生异常",
extra={"symbol": symbol, "timeframe": timeframe},
)
await asyncio.sleep(interval)
continue
candles, _ = normalize_candles_for_timeframe(candles, tf_ms)
if candles:
closed_flags = [True] * len(candles)
logger.info(
"轮询拉取基础周期完成",
extra={
"symbol": symbol,
"timeframe": timeframe,
"count": len(candles),
"since": since,
"mode": "rest_poll",
},
)
try:
await process_candles(
symbol,
timeframe,
candles,
derived_timeframes,
tf_ms,
finalized=True,
closed_flags=closed_flags,
)
except asyncio.CancelledError:
raise
except Exception as exc:
logger.exception(
"处理基础周期 K 线失败 [%s %s]",
symbol,
timeframe,
)
state = fetch_states.get(state_key)
if state:
state.last_error = str(exc)
state.consecutive_errors += 1
await asyncio.sleep(interval)
continue
await asyncio.sleep(interval)
except asyncio.CancelledError:
raise
finally:
with suppress(Exception):
await exchange.close()
logger.info("轮询任务退出", extra={"symbol": symbol, "timeframe": timeframe})
async def stream_loop(symbol: str, timeframe: str, derived_timeframes: List[str], tf_ms: int):
state_key = (symbol, timeframe)
url = build_stream_url(symbol, timeframe)
while True:
try:
async with websockets.connect(url, ping_interval=20, ping_timeout=20) as ws:
logger.info("WebSocket 已连接", extra={"symbol": symbol, "timeframe": timeframe, "url": url})
async for message in ws:
data = json.loads(message)
kline = data.get("k")
if not kline:
continue
is_closed = bool(kline.get("x"))
row: CandleRow = [
int(kline["t"]),
float(kline["o"]),
float(kline["h"]),
float(kline["l"]),
float(kline["c"]),
float(kline["v"]),
]
await process_candles(
symbol,
timeframe,
[row],
derived_timeframes,
tf_ms,
finalized=is_closed,
closed_flags=[is_closed],
)
except asyncio.CancelledError:
logger.info("取消 WebSocket 任务", extra={"symbol": symbol, "timeframe": timeframe})
raise
except Exception as exc:
logger.warning(
"WebSocket 连接异常,准备重连",
extra={"symbol": symbol, "timeframe": timeframe, "error": str(exc)},
)
state = fetch_states.get(state_key)
start_since = None
if state and state.last_candle_ts:
start_since = state.last_candle_ts + tf_ms
if start_since:
await rest_catchup(symbol, timeframe, derived_timeframes, tf_ms, start_since)
await asyncio.sleep(5.0)
def build_exchange():
if EXCHANGE.lower() == "binance":
return ccxt_async.binance(
{
"enableRateLimit": True,
"timeout": 20_000,
"options": {
"adjustForTimeDifference": True,
"defaultType": "future",
"defaultSubType": "linear",
"defaultMarket": "future",
"defaultSettle": "USDT",
},
}
)
raise RuntimeError(f"Unsupported EXCHANGE: {EXCHANGE}")
async def fetch_loop(symbol: str, timeframe: str):
"""初次通过 REST 补齐历史,随后持续轮询/流式拉取增量。"""
derived_timeframes = AGGREGATION_TARGETS.get(timeframe, [])
global RESAMPLE_WARNING_EMITTED
if derived_timeframes and not RESAMPLE_AVAILABLE and not RESAMPLE_WARNING_EMITTED:
logger.warning(
"缺少 technical.util.resample_to_interval 模块,聚合时间周期生成已跳过",
extra={"timeframe": timeframe},
)
RESAMPLE_WARNING_EMITTED = True
tf_ms = tf_to_ms(timeframe)
state_key = (symbol, timeframe)
last_ts = cache_get_last_timestamp(symbol, timeframe)
fetch_states[state_key] = FetchState(symbol=symbol, timeframe=timeframe, last_candle_ts=last_ts)
initial_sync_flushed = False
start_from = parse_start_from_ms(START_FROM)
backoff = 1.0
try:
while True:
state = fetch_states[state_key]
state.started_at = datetime.utcnow()
if state.last_candle_ts is not None:
rewind_since = max(0, state.last_candle_ts - tf_ms)
initial_since = max(start_from, rewind_since)
else:
initial_since = start_from
logger.info(
"启动拉取任务 [%s %s] since=%s",
symbol,
timeframe,
initial_since,
)
try:
await rest_catchup(symbol, timeframe, derived_timeframes, tf_ms, initial_since)
if not initial_sync_flushed:
await flush_dirty_base_snapshots(force_all=True)
logger.info(
"初次同步完成,基础周期数据已写盘",
extra={"symbol": symbol, "timeframe": timeframe},
)
initial_sync_flushed = True
if WS_ENABLED:
await stream_loop(symbol, timeframe, derived_timeframes, tf_ms)
else:
await rest_poll_loop(symbol, timeframe, derived_timeframes, tf_ms)
except asyncio.CancelledError:
logger.info("取消拉取任务 [%s %s]", symbol, timeframe)
state.last_error = "cancelled"
raise
except Exception as exc:
state.last_error = str(exc)
state.consecutive_errors += 1
logger.exception(
"拉取任务异常 [%s %s]%.1f 秒后重启",
symbol,
timeframe,
backoff,
)
await asyncio.sleep(backoff)
backoff = min(backoff * BACKOFF_BASE, BACKOFF_MAX)
continue
else:
backoff = 1.0
logger.warning("拉取循环提前结束 [%s %s]1 秒后重启", symbol, timeframe)
await asyncio.sleep(1.0)
finally:
logger.info("拉取任务退出", extra={"symbol": symbol, "timeframe": timeframe})
@app.on_event("startup")
async def on_start():
logger.info("API 服务启动完成")
if IS_ENGINE_PROCESS:
global engine_runner_stop, engine_runner_task
if engine_runner_task is None or engine_runner_task.done():
engine_runner_stop = Event()
engine_runner_task = asyncio.create_task(run_engine(engine_runner_stop))
@app.on_event("shutdown")
async def on_shutdown():
logger.info("API 服务准备退出")
if IS_ENGINE_PROCESS:
global engine_runner_stop, engine_runner_task
if engine_runner_stop is not None:
engine_runner_stop.set()
if engine_runner_task is not None:
with suppress(Exception):
await engine_runner_task
engine_runner_task = None
engine_runner_stop = None
@app.get("/health")
async def health():
now = datetime.utcnow().replace(microsecond=0).isoformat() + "Z"
engine_status = load_engine_status_snapshot() or {"updated_at": None, "tasks": [], "queues": {}}
return {
"status": "ok",
"time": now,
"exchange": EXCHANGE,
"symbols": SYMBOLS,
"base_timeframes": FETCH_TIMEFRAMES,
"derived_timeframes": DERIVED_TIMEFRAMES,
"timeframes": AVAILABLE_TIMEFRAMES,
"engine": engine_status,
"tasks": engine_status.get("tasks", []),
}
@app.get("/api/candles")
def api_candles(
symbol: str = Query(..., description="如 BTC/USDT:USDT"),
tf: str = Query("1m", description="时间周期"),
start: Optional[int] = Query(None, description="开始时间戳(ms)"),
end: Optional[int] = Query(None, description="结束时间戳(ms)"),
):
try:
ensure_symbol_timeframe(symbol, tf)
df = cache_get(symbol, tf, start, end)
records = df.to_dict("records") if not df.empty else []
return JSONResponse(records)
except Exception as e:
return JSONResponse({"error": str(e)}, status_code=500)
@app.websocket("/ws")
async def ws_endpoint(websocket: WebSocket, symbol: str, tf: str, since: Optional[int] = None):
if symbol not in VALID_SYMBOLS or tf not in VALID_TIMEFRAMES:
await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="invalid symbol/timeframe")
return
await hub.subscribe(websocket, symbol, tf)
try:
snap = cache_get(symbol, tf, since, None)
await websocket.send_text(
json.dumps(
{
"topic": f"candles.{symbol}.{tf}",
"type": "snapshot",
"data": [
{"t": int(r["timestamp"]), "o": r["open"], "h": r["high"], "l": r["low"], "c": r["close"], "v": r["volume"]}
for _, r in snap.iterrows()
],
},
ensure_ascii=False,
)
)
except Exception:
pass
try:
while True:
await asyncio.sleep(30)
await websocket.send_text(json.dumps({"type": "ping", "ts": int(datetime.utcnow().timestamp() * 1000)}))
except WebSocketDisconnect:
return
@app.get("/")
def root():
return {
"service": "Local Data Service",
"exchange": EXCHANGE,
"symbols": SYMBOLS,
"base_timeframes": FETCH_TIMEFRAMES,
"derived_timeframes": DERIVED_TIMEFRAMES,
"timeframes": AVAILABLE_TIMEFRAMES,
}
def _fallback_resample_to_interval(df: pd.DataFrame, minutes: int) -> pd.DataFrame:
if df.empty or minutes <= 0:
return pd.DataFrame(columns=CANDLE_COLUMNS)
working = df.copy()
if "timestamp" not in working.columns:
return pd.DataFrame(columns=CANDLE_COLUMNS)
working["date"] = pd.to_datetime(working["timestamp"], unit="ms", utc=True)
working = working.set_index("date", drop=True)
columns = ["open", "high", "low", "close", "volume"]
for column in columns:
if column not in working.columns:
working[column] = 0.0
working = working[columns]
rule = f"{minutes}T"
aggregated = working.resample(rule, label="left", closed="left").agg(
{
"open": "first",
"high": "max",
"low": "min",
"close": "last",
"volume": "sum",
}
)
aggregated = aggregated.dropna(subset=["open", "high", "low", "close"]).reset_index()
aggregated["timestamp"] = (aggregated["date"].astype("int64") // 1_000_000)
aggregated = aggregated.drop(columns=["date"], errors="ignore")
aggregated = aggregated.dropna(subset=["timestamp"]).reset_index(drop=True)
aggregated["timestamp"] = aggregated["timestamp"].astype("int64")
return aggregated[CANDLE_COLUMNS]