176 lines
6.0 KiB
Python
176 lines
6.0 KiB
Python
import os
|
|
from dataclasses import replace
|
|
from datetime import UTC, date, datetime, timedelta
|
|
from decimal import Decimal
|
|
from pathlib import Path
|
|
|
|
import psycopg
|
|
import pytest
|
|
from alembic import command
|
|
from alembic.config import Config
|
|
|
|
from zhixing_server.bootstrap.config import 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("%", "%%"))
|
|
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"),
|
|
),
|
|
)
|
|
).inserted
|
|
== 1
|
|
)
|
|
|
|
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,),
|
|
)
|