rebuild(stage-5): establish market data gateway and charts
This commit is contained in:
@@ -0,0 +1,159 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections.abc import Callable
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from backend.data.contracts import (
|
||||
DataSource,
|
||||
DataUsage,
|
||||
ObservationMetadata,
|
||||
ProviderResult,
|
||||
SnapshotState,
|
||||
)
|
||||
from backend.data.providers.base import ProviderError
|
||||
|
||||
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
|
||||
|
||||
class TushareProvider:
|
||||
source = DataSource.TUSHARE
|
||||
url = "http://api.tushare.pro"
|
||||
|
||||
def __init__(self, token: str | None | Callable[[], str | None], timeout: int = 20) -> None:
|
||||
self._token_provider = token if callable(token) else lambda: token
|
||||
self._timeout = timeout
|
||||
|
||||
@property
|
||||
def configured(self) -> bool:
|
||||
return bool(self._token())
|
||||
|
||||
def calendar(self, start_date: str, end_date: str) -> ProviderResult:
|
||||
result = self._query(
|
||||
"trade_cal",
|
||||
{"exchange": "SSE", "start_date": _compact(start_date), "end_date": _compact(end_date)},
|
||||
"cal_date,is_open,pretrade_date",
|
||||
unit="calendar_day",
|
||||
)
|
||||
days = (
|
||||
datetime.fromisoformat(end_date).date() - datetime.fromisoformat(start_date).date()
|
||||
).days + 1
|
||||
return ProviderResult(
|
||||
result.rows,
|
||||
replace(result.metadata, coverage=min(len(result.rows) / max(days, 1), 1)),
|
||||
)
|
||||
|
||||
def entities(self) -> ProviderResult:
|
||||
rows: list[dict[str, Any]] = []
|
||||
for status in ("L", "P", "D"):
|
||||
result = self._query(
|
||||
"stock_basic",
|
||||
{"exchange": "", "list_status": status},
|
||||
"ts_code,symbol,name,industry,list_status,list_date,delist_date",
|
||||
unit="entity",
|
||||
)
|
||||
rows.extend(result.rows)
|
||||
return ProviderResult(
|
||||
tuple(rows), _metadata(self.source, "entity", min(len(rows) / 5300, 1))
|
||||
)
|
||||
|
||||
def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult:
|
||||
api_name = "index_daily" if entity_type == "index" else "daily"
|
||||
if entity_type in {"sector", "theme"}:
|
||||
api_name = "ths_daily"
|
||||
end = datetime.strptime(_compact(end_date), "%Y%m%d")
|
||||
start = (end - timedelta(days=380)).strftime("%Y%m%d")
|
||||
return self._query(
|
||||
api_name,
|
||||
{"ts_code": identifier, "start_date": start, "end_date": end.strftime("%Y%m%d")},
|
||||
"ts_code,trade_date,open,high,low,close,vol,amount,pct_chg",
|
||||
unit="yuan/share",
|
||||
adjustment="unadjusted",
|
||||
)
|
||||
|
||||
def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult:
|
||||
if entity_type != "stock":
|
||||
raise ProviderError("Tushare minute charts only support stocks")
|
||||
date = _display(trade_date)
|
||||
return self._query(
|
||||
"stk_mins",
|
||||
{
|
||||
"ts_code": identifier,
|
||||
"freq": "1min",
|
||||
"start_date": f"{date} 09:30:00",
|
||||
"end_date": f"{date} 15:00:00",
|
||||
},
|
||||
"ts_code,trade_time,open,high,low,close,vol,amount",
|
||||
unit="yuan/share",
|
||||
)
|
||||
|
||||
def _query(
|
||||
self,
|
||||
api_name: str,
|
||||
params: dict[str, Any],
|
||||
fields: str,
|
||||
*,
|
||||
unit: str,
|
||||
adjustment: str = "not_applicable",
|
||||
) -> ProviderResult:
|
||||
if not self.configured:
|
||||
raise ProviderError("行情服务尚未配置")
|
||||
token = self._token()
|
||||
if not token:
|
||||
raise ProviderError("行情服务尚未配置")
|
||||
body = json.dumps(
|
||||
{"api_name": api_name, "token": token, "params": params, "fields": fields},
|
||||
ensure_ascii=False,
|
||||
).encode("utf-8")
|
||||
request = urllib.request.Request(
|
||||
self.url,
|
||||
data=body,
|
||||
headers={"Content-Type": "application/json", "User-Agent": "XiaobaiReview/2"},
|
||||
method="POST",
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(request, timeout=self._timeout) as response:
|
||||
payload = json.loads(response.read().decode("utf-8"))
|
||||
except (urllib.error.URLError, TimeoutError, OSError, json.JSONDecodeError) as exc:
|
||||
raise ProviderError("行情服务请求失败") from exc
|
||||
if payload.get("code") not in (None, 0):
|
||||
raise ProviderError(str(payload.get("msg") or "行情服务拒绝请求"))
|
||||
data = payload.get("data") or {}
|
||||
columns = data.get("fields") or []
|
||||
rows = tuple(dict(zip(columns, item, strict=False)) for item in data.get("items") or [])
|
||||
return ProviderResult(rows, _metadata(self.source, unit, 1 if rows else 0, adjustment))
|
||||
|
||||
def _token(self) -> str:
|
||||
return str(self._token_provider() or "").strip()
|
||||
|
||||
|
||||
def _metadata(
|
||||
source: DataSource, unit: str, coverage: float, adjustment: str = "not_applicable"
|
||||
) -> ObservationMetadata:
|
||||
return ObservationMetadata(
|
||||
source=source,
|
||||
observed_at=datetime.now(SHANGHAI),
|
||||
unit=unit,
|
||||
adjustment=adjustment,
|
||||
freshness_seconds=0,
|
||||
coverage=coverage,
|
||||
state=SnapshotState.ARCHIVE,
|
||||
usage=DataUsage.CALCULATION,
|
||||
)
|
||||
|
||||
|
||||
def _compact(value: str) -> str:
|
||||
normalized = value.replace("-", "")
|
||||
if len(normalized) != 8 or not normalized.isdigit():
|
||||
raise ProviderError("日期格式无效")
|
||||
return normalized
|
||||
|
||||
|
||||
def _display(value: str) -> str:
|
||||
compact = _compact(value)
|
||||
return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"
|
||||
Reference in New Issue
Block a user