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:
@@ -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")]
|
||||
Reference in New Issue
Block a user