feat(selection): expose run sector aggregates and sector filter on results API

- sector_radar: add batch sector-count aggregation and sector member lookup
  over the strict last-good membership snapshot (postgres + in-memory fakes)
- selection: add SelectionSectorReader port, list_sector_counts use case,
  and sector_stock_codes filtering via run identity resolution; queries stay
  inside the selection context per ADR 0001
- http: add GET /api/v1/selection/sectors and forward sector param on
  /results and /runs/{run_id}
- fix stale positional args in pattern-scoring run tests; cover new behavior
  with read-service, application, and HTTP contract tests
This commit is contained in:
yuxuanhui
2026-09-05 19:48:30 +08:00
parent 41a9b4eb9a
commit 7e0f13d678
13 changed files with 1135 additions and 36 deletions
@@ -27,6 +27,7 @@ from zhixing_server.modules.sector_radar.domain.models import (
from zhixing_server.modules.sector_radar.domain.persistence import (
MembershipRecord,
RankingRecord,
SectorCountEntry,
)
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
from zhixing_server.modules.sector_radar.infrastructure.memory import (
@@ -315,3 +316,89 @@ def test_stock_membership_query_rejects_invalid_values() -> None:
StockSectorQuery(trade_date=TARGET_DATE, ts_code="x" * 13)
with pytest.raises(ValueError):
StockSectorQuery(trade_date=TARGET_DATE, ts_code="000001.SZ", concept_limit=0)
def test_sector_counts_aggregates_concepts_and_orders_by_count_then_code() -> None:
repository = _membership_repository()
repository.save_memberships(
(
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
_membership_record("000003.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0003.DC", "数字经济"),
)
)
reader = ReadSectorRadar(repository)
snapshot = reader.sector_counts(["000001.SZ", "000002.SZ", "000003.SZ"], TARGET_DATE)
assert snapshot.status == "success"
assert snapshot.trade_date == TARGET_DATE
assert snapshot.sector_type is SectorType.CONCEPT
assert snapshot.counts == (
SectorCountEntry(sector_code="BK0001.DC", sector_name="机器人", stock_count=3),
SectorCountEntry(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2),
SectorCountEntry(sector_code="BK0002.DC", sector_name="人工智能", stock_count=1),
)
def test_sector_counts_supports_industry_type() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_counts(
["000001.SZ", "000002.SZ"],
TARGET_DATE,
sector_type=SectorType.INDUSTRY,
)
assert snapshot.status == "success"
assert snapshot.sector_type is SectorType.INDUSTRY
assert snapshot.counts == (
SectorCountEntry(sector_code="BK0901.DC", sector_name="银行", stock_count=1),
SectorCountEntry(sector_code="BK0902.DC", sector_name="房地产", stock_count=1),
)
def test_sector_counts_without_publication_is_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
snapshot = reader.sector_counts(["000001.SZ"], TARGET_DATE)
assert snapshot.status == "no_data"
assert snapshot.trade_date is None
assert snapshot.counts == ()
def test_sector_counts_with_empty_stock_set_skips_publication_lookup() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_counts([], TARGET_DATE)
assert snapshot.status == "no_data"
assert snapshot.counts == ()
def test_sector_member_codes_returns_available_members() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
assert snapshot.status == "success"
assert snapshot.trade_date == TARGET_DATE
assert snapshot.stock_codes == ("000001.SZ",)
def test_sector_member_codes_without_publication_is_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
assert snapshot.status == "no_data"
assert snapshot.trade_date is None
assert snapshot.stock_codes == ()
def test_sector_member_codes_rejects_blank_sector_code() -> None:
reader = ReadSectorRadar(_membership_repository())
with pytest.raises(ValueError):
reader.sector_member_codes(TARGET_DATE, " ")
@@ -136,7 +136,13 @@ class FakeStore:
kwargs["error_message"] = error_message
self.finished = (run_id, status, kwargs)
def get_run(self, run_id: str, *, query: SelectionResultQuery | None = None):
def get_run(
self,
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
):
return None
def get_latest_run(
@@ -145,9 +151,16 @@ class FakeStore:
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
):
return None
def get_run_identity(self, run_id: str):
return None
def get_latest_run_identity(self, strategy: str, target_trade_date: date | None = None):
return None
class BatchStore(FakeStore):
def __init__(self) -> None:
@@ -427,8 +440,8 @@ def test_execute_loads_cases_once_and_scores_only_selected_stocks() -> None:
reader,
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
batch_size=1,
)
@@ -463,8 +476,8 @@ def test_execute_isolates_pattern_scoring_failure_from_selection_status() -> Non
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
)
@@ -497,8 +510,8 @@ def test_execute_skips_pattern_dependencies_when_feature_flag_is_disabled() -> N
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=False,
)
@@ -535,8 +548,8 @@ def test_execute_marks_scores_failed_when_case_library_is_unavailable() -> None:
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
)
@@ -0,0 +1,264 @@
"""Sector-filter behavior for selection result reads and sector aggregates."""
from collections.abc import Sequence
from dataclasses import replace
from datetime import date
from decimal import Decimal
from zhixing_server.modules.selection.application.run import (
RunZhixingB1,
SelectionSectorAggregates,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionExecutionSource,
SelectionResultQuery,
SelectionRun,
SelectionRunIdentity,
SelectionRunItem,
SelectionRunStatus,
SelectionSectorCount,
SelectionSectorMembership,
)
TARGET = date(2026, 8, 8)
class FakeReader:
def load_execution_source(
self, strategy: str, target_trade_date: date
) -> SelectionExecutionSource:
raise AssertionError("sector reads must not load execution sources")
def load_history(self, ts_code: str, target_trade_date: date):
raise AssertionError("sector reads must not load histories")
class FakeSectorReader:
def __init__(
self,
counts: tuple[SelectionSectorCount, ...] = (),
member_codes: tuple[str, ...] = (),
) -> None:
self.counts = counts
self.member_codes = member_codes
self.count_calls: list[tuple[tuple[str, ...], date, str]] = []
self.member_calls: list[tuple[date, str, str]] = []
def sector_counts(
self,
stock_codes: Sequence[str],
target_trade_date: date,
*,
sector_type: str = "concept",
) -> SelectionSectorMembership:
self.count_calls.append((tuple(stock_codes), target_trade_date, sector_type))
return SelectionSectorMembership(
snapshot_trade_date=target_trade_date,
sector_counts=self.counts,
)
def sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
*,
sector_type: str = "concept",
) -> tuple[str, ...]:
self.member_calls.append((target_trade_date, sector_code, sector_type))
return self.member_codes
class FakeStore:
def __init__(self, run: SelectionRun | None) -> None:
self.run = run
self.sector_codes_seen: Sequence[str] | None = None
self.queries: list[SelectionResultQuery] = []
def prepare_run(self, *args: object, **kwargs: object):
raise AssertionError("sector reads must not prepare runs")
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
raise AssertionError("sector reads must not record items")
def finish_run(self, *args: object, **kwargs: object) -> None:
raise AssertionError("sector reads must not finish runs")
def get_run(
self,
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
self.queries.append(query or SelectionResultQuery())
self.sector_codes_seen = sector_stock_codes
if self.run is not None and self.run.id == run_id:
return self.run
return None
def get_latest_run(
self,
strategy: str,
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
self.queries.append(query or SelectionResultQuery())
self.sector_codes_seen = sector_stock_codes
return self.run
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
if self.run is None or self.run.id != run_id:
return None
return SelectionRunIdentity(
run_id=self.run.id,
target_trade_date=self.run.target_trade_date,
)
def get_latest_run_identity(
self,
strategy: str,
target_trade_date: date | None = None,
) -> SelectionRunIdentity | None:
if self.run is None:
return None
return SelectionRunIdentity(
run_id=self.run.id,
target_trade_date=self.run.target_trade_date,
)
def _run(status: SelectionRunStatus = "success") -> SelectionRun:
return SelectionRun(
id="run-1",
strategy="zhixing_b1",
target_trade_date=TARGET,
market_sync_batch_id="market-run-1",
status=status,
target_count=3,
eligible_count=3,
evaluated_count=3,
selected_stock_count=2,
signal_count=2,
failed_count=0,
coverage=Decimal(1),
items=(
SelectionRunItem(ts_code="000001.SZ", name="A", status="selected", signal_count=1),
SelectionRunItem(ts_code="000002.SZ", name="B", status="selected", signal_count=1),
SelectionRunItem(ts_code="000003.SZ", name="C", status="no_signal", signal_count=0),
),
)
def test_list_sector_counts_aggregates_only_selected_stocks() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(
counts=(SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=2),)
)
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1")
assert isinstance(aggregates, SelectionSectorAggregates)
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "concept")]
assert aggregates.snapshot_trade_date == TARGET
assert aggregates.sector_type == "concept"
assert aggregates.sectors == sector_reader.counts
assert aggregates.run is store.run
def test_list_sector_counts_without_run_returns_none() -> None:
service = RunZhixingB1(FakeReader(), FakeStore(None), sector_reader=FakeSectorReader())
assert service.list_sector_counts("zhixing_b1") is None
def test_list_sector_counts_without_selected_stocks_skips_reader() -> None:
run = replace(_run(), items=())
sector_reader = FakeSectorReader()
service = RunZhixingB1(FakeReader(), FakeStore(run), sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1")
assert sector_reader.count_calls == []
assert aggregates is not None
assert aggregates.snapshot_trade_date is None
assert aggregates.sectors == ()
def test_get_latest_resolves_sector_filter_against_run_snapshot() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(member_codes=("000001.SZ", "000002.SZ"))
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
run = service.get_latest(
"zhixing_b1",
query=SelectionResultQuery(page=2, page_size=20, sector="BK0001.DC"),
)
assert run is store.run
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
assert store.sector_codes_seen == ("000001.SZ", "000002.SZ")
assert store.queries[-1].sector is None
assert store.queries[-1].page == 2
def test_get_latest_with_empty_sector_members_returns_empty_codes_page() -> None:
store = FakeStore(_run())
service = RunZhixingB1(
FakeReader(),
store,
sector_reader=FakeSectorReader(member_codes=()),
)
service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK9999.DC"))
assert store.sector_codes_seen == ()
def test_get_latest_without_sector_reader_ignores_the_filter() -> None:
store = FakeStore(_run())
service = RunZhixingB1(FakeReader(), store)
run = service.get_latest(
"zhixing_b1",
query=SelectionResultQuery(sector="BK0001.DC"),
)
assert run is store.run
assert store.sector_codes_seen is None
assert store.queries[-1].sector == "BK0001.DC"
def test_get_run_resolves_sector_filter_against_run_snapshot() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(member_codes=("000001.SZ",))
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
run = service.get_run("run-1", query=SelectionResultQuery(sector="BK0001.DC"))
assert run is store.run
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
assert store.sector_codes_seen == ("000001.SZ",)
def test_get_latest_without_run_and_sector_filter_returns_none() -> None:
service = RunZhixingB1(
FakeReader(),
FakeStore(None),
sector_reader=FakeSectorReader(member_codes=()),
)
assert service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK0001.DC")) is None
def test_list_sector_counts_accepts_industry_type() -> None:
sector_reader = FakeSectorReader()
service = RunZhixingB1(FakeReader(), FakeStore(_run()), sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1", sector_type="industry")
assert aggregates is not None
assert aggregates.sector_type == "industry"
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "industry")]