feat(market-data): 迁移 Tushare 同步到 PostgreSQL
This commit is contained in:
@@ -0,0 +1,165 @@
|
||||
from collections.abc import Generator, Iterable, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
|
||||
from zhixing_server.modules.market_data.application.sync import (
|
||||
SyncMarketData,
|
||||
SyncMarketDataCommand,
|
||||
)
|
||||
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from zhixing_server.modules.market_data.domain.ports import WriteResult
|
||||
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
|
||||
|
||||
|
||||
class FakeSource:
|
||||
def __init__(self, target: date) -> None:
|
||||
self.target = target
|
||||
self.stock = Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L")
|
||||
|
||||
def fetch_stocks(self) -> Sequence[Stock]:
|
||||
return (self.stock,)
|
||||
|
||||
def fetch_open_dates(self, start: date, end: date) -> Sequence[date]:
|
||||
return (self.target,) if start <= self.target <= end else ()
|
||||
|
||||
def fetch_bars(self, ts_code: str, window: SyncWindow) -> Sequence[Bar]:
|
||||
return (Bar(ts_code, self.target, close=Decimal("10")),)
|
||||
|
||||
def fetch_daily_basic(self, trade_date: date) -> Sequence[DailyBasic]:
|
||||
return (DailyBasic("000001.SZ", trade_date, close=Decimal("10")),)
|
||||
|
||||
|
||||
class InMemoryRepository:
|
||||
def __init__(self) -> None:
|
||||
self.bars: dict[tuple[str, date], Bar] = {}
|
||||
self.daily_basic: dict[tuple[str, date], DailyBasic] = {}
|
||||
self.batch_status: str | None = None
|
||||
self.batch_counter = 0
|
||||
|
||||
@contextmanager
|
||||
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
|
||||
yield True
|
||||
|
||||
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult:
|
||||
records = tuple(rows)
|
||||
return WriteResult(inserted=len(records))
|
||||
|
||||
def upsert_bars(
|
||||
self,
|
||||
rows: Iterable[Bar],
|
||||
window: SyncWindow,
|
||||
*,
|
||||
full_snapshot: bool,
|
||||
) -> WriteResult:
|
||||
records = tuple(rows)
|
||||
inserted = 0
|
||||
updated = 0
|
||||
unchanged = 0
|
||||
code = records[0].ts_code
|
||||
if full_snapshot:
|
||||
for key in tuple(self.bars):
|
||||
if (
|
||||
key[0] == code
|
||||
and window.contains(key[1])
|
||||
and key[1] not in {row.trade_date for row in records}
|
||||
):
|
||||
del self.bars[key]
|
||||
for row in records:
|
||||
key = (row.ts_code, row.trade_date)
|
||||
previous = self.bars.get(key)
|
||||
if previous is None:
|
||||
inserted += 1
|
||||
elif previous == row:
|
||||
unchanged += 1
|
||||
else:
|
||||
updated += 1
|
||||
self.bars[key] = row
|
||||
return WriteResult(inserted=inserted, updated=updated, unchanged=unchanged)
|
||||
|
||||
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult:
|
||||
inserted = 0
|
||||
updated = 0
|
||||
unchanged = 0
|
||||
for row in rows:
|
||||
key = (row.ts_code, row.trade_date)
|
||||
previous = self.daily_basic.get(key)
|
||||
if previous is None:
|
||||
inserted += 1
|
||||
elif previous == row:
|
||||
unchanged += 1
|
||||
else:
|
||||
updated += 1
|
||||
self.daily_basic[key] = row
|
||||
return WriteResult(inserted=inserted, updated=updated, unchanged=unchanged)
|
||||
|
||||
def purge_before(self, window: SyncWindow) -> None:
|
||||
self.bars = {key: row for key, row in self.bars.items() if key[1] >= window.start}
|
||||
self.daily_basic = {
|
||||
key: row for key, row in self.daily_basic.items() if key[1] >= window.start
|
||||
}
|
||||
|
||||
def has_bar(self, ts_code: str, trade_date: date) -> bool:
|
||||
return (ts_code, trade_date) in self.bars
|
||||
|
||||
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
|
||||
return (ts_code, trade_date) in self.daily_basic
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
target_trade_date: date,
|
||||
window: SyncWindow,
|
||||
mode: str,
|
||||
parent_batch_id: str | None,
|
||||
target_count: int,
|
||||
) -> str:
|
||||
self.batch_counter += 1
|
||||
return f"batch-{self.batch_counter}"
|
||||
|
||||
def record_batch(
|
||||
self,
|
||||
batch_id: str,
|
||||
status: str,
|
||||
valid_count: int,
|
||||
coverage: Decimal,
|
||||
strategy_eligible: bool,
|
||||
) -> None:
|
||||
self.batch_status = status
|
||||
|
||||
def record_item(
|
||||
self,
|
||||
batch_id: str,
|
||||
item_kind: str,
|
||||
item_key: str,
|
||||
status: str,
|
||||
result: WriteResult,
|
||||
fingerprint: str | None = None,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
|
||||
return ()
|
||||
|
||||
|
||||
def test_sync_is_idempotent_and_reports_coverage(tmp_path: Path) -> None:
|
||||
target = date(2024, 1, 2)
|
||||
repository = InMemoryRepository()
|
||||
use_case = SyncMarketData(
|
||||
FakeSource(target),
|
||||
CsvSnapshotStore(tmp_path),
|
||||
repository,
|
||||
today=target,
|
||||
)
|
||||
|
||||
first = use_case.execute(SyncMarketDataCommand(target_trade_date=target))
|
||||
second = use_case.execute(SyncMarketDataCommand(target_trade_date=target))
|
||||
|
||||
assert first.status == "success"
|
||||
assert first.coverage == Decimal("1")
|
||||
assert first.strategy_eligible
|
||||
assert second.status == "success"
|
||||
assert second.inserted_count == 1
|
||||
assert second.unchanged_count >= 2
|
||||
Reference in New Issue
Block a user