diff --git a/next/backend/bootstrap/container.py b/next/backend/bootstrap/container.py index de3a1d1..0a7b1b1 100644 --- a/next/backend/bootstrap/container.py +++ b/next/backend/bootstrap/container.py @@ -3,6 +3,10 @@ from __future__ import annotations from dataclasses import dataclass from backend.bootstrap.settings import Settings +from backend.data.gateway import DataGateway +from backend.data.policy import DataSourcePolicy +from backend.data.providers import EastmoneyProvider, IfindProvider, TushareProvider +from backend.data.repository import MarketRepository from backend.database.connection import Database from backend.database.repositories.status import DatabaseStatusRepository from backend.features.accounts.credentials import SystemCredentialService @@ -12,6 +16,7 @@ from backend.features.accounts.service import ( AccountService, MembershipService, ) +from backend.features.market import MarketService from backend.security import PasswordHasher, load_or_create_cipher @@ -24,6 +29,7 @@ class ApplicationContainer: memberships: MembershipService system_credentials: SystemCredentialService model_pool: ModelPoolService + market: MarketService def build_container(settings: Settings) -> ApplicationContainer: @@ -32,12 +38,27 @@ def build_container(settings: Settings) -> ApplicationContainer: credential_repository = SystemCredentialRepository() model_pool_repository = ModelPoolRepository() cipher = load_or_create_cipher(settings) + credentials = SystemCredentialService(database, credential_repository, cipher) + gateway = DataGateway( + database, + MarketRepository(), + ( + TushareProvider(lambda: credentials.get("tushare_token")), + IfindProvider( + lambda: credentials.get("ifind_refresh_token"), + lambda: credentials.get("ifind_access_token"), + ), + EastmoneyProvider(), + ), + DataSourcePolicy(), + ) return ApplicationContainer( settings=settings, database=database, database_status=DatabaseStatusRepository(database), accounts=AccountService(database, account_repository, PasswordHasher(), cipher), memberships=MembershipService(database, account_repository), - system_credentials=SystemCredentialService(database, credential_repository, cipher), + system_credentials=credentials, model_pool=ModelPoolService(database, model_pool_repository, cipher), + market=MarketService(gateway), ) diff --git a/next/backend/data/__init__.py b/next/backend/data/__init__.py new file mode 100644 index 0000000..dc2c18b --- /dev/null +++ b/next/backend/data/__init__.py @@ -0,0 +1,3 @@ +from backend.data.gateway import DataGateway + +__all__ = ["DataGateway"] diff --git a/next/backend/data/contracts.py b/next/backend/data/contracts.py new file mode 100644 index 0000000..0df6ac0 --- /dev/null +++ b/next/backend/data/contracts.py @@ -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 diff --git a/next/backend/data/gateway.py b/next/backend/data/gateway.py new file mode 100644 index 0000000..bc2d7bd --- /dev/null +++ b/next/backend/data/gateway.py @@ -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 diff --git a/next/backend/data/policy.py b/next/backend/data/policy.py new file mode 100644 index 0000000..23fa8cf --- /dev/null +++ b/next/backend/data/policy.py @@ -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), ()) diff --git a/next/backend/data/providers/__init__.py b/next/backend/data/providers/__init__.py new file mode 100644 index 0000000..b049656 --- /dev/null +++ b/next/backend/data/providers/__init__.py @@ -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"] diff --git a/next/backend/data/providers/base.py b/next/backend/data/providers/base.py new file mode 100644 index 0000000..d68934a --- /dev/null +++ b/next/backend/data/providers/base.py @@ -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: ... diff --git a/next/backend/data/providers/eastmoney.py b/next/backend/data/providers/eastmoney.py new file mode 100644 index 0000000..b012e88 --- /dev/null +++ b/next/backend/data/providers/eastmoney.py @@ -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("该标的暂无展示分时数据") diff --git a/next/backend/data/providers/ifind.py b/next/backend/data/providers/ifind.py new file mode 100644 index 0000000..5ba8b56 --- /dev/null +++ b/next/backend/data/providers/ifind.py @@ -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 diff --git a/next/backend/data/providers/tushare.py b/next/backend/data/providers/tushare.py new file mode 100644 index 0000000..e737582 --- /dev/null +++ b/next/backend/data/providers/tushare.py @@ -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:]}" diff --git a/next/backend/data/quality.py b/next/backend/data/quality.py new file mode 100644 index 0000000..668ce29 --- /dev/null +++ b/next/backend/data/quality.py @@ -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 diff --git a/next/backend/data/repository.py b/next/backend/data/repository.py new file mode 100644 index 0000000..d249628 --- /dev/null +++ b/next/backend/data/repository.py @@ -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 "")), + ), + ) + ) diff --git a/next/backend/database/migrations/m0003_market_foundation.py b/next/backend/database/migrations/m0003_market_foundation.py new file mode 100644 index 0000000..e767d28 --- /dev/null +++ b/next/backend/database/migrations/m0003_market_foundation.py @@ -0,0 +1,114 @@ +from __future__ import annotations + +import sqlite3 + +from backend.database.migrations.runner import Migration + + +def upgrade(connection: sqlite3.Connection) -> None: + connection.execute( + """ + CREATE TABLE trading_days ( + trade_date TEXT PRIMARY KEY, + is_open INTEGER NOT NULL CHECK (is_open IN (0, 1)), + previous_open_date TEXT, + source TEXT NOT NULL, + observed_at TEXT NOT NULL + ) + """ + ) + connection.execute( + """ + CREATE TABLE market_entities ( + entity_type TEXT NOT NULL, + identifier TEXT NOT NULL, + code TEXT NOT NULL, + name TEXT NOT NULL, + search_key TEXT NOT NULL, + sector TEXT, + active INTEGER NOT NULL DEFAULT 1 CHECK (active IN (0, 1)), + source TEXT NOT NULL, + observed_at TEXT NOT NULL, + PRIMARY KEY (entity_type, identifier) + ) + """ + ) + connection.execute( + """ + CREATE INDEX market_entities_search_idx + ON market_entities(active, search_key, entity_type) + """ + ) + connection.execute( + """ + INSERT INTO market_entities ( + entity_type, identifier, code, name, search_key, + sector, active, source, observed_at + ) VALUES + ( + 'index', '000001.SH', '000001', '上证指数', + '000001 上证指数 shanghai', NULL, 1, 'local', '2026-07-30T00:00:00+08:00' + ), + ( + 'index', '399001.SZ', '399001', '深证成指', + '399001 深证成指 shenzhen', NULL, 1, 'local', '2026-07-30T00:00:00+08:00' + ), + ( + 'index', '399006.SZ', '399006', '创业板指', + '399006 创业板指 chinext', NULL, 1, 'local', '2026-07-30T00:00:00+08:00' + ) + """ + ) + connection.execute( + """ + CREATE TABLE market_summaries ( + trade_date TEXT PRIMARY KEY, + observed_at TEXT NOT NULL, + state TEXT NOT NULL CHECK (state IN ('realtime', 'final', 'archive')), + source TEXT NOT NULL, + coverage REAL NOT NULL CHECK (coverage >= 0 AND coverage <= 1), + payload_json TEXT NOT NULL, + created_at TEXT NOT NULL + ) + """ + ) + connection.execute( + """ + CREATE TABLE chart_series ( + entity_type TEXT NOT NULL, + identifier TEXT NOT NULL, + interval TEXT NOT NULL CHECK (interval IN ('day', 'minute')), + trade_date TEXT NOT NULL, + observed_at TEXT NOT NULL, + source TEXT NOT NULL, + usage TEXT NOT NULL CHECK (usage IN ('display', 'calculation')), + adjustment TEXT NOT NULL, + coverage REAL NOT NULL CHECK (coverage >= 0 AND coverage <= 1), + payload_json TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (entity_type, identifier, interval, trade_date) + ) + """ + ) + connection.execute( + """ + CREATE INDEX chart_series_latest_idx + ON chart_series(entity_type, identifier, interval, trade_date DESC) + """ + ) + + +def downgrade(connection: sqlite3.Connection) -> None: + connection.execute("DROP TABLE chart_series") + connection.execute("DROP TABLE market_summaries") + connection.execute("DROP TABLE market_entities") + connection.execute("DROP TABLE trading_days") + + +MIGRATION = Migration( + version=3, + name="create_market_foundation", + signature="market:v1:calendar-entities-summary-chart-provenance", + upgrade=upgrade, + downgrade=downgrade, +) diff --git a/next/backend/database/migrations/registry.py b/next/backend/database/migrations/registry.py index e0ddd26..7192df9 100644 --- a/next/backend/database/migrations/registry.py +++ b/next/backend/database/migrations/registry.py @@ -1,5 +1,6 @@ from backend.database.migrations.m0001_accounts import MIGRATION as ACCOUNTS from backend.database.migrations.m0002_model_pool import MIGRATION as MODEL_POOL +from backend.database.migrations.m0003_market_foundation import MIGRATION as MARKET_FOUNDATION from backend.database.migrations.runner import Migration -MIGRATIONS: tuple[Migration, ...] = (ACCOUNTS, MODEL_POOL) +MIGRATIONS: tuple[Migration, ...] = (ACCOUNTS, MODEL_POOL, MARKET_FOUNDATION) diff --git a/next/backend/features/market/__init__.py b/next/backend/features/market/__init__.py new file mode 100644 index 0000000..ebf2799 --- /dev/null +++ b/next/backend/features/market/__init__.py @@ -0,0 +1,3 @@ +from backend.features.market.service import MarketService + +__all__ = ["MarketService"] diff --git a/next/backend/features/market/routes.py b/next/backend/features/market/routes.py new file mode 100644 index 0000000..fbc25cf --- /dev/null +++ b/next/backend/features/market/routes.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from typing import Annotated, Literal + +from fastapi import APIRouter, Path, Query, Request + +from backend.features.accounts.auth import AdminWritePrincipal, AuthenticatedPrincipal +from backend.features.market.schemas import ( + ChartResponse, + MarketSummaryResponse, + ReferenceSyncResponse, + SearchResponse, + TradeContextResponse, +) + +router = APIRouter(prefix="/market", tags=["market"]) + + +@router.get("/context", response_model=TradeContextResponse) +def context( + request: Request, + _principal: AuthenticatedPrincipal, + requested_date: Annotated[str | None, Query(alias="date")] = None, +) -> dict: + return request.app.state.container.market.context(requested_date) + + +@router.get("/summary", response_model=MarketSummaryResponse) +def summary( + request: Request, + _principal: AuthenticatedPrincipal, + requested_date: Annotated[str | None, Query(alias="date")] = None, +) -> dict: + return request.app.state.container.market.summary(requested_date) + + +@router.get("/search", response_model=SearchResponse) +def search( + request: Request, + _principal: AuthenticatedPrincipal, + query: Annotated[str, Query(alias="q", max_length=80)] = "", +) -> dict: + return request.app.state.container.market.search(query) + + +@router.get("/entities/{entity_type}/{identifier}/charts/{interval}", response_model=ChartResponse) +def chart( + request: Request, + _principal: AuthenticatedPrincipal, + entity_type: Annotated[Literal["stock", "sector", "theme", "index"], Path()], + identifier: Annotated[str, Path(min_length=1, max_length=40)], + interval: Annotated[Literal["day", "minute"], Path()], +) -> dict: + return request.app.state.container.market.chart(entity_type, identifier, interval) + + +@router.post("/reference-sync", response_model=ReferenceSyncResponse) +def refresh_reference(request: Request, _principal: AdminWritePrincipal) -> dict[str, int | str]: + return request.app.state.container.market.refresh_reference() diff --git a/next/backend/features/market/schemas.py b/next/backend/features/market/schemas.py new file mode 100644 index 0000000..6019d6f --- /dev/null +++ b/next/backend/features/market/schemas.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Any, Literal + +from pydantic import BaseModel, Field + + +class TradeContextResponse(BaseModel): + requested_date: str + actual_date: str | None + previous_date: str | None + observed_at: datetime | None + state: str | None + carried_forward: bool + message: str + + +class MarketSummaryResponse(BaseModel): + context: TradeContextResponse + values: dict[str, Any] | None + + +class SearchResultResponse(BaseModel): + entity_type: Literal["stock", "sector", "theme", "index"] + identifier: str + code: str + name: str + sector: str | None + + +class SearchGroupResponse(BaseModel): + entity_type: Literal["stock", "sector", "theme", "index"] + label: str + items: list[SearchResultResponse] + + +class SearchResponse(BaseModel): + query: str + groups: list[SearchGroupResponse] + + +class ChartPointResponse(BaseModel): + time: str + open: float + high: float + low: float + close: float + volume: float + amount: float + average: float | None + + +class ChartResponse(BaseModel): + entity_type: str + identifier: str + code: str + name: str + interval: Literal["day", "minute"] + trade_date: str + observed_at: datetime + previous_close: float | None + range_start: str | None + range_end: str | None + points: list[ChartPointResponse] + + +class ReferenceSyncResponse(BaseModel): + calendar_days: int = Field(ge=1) + entities: int = Field(ge=1) + observed_at: datetime diff --git a/next/backend/features/market/service.py b/next/backend/features/market/service.py new file mode 100644 index 0000000..b980d86 --- /dev/null +++ b/next/backend/features/market/service.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +from typing import Any + +from backend.data.gateway import DataGateway, MarketDataUnavailable +from backend.data.providers.base import ProviderError +from backend.data.quality import DataQualityError +from backend.http.errors import AppError + + +class MarketService: + def __init__(self, gateway: DataGateway) -> None: + self._gateway = gateway + + def context(self, requested_date: str | None = None) -> dict[str, Any]: + context = self._call(self._gateway.trade_context, requested_date) + return _context(context) + + def summary(self, requested_date: str | None = None) -> dict[str, Any]: + result = self._call(self._gateway.summary, requested_date) + return {"context": _context(result["context"]), "values": result["values"]} + + def search(self, query: str) -> dict[str, Any]: + normalized = " ".join(query.split()) + items = self._gateway.search(normalized) if normalized else () + labels = {"stock": "股票", "sector": "板块", "theme": "题材", "index": "指数"} + groups = [] + for entity_type in ("stock", "sector", "theme", "index"): + groups.append( + { + "entity_type": entity_type, + "label": labels[entity_type], + "items": [ + { + "entity_type": item.entity_type, + "identifier": item.identifier, + "code": item.code, + "name": item.name, + "sector": item.sector, + } + for item in items + if item.entity_type == entity_type + ], + } + ) + return {"query": normalized, "groups": groups} + + def chart(self, entity_type: str, identifier: str, interval: str) -> dict[str, Any]: + series = self._call(self._gateway.chart, entity_type, identifier, interval) + return { + "entity_type": series.entity.entity_type, + "identifier": series.entity.identifier, + "code": series.entity.code, + "name": series.entity.name, + "interval": series.interval, + "trade_date": series.trade_date, + "observed_at": series.metadata.observed_at, + "previous_close": series.previous_close, + "range_start": "09:30" if interval == "minute" else None, + "range_end": "15:00" if interval == "minute" else None, + "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 + ], + } + + def refresh_reference(self) -> dict[str, int | str]: + return self._call(self._gateway.refresh_reference) + + @staticmethod + def _call(function, *args): + try: + return function(*args) + except (MarketDataUnavailable, ProviderError, DataQualityError) as exc: + raise AppError("market_data_unavailable", str(exc), 503) from exc + + +def _context(context) -> dict[str, Any]: + return { + "requested_date": context.requested_date, + "actual_date": context.actual_date, + "previous_date": context.previous_date, + "observed_at": context.observed_at, + "state": context.state.value if context.state else None, + "carried_forward": context.carried_forward, + "message": context.message, + } diff --git a/next/backend/http/router.py b/next/backend/http/router.py index ffb17bc..b08e747 100644 --- a/next/backend/http/router.py +++ b/next/backend/http/router.py @@ -1,8 +1,10 @@ from fastapi import APIRouter from backend.features.accounts.routes import router as accounts_router +from backend.features.market.routes import router as market_router from backend.http.routes.health import router as health_router api_router = APIRouter() api_router.include_router(health_router) api_router.include_router(accounts_router) +api_router.include_router(market_router) diff --git a/next/docs/evidence/stage-4/shell-dark-1920x1080.jpg b/next/docs/evidence/stage-4/shell-dark-1920x1080.jpg index 967c40b..9d45d97 100644 Binary files a/next/docs/evidence/stage-4/shell-dark-1920x1080.jpg and b/next/docs/evidence/stage-4/shell-dark-1920x1080.jpg differ diff --git a/next/docs/evidence/stage-4/shell-dark-390x844.jpg b/next/docs/evidence/stage-4/shell-dark-390x844.jpg index 1748058..9675249 100644 Binary files a/next/docs/evidence/stage-4/shell-dark-390x844.jpg and b/next/docs/evidence/stage-4/shell-dark-390x844.jpg differ diff --git a/next/docs/evidence/stage-4/shell-light-1920x1080.jpg b/next/docs/evidence/stage-4/shell-light-1920x1080.jpg index 1761c1b..7a5a32b 100644 Binary files a/next/docs/evidence/stage-4/shell-light-1920x1080.jpg and b/next/docs/evidence/stage-4/shell-light-1920x1080.jpg differ diff --git a/next/docs/evidence/stage-4/system-dark-3840x2160.jpg b/next/docs/evidence/stage-4/system-dark-3840x2160.jpg index fd06459..d0608c8 100644 Binary files a/next/docs/evidence/stage-4/system-dark-3840x2160.jpg and b/next/docs/evidence/stage-4/system-dark-3840x2160.jpg differ diff --git a/next/docs/evidence/stage-5/README.md b/next/docs/evidence/stage-5/README.md new file mode 100644 index 0000000..b5ff506 --- /dev/null +++ b/next/docs/evidence/stage-5/README.md @@ -0,0 +1,49 @@ +# 阶段5验收记录 + +## 范围 + +本阶段建立唯一数据网关、数据源准入策略、交易日期与真实快照契约、标的目录、全局搜索、日K和分时图基础。情绪计算、五类股池和完整个股业务详情仍分别属于阶段6及后续阶段,本阶段不提前复制旧业务。 + +## 数据真相 + +- 每条可持久化图表序列保留标的、真实交易日、观测时间、来源、用途、单位、复权、覆盖率和实时/最终/归档状态。 +- Tushare负责交易日历、股票目录和日线;iFinD优先用于展示图表;东方财富只允许进入展示分时兜底,策略和情绪计算会被策略层拒绝。 +- 今天没有有效开高低收、成交量和成交额时,日K删除空的今日柱,继续显示最后真实交易日。 +- 历史快照只按实际交易日返回;沿用旧快照时,接口同时返回请求日期、实际数据日期和沿用说明。 +- 计算数据覆盖率不达标时失败关闭;从未同步时显示等待首次同步,不产生模拟数据。 + +## 搜索与图表 + +- `Ctrl+K`搜索经唯一浏览器API出口请求,160毫秒防抖,按股票、板块、题材、指数固定分组。 +- 支持键盘上下循环、Enter打开、Escape关闭,并区分无输入、加载、空结果和失败。 +- 桌面搜索结果提供最新日K预览,可切换分时;点击进入共用行情详情基础页。移动端点击进入详情,不依赖Hover。 +- 日K上涨柱为空心且影线不穿实体,下跌柱为实心;分时提供昨收零轴和均价线,横轴契约固定09:30至15:00。 +- 普通用户响应不包含供应商名称;行情管理保留管理员可见的凭据与基础资料同步入口。 + +## 减法证据 + +- 未复制旧`TushareClient`、`IfindHttpClient`或`MarketChartClient`;只迁移通用协议、授权刷新和必要归一化规则。 +- 所有外部源通过一个`DataGateway`和一份`DataSourcePolicy`进入系统;页面没有直接访问供应商。 +- 搜索预览与详情复用同一个图表组件和同一接口,不创建股票、板块、题材、指数四套图表实现。 +- 新增后端最长文件为379行,未超过章程400行目标;CSS色值只存在于`tokens.css`。 +- 未将测试行情写入产品数据库,浏览器视觉样本由Playwright路由固定,仅用于前端验收。 + +## 自动验证 + +- Ruff:通过。 +- pytest:46项通过,覆盖数据源准入、快照沿用、盘前空K线、iFinD顶层表格与过期授权刷新、搜索权限与固定分组、图表响应不泄露来源。 +- Vue类型检查:通过。 +- Vitest:2个文件、5项通过。 +- Vite生产构建:通过。 +- Playwright:3项通过,覆盖既有Shell回归,以及摘要、搜索、日K预览、详情、分时零轴、日夜主题、1920×1080和390×844视口。 +- 密钥扫描:已提供账号密码与令牌均未进入`next/`;组件CSS未发现令牌外色值。 + +## 截图 + +- [日间搜索与日K预览 1920×1080](search-preview-light-1920x1080.jpg) +- [夜间分时详情 1920×1080](entity-detail-dark-1920x1080.jpg) +- [夜间移动分时详情 390×844](entity-detail-dark-390x844.jpg) + +## 后续入口 + +阶段6将使用本阶段的交易日期、摘要存储和数据质量门实现情绪周期与五类股池。只有阶段6的确定性行情任务可以写入正式市场摘要,展示兜底源不能借图表接口进入计算。 diff --git a/next/docs/evidence/stage-5/entity-detail-dark-1920x1080.jpg b/next/docs/evidence/stage-5/entity-detail-dark-1920x1080.jpg new file mode 100644 index 0000000..9d3f2f6 Binary files /dev/null and b/next/docs/evidence/stage-5/entity-detail-dark-1920x1080.jpg differ diff --git a/next/docs/evidence/stage-5/entity-detail-dark-390x844.jpg b/next/docs/evidence/stage-5/entity-detail-dark-390x844.jpg new file mode 100644 index 0000000..adf20ea Binary files /dev/null and b/next/docs/evidence/stage-5/entity-detail-dark-390x844.jpg differ diff --git a/next/docs/evidence/stage-5/search-preview-light-1920x1080.jpg b/next/docs/evidence/stage-5/search-preview-light-1920x1080.jpg new file mode 100644 index 0000000..7ff958c Binary files /dev/null and b/next/docs/evidence/stage-5/search-preview-light-1920x1080.jpg differ diff --git a/next/frontend/src/app/router.ts b/next/frontend/src/app/router.ts index 8571854..3482f6b 100644 --- a/next/frontend/src/app/router.ts +++ b/next/frontend/src/app/router.ts @@ -2,6 +2,7 @@ import { createRouter, createWebHistory } from "vue-router"; import SystemManagementView from "./views/SystemManagementView.vue"; import WorkspaceView from "./views/WorkspaceView.vue"; +import EntityDetailView from "./views/EntityDetailView.vue"; import { findWorkspace } from "./workspaceRegistry"; export default createRouter({ @@ -15,6 +16,11 @@ export default createRouter({ beforeEnter: (to) => (findWorkspace(String(to.params.workspace)) ? true : "/workspace/emotion"), }, { path: "/system", name: "system", component: SystemManagementView }, + { + path: "/market/:entityType/:identifier", + name: "entity-detail", + component: EntityDetailView, + }, { path: "/:pathMatch(.*)*", redirect: "/workspace/emotion" }, ], }); diff --git a/next/frontend/src/app/shell/AppShell.vue b/next/frontend/src/app/shell/AppShell.vue index cac0a86..5bcf26e 100644 --- a/next/frontend/src/app/shell/AppShell.vue +++ b/next/frontend/src/app/shell/AppShell.vue @@ -4,6 +4,7 @@ import { onBeforeUnmount, onMounted } from "vue"; import DialogHost from "../../shared/components/DialogHost.vue"; import ToastHost from "../../shared/components/ToastHost.vue"; import { useUiStore } from "../../shared/stores/ui"; +import { useMarketStore } from "../../shared/stores/market"; import DesktopSidebar from "./DesktopSidebar.vue"; import MarketStrip from "./MarketStrip.vue"; import MobileNav from "./MobileNav.vue"; @@ -11,6 +12,7 @@ import StatusBar from "./StatusBar.vue"; import TopBar from "./TopBar.vue"; const ui = useUiStore(); +const market = useMarketStore(); function globalShortcut(event: KeyboardEvent): void { if (!(event.ctrlKey || event.metaKey) || event.key.toLowerCase() !== "k") return; @@ -22,7 +24,10 @@ function globalShortcut(event: KeyboardEvent): void { ui.openDialog("search"); } -onMounted(() => window.addEventListener("keydown", globalShortcut)); +onMounted(() => { + window.addEventListener("keydown", globalShortcut); + void market.load(); +}); onBeforeUnmount(() => window.removeEventListener("keydown", globalShortcut)); diff --git a/next/frontend/src/app/shell/MarketStrip.vue b/next/frontend/src/app/shell/MarketStrip.vue index ce47dc1..3b79138 100644 --- a/next/frontend/src/app/shell/MarketStrip.vue +++ b/next/frontend/src/app/shell/MarketStrip.vue @@ -1,28 +1,68 @@