189 lines
5.6 KiB
Python
189 lines
5.6 KiB
Python
from datetime import UTC, date, datetime
|
|
from decimal import Decimal
|
|
|
|
from zhixing_server.modules.sector_radar.domain.models import (
|
|
MembershipStatus,
|
|
SectorType,
|
|
StockFactStatus,
|
|
)
|
|
from zhixing_server.modules.sector_radar.domain.normalize import (
|
|
normalize_memberships,
|
|
normalize_stock_facts,
|
|
)
|
|
from zhixing_server.modules.sector_radar.domain.source import (
|
|
DailyRow,
|
|
MoneyflowDcRow,
|
|
SectorIndexRow,
|
|
SectorMemberRow,
|
|
SourceResult,
|
|
StockBasicRow,
|
|
SuspendRow,
|
|
build_source_snapshot,
|
|
)
|
|
|
|
TARGET_DATE = date(2026, 8, 28)
|
|
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
|
|
|
|
|
def result[T](api_name: str, rows: tuple[T, ...]) -> SourceResult[T]:
|
|
snapshot = build_source_snapshot(
|
|
api_name=api_name,
|
|
params={"trade_date": "20260828"},
|
|
rows=(),
|
|
target_trade_date=TARGET_DATE,
|
|
observed_at=OBSERVED_AT,
|
|
)
|
|
return SourceResult((snapshot,), rows)
|
|
|
|
|
|
def basic(ts_code: str, *, list_date: date = date(2020, 1, 1)) -> StockBasicRow:
|
|
return StockBasicRow(
|
|
ts_code=ts_code,
|
|
symbol=ts_code.split(".")[0],
|
|
name=ts_code,
|
|
market="主板",
|
|
exchange="SZSE",
|
|
list_status="L",
|
|
list_date=list_date,
|
|
delist_date=None,
|
|
)
|
|
|
|
|
|
def daily(ts_code: str, amount: Decimal | None) -> DailyRow:
|
|
return DailyRow(
|
|
ts_code=ts_code,
|
|
trade_date=TARGET_DATE,
|
|
close=Decimal("10"),
|
|
pre_close=Decimal("10"),
|
|
pct_chg=Decimal(0),
|
|
volume=Decimal(0),
|
|
amount_thousand_yuan=amount,
|
|
)
|
|
|
|
|
|
def moneyflow(ts_code: str, amount: Decimal | None) -> MoneyflowDcRow:
|
|
return MoneyflowDcRow(
|
|
trade_date=TARGET_DATE,
|
|
ts_code=ts_code,
|
|
name=ts_code,
|
|
net_amount_ten_thousand_yuan=amount,
|
|
net_amount_rate=Decimal(0),
|
|
pct_change=Decimal(0),
|
|
close=Decimal("10"),
|
|
)
|
|
|
|
|
|
def test_membership_normalization_persists_an_explicit_unknown_sector() -> None:
|
|
indices = (
|
|
SectorIndexRow(
|
|
TARGET_DATE,
|
|
SectorType.CONCEPT,
|
|
"BK0001.DC",
|
|
"机器人",
|
|
"一级",
|
|
Decimal(1),
|
|
None,
|
|
),
|
|
SectorIndexRow(
|
|
TARGET_DATE,
|
|
SectorType.CONCEPT,
|
|
"BK0002.DC",
|
|
"低空经济",
|
|
"一级",
|
|
Decimal(1),
|
|
None,
|
|
),
|
|
)
|
|
member = SectorMemberRow(
|
|
TARGET_DATE,
|
|
"BK0001.DC",
|
|
"000001.SZ",
|
|
"平安银行",
|
|
)
|
|
all_snapshot = build_source_snapshot(
|
|
api_name="dc_member",
|
|
params={"trade_date": "20260828"},
|
|
rows=(
|
|
{
|
|
"trade_date": "20260828",
|
|
"ts_code": member.sector_code,
|
|
"con_code": member.stock_code,
|
|
"name": member.stock_name,
|
|
},
|
|
),
|
|
target_trade_date=TARGET_DATE,
|
|
partition_key="all",
|
|
observed_at=OBSERVED_AT,
|
|
)
|
|
empty_partition = build_source_snapshot(
|
|
api_name="dc_member",
|
|
params={"trade_date": "20260828", "ts_code": "BK0002.DC"},
|
|
rows=(),
|
|
target_trade_date=TARGET_DATE,
|
|
partition_key="BK0002.DC",
|
|
observed_at=OBSERVED_AT,
|
|
)
|
|
|
|
records = normalize_memberships(
|
|
indices,
|
|
SourceResult((all_snapshot, empty_partition), (member,)),
|
|
)
|
|
|
|
assert records[0].status is MembershipStatus.AVAILABLE
|
|
assert records[0].stock_code == "000001.SZ"
|
|
assert records[1].status is MembershipStatus.UNKNOWN
|
|
assert records[1].stock_code is None
|
|
assert records[1].membership_key == "__membership_unknown__"
|
|
|
|
|
|
def test_stock_fact_normalization_preserves_all_missing_and_zero_states() -> None:
|
|
codes = tuple(f"00000{index}.SZ" for index in range(1, 9))
|
|
basics = tuple(
|
|
basic(code, list_date=date(2027, 1, 1) if code == codes[7] else date(2020, 1, 1))
|
|
for code in codes
|
|
)
|
|
daily_rows = (
|
|
daily(codes[0], Decimal("1")),
|
|
daily(codes[3], None),
|
|
daily(codes[4], Decimal("1")),
|
|
daily(codes[5], Decimal("1")),
|
|
daily(codes[6], Decimal("0")),
|
|
daily(codes[7], Decimal("1")),
|
|
)
|
|
moneyflow_rows = (
|
|
moneyflow(codes[0], Decimal("0")),
|
|
moneyflow(codes[3], Decimal("1")),
|
|
moneyflow(codes[5], None),
|
|
moneyflow(codes[6], Decimal("0")),
|
|
moneyflow(codes[7], Decimal("1")),
|
|
)
|
|
suspensions = (
|
|
SuspendRow(
|
|
ts_code=codes[1],
|
|
trade_date=TARGET_DATE,
|
|
suspend_timing="09:30",
|
|
suspend_type="停牌",
|
|
),
|
|
)
|
|
|
|
facts = normalize_stock_facts(
|
|
target_trade_date=TARGET_DATE,
|
|
candidate_codes=codes,
|
|
stock_basics=result("stock_basic", basics),
|
|
suspensions=result("suspend_d", suspensions),
|
|
daily=result("daily", daily_rows),
|
|
moneyflow=result("moneyflow_dc", moneyflow_rows),
|
|
)
|
|
by_code = {fact.ts_code: fact for fact in facts}
|
|
|
|
assert by_code[codes[0]].status is StockFactStatus.AVAILABLE
|
|
assert by_code[codes[0]].turnover_yuan == Decimal("1000")
|
|
assert by_code[codes[0]].net_amount_yuan == Decimal("0")
|
|
assert by_code[codes[1]].status is StockFactStatus.SUSPENDED
|
|
assert by_code[codes[2]].status is StockFactStatus.MISSING_DAILY
|
|
assert by_code[codes[3]].status is StockFactStatus.NULL_DAILY_AMOUNT
|
|
assert by_code[codes[4]].status is StockFactStatus.MISSING_MONEYFLOW
|
|
assert by_code[codes[5]].status is StockFactStatus.NULL_MONEYFLOW
|
|
assert by_code[codes[6]].status is StockFactStatus.LOW_LIQUIDITY
|
|
assert by_code[codes[7]].status is StockFactStatus.LIFECYCLE_INVALID
|