feat(sector-radar): add weighted scores and rank-change views
This commit is contained in:
@@ -0,0 +1,160 @@
|
||||
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
|
||||
)
|
||||
Reference in New Issue
Block a user