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

204 lines
6.6 KiB
Python
Raw Normal View History

import logging
from collections.abc import Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from pathlib import Path
import pytest
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
def test_sync_logs_progress_for_initialize_and_daily_update(
tmp_path: Path,
caplog: pytest.LogCaptureFixture,
) -> None:
target = date(2024, 1, 2)
repository = InMemoryRepository()
use_case = SyncMarketData(
FakeSource(target),
CsvSnapshotStore(tmp_path),
repository,
today=target,
)
caplog.set_level(logging.INFO, logger="zhixing_server.modules.market_data.application.sync")
initialize = use_case.execute(
SyncMarketDataCommand(mode="initialize", target_trade_date=target)
)
initialize_messages = [record.getMessage() for record in caplog.records]
assert initialize.status == "success"
assert any("stage=daily_basic" in message for message in initialize_messages)
assert any(
"stage=bar" in message and "progress=1/1" in message
for message in initialize_messages
)
caplog.clear()
daily = use_case.execute(SyncMarketDataCommand(mode="daily", target_trade_date=target))
daily_messages = [record.getMessage() for record in caplog.records]
assert daily.status == "success"
assert any("market_data_sync_started mode=daily" in message for message in daily_messages)
assert any("market_data_sync_finished" in message for message in daily_messages)