Files
zhixing-system/zhixing-server/tests/unit/market_data/test_sync.py
T

166 lines
5.3 KiB
Python

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