feat(market-data): 优化同步并增加完整性检查
This commit is contained in:
@@ -33,6 +33,8 @@ def test_postgres_migration_creates_market_data_contract(
|
||||
"market_daily_basic",
|
||||
"market_sync_batch",
|
||||
"market_sync_item",
|
||||
"market_integrity_check",
|
||||
"market_integrity_issue",
|
||||
"selection_run",
|
||||
"selection_run_item",
|
||||
"selection_signal",
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
from zhixing_server.modules.market_data.application.sync import SyncFailure, SyncItemOutcome
|
||||
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, SyncWindow
|
||||
from zhixing_server.modules.market_data.domain.ports import MarketDataRepositoryError, WriteResult
|
||||
from zhixing_server.modules.market_data.infrastructure.postgres import (
|
||||
PostgresMarketDataRepository,
|
||||
)
|
||||
|
||||
TEST_CODES = ("991901.SZ", "991902.SZ")
|
||||
TARGET = date(2024, 1, 2)
|
||||
WINDOW = SyncWindow(start=TARGET, end=TARGET)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _database(database_url: str) -> Any:
|
||||
with psycopg.connect(database_url) as connection, connection.transaction():
|
||||
yield connection
|
||||
|
||||
|
||||
def _cleanup(database_url: str) -> None:
|
||||
with _database(database_url) as connection:
|
||||
connection.execute(
|
||||
"DELETE FROM market_sync_item WHERE batch_id IN "
|
||||
"(SELECT id FROM market_sync_batch WHERE id LIKE '9919%')"
|
||||
)
|
||||
connection.execute("DELETE FROM market_sync_batch WHERE id LIKE '9919%'")
|
||||
connection.execute(
|
||||
"DELETE FROM market_daily_basic WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM market_daily_bar WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM market_stock WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
|
||||
|
||||
def _prepare_database(database_url: str) -> None:
|
||||
server_root = Path(__file__).parents[2]
|
||||
config = Config(str(server_root / "alembic.ini"))
|
||||
config.set_main_option(
|
||||
"sqlalchemy.url",
|
||||
database_url.replace("%", "%%").replace("postgresql://", "postgresql+psycopg://"),
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
_cleanup(database_url)
|
||||
with _database(database_url) as connection:
|
||||
connection.executemany(
|
||||
"""
|
||||
INSERT INTO market_stock
|
||||
(ts_code, name, market, exchange, list_status, is_active)
|
||||
VALUES (%s, %s, '主板', 'SZSE', 'L', true)
|
||||
""",
|
||||
[(code, f"测试{code}") for code in TEST_CODES],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_pool_upsert_rollback_batch_audit_and_set_coverage() -> 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")
|
||||
_prepare_database(database_url)
|
||||
repository = PostgresMarketDataRepository(database_url, max_connections=4)
|
||||
batch_id: str | None = None
|
||||
try:
|
||||
bars = {code: Bar(code, TARGET, close=Decimal("10")) for code in TEST_CODES}
|
||||
with repository:
|
||||
|
||||
def upsert(code: str) -> WriteResult:
|
||||
return repository.upsert_bars((bars[code],), WINDOW, full_snapshot=True)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
results = tuple(executor.map(upsert, TEST_CODES))
|
||||
assert [result.inserted for result in results] == [1, 1]
|
||||
|
||||
with pytest.raises(MarketDataRepositoryError):
|
||||
repository.upsert_bars(
|
||||
(
|
||||
Bar(TEST_CODES[0], TARGET, close=Decimal("11")),
|
||||
Bar(TEST_CODES[0], TARGET, close=Decimal("12")),
|
||||
),
|
||||
WINDOW,
|
||||
full_snapshot=True,
|
||||
)
|
||||
|
||||
unchanged = repository.upsert_daily_basic(
|
||||
(DailyBasic(TEST_CODES[0], TARGET, close=Decimal("10")),),
|
||||
WINDOW,
|
||||
)
|
||||
assert unchanged.inserted == 1
|
||||
assert repository.count_valid_stocks(TARGET) == 1
|
||||
|
||||
batch_id = "9919-pool-test"
|
||||
outcomes = (
|
||||
SyncItemOutcome("bar", TEST_CODES[0], "success", WriteResult(inserted=1)),
|
||||
SyncItemOutcome(
|
||||
"bar",
|
||||
TEST_CODES[1],
|
||||
"failed",
|
||||
failure=SyncFailure("bar", TEST_CODES[1], "source_error", "safe failure"),
|
||||
),
|
||||
)
|
||||
# Use the real schema row directly because the test repository's
|
||||
# create_batch API generates a UUID for normal production calls.
|
||||
with _database(database_url) as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO market_sync_batch
|
||||
(id, target_trade_date, window_start, mode, status, target_count)
|
||||
VALUES (%s, %s, %s, 'daily', 'running', 2)
|
||||
""",
|
||||
(batch_id, TARGET, TARGET),
|
||||
)
|
||||
repository.record_items(batch_id, outcomes)
|
||||
with _database(database_url) as connection:
|
||||
count = connection.execute(
|
||||
"SELECT count(*) FROM market_sync_item WHERE batch_id = %s",
|
||||
(batch_id,),
|
||||
).fetchone()[0]
|
||||
assert count == 2
|
||||
finally:
|
||||
repository.close()
|
||||
_cleanup(database_url)
|
||||
Reference in New Issue
Block a user