203 lines
6.6 KiB
Python
203 lines
6.6 KiB
Python
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)
|