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)