feat(market-data): 迁移 Tushare 同步到 PostgreSQL
This commit is contained in:
@@ -0,0 +1,46 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic
|
||||
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
|
||||
|
||||
|
||||
def make_bar(trade_date: date, close: str = "10") -> Bar:
|
||||
return Bar(
|
||||
ts_code="000001.SZ",
|
||||
trade_date=trade_date,
|
||||
close=Decimal(close),
|
||||
)
|
||||
|
||||
|
||||
def make_basic(trade_date: date) -> DailyBasic:
|
||||
return DailyBasic(
|
||||
ts_code="000001.SZ",
|
||||
trade_date=trade_date,
|
||||
close=Decimal("10"),
|
||||
total_mv=Decimal("100000"),
|
||||
)
|
||||
|
||||
|
||||
def test_bar_snapshot_is_atomic_and_failed_publish_preserves_old_file(tmp_path: Path) -> None:
|
||||
store = CsvSnapshotStore(tmp_path)
|
||||
original = (make_bar(date(2024, 1, 2)),)
|
||||
staged = store.stage_bars("000001.SZ", original)
|
||||
store.publish(staged)
|
||||
|
||||
replacement = store.stage_bars("000001.SZ", (make_bar(date(2024, 1, 2), "11"),))
|
||||
store.discard(replacement)
|
||||
|
||||
assert store.read_bars("000001.SZ") == original
|
||||
assert not replacement.temporary_path.exists()
|
||||
|
||||
|
||||
def test_daily_basic_snapshot_rejects_duplicate_codes(tmp_path: Path) -> None:
|
||||
store = CsvSnapshotStore(tmp_path)
|
||||
target = date(2024, 1, 2)
|
||||
|
||||
with pytest.raises(ValueError, match="duplicate"):
|
||||
store.stage_daily_basic(target, (make_basic(target), make_basic(target)))
|
||||
@@ -0,0 +1,73 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from zhixing_server.modules.market_data.domain.fingerprint import (
|
||||
SnapshotChange,
|
||||
compare_snapshots,
|
||||
snapshot_fingerprint,
|
||||
)
|
||||
from zhixing_server.modules.market_data.domain.models import Bar, Stock, SyncWindow
|
||||
from zhixing_server.modules.market_data.domain.rules import filter_current_hs_a_stocks
|
||||
|
||||
|
||||
def make_bar(trade_date: date, close: str = "10") -> Bar:
|
||||
return Bar(
|
||||
ts_code="000001.SZ",
|
||||
trade_date=trade_date,
|
||||
open=Decimal("9"),
|
||||
high=Decimal(close),
|
||||
low=Decimal("8"),
|
||||
close=Decimal(close),
|
||||
pre_close=Decimal("9"),
|
||||
change=Decimal("1"),
|
||||
pct_chg=Decimal("11.11"),
|
||||
vol=Decimal("100"),
|
||||
amount=Decimal("1000"),
|
||||
)
|
||||
|
||||
|
||||
def test_window_uses_inclusive_calendar_boundary() -> None:
|
||||
window = SyncWindow.from_target(date(2024, 2, 29))
|
||||
|
||||
assert window.start == date(2018, 2, 28)
|
||||
assert window.end == date(2024, 2, 29)
|
||||
assert window.contains(date(2018, 2, 28))
|
||||
|
||||
|
||||
def test_universe_keeps_current_non_st_hs_a_stocks() -> None:
|
||||
stocks = (
|
||||
Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L"),
|
||||
Stock("600000.SH", "浦发银行", exchange="SSE", list_status="L"),
|
||||
Stock("300001.SZ", "特锐德", exchange="SZSE", list_status="L"),
|
||||
Stock("600001.SH", "*ST风险", exchange="SSE", list_status="L"),
|
||||
Stock("830001.BJ", "北交所", exchange="BSE", list_status="L"),
|
||||
)
|
||||
|
||||
assert [stock.ts_code for stock in filter_current_hs_a_stocks(stocks)] == [
|
||||
"000001.SZ",
|
||||
"300001.SZ",
|
||||
"600000.SH",
|
||||
]
|
||||
|
||||
|
||||
def test_fingerprint_is_order_independent_and_excludes_runtime_metadata() -> None:
|
||||
rows = (make_bar(date(2024, 1, 2)), make_bar(date(2024, 1, 3)))
|
||||
|
||||
assert snapshot_fingerprint(rows) == snapshot_fingerprint(tuple(reversed(rows)))
|
||||
assert compare_snapshots(rows, rows).change is SnapshotChange.UNCHANGED
|
||||
assert (
|
||||
compare_snapshots(rows, (*rows, make_bar(date(2024, 1, 4)))).change
|
||||
is SnapshotChange.NEW_DATES
|
||||
)
|
||||
|
||||
|
||||
def test_fingerprint_detects_repairs_missing_rows_and_earlier_rows() -> None:
|
||||
old = (make_bar(date(2024, 1, 2)), make_bar(date(2024, 1, 3)))
|
||||
|
||||
changed = (make_bar(date(2024, 1, 2), close="11"), make_bar(date(2024, 1, 3)))
|
||||
missing = (make_bar(date(2024, 1, 2)),)
|
||||
earlier = (make_bar(date(2024, 1, 1)), *old)
|
||||
|
||||
assert compare_snapshots(old, changed).change is SnapshotChange.CHANGED
|
||||
assert compare_snapshots(old, missing).change is SnapshotChange.CHANGED
|
||||
assert compare_snapshots(old, earlier).change is SnapshotChange.CHANGED
|
||||
@@ -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