1458 lines
52 KiB
Python
1458 lines
52 KiB
Python
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-DD(UTC 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]
|
||
|
||
|