更新data_provider逻辑,能够更快开始提供服务,添加说明
This commit is contained in:
+41
-17
@@ -19,6 +19,7 @@ import ccxt # type: ignore
|
||||
import pandas as pd # type: ignore
|
||||
from fastapi import FastAPI, HTTPException, Query, WebSocket, WebSocketDisconnect
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import HTMLResponse
|
||||
import uvicorn
|
||||
from technical.util import resample_to_interval
|
||||
|
||||
@@ -405,11 +406,23 @@ class DataProvider:
|
||||
return results
|
||||
|
||||
def initialize(self) -> None:
|
||||
"""阻塞式启动:加载本地、从倒数第二根或配置起点补历史、写盘并 set _ready。"""
|
||||
logger.info("开始初始化数据提供商")
|
||||
"""快速启动:仅加载本地磁盘已有数据到内存,然后立即标记就绪,不阻塞服务。"""
|
||||
logger.info("开始加载本地数据")
|
||||
for symbol in self.symbols:
|
||||
for timeframe in self.timeframes:
|
||||
existing = self._load_local(symbol, timeframe)
|
||||
with self._lock:
|
||||
self.data.setdefault(symbol, {})[timeframe] = existing
|
||||
self._ready.set()
|
||||
logger.info("本地数据加载完成,服务已就绪")
|
||||
|
||||
def run_initial_history_fetch(self) -> None:
|
||||
"""后台一次性拉取所有 symbol/tf 的历史数据(从本地末根或配置起点到当前),然后落盘。"""
|
||||
logger.info("开始后台历史数据拉取")
|
||||
for symbol in self.symbols:
|
||||
for timeframe in self.timeframes:
|
||||
with self._lock:
|
||||
existing = list(self.data.get(symbol, {}).get(timeframe, []))
|
||||
tf_ms = TIMEFRAME_TO_MS[timeframe]
|
||||
last_ts = existing[-1]["timestamp"] if existing else None
|
||||
if last_ts is not None:
|
||||
@@ -433,10 +446,10 @@ class DataProvider:
|
||||
history = self._fetch_history(symbol, timeframe, fetch_since)
|
||||
merged = self._merge_candles(timeframe, existing, history)
|
||||
with self._lock:
|
||||
self.data.setdefault(symbol, {})[timeframe] = merged
|
||||
self.data[symbol][timeframe] = merged
|
||||
self._write_to_disk(symbol, timeframe, merged)
|
||||
self._ready.set()
|
||||
logger.info("数据初始化完成")
|
||||
self._notify_update(symbol, timeframe)
|
||||
logger.info("后台历史数据拉取完成")
|
||||
|
||||
def resample_df(self, df: pd.DataFrame, interval: int) -> pd.DataFrame:
|
||||
"""将基础周期 DataFrame 聚合为 interval 分钟周期(freqtrade technical.util)。"""
|
||||
@@ -788,8 +801,15 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
loop = asyncio.get_running_loop()
|
||||
ws_manager.set_loop(loop)
|
||||
provider.on_update(_on_data_update)
|
||||
# 快速加载本地数据后立即就绪,不阻塞服务启动
|
||||
await loop.run_in_executor(None, provider.initialize)
|
||||
provider.start_background_workers()
|
||||
# 后台拉取历史数据补齐(不阻塞 HTTP/WS 服务)
|
||||
threading.Thread(
|
||||
target=provider.run_initial_history_fetch,
|
||||
name="initial-history-fetch",
|
||||
daemon=True,
|
||||
).start()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
@@ -840,18 +860,15 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
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,
|
||||
"symbols": provider.symbols,
|
||||
"base_timeframes": provider.timeframes,
|
||||
"derived_timeframes": provider.get_derived_timeframes(),
|
||||
"timeframes": provider.get_available_timeframes(),
|
||||
"ready": provider.is_ready(),
|
||||
}
|
||||
homepage_path = Path(__file__).resolve().parent / "homepage.html"
|
||||
docs_path = Path(__file__).resolve().parent / "api_docs.html"
|
||||
|
||||
@app.get("/", response_class=HTMLResponse)
|
||||
async def root():
|
||||
"""服务主页。"""
|
||||
if homepage_path.exists():
|
||||
return HTMLResponse(content=homepage_path.read_text(encoding="utf-8"))
|
||||
return HTMLResponse(content="<h1>主页页面未找到</h1>", status_code=404)
|
||||
|
||||
@app.websocket("/ws")
|
||||
async def websocket_endpoint(ws: WebSocket):
|
||||
@@ -921,6 +938,13 @@ def create_app(provider: DataProvider) -> FastAPI:
|
||||
finally:
|
||||
await ws_manager.disconnect(ws)
|
||||
|
||||
@app.get("/api/docs", response_class=HTMLResponse, include_in_schema=False)
|
||||
async def api_docs():
|
||||
"""返回自定义 API 文档页面。"""
|
||||
if docs_path.exists():
|
||||
return HTMLResponse(content=docs_path.read_text(encoding="utf-8"))
|
||||
return HTMLResponse(content="<h1>API 文档页面未找到</h1>", status_code=404)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user