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 )