feat(market-data): 迁移 Tushare 同步到 PostgreSQL

This commit is contained in:
yuxuanhui
2026-08-05 17:57:32 +08:00
parent e69bcbe821
commit d7d5ed8979
34 changed files with 3230 additions and 3 deletions
@@ -0,0 +1,31 @@
import os
from pathlib import Path
import pytest
from alembic import command
from alembic.config import Config
from sqlalchemy import Engine, create_engine, inspect
@pytest.mark.integration
def test_postgres_migration_creates_market_data_contract() -> None:
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
server_root = Path(__file__).parents[2]
config = Config(str(server_root / "alembic.ini"))
config.set_main_option("sqlalchemy.url", database_url.replace("%", "%%"))
engine: Engine = create_engine(database_url)
command.upgrade(config, "head")
try:
tables = set(inspect(engine).get_table_names())
assert {
"market_stock",
"market_daily_bar",
"market_daily_basic",
"market_sync_batch",
"market_sync_item",
} <= tables
finally:
engine.dispose()
@@ -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