rebuild(stage-5): establish market data gateway and charts
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from backend.data.gateway import DataGateway
|
||||
|
||||
__all__ = ["DataGateway"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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), ())
|
||||
@@ -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:]}"
|
||||
@@ -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
|
||||
@@ -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 "")),
|
||||
),
|
||||
)
|
||||
)
|
||||
Reference in New Issue
Block a user