160 lines
5.5 KiB
Python
160 lines
5.5 KiB
Python
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:]}"
|