refactor: route market providers through data gateway
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from .gateway import DataGateway, build_data_gateway
|
||||
from .policy import DataPolicyError, DataSourcePolicy
|
||||
|
||||
__all__ = ["DataGateway", "DataPolicyError", "DataSourcePolicy", "build_data_gateway"]
|
||||
@@ -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)
|
||||
@@ -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(),
|
||||
)
|
||||
@@ -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
|
||||
@@ -0,0 +1,4 @@
|
||||
from .ifind import IfindProvider
|
||||
from .tushare import TushareProvider
|
||||
|
||||
__all__ = ["IfindProvider", "TushareProvider"]
|
||||
@@ -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)
|
||||
@@ -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())
|
||||
Reference in New Issue
Block a user