import os from dataclasses import replace from datetime import UTC, date, datetime, timedelta from decimal import Decimal from pathlib import Path from unittest.mock import patch import psycopg import pytest from alembic import command from alembic.config import Config from zhixing_server.bootstrap.config import Settings, sqlalchemy_database_url from zhixing_server.modules.sector_radar.domain.models import ( MembershipStatus, PublicationStatus, RadarPublication, SectorType, StockFactStatus, ) from zhixing_server.modules.sector_radar.domain.persistence import ( MembershipRecord, StockFactRecord, ) from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot from zhixing_server.modules.sector_radar.infrastructure.postgres import ( PostgresSectorRadarRepository, ) TARGET_DATE = date(2099, 1, 4) STARTED_AT = datetime(2099, 1, 4, 17, 30, tzinfo=UTC) def prepare_database(database_url: str) -> None: server_root = Path(__file__).parents[2] config = Config(str(server_root / "alembic.ini")) sqlalchemy_url = sqlalchemy_database_url(database_url) config.set_main_option("sqlalchemy.url", sqlalchemy_url.replace("%", "%%")) config.config_file_name = None with patch( "zhixing_server.bootstrap.config.get_settings", return_value=Settings(database_url=database_url), ): command.upgrade(config, "head") @pytest.mark.integration def test_postgres_sector_radar_revisions_and_last_good() -> 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) snapshot = build_source_snapshot( api_name="dc_member", params={"trade_date": "20990104"}, rows=( { "trade_date": "20990104", "ts_code": "BKTEST.DC", "con_code": "000001.SZ", "name": "测试股票", }, ), target_trade_date=TARGET_DATE, observed_at=STARTED_AT, ) fact_revision = "b" * 64 publication_ids = ("test-sector-radar-success", "test-sector-radar-failed") with psycopg.connect(database_url) as connection, connection.transaction(): connection.execute( "DELETE FROM sector_radar_publication WHERE id = ANY(%s)", (list(publication_ids),), ) connection.execute( "DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s", (fact_revision,), ) connection.execute( "DELETE FROM sector_radar_source_snapshot WHERE id = %s", (snapshot.snapshot_id,), ) repository = PostgresSectorRadarRepository(database_url, max_connections=2) try: assert repository.save_source_snapshots((snapshot,)).inserted == 1 assert ( repository.save_source_snapshots( (replace(snapshot, observed_at=STARTED_AT + timedelta(minutes=1)),) ).unchanged == 1 ) assert ( repository.save_memberships( ( MembershipRecord( source_snapshot_id=snapshot.snapshot_id, trade_date=TARGET_DATE, sector_type=SectorType.CONCEPT, sector_code="BKTEST.DC", sector_name="测试概念", stock_code="000001.SZ", stock_name="测试股票", status=MembershipStatus.AVAILABLE, ), ) ).inserted == 1 ) assert ( repository.save_stock_facts( ( StockFactRecord( fact_revision=fact_revision, source_snapshot_ids=(snapshot.snapshot_id,), trade_date=TARGET_DATE, ts_code="000001.SZ", status=StockFactStatus.AVAILABLE, turnover_yuan=Decimal("1000"), net_amount_yuan=Decimal("100"), pct_change=Decimal("1.25"), active_buy_net_amount_yuan=Decimal("-25000"), ), ) ).inserted == 1 ) with psycopg.connect(database_url) as connection: detail_fact = connection.execute( "SELECT pct_change, active_buy_net_amount_yuan " "FROM sector_radar_stock_fact WHERE fact_revision = %s", (fact_revision,), ).fetchone() assert detail_fact == (Decimal("1.25"), Decimal("-25000")) running = RadarPublication( publication_id=publication_ids[0], target_trade_date=TARGET_DATE, status=PublicationStatus.RUNNING, source_version="tushare-pro-v1", universe_version=snapshot.content_sha256, metric_versions=("zhixing_amount_net_bn_v1",), input_hash=None, coverage=Decimal(0), started_at=STARTED_AT, ) repository.create_publication(running) repository.finish_publication( replace( running, status=PublicationStatus.SUCCESS, input_hash="a" * 64, coverage=Decimal(1), finished_at=STARTED_AT + timedelta(minutes=5), ) ) failed = replace( running, publication_id=publication_ids[1], started_at=STARTED_AT + timedelta(minutes=6), ) repository.create_publication(failed) repository.finish_publication( replace( failed, status=PublicationStatus.FAILED, coverage=Decimal("0.8"), finished_at=STARTED_AT + timedelta(minutes=7), error_summary="safe_error", ) ) last_good = repository.get_last_good_publication(TARGET_DATE) assert last_good is not None assert last_good.publication_id == publication_ids[0] finally: repository.close() with psycopg.connect(database_url) as connection, connection.transaction(): connection.execute( "DELETE FROM sector_radar_publication WHERE id = ANY(%s)", (list(publication_ids),), ) connection.execute( "DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s", (fact_revision,), ) connection.execute( "DELETE FROM sector_radar_source_snapshot WHERE id = %s", (snapshot.snapshot_id,), )