rebuild(stage-5): establish market data gateway and charts
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
from backend.data.providers.eastmoney import EastmoneyProvider
|
||||
from backend.data.providers.ifind import IfindProvider
|
||||
from backend.data.providers.tushare import TushareProvider
|
||||
|
||||
__all__ = ["EastmoneyProvider", "IfindProvider", "TushareProvider"]
|
||||
@@ -0,0 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from backend.data.contracts import DataSource, ProviderResult
|
||||
|
||||
|
||||
class ProviderError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class MarketDataProvider(Protocol):
|
||||
source: DataSource
|
||||
|
||||
@property
|
||||
def configured(self) -> bool: ...
|
||||
|
||||
def calendar(self, start_date: str, end_date: str) -> ProviderResult: ...
|
||||
|
||||
def entities(self) -> ProviderResult: ...
|
||||
|
||||
def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult: ...
|
||||
|
||||
def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult: ...
|
||||
@@ -0,0 +1,105 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from datetime import datetime
|
||||
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")
|
||||
INDEX_CODES = {"000001.SH": "1.000001", "399001.SZ": "0.399001", "399006.SZ": "0.399006"}
|
||||
|
||||
|
||||
class EastmoneyProvider:
|
||||
source = DataSource.EASTMONEY
|
||||
url = "https://push2delay.eastmoney.com/api/qt/stock/trends2/get"
|
||||
|
||||
def __init__(self, timeout: int = 6) -> None:
|
||||
self._timeout = timeout
|
||||
|
||||
@property
|
||||
def configured(self) -> bool:
|
||||
return True
|
||||
|
||||
def calendar(self, start_date: str, end_date: str) -> ProviderResult:
|
||||
raise ProviderError("The display provider is not a calendar authority")
|
||||
|
||||
def entities(self) -> ProviderResult:
|
||||
raise ProviderError("The display provider is not an entity authority")
|
||||
|
||||
def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult:
|
||||
raise ProviderError("The display provider does not supply canonical daily bars")
|
||||
|
||||
def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult:
|
||||
secid = self._secid(entity_type, identifier)
|
||||
params = urllib.parse.urlencode(
|
||||
{
|
||||
"secid": secid,
|
||||
"fields1": "f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f11,f12,f13",
|
||||
"fields2": "f51,f52,f53,f54,f55,f56,f57,f58",
|
||||
"iscr": "0",
|
||||
"ndays": "1",
|
||||
}
|
||||
)
|
||||
request = urllib.request.Request(
|
||||
f"{self.url}?{params}",
|
||||
headers={"Accept": "application/json", "User-Agent": "XiaobaiReview/2"},
|
||||
)
|
||||
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
|
||||
data = payload.get("data") or {}
|
||||
rows = []
|
||||
for raw in data.get("trends") or []:
|
||||
fields = str(raw).split(",")
|
||||
if len(fields) < 8 or " " not in fields[0]:
|
||||
continue
|
||||
date, time = fields[0].split(" ", 1)
|
||||
if date != trade_date or not "09:30" <= time[:5] <= "15:00":
|
||||
continue
|
||||
rows.append(
|
||||
{
|
||||
"time": fields[0],
|
||||
"open": fields[1],
|
||||
"close": fields[2],
|
||||
"high": fields[3],
|
||||
"low": fields[4],
|
||||
"volume": fields[5],
|
||||
"amount": fields[6],
|
||||
"avgPrice": fields[7],
|
||||
"preClose": data.get("preClose"),
|
||||
}
|
||||
)
|
||||
metadata = ObservationMetadata(
|
||||
source=self.source,
|
||||
observed_at=datetime.now(SHANGHAI),
|
||||
unit="yuan/share",
|
||||
adjustment="unadjusted",
|
||||
freshness_seconds=0,
|
||||
coverage=1 if rows else 0,
|
||||
state=SnapshotState.REALTIME,
|
||||
usage=DataUsage.DISPLAY,
|
||||
)
|
||||
return ProviderResult(tuple(rows), metadata)
|
||||
|
||||
@staticmethod
|
||||
def _secid(entity_type: str, identifier: str) -> str:
|
||||
if entity_type == "index" and identifier in INDEX_CODES:
|
||||
return INDEX_CODES[identifier]
|
||||
code = identifier.split(".")[0]
|
||||
if entity_type == "stock" and len(code) == 6 and code.isdigit():
|
||||
market = "1" if code.startswith(("5", "6", "9")) else "0"
|
||||
return f"{market}.{code}"
|
||||
raise ProviderError("该标的暂无展示分时数据")
|
||||
@@ -0,0 +1,233 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
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 IfindProvider:
|
||||
source = DataSource.IFIND
|
||||
base_url = "https://quantapi.51ifind.com/api/v1"
|
||||
auth_error_codes = {-1302, -1303, -1304, -4302, -4303}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
refresh_token: str | None | Callable[[], str | None],
|
||||
access_token: str | None | Callable[[], str | None],
|
||||
timeout: int = 15,
|
||||
) -> None:
|
||||
self._refresh_provider = refresh_token if callable(refresh_token) else lambda: refresh_token
|
||||
self._access_provider = access_token if callable(access_token) else lambda: access_token
|
||||
self._issued_access = ""
|
||||
self._timeout = timeout
|
||||
self._lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def configured(self) -> bool:
|
||||
return bool(self._refresh() or self._configured_access())
|
||||
|
||||
def calendar(self, start_date: str, end_date: str) -> ProviderResult:
|
||||
raise ProviderError("iFinD is not the calendar authority")
|
||||
|
||||
def entities(self) -> ProviderResult:
|
||||
raise ProviderError("iFinD is not the entity-directory authority")
|
||||
|
||||
def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult:
|
||||
end = datetime.strptime(_compact(end_date), "%Y%m%d")
|
||||
start = end.replace(year=end.year - 1).strftime("%Y-%m-%d")
|
||||
payload = self._request(
|
||||
"cmd_history_quotation",
|
||||
{
|
||||
"codes": identifier,
|
||||
"indicators": "open,high,low,close,volume,amount",
|
||||
"startdate": start,
|
||||
"enddate": end.strftime("%Y-%m-%d"),
|
||||
"functionpara": {"CPS": "forward1", "Fill": "Omit"},
|
||||
},
|
||||
)
|
||||
return _result(payload, "yuan/share", "forward", SnapshotState.ARCHIVE)
|
||||
|
||||
def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult:
|
||||
date = _display(trade_date)
|
||||
payload = self._request(
|
||||
"high_frequency",
|
||||
{
|
||||
"codes": identifier,
|
||||
"indicators": "open,high,low,close,volume,amount,avgPrice",
|
||||
"starttime": f"{date} 09:30:00",
|
||||
"endtime": f"{date} 15:00:00",
|
||||
"functionpara": {
|
||||
"CPS": "forward1",
|
||||
"Fill": "Previous",
|
||||
"Timeformat": "LocalTime",
|
||||
"Interval": "1",
|
||||
"Limitstart": "09:30:00",
|
||||
"Limitend": "15:00:00",
|
||||
},
|
||||
},
|
||||
)
|
||||
return _result(payload, "yuan/share", "forward", SnapshotState.REALTIME)
|
||||
|
||||
def _request(self, endpoint: str, body: dict[str, Any]) -> dict[str, Any]:
|
||||
if not self.configured:
|
||||
raise ProviderError("实时行情服务尚未配置")
|
||||
payload = self._post(endpoint, body, self._access())
|
||||
code = _error_code(payload)
|
||||
if self._auth_error(payload) and self._refresh():
|
||||
payload = self._post(endpoint, body, self._refresh_access())
|
||||
code = _error_code(payload)
|
||||
if code != 0:
|
||||
raise ProviderError(
|
||||
str(payload.get("errmsg") or payload.get("message") or "实时行情服务拒绝请求")
|
||||
)
|
||||
return payload
|
||||
|
||||
def _access(self) -> str:
|
||||
with self._lock:
|
||||
if self._issued_access:
|
||||
return self._issued_access
|
||||
configured = self._configured_access()
|
||||
if configured:
|
||||
return configured
|
||||
refresh = self._refresh()
|
||||
if not refresh:
|
||||
raise ProviderError("实时行情服务尚未配置")
|
||||
payload = self._post("get_access_token", {}, "", refresh)
|
||||
token = str((payload.get("data") or {}).get("access_token") or "").strip()
|
||||
if not token:
|
||||
raise ProviderError("实时行情服务授权失败")
|
||||
self._issued_access = token
|
||||
return token
|
||||
|
||||
def _refresh_access(self) -> str:
|
||||
refresh = self._refresh()
|
||||
if not refresh:
|
||||
raise ProviderError("实时行情服务授权失败")
|
||||
with self._lock:
|
||||
payload = self._post("get_access_token", {}, "", refresh)
|
||||
token = str((payload.get("data") or {}).get("access_token") or "").strip()
|
||||
if not token:
|
||||
raise ProviderError("实时行情服务授权失败")
|
||||
self._issued_access = token
|
||||
return token
|
||||
|
||||
def _auth_error(self, payload: dict[str, Any]) -> bool:
|
||||
message = str(payload.get("errmsg") or payload.get("message") or "").casefold()
|
||||
return (
|
||||
_error_code(payload) in self.auth_error_codes
|
||||
or "token" in message
|
||||
or "鉴权" in message
|
||||
)
|
||||
|
||||
def _refresh(self) -> str:
|
||||
return str(self._refresh_provider() or "").strip()
|
||||
|
||||
def _configured_access(self) -> str:
|
||||
return str(self._access_provider() or "").strip()
|
||||
|
||||
def _post(
|
||||
self, endpoint: str, body: dict[str, Any], access: str, refresh: str = ""
|
||||
) -> dict[str, Any]:
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": "XiaobaiReview/2",
|
||||
"ifindlang": "cn",
|
||||
}
|
||||
if access:
|
||||
headers["access_token"] = access
|
||||
if refresh:
|
||||
headers["refresh_token"] = refresh
|
||||
request = urllib.request.Request(
|
||||
f"{self.base_url}/{endpoint}",
|
||||
data=json.dumps(body, ensure_ascii=False).encode("utf-8"),
|
||||
headers=headers,
|
||||
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 not isinstance(payload, dict):
|
||||
raise ProviderError("实时行情服务返回格式无效")
|
||||
return payload
|
||||
|
||||
|
||||
def _result(
|
||||
payload: dict[str, Any], unit: str, adjustment: str, state: SnapshotState
|
||||
) -> ProviderResult:
|
||||
tables = payload.get("tables") or (payload.get("data") or {}).get("tables") or []
|
||||
if isinstance(tables, dict):
|
||||
tables = [tables]
|
||||
rows: list[dict[str, Any]] = []
|
||||
for block in tables:
|
||||
columns = block.get("table") or {}
|
||||
if not columns:
|
||||
continue
|
||||
times = block.get("time") or []
|
||||
codes = block.get("thscode") or block.get("thscodes") or []
|
||||
if isinstance(codes, str):
|
||||
codes = [codes]
|
||||
size = max(
|
||||
(len(value) for value in columns.values() if isinstance(value, list)),
|
||||
default=len(times) if isinstance(times, list) else 1,
|
||||
)
|
||||
for index in range(size):
|
||||
row = {
|
||||
key: values[index]
|
||||
if isinstance(values, list) and index < len(values)
|
||||
else values if index == 0 else None
|
||||
for key, values in columns.items()
|
||||
}
|
||||
if isinstance(times, list) and index < len(times):
|
||||
row["time"] = times[index]
|
||||
if codes:
|
||||
row["thscode"] = codes[index] if index < len(codes) else codes[0]
|
||||
rows.append(row)
|
||||
metadata = ObservationMetadata(
|
||||
source=DataSource.IFIND,
|
||||
observed_at=datetime.now(SHANGHAI),
|
||||
unit=unit,
|
||||
adjustment=adjustment,
|
||||
freshness_seconds=0,
|
||||
coverage=1 if rows else 0,
|
||||
state=state,
|
||||
usage=DataUsage.DISPLAY,
|
||||
)
|
||||
return ProviderResult(tuple(rows), metadata)
|
||||
|
||||
|
||||
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:]}"
|
||||
|
||||
|
||||
def _error_code(payload: dict[str, Any]) -> int:
|
||||
try:
|
||||
return int(payload.get("errorcode", payload.get("code", 0)) or 0)
|
||||
except (TypeError, ValueError):
|
||||
return -1
|
||||
@@ -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