feat(sector-radar): add weighted scores and rank-change views
This commit is contained in:
@@ -599,10 +599,14 @@ def test_tenth_trading_day_publishes_swing_and_five_rank_changes() -> None:
|
||||
swing = tuple(
|
||||
ranking
|
||||
for ranking in current
|
||||
if ranking.observation.metric_version == "zhixing_swing_equal_3_10_v1"
|
||||
if ranking.observation.metric_version == "zhixing_swing_weighted_v2"
|
||||
)
|
||||
assert len(swing) == 2
|
||||
assert all(ranking.observation.value == Decimal("0.03") for ranking in swing)
|
||||
assert all(
|
||||
ranking.observation.value == pytest.approx(Decimal(150000) / Decimal(5000100))
|
||||
for ranking in swing
|
||||
)
|
||||
assert all(ranking.observation.weighted_score is not None for ranking in swing)
|
||||
assert all(
|
||||
tuple(change.value for change in ranking.rank_changes) == (None, None, None, None, None)
|
||||
for ranking in swing
|
||||
@@ -786,7 +790,7 @@ def test_detail_history_and_ranking_extras_http_use_the_same_publication() -> No
|
||||
row = ranking.json()["rows"][0]
|
||||
assert row["pct_change"] == "1"
|
||||
assert Decimal(row["daily_net_amount_yuan"]) == 150000
|
||||
assert Decimal(row["daily_ratio"]) == Decimal("0.03")
|
||||
assert Decimal(row["daily_ratio"]) == Decimal(150000) / Decimal(5000100)
|
||||
assert row["on_list_count"] == row["history_available_days"] == 1
|
||||
assert client.get(base + "/detail").status_code == 422
|
||||
absent = client.get(base + "/detail", params={"trade_date": "2020-01-01"})
|
||||
@@ -805,6 +809,13 @@ def test_postgres_detail_migration_and_build_roundtrip(monkeypatch: pytest.Monke
|
||||
|
||||
from zhixing_server.bootstrap.config import sqlalchemy_database_url
|
||||
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
|
||||
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
|
||||
from zhixing_server.modules.sector_radar.domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.models import MetricKind
|
||||
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
|
||||
PostgresSectorRadarRepository,
|
||||
)
|
||||
@@ -843,7 +854,7 @@ def test_postgres_detail_migration_and_build_roundtrip(monkeypatch: pytest.Monke
|
||||
assert detail.summary["amount"].metric_value == Decimal("0.0015")
|
||||
with psycopg.connect(database_url) as connection:
|
||||
assert connection.execute("SELECT version_num FROM alembic_version").fetchone() == (
|
||||
"0009_radar_sector_detail",
|
||||
"0010_radar_weighted_score",
|
||||
)
|
||||
row = connection.execute(
|
||||
"SELECT pct_change, leading_code FROM sector_radar_daily_aggregate "
|
||||
@@ -863,5 +874,140 @@ def test_postgres_detail_migration_and_build_roundtrip(monkeypatch: pytest.Monke
|
||||
"SET active_buy_net_amount_yuan = 'NaN'::numeric WHERE trade_date = %s",
|
||||
(target,),
|
||||
)
|
||||
# Exercise every ranking projection against persisted v2 scores, while
|
||||
# keeping the original source publication and its NULL score readable.
|
||||
start, end = target + timedelta(days=31), target + timedelta(days=46)
|
||||
interval = BuildSectorRadarCommand(start_date=start, end_date=end)
|
||||
source = FakeRadarSource()
|
||||
legacy = BuildSectorRadar(
|
||||
source,
|
||||
repository,
|
||||
now_fn=lambda: NOW,
|
||||
strategies=(
|
||||
AmountNetStrategy(),
|
||||
RatioTurnoverStrategy(),
|
||||
SwingEqualThreeToTenStrategy(),
|
||||
),
|
||||
).execute(interval)
|
||||
assert legacy.status in ("success", "unchanged")
|
||||
old_id = legacy.outcomes[-1].publication_id
|
||||
assert old_id is not None
|
||||
old_rows = tuple(repository.load_rankings(old_id))
|
||||
assert all(row.observation.weighted_score is None for row in old_rows)
|
||||
source.calls.clear()
|
||||
rescore = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1))
|
||||
updated = rescore.execute(interval)
|
||||
assert updated.status in ("success", "unchanged")
|
||||
updated_id = updated.outcomes[-1].publication_id
|
||||
assert updated_id is not None and updated_id != old_id
|
||||
assert source.calls == []
|
||||
assert tuple(repository.load_rankings(old_id)) == old_rows
|
||||
current_rows = tuple(repository.load_rankings(updated_id))
|
||||
scored = next(
|
||||
row
|
||||
for row in current_rows
|
||||
if row.observation.metric_kind is MetricKind.SWING
|
||||
and row.observation.sector_type is SectorType.CONCEPT
|
||||
)
|
||||
expected = ((Decimal(5000000) + 1).log10() * 100).quantize(Decimal("0.000000000001"))
|
||||
assert scored.observation.weighted_score == expected
|
||||
assert scored.rank_change(5) == 0
|
||||
assert dict(repository.load_publication_rankings((updated_id,)))[updated_id] == current_rows
|
||||
previous = dict(repository.load_previous_rankings(end + timedelta(days=1), limit_dates=1))
|
||||
assert previous[end] == current_rows
|
||||
historical = repository.load_ranked_history((updated_id,), ("BK0001.DC",))
|
||||
assert any(row.ranking == scored for row in historical)
|
||||
detail = ReadRadarDetails(repository).detail(end, SectorType.CONCEPT, "BK0001.DC")
|
||||
assert detail.summary["swing"].weighted_score == expected
|
||||
assert detail.pct_change == 1 and len(detail.members) == 5
|
||||
assert rescore.execute(interval).status == "unchanged"
|
||||
with (
|
||||
psycopg.connect(database_url) as connection,
|
||||
pytest.raises(psycopg.errors.CheckViolation),
|
||||
connection.transaction(),
|
||||
):
|
||||
connection.execute(
|
||||
"UPDATE sector_radar_ranking SET weighted_score = 'NaN'::numeric "
|
||||
"WHERE publication_id = %s",
|
||||
(updated_id,),
|
||||
)
|
||||
finally:
|
||||
repository.close()
|
||||
|
||||
|
||||
def test_offline_rescore_is_idempotent_preserves_old_publications_and_details() -> None:
|
||||
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
|
||||
from zhixing_server.modules.sector_radar.application.read import (
|
||||
RadarQuery,
|
||||
RadarView,
|
||||
ReadSectorRadar,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
|
||||
from zhixing_server.modules.sector_radar.domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
|
||||
source = FakeRadarSource()
|
||||
repository = InMemorySectorRadarRepository()
|
||||
end = TARGET_DATE + timedelta(days=15)
|
||||
command = BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
|
||||
initial = BuildSectorRadar(
|
||||
source,
|
||||
repository,
|
||||
now_fn=lambda: NOW,
|
||||
strategies=(AmountNetStrategy(), RatioTurnoverStrategy(), SwingEqualThreeToTenStrategy()),
|
||||
).execute(command)
|
||||
assert initial.status == "success"
|
||||
old_publications = repository.publications.copy()
|
||||
old_rankings = repository.rankings.copy()
|
||||
snapshot_count = len(repository.source_snapshots)
|
||||
source.calls.clear()
|
||||
source.fail_daily = True
|
||||
service = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1))
|
||||
result = service.execute(command)
|
||||
assert result.status == "success"
|
||||
assert all(repository.publications[key] == value for key, value in old_publications.items())
|
||||
assert all(repository.rankings[key] == value for key, value in old_rankings.items())
|
||||
assert len(repository.source_snapshots) == snapshot_count
|
||||
assert source.calls == []
|
||||
latest = repository.get_last_good_publication(end)
|
||||
assert latest is not None
|
||||
assert "zhixing_swing_weighted_v2" in latest.metric_versions
|
||||
count = len(repository.publications)
|
||||
assert service.execute(command).status == "unchanged"
|
||||
assert len(repository.publications) == count
|
||||
reader = ReadSectorRadar(repository)
|
||||
page = reader.query(RadarQuery(trade_date=end, view=RadarView.SWING))
|
||||
assert len(page.rows) == 1
|
||||
assert page.rows[0].observation.weighted_score is not None
|
||||
assert page.rows[0].rank_change(5) == 0
|
||||
detail = ReadRadarDetails(repository).detail(end, SectorType.CONCEPT, "BK0001.DC")
|
||||
assert len(detail.members) == 5
|
||||
assert detail.pct_change == 1
|
||||
assert detail.summary["swing"].weighted_score == page.rows[0].observation.weighted_score
|
||||
assert detail.summary["swing"].rank_position == page.rows[0].rank_position
|
||||
|
||||
|
||||
def test_offline_rescore_failure_does_not_replace_last_good(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from unittest.mock import Mock
|
||||
|
||||
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
|
||||
|
||||
repository = InMemorySectorRadarRepository()
|
||||
command = BuildSectorRadarCommand(trade_date=TARGET_DATE)
|
||||
BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(command)
|
||||
before = repository.get_last_good_publication()
|
||||
monkeypatch.setattr(
|
||||
repository, "finalize_publication", Mock(side_effect=RuntimeError("injected failure"))
|
||||
)
|
||||
result = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1)).execute(
|
||||
command
|
||||
)
|
||||
assert result.status == "failed"
|
||||
assert repository.get_last_good_publication() == before
|
||||
latest = repository.get_latest_publication()
|
||||
assert latest is not None and latest.status is PublicationStatus.FAILED
|
||||
|
||||
@@ -61,6 +61,7 @@ class FakeConnection:
|
||||
1,
|
||||
Decimal(100),
|
||||
{"1": 3, "2": None},
|
||||
None,
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -161,6 +162,7 @@ def test_load_rankings_reconstructs_values_and_rank_changes() -> None:
|
||||
ranking = rankings[0]
|
||||
assert ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
|
||||
assert ranking.observation.value == Decimal("12.5")
|
||||
assert ranking.observation.weighted_score is None
|
||||
assert ranking.rank_position == 1
|
||||
assert ranking.rank_change(1) == 3
|
||||
assert ranking.rank_change(2) is None
|
||||
|
||||
@@ -158,7 +158,45 @@ def test_percentile_side_is_selected_before_search_and_pagination() -> None:
|
||||
|
||||
|
||||
def test_rank_change_uses_selected_metric_days_and_pool_sides() -> None:
|
||||
reader = ReadSectorRadar(_published_repository())
|
||||
from zhixing_server.modules.sector_radar.domain.persistence import (
|
||||
PublicationSourceGroup,
|
||||
PublicationSourceRecord,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
|
||||
|
||||
repository = _published_repository()
|
||||
dates = [date(2026, 8, day) for day in [21, 24, 25, 26, 27, 28]]
|
||||
calendar = build_source_snapshot(
|
||||
api_name="trade_cal",
|
||||
params={},
|
||||
observed_at=NOW,
|
||||
target_trade_date=TARGET_DATE,
|
||||
rows=tuple({"exchange": "SSE", "cal_date": day.isoformat(), "is_open": 1} for day in dates),
|
||||
)
|
||||
repository.save_source_snapshots((calendar,))
|
||||
repository.save_publication_sources(
|
||||
(
|
||||
PublicationSourceRecord(
|
||||
"publication-success", PublicationSourceGroup.CALENDAR, 0, calendar
|
||||
),
|
||||
)
|
||||
)
|
||||
old = _running("past-publication", dates[0])
|
||||
repository.create_publication(old)
|
||||
past = tuple(
|
||||
replace(
|
||||
row.observation,
|
||||
trade_date=dates[0],
|
||||
value=Decimal(index),
|
||||
sector_code="BK9999.DC" if index == 5 else row.observation.sector_code,
|
||||
)
|
||||
for index, row in enumerate(_amount_rankings(), 1)
|
||||
)
|
||||
repository.save_rankings(
|
||||
RankingRecord(old.publication_id, row) for row in rank_metric_observations(past)
|
||||
)
|
||||
repository.finish_publication(_finish(old, PublicationStatus.SUCCESS))
|
||||
reader = ReadSectorRadar(repository)
|
||||
query = RadarQuery(
|
||||
view=RadarView.RANK_CHANGE,
|
||||
rank_change_metric=MetricKind.AMOUNT,
|
||||
@@ -170,12 +208,14 @@ def test_rank_change_uses_selected_metric_days_and_pool_sides() -> None:
|
||||
all_rows = reader.query(query)
|
||||
|
||||
assert top.total == 1
|
||||
assert top.rows[0].rank_change(5) == 5
|
||||
assert top.rows[0].rank_change(5) == 9
|
||||
assert bottom.total == 1
|
||||
assert bottom.rows[0].rank_change(5) == -4
|
||||
assert bottom.rows[0].rank_change(5) == -9
|
||||
assert all_rows.total == 10
|
||||
assert all_rows.rows[-1].observation.sector_code == "BK0005.DC"
|
||||
assert all_rows.rows[-1].rank_change(5) is None
|
||||
assert top.comparison_trade_date == dates[0]
|
||||
assert top.rank_change_values["BK0001.DC"][MetricKind.AMOUNT] == 9
|
||||
|
||||
|
||||
def test_latest_partial_attempt_is_visible_but_does_not_replace_last_good() -> None:
|
||||
|
||||
@@ -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