rebuild(stage-5): establish market data gateway and charts

This commit is contained in:
leefer
2026-07-30 02:35:42 +08:00
parent 40ad5d6836
commit cf0ab7026f
45 changed files with 2701 additions and 46 deletions
+3
View File
@@ -0,0 +1,3 @@
from backend.data.gateway import DataGateway
__all__ = ["DataGateway"]
+91
View File
@@ -0,0 +1,91 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime
from enum import StrEnum
from typing import Any
class DataSource(StrEnum):
TUSHARE = "tushare"
IFIND = "ifind"
EASTMONEY = "eastmoney"
TENCENT = "tencent"
LOCAL = "local"
class DataUsage(StrEnum):
DISPLAY = "display"
CALCULATION = "calculation"
class SnapshotState(StrEnum):
REALTIME = "realtime"
FINAL = "final"
ARCHIVE = "archive"
@dataclass(frozen=True, slots=True)
class ObservationMetadata:
source: DataSource
observed_at: datetime
unit: str
adjustment: str
freshness_seconds: int
coverage: float
state: SnapshotState
usage: DataUsage
def __post_init__(self) -> None:
if not 0 <= self.coverage <= 1:
raise ValueError("coverage must be between zero and one")
if self.freshness_seconds < 0:
raise ValueError("freshness_seconds cannot be negative")
@dataclass(frozen=True, slots=True)
class ProviderResult:
rows: tuple[dict[str, Any], ...]
metadata: ObservationMetadata
@dataclass(frozen=True, slots=True)
class TradeContext:
requested_date: str
actual_date: str | None
previous_date: str | None
observed_at: datetime | None
state: SnapshotState | None
carried_forward: bool
message: str
@dataclass(frozen=True, slots=True)
class MarketEntity:
entity_type: str
identifier: str
code: str
name: str
sector: str | None = None
@dataclass(frozen=True, slots=True)
class ChartPoint:
time: str
open: float
high: float
low: float
close: float
volume: float
amount: float
average: float | None = None
@dataclass(frozen=True, slots=True)
class ChartSeries:
entity: MarketEntity
interval: str
trade_date: str
previous_close: float | None
points: tuple[ChartPoint, ...]
metadata: ObservationMetadata
+379
View File
@@ -0,0 +1,379 @@
from __future__ import annotations
import json
from datetime import date, datetime, time, timedelta
from typing import Any
from zoneinfo import ZoneInfo
from backend.data.contracts import (
ChartPoint,
ChartSeries,
DataSource,
DataUsage,
MarketEntity,
ObservationMetadata,
SnapshotState,
TradeContext,
)
from backend.data.policy import DataSourcePolicy
from backend.data.providers.base import MarketDataProvider, ProviderError
from backend.data.quality import DataQualityError, require_quality
from backend.data.repository import MarketRepository
from backend.database.connection import Database
SHANGHAI = ZoneInfo("Asia/Shanghai")
class MarketDataUnavailable(RuntimeError):
pass
class DataGateway:
def __init__(
self,
database: Database,
repository: MarketRepository,
providers: tuple[MarketDataProvider, ...],
policy: DataSourcePolicy,
) -> None:
self._database = database
self._repository = repository
self._providers = {provider.source: provider for provider in providers}
self._policy = policy
def refresh_reference(self, now: datetime | None = None) -> dict[str, int | str]:
clock = now or datetime.now(SHANGHAI)
provider = self._provider(DataSource.TUSHARE)
start = (clock.date() - timedelta(days=370)).isoformat()
end = (clock.date() + timedelta(days=40)).isoformat()
calendar = require_quality("calendar", provider.calendar(start, end))
entities = require_quality("entities", provider.entities())
observed_at = clock.isoformat(timespec="seconds")
with self._database.transaction() as connection:
self._repository.replace_calendar(
connection, calendar.rows, calendar.metadata.source.value, observed_at
)
self._repository.replace_stocks(
connection, entities.rows, entities.metadata.source.value, observed_at
)
return {
"calendar_days": len(calendar.rows),
"entities": len(entities.rows),
"observed_at": observed_at,
}
def trade_context(
self, requested_date: str | None = None, now: datetime | None = None
) -> TradeContext:
clock = now or datetime.now(SHANGHAI)
requested = _date(requested_date or clock.date().isoformat())
with self._database.read() as connection:
summary = self._repository.latest_summary(connection, requested)
if summary is None:
return TradeContext(
requested_date=requested,
actual_date=None,
previous_date=None,
observed_at=None,
state=None,
carried_forward=False,
message="等待管理员首次同步真实行情",
)
actual = str(summary["trade_date"])
observed_at = datetime.fromisoformat(str(summary["observed_at"]))
state = SnapshotState(str(summary["state"]))
carried = actual != requested
with self._database.read() as connection:
dates = self._repository.open_dates(connection, actual, 2)
return TradeContext(
requested_date=requested,
actual_date=actual,
previous_date=dates[1] if len(dates) > 1 else None,
observed_at=observed_at,
state=state,
carried_forward=carried,
message="沿用最近真实收盘快照" if carried else "",
)
def latest_available_date(self, now: datetime | None = None) -> str:
context = self.trade_context(now=now)
return context.actual_date or (now or datetime.now(SHANGHAI)).date().isoformat()
def summary(self, requested_date: str | None = None) -> dict[str, Any]:
context = self.trade_context(requested_date)
if context.actual_date is None:
return {"context": context, "values": None}
with self._database.read() as connection:
row = self._repository.latest_summary(connection, context.actual_date)
return {"context": context, "values": json.loads(str(row["payload_json"])) if row else None}
def search(self, query: str) -> tuple[MarketEntity, ...]:
with self._database.read() as connection:
return self._repository.search(connection, query)
def chart(
self,
entity_type: str,
identifier: str,
interval: str,
now: datetime | None = None,
) -> ChartSeries:
if entity_type not in {"stock", "sector", "theme", "index"}:
raise MarketDataUnavailable("不支持的行情标的类型")
if interval not in {"day", "minute"}:
raise MarketDataUnavailable("不支持的行情周期")
entity = self._resolve_entity(entity_type, identifier)
clock = now or datetime.now(SHANGHAI)
with self._database.read() as connection:
cached = self._repository.chart(connection, entity_type, entity.identifier, interval)
if cached and not _chart_stale(cached, interval, clock):
return _stored_chart(entity, cached)
try:
series = self._load_chart(entity, interval, clock)
except (ProviderError, DataQualityError) as exc:
if cached:
return _stored_chart(entity, cached)
raise MarketDataUnavailable("当前没有可用的真实行情数据") from exc
self._save_chart(series)
return series
def _load_chart(self, entity: MarketEntity, interval: str, clock: datetime) -> ChartSeries:
dataset = "daily_chart" if interval == "day" else "minute_chart"
target_dates = self._chart_dates(clock)
errors: list[Exception] = []
for source in self._policy.candidates(dataset, DataUsage.DISPLAY):
provider = self._providers.get(source)
if provider is None or not provider.configured:
continue
self._policy.assert_allowed(source, DataUsage.DISPLAY)
for target in target_dates:
try:
raw = (
provider.daily(entity.entity_type, entity.identifier, target)
if interval == "day"
else provider.minute(entity.entity_type, entity.identifier, target)
)
require_quality(dataset, raw)
normalized = _normalize_chart(entity, interval, raw.rows, raw.metadata, clock)
if normalized.points:
return normalized
except (ProviderError, DataQualityError, ValueError) as exc:
errors.append(exc)
if interval == "day":
break
raise MarketDataUnavailable("当前没有可用的真实行情数据") from (
errors[-1] if errors else None
)
def _chart_dates(self, clock: datetime) -> tuple[str, ...]:
today = clock.date().isoformat()
with self._database.read() as connection:
dates = self._repository.open_dates(connection, today, 8)
if dates:
return dates
return tuple((clock.date() - timedelta(days=offset)).isoformat() for offset in range(8))
def _resolve_entity(self, entity_type: str, identifier: str) -> MarketEntity:
normalized = identifier.strip().upper()
with self._database.read() as connection:
entity = self._repository.entity(connection, entity_type, normalized)
if entity is None and entity_type == "stock" and normalized.isdigit():
matches = self._repository.search(connection, normalized, 4)
entity = next(
(
item
for item in matches
if item.entity_type == "stock" and item.code == normalized
),
None,
)
if entity:
return entity
if entity_type == "stock" and len(normalized) == 6 and normalized.isdigit():
suffix = (
"BJ"
if normalized.startswith(("4", "8", "9"))
else "SH"
if normalized.startswith("6")
else "SZ"
)
return MarketEntity("stock", f"{normalized}.{suffix}", normalized, normalized)
raise MarketDataUnavailable("未找到该行情标的")
def _save_chart(self, series: ChartSeries) -> None:
payload = {
"previous_close": series.previous_close,
"points": [
{
"time": point.time,
"open": point.open,
"high": point.high,
"low": point.low,
"close": point.close,
"volume": point.volume,
"amount": point.amount,
"average": point.average,
}
for point in series.points
],
}
with self._database.transaction() as connection:
self._repository.save_chart(
connection,
entity_type=series.entity.entity_type,
identifier=series.entity.identifier,
interval=series.interval,
trade_date=series.trade_date,
observed_at=series.metadata.observed_at.isoformat(timespec="seconds"),
source=series.metadata.source.value,
usage=series.metadata.usage.value,
adjustment=series.metadata.adjustment,
coverage=series.metadata.coverage,
payload=payload,
)
def _provider(self, source: DataSource) -> MarketDataProvider:
provider = self._providers.get(source)
if provider is None or not provider.configured:
raise MarketDataUnavailable("所需行情服务尚未配置")
return provider
def _normalize_chart(
entity: MarketEntity,
interval: str,
rows: tuple[dict[str, Any], ...],
metadata: ObservationMetadata,
clock: datetime,
) -> ChartSeries:
parsed: list[tuple[str, ChartPoint, float | None]] = []
for row in rows:
stamp = str(row.get("trade_date") or row.get("time") or row.get("trade_time") or "")
trade_date = _row_date(stamp)
point_time = trade_date if interval == "day" else _row_time(stamp)
close = _number(row.get("close"))
open_price = _number(row.get("open"))
high = _number(row.get("high"))
low = _number(row.get("low"))
if not trade_date or close <= 0 or open_price <= 0 or high <= 0 or low <= 0:
continue
if interval == "minute" and not "09:30" <= point_time <= "15:00":
continue
volume = _number(row.get("volume", row.get("vol")))
amount = _number(row.get("amount"))
if metadata.source is DataSource.TUSHARE:
volume *= 100
amount *= 1000
parsed.append(
(
trade_date,
ChartPoint(
time=point_time,
open=open_price,
high=high,
low=low,
close=close,
volume=volume,
amount=amount,
average=_optional_number(row.get("avgPrice", row.get("average"))),
),
_optional_number(row.get("preClose", row.get("pre_close"))),
)
)
parsed.sort(key=lambda item: (item[0], item[1].time))
if not parsed:
raise DataQualityError("chart contains no valid points")
if interval == "minute":
latest = parsed[-1][0]
parsed = [item for item in parsed if item[0] == latest]
else:
today = clock.date().isoformat()
if parsed[-1][0] == today and not _valid_today_bar(parsed[-1][1], clock):
parsed.pop()
if not parsed:
raise DataQualityError("chart contains no completed bar")
parsed = parsed[-90:]
previous = parsed[0][2]
if previous is None and interval == "day" and len(parsed) > 1:
previous = parsed[-2][1].close
return ChartSeries(
entity=entity,
interval=interval,
trade_date=parsed[-1][0],
previous_close=previous,
points=tuple(item[1] for item in parsed),
metadata=metadata,
)
def _valid_today_bar(point: ChartPoint, clock: datetime) -> bool:
if clock.time() < time(9, 30):
return False
return (
point.volume > 0
and point.amount > 0
and point.high >= max(point.open, point.close)
and point.low <= min(point.open, point.close)
)
def _chart_stale(row: Any, interval: str, clock: datetime) -> bool:
observed = datetime.fromisoformat(str(row["observed_at"]))
if interval == "day":
return row["trade_date"] < clock.date().isoformat() and clock.time() >= time(15, 5)
return (clock - observed.astimezone(SHANGHAI)).total_seconds() > 30
def _stored_chart(entity: MarketEntity, row: Any) -> ChartSeries:
payload = json.loads(str(row["payload_json"]))
points = tuple(ChartPoint(**point) for point in payload.get("points") or [])
metadata = ObservationMetadata(
source=DataSource(str(row["source"])),
observed_at=datetime.fromisoformat(str(row["observed_at"])),
unit="yuan/share",
adjustment=str(row["adjustment"]),
freshness_seconds=0,
coverage=float(row["coverage"]),
state=SnapshotState.ARCHIVE,
usage=DataUsage(str(row["usage"])),
)
return ChartSeries(
entity,
str(row["interval"]),
str(row["trade_date"]),
payload.get("previous_close"),
points,
metadata,
)
def _date(value: str) -> str:
try:
return date.fromisoformat(value).isoformat()
except ValueError as exc:
raise MarketDataUnavailable("日期格式无效") from exc
def _row_date(value: str) -> str:
compact = value[:10].replace("-", "")
if len(compact) != 8 or not compact.isdigit():
return ""
return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"
def _row_time(value: str) -> str:
if " " in value:
return value.split(" ", 1)[1][:5]
return value[-8:-3] if len(value) >= 8 else value[:5]
def _number(value: Any) -> float:
try:
return float(value or 0)
except (TypeError, ValueError):
return 0.0
def _optional_number(value: Any) -> float | None:
number = _number(value)
return number if number > 0 else None
+34
View File
@@ -0,0 +1,34 @@
from __future__ import annotations
from dataclasses import dataclass
from backend.data.contracts import DataSource, DataUsage
class DataPolicyError(RuntimeError):
pass
@dataclass(frozen=True, slots=True)
class DataSourcePolicy:
calculation_sources: frozenset[DataSource] = frozenset(
{DataSource.TUSHARE, DataSource.IFIND, DataSource.LOCAL}
)
display_sources: frozenset[DataSource] = frozenset(DataSource)
def assert_allowed(self, source: DataSource, usage: DataUsage) -> None:
allowed = (
self.calculation_sources if usage is DataUsage.CALCULATION else self.display_sources
)
if source not in allowed:
raise DataPolicyError(f"{source.value} cannot be used for {usage.value}")
def candidates(self, dataset: str, usage: DataUsage) -> tuple[DataSource, ...]:
routes = {
("calendar", DataUsage.CALCULATION): (DataSource.TUSHARE,),
("entities", DataUsage.CALCULATION): (DataSource.TUSHARE,),
("daily_chart", DataUsage.DISPLAY): (DataSource.IFIND, DataSource.TUSHARE),
("minute_chart", DataUsage.DISPLAY): (DataSource.IFIND, DataSource.EASTMONEY),
("realtime_quote", DataUsage.CALCULATION): (DataSource.IFIND, DataSource.TUSHARE),
}
return routes.get((dataset, usage), ())
+5
View File
@@ -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"]
+24
View File
@@ -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: ...
+105
View File
@@ -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("该标的暂无展示分时数据")
+233
View File
@@ -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
+159
View File
@@ -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:]}"
+26
View File
@@ -0,0 +1,26 @@
from __future__ import annotations
from backend.data.contracts import ProviderResult
class DataQualityError(RuntimeError):
pass
MINIMUM_COVERAGE = {
"calendar": 1.0,
"entities": 0.98,
"daily_chart": 1.0,
"minute_chart": 1.0,
}
def require_quality(dataset: str, result: ProviderResult) -> ProviderResult:
minimum = MINIMUM_COVERAGE.get(dataset, 1.0)
if not result.rows:
raise DataQualityError(f"{dataset} returned no real observations")
if result.metadata.coverage < minimum:
raise DataQualityError(
f"{dataset} coverage {result.metadata.coverage:.3f} is below {minimum:.3f}"
)
return result
+215
View File
@@ -0,0 +1,215 @@
from __future__ import annotations
import json
import re
import sqlite3
from datetime import datetime
from typing import Any
from backend.data.contracts import MarketEntity
class MarketRepository:
def replace_calendar(
self,
connection: sqlite3.Connection,
rows: tuple[dict[str, Any], ...],
source: str,
observed_at: str,
) -> None:
connection.executemany(
"""
INSERT INTO trading_days (trade_date, is_open, previous_open_date, source, observed_at)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(trade_date) DO UPDATE SET
is_open = excluded.is_open,
previous_open_date = excluded.previous_open_date,
source = excluded.source,
observed_at = excluded.observed_at
""",
[
(
_display(str(row.get("cal_date") or "")),
1 if int(row.get("is_open") or 0) == 1 else 0,
_display(str(row.get("pretrade_date") or "")) or None,
source,
observed_at,
)
for row in rows
if _display(str(row.get("cal_date") or ""))
],
)
def replace_stocks(
self,
connection: sqlite3.Connection,
rows: tuple[dict[str, Any], ...],
source: str,
observed_at: str,
) -> None:
connection.executemany(
"""
INSERT INTO market_entities (
entity_type, identifier, code, name, search_key,
sector, active, source, observed_at
)
VALUES ('stock', ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(entity_type, identifier) DO UPDATE SET
code = excluded.code,
name = excluded.name,
search_key = excluded.search_key,
sector = excluded.sector,
active = excluded.active,
source = excluded.source,
observed_at = excluded.observed_at
""",
[
(
str(row.get("ts_code") or "").upper(),
str(row.get("symbol") or ""),
str(row.get("name") or "").strip(),
_search_key(row),
str(row.get("industry") or "").strip() or None,
0 if row.get("list_status") == "D" else 1,
source,
observed_at,
)
for row in rows
if row.get("ts_code") and row.get("symbol") and row.get("name")
],
)
def search(
self, connection: sqlite3.Connection, query: str, limit: int = 32
) -> tuple[MarketEntity, ...]:
normalized = _normalize(query)
if not normalized:
return ()
rows = connection.execute(
"""
SELECT entity_type, identifier, code, name, sector
FROM market_entities
WHERE active = 1 AND search_key LIKE ?
ORDER BY
CASE WHEN code = ? THEN 0 WHEN name = ? THEN 1 WHEN code LIKE ? THEN 2 ELSE 3 END,
entity_type, name
LIMIT ?
""",
(f"%{normalized}%", normalized, query.strip(), f"{normalized}%", limit),
).fetchall()
return tuple(MarketEntity(**dict(row)) for row in rows)
def entity(
self, connection: sqlite3.Connection, entity_type: str, identifier: str
) -> MarketEntity | None:
row = connection.execute(
"""
SELECT entity_type, identifier, code, name, sector
FROM market_entities WHERE entity_type = ? AND identifier = ? AND active = 1
""",
(entity_type, identifier),
).fetchone()
return MarketEntity(**dict(row)) if row else None
def open_dates(
self, connection: sqlite3.Connection, through: str, limit: int = 12
) -> tuple[str, ...]:
return tuple(
str(row["trade_date"])
for row in connection.execute(
"""
SELECT trade_date FROM trading_days
WHERE is_open = 1 AND trade_date <= ?
ORDER BY trade_date DESC LIMIT ?
""",
(through, limit),
)
)
def latest_summary(self, connection: sqlite3.Connection, through: str) -> sqlite3.Row | None:
return connection.execute(
"SELECT * FROM market_summaries WHERE trade_date <= ? ORDER BY trade_date DESC LIMIT 1",
(through,),
).fetchone()
def save_chart(
self,
connection: sqlite3.Connection,
*,
entity_type: str,
identifier: str,
interval: str,
trade_date: str,
observed_at: str,
source: str,
usage: str,
adjustment: str,
coverage: float,
payload: dict[str, Any],
) -> None:
connection.execute(
"""
INSERT INTO chart_series
(entity_type, identifier, interval, trade_date, observed_at, source, usage,
adjustment, coverage, payload_json, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(entity_type, identifier, interval, trade_date) DO UPDATE SET
observed_at = excluded.observed_at,
source = excluded.source,
usage = excluded.usage,
adjustment = excluded.adjustment,
coverage = excluded.coverage,
payload_json = excluded.payload_json,
created_at = excluded.created_at
""",
(
entity_type,
identifier,
interval,
trade_date,
observed_at,
source,
usage,
adjustment,
coverage,
json.dumps(payload, ensure_ascii=False, separators=(",", ":")),
datetime.now().astimezone().isoformat(timespec="seconds"),
),
)
def chart(
self, connection: sqlite3.Connection, entity_type: str, identifier: str, interval: str
) -> sqlite3.Row | None:
return connection.execute(
"""
SELECT * FROM chart_series
WHERE entity_type = ? AND identifier = ? AND interval = ?
ORDER BY trade_date DESC LIMIT 1
""",
(entity_type, identifier, interval),
).fetchone()
def _display(value: str) -> str:
compact = value.replace("-", "")
if len(compact) != 8 or not compact.isdigit():
return ""
return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"
def _normalize(value: str) -> str:
return re.sub(r"\s+", "", value).casefold()
def _search_key(row: dict[str, Any]) -> str:
return " ".join(
filter(
None,
(
_normalize(str(row.get("symbol") or "")),
_normalize(str(row.get("ts_code") or "")),
_normalize(str(row.get("name") or "")),
_normalize(str(row.get("industry") or "")),
),
)
)