161 lines
5.9 KiB
Python
161 lines
5.9 KiB
Python
from dataclasses import replace
|
|
from datetime import date, timedelta
|
|
from decimal import Decimal
|
|
|
|
import pytest
|
|
|
|
from zhixing_server.modules.sector_radar.application.scoring import (
|
|
calculate_rankings,
|
|
calendar_rank_changes,
|
|
comparison_dates,
|
|
)
|
|
from zhixing_server.modules.sector_radar.domain.models import (
|
|
MetricKind,
|
|
RankedMetric,
|
|
RankSide,
|
|
SectorDailyAggregate,
|
|
SectorType,
|
|
)
|
|
from zhixing_server.modules.sector_radar.domain.ranking import select_percentile_side
|
|
from zhixing_server.modules.sector_radar.domain.weighted import (
|
|
RATIO_WEIGHTED_VERSION,
|
|
SWING_WEIGHTED_VERSION,
|
|
resolve_metric_version,
|
|
)
|
|
|
|
DAYS = tuple(
|
|
date(2026, 9, 1) + timedelta(days=n)
|
|
for n in range(21)
|
|
if (date(2026, 9, 1) + timedelta(days=n)).weekday() < 5
|
|
)
|
|
|
|
|
|
def aggregate(
|
|
code: str,
|
|
day: date,
|
|
ratio: str,
|
|
turnover: str = "9999999999",
|
|
kind: SectorType = SectorType.CONCEPT,
|
|
) -> SectorDailyAggregate:
|
|
amount = Decimal(turnover)
|
|
return SectorDailyAggregate(
|
|
day,
|
|
kind,
|
|
code,
|
|
code,
|
|
10,
|
|
10,
|
|
Decimal(ratio) * (amount + 100),
|
|
amount,
|
|
Decimal(1),
|
|
Decimal(1),
|
|
)
|
|
|
|
|
|
def scores(
|
|
rows: list[SectorDailyAggregate], target: date = DAYS[9]
|
|
) -> dict[tuple[SectorType, str, MetricKind], RankedMetric]:
|
|
return {
|
|
(row.observation.sector_type, row.observation.sector_code, row.observation.metric_kind): row
|
|
for row in calculate_rankings(
|
|
target,
|
|
[item for item in rows if item.trade_date == target],
|
|
[item for item in rows if item.trade_date != target],
|
|
DAYS,
|
|
)
|
|
}
|
|
|
|
|
|
def test_liquidity_weight_can_reverse_raw_ratio_order_and_separates_pools() -> None:
|
|
rows = [
|
|
aggregate(code, day, ratio, turnover, kind)
|
|
for day in DAYS[:10]
|
|
for code, ratio, turnover, kind in [
|
|
("small", ".3", "999999", SectorType.CONCEPT),
|
|
("large", ".2", "999999999999", SectorType.CONCEPT),
|
|
("medium", ".1", "9999999999", SectorType.CONCEPT),
|
|
("industry", "-.5", "9999999999", SectorType.INDUSTRY),
|
|
]
|
|
]
|
|
ranked = scores(rows)
|
|
small = ranked[(SectorType.CONCEPT, "small", MetricKind.RATIO)]
|
|
large = ranked[(SectorType.CONCEPT, "large", MetricKind.RATIO)]
|
|
assert small.observation.value == Decimal(".3")
|
|
assert small.observation.weighted_score == 600
|
|
assert large.observation.weighted_score == 800
|
|
assert large.rank_position == 1 and small.rank_position == 2
|
|
assert (
|
|
ranked[(SectorType.INDUSTRY, "industry", MetricKind.RATIO)].observation.weighted_score
|
|
== 1000
|
|
)
|
|
|
|
|
|
def test_swing_ranks_each_window_before_combining_not_the_averaged_raw_ratio() -> None:
|
|
rows = [
|
|
aggregate(code, day, ratio)
|
|
for index, day in enumerate(DAYS[:10])
|
|
for code, ratio in [("A", ".1" if index >= 7 else "-.1"), ("B", ".03"), ("C", ".05")]
|
|
]
|
|
ranked = scores(rows)
|
|
a = ranked[(SectorType.CONCEPT, "A", MetricKind.SWING)]
|
|
b = ranked[(SectorType.CONCEPT, "B", MetricKind.SWING)]
|
|
assert a.observation.value == b.observation.value == Decimal(".03")
|
|
assert a.observation.weighted_score == pytest.approx(Decimal("666.6666666666667"))
|
|
assert b.observation.weighted_score == 500
|
|
assert a.rank_position == 2 and b.rank_position == 3
|
|
|
|
|
|
def test_tied_features_use_average_percentiles_and_code_breaks_final_ties() -> None:
|
|
ranked = scores([aggregate(code, day, "0") for day in DAYS[:10] for code in ["B", "A"]])
|
|
assert ranked[(SectorType.CONCEPT, "A", MetricKind.RATIO)].observation.weighted_score == 750
|
|
assert ranked[(SectorType.CONCEPT, "B", MetricKind.RATIO)].observation.weighted_score == 750
|
|
assert ranked[(SectorType.CONCEPT, "A", MetricKind.RATIO)].rank_position == 1
|
|
|
|
|
|
def test_missing_calendar_session_is_not_replaced_by_an_older_success() -> None:
|
|
rows = [aggregate("A", day, ".2") for day in DAYS[:10] if day != DAYS[7]]
|
|
ranked = scores(rows)
|
|
daily = ranked[(SectorType.CONCEPT, "A", MetricKind.RATIO)]
|
|
assert daily.observation.value == Decimal(".2")
|
|
assert daily.observation.weighted_score is None and daily.rank_position is None
|
|
assert ranked[(SectorType.CONCEPT, "A", MetricKind.SWING)].observation.value is None
|
|
|
|
|
|
def test_future_inputs_do_not_affect_scores_and_bottom_uses_score_order() -> None:
|
|
rows = [
|
|
aggregate(f"C{n}", day, str(n), str(10 ** (6 + n) - 1))
|
|
for day in DAYS[:10]
|
|
for n in range(1, 11)
|
|
]
|
|
original = scores(rows)
|
|
future = scores(rows + [aggregate("C1", DAYS[10], "1000000")])
|
|
assert future == original
|
|
ratio_rows = [
|
|
row for row in original.values() if row.observation.metric_kind is MetricKind.RATIO
|
|
]
|
|
bottom = select_percentile_side(ratio_rows, RankSide.BOTTOM)
|
|
assert [row.observation.sector_code for row in bottom] == ["C1"]
|
|
|
|
|
|
def test_rank_changes_resolve_trading_dates_and_do_not_mix_versions() -> None:
|
|
rows = [aggregate("A", day, ".1") for day in DAYS[:11]]
|
|
before = tuple(scores(rows, DAYS[9]).values())
|
|
current = tuple(scores(rows, DAYS[10]).values())
|
|
assert comparison_dates(DAYS, DAYS[10])[1] == DAYS[9]
|
|
missing = calendar_rank_changes(current, {DAYS[8]: before}, DAYS, DAYS[10])
|
|
assert all(row.rank_change(1) is None for row in missing)
|
|
older = tuple(
|
|
replace(row, observation=replace(row.observation, metric_version="legacy"))
|
|
for row in before
|
|
)
|
|
incompatible = calendar_rank_changes(current, {DAYS[9]: older}, DAYS, DAYS[10])
|
|
assert all(row.rank_change(1) is None for row in incompatible)
|
|
comparable = calendar_rank_changes(current, {DAYS[9]: before}, DAYS, DAYS[10])
|
|
assert all(row.rank_change(1) == 0 for row in comparable)
|
|
assert (
|
|
resolve_metric_version(MetricKind.RATIO, [RATIO_WEIGHTED_VERSION]) == RATIO_WEIGHTED_VERSION
|
|
)
|
|
assert (
|
|
resolve_metric_version(MetricKind.SWING, [SWING_WEIGHTED_VERSION]) == SWING_WEIGHTED_VERSION
|
|
)
|