refactor: route market providers through data gateway

This commit is contained in:
leefer
2026-07-29 17:25:21 +08:00
parent 3994387935
commit 7c8b8ca21e
12 changed files with 318 additions and 30 deletions
+10 -8
View File
@@ -2,9 +2,11 @@ from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from collections.abc import Callable
from alert_service import AlertService
from chart_data_provider import EastmoneyChartClient, MarketChartClient
from backend.data import DataGateway, build_data_gateway
from chart_data_provider import MarketChartClient
from database import ReviewDatabase
from ifind_client import IfindHttpClient
from mentor_agent import MentorSkillRegistry
@@ -17,6 +19,7 @@ from trade_journal import TradeJournalService
@dataclass(frozen=True)
class ApplicationContainer:
database: ReviewDatabase
data_gateway: DataGateway
ifind: IfindHttpClient
screener: ScreenerEngine
strategy_tracking: StrategyTrackingService
@@ -32,19 +35,18 @@ def build_application_container(
credentials: dict[str, object],
mentor_skills_dir: Path,
private_mentor_skills_dir: Path,
tushare_token_supplier: Callable[[], str] | None = None,
) -> ApplicationContainer:
ifind = IfindHttpClient(
str(credentials.get("ifind_refresh_token") or ""),
str(credentials.get("ifind_access_token") or ""),
)
data_gateway = build_data_gateway(credentials, tushare_token_supplier)
return ApplicationContainer(
database=database,
ifind=ifind,
data_gateway=data_gateway,
ifind=data_gateway.ifind,
screener=ScreenerEngine(database),
strategy_tracking=StrategyTrackingService(database),
alert_service=AlertService(database),
trade_journal=TradeJournalService(database),
mentor_skills=MentorSkillRegistry(mentor_skills_dir, private_mentor_skills_dir),
realtime_aggregator=WebRealtimeAggregator(),
chart_data=MarketChartClient(ifind, EastmoneyChartClient()),
realtime_aggregator=data_gateway.realtime_observer,
chart_data=data_gateway.chart_data,
)
+4
View File
@@ -0,0 +1,4 @@
from .gateway import DataGateway, build_data_gateway
from .policy import DataPolicyError, DataSourcePolicy
__all__ = ["DataGateway", "DataPolicyError", "DataSourcePolicy", "build_data_gateway"]
+29
View File
@@ -0,0 +1,29 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
DataUsage = Literal["display", "calculation"]
@dataclass(frozen=True)
class ProviderContract:
id: str
provider_class: str
calculation_allowed: bool
@dataclass(frozen=True)
class DatasetContract:
id: str
entity: str
frequency: str
primary: str
fallbacks: tuple[str, ...]
usage: str
fields: tuple[str, ...]
@property
def providers(self) -> tuple[str, ...]:
return (self.primary, *self.fallbacks)
+57
View File
@@ -0,0 +1,57 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from backend.data.contracts import DataUsage
from backend.data.policy import DataSourcePolicy
from backend.data.providers import IfindProvider, TushareProvider
from chart_data_provider import EastmoneyChartClient, MarketChartClient
from ifind_client import IfindHttpClient
from realtime_aggregator import WebRealtimeAggregator
from tushare_client import TushareClient
@dataclass(frozen=True)
class DataGateway:
policy: DataSourcePolicy
tushare_provider: TushareProvider
ifind_provider: IfindProvider
chart_data: MarketChartClient
realtime_observer: WebRealtimeAggregator
@property
def ifind(self) -> IfindHttpClient:
return self.ifind_provider.client
def tushare(
self,
dataset_id: str = "",
usage: DataUsage = "calculation",
) -> TushareClient:
if dataset_id:
self.policy.assert_allowed(dataset_id, "tushare", usage)
return self.tushare_provider.client()
def assert_source(self, dataset_id: str, provider_id: str, usage: DataUsage) -> None:
self.policy.assert_allowed(dataset_id, provider_id, usage)
def build_data_gateway(
credentials: dict[str, object],
tushare_token_supplier: Callable[[], str] | None = None,
) -> DataGateway:
ifind = IfindHttpClient(
str(credentials.get("ifind_refresh_token") or ""),
str(credentials.get("ifind_access_token") or ""),
)
token_supplier = tushare_token_supplier or (
lambda: str(credentials.get("tushare_token") or "")
)
return DataGateway(
policy=DataSourcePolicy.load(),
tushare_provider=TushareProvider(token_supplier),
ifind_provider=IfindProvider(ifind),
chart_data=MarketChartClient(ifind, EastmoneyChartClient()),
realtime_observer=WebRealtimeAggregator(),
)
+77
View File
@@ -0,0 +1,77 @@
from __future__ import annotations
import json
from pathlib import Path
from app_config import APP_DIR
from backend.data.contracts import DataUsage, DatasetContract, ProviderContract
class DataPolicyError(RuntimeError):
pass
class DataSourcePolicy:
def __init__(
self,
providers: dict[str, ProviderContract],
datasets: dict[str, DatasetContract],
) -> None:
self.providers = dict(providers)
self.datasets = dict(datasets)
@classmethod
def load(cls, path: Path | None = None) -> "DataSourcePolicy":
config_path = path or APP_DIR / "config" / "data-fields.config.json"
payload = json.loads(config_path.read_text(encoding="utf-8"))
providers = {
provider_id: ProviderContract(
id=provider_id,
provider_class=str(item["class"]),
calculation_allowed=bool(item["calculation_allowed"]),
)
for provider_id, item in payload["providers"].items()
}
datasets = {
item["id"]: DatasetContract(
id=str(item["id"]),
entity=str(item["entity"]),
frequency=str(item["frequency"]),
primary=str(item["primary"]),
fallbacks=tuple(str(value) for value in item.get("fallbacks", [])),
usage=str(item["usage"]),
fields=tuple(str(value) for value in item.get("fields", [])),
)
for item in payload["datasets"]
}
return cls(providers, datasets)
def dataset(self, dataset_id: str) -> DatasetContract:
try:
return self.datasets[dataset_id]
except KeyError as exc:
raise DataPolicyError(f"Unregistered dataset: {dataset_id}") from exc
def assert_allowed(
self,
dataset_id: str,
provider_id: str,
usage: DataUsage,
) -> DatasetContract:
dataset = self.dataset(dataset_id)
if dataset.usage == "blocked":
raise DataPolicyError(f"Dataset is blocked: {dataset_id}")
if provider_id not in dataset.providers:
raise DataPolicyError(
f"Provider {provider_id} is not registered for dataset {dataset_id}"
)
try:
provider = self.providers[provider_id]
except KeyError as exc:
raise DataPolicyError(f"Unregistered provider: {provider_id}") from exc
if usage == "calculation":
if dataset.usage != "calculation" or not provider.calculation_allowed:
raise DataPolicyError(
f"Provider {provider_id} cannot calculate dataset {dataset_id}"
)
return dataset
+4
View File
@@ -0,0 +1,4 @@
from .ifind import IfindProvider
from .tushare import TushareProvider
__all__ = ["IfindProvider", "TushareProvider"]
+11
View File
@@ -0,0 +1,11 @@
from __future__ import annotations
from ifind_client import IfindHttpClient
class IfindProvider:
def __init__(self, client: IfindHttpClient) -> None:
self.client = client
def set_credentials(self, refresh_token: str, access_token: str = "") -> None:
self.client.set_credentials(refresh_token, access_token)
+18
View File
@@ -0,0 +1,18 @@
from __future__ import annotations
from collections.abc import Callable
from tushare_client import TushareClient
class TushareProvider:
def __init__(
self,
token_supplier: Callable[[], str],
client_factory: Callable[[str], TushareClient] = TushareClient,
) -> None:
self._token_supplier = token_supplier
self._client_factory = client_factory
def client(self) -> TushareClient:
return self._client_factory(str(self._token_supplier() or "").strip())