From 7e0f13d678e0ec1a70f02e9bf38b7bad4dd274c3 Mon Sep 17 00:00:00 2001 From: yuxuanhui Date: Sat, 5 Sep 2026 19:48:30 +0800 Subject: [PATCH] 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 --- .../modules/sector_radar/application/read.py | 102 ++++++- .../sector_radar/domain/persistence.py | 31 ++ .../sector_radar/infrastructure/memory.py | 61 +++- .../sector_radar/infrastructure/postgres.py | 55 ++++ .../modules/selection/application/run.py | 106 ++++++- .../modules/selection/domain/runs.py | 56 ++++ .../selection/infrastructure/postgres_runs.py | 91 +++++- .../infrastructure/sector_membership.py | 84 ++++++ .../modules/selection/presentation/http.py | 95 ++++++- zhixing-server/tests/test_selection_http.py | 108 ++++++- .../tests/unit/sector_radar/test_read.py | 87 ++++++ .../tests/unit/selection/test_run.py | 31 +- .../unit/selection/test_sector_filter.py | 264 ++++++++++++++++++ 13 files changed, 1135 insertions(+), 36 deletions(-) create mode 100644 zhixing-server/src/zhixing_server/modules/selection/infrastructure/sector_membership.py create mode 100644 zhixing-server/tests/unit/selection/test_sector_filter.py diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/application/read.py b/zhixing-server/src/zhixing_server/modules/sector_radar/application/read.py index 4c69780..c6ec4fe 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/application/read.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/application/read.py @@ -2,7 +2,7 @@ from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Iterator, Sequence from dataclasses import dataclass from datetime import date from enum import StrEnum @@ -21,7 +21,11 @@ from ..domain.models import ( RankSide, SectorType, ) -from ..domain.persistence import SectorRadarRepository, StockMembershipEntry +from ..domain.persistence import ( + SectorCountEntry, + SectorRadarRepository, + StockMembershipEntry, +) from ..domain.ranking import select_percentile_side, select_rank_change_side ReadStatus = Literal["success", "no_data"] @@ -148,6 +152,26 @@ class StockSectorMembership: concept_total: int +@dataclass(frozen=True, slots=True) +class SectorCountsSnapshot: + """Sector counts for a bounded stock set on one last-good snapshot date.""" + + status: ReadStatus + sector_type: SectorType + trade_date: date | None + counts: tuple[SectorCountEntry, ...] = () + + +@dataclass(frozen=True, slots=True) +class SectorMembersSnapshot: + """One sector's member stock codes on one last-good snapshot date.""" + + status: ReadStatus + sector_type: SectorType + trade_date: date | None + stock_codes: tuple[str, ...] = () + + _METRIC_DEFINITIONS = { MetricKind.AMOUNT: RadarMetricDefinition( metric_kind=MetricKind.AMOUNT, @@ -290,6 +314,78 @@ class ReadSectorRadar: concept_total=len(concepts), ) + def sector_counts( + self, + stock_codes: Sequence[str], + target_trade_date: date, + *, + sector_type: SectorType = SectorType.CONCEPT, + ) -> SectorCountsSnapshot: + """Count claimed stocks per sector from the strict last-good snapshot.""" + + if not stock_codes: + return SectorCountsSnapshot( + status="no_data", + sector_type=sector_type, + trade_date=None, + counts=(), + ) + publication = self.repository.get_last_good_publication(target_trade_date) + if publication is None: + return SectorCountsSnapshot( + status="no_data", + sector_type=sector_type, + trade_date=None, + counts=(), + ) + counts = tuple( + self.repository.load_sector_counts( + publication.target_trade_date, + stock_codes, + sector_type, + ) + ) + return SectorCountsSnapshot( + status="success", + sector_type=sector_type, + trade_date=publication.target_trade_date, + counts=counts, + ) + + def sector_member_codes( + self, + target_trade_date: date, + sector_code: str, + *, + sector_type: SectorType = SectorType.CONCEPT, + ) -> SectorMembersSnapshot: + """Return one sector's members from the strict last-good snapshot.""" + + normalized_code = sector_code.strip() + if not normalized_code: + raise ValueError("sector_code must not be empty") + publication = self.repository.get_last_good_publication(target_trade_date) + if publication is None: + return SectorMembersSnapshot( + status="no_data", + sector_type=sector_type, + trade_date=None, + stock_codes=(), + ) + codes = tuple( + self.repository.load_sector_member_codes( + publication.target_trade_date, + normalized_code, + sector_type, + ) + ) + return SectorMembersSnapshot( + status="success", + sector_type=sector_type, + trade_date=publication.target_trade_date, + stock_codes=codes, + ) + def _sector_refs(entries: Iterator[StockMembershipEntry]) -> tuple[SectorRef, ...]: """Project membership entries into ordered public sector references.""" @@ -309,6 +405,8 @@ __all__ = [ "RadarView", "RankingPage", "ReadSectorRadar", + "SectorCountsSnapshot", + "SectorMembersSnapshot", "SectorRef", "StockSectorMembership", "StockSectorQuery", diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py index df20cb6..831df77 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py @@ -81,6 +81,23 @@ class StockMembershipEntry: raise ValueError("membership sector identity fields must not be empty") +@dataclass(frozen=True, slots=True) +class SectorCountEntry: + """One sector and the number of stocks it claimed on one snapshot date.""" + + sector_code: str + sector_name: str + stock_count: int + + def __post_init__(self) -> None: + """Reject empty sector identity and non-positive counts.""" + + if not self.sector_code.strip() or not self.sector_name.strip(): + raise ValueError("sector count identity fields must not be empty") + if self.stock_count < 1: + raise ValueError("sector stock_count must be positive") + + @dataclass(frozen=True, slots=True) class StockFactRecord: """One normalized stock fact revision with all contributing raw snapshots.""" @@ -217,6 +234,20 @@ class SectorRadarRepository(Protocol): self, trade_date: date, stock_code: str ) -> Sequence[StockMembershipEntry]: ... + def load_sector_counts( + self, + trade_date: date, + stock_codes: Sequence[str], + sector_type: SectorType, + ) -> Sequence[SectorCountEntry]: ... + + def load_sector_member_codes( + self, + trade_date: date, + sector_code: str, + sector_type: SectorType, + ) -> Sequence[str]: ... + def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ... def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ... diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py index 1d5ee0d..10abaa1 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py @@ -7,13 +7,20 @@ from contextlib import contextmanager from dataclasses import replace from datetime import date, datetime -from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate +from ..domain.models import ( + PublicationStatus, + RadarPublication, + RankedMetric, + SectorDailyAggregate, + SectorType, +) from ..domain.persistence import ( DailyAggregateRecord, MembershipRecord, PublicationSourceGroup, PublicationSourceRecord, RankingRecord, + SectorCountEntry, StockFactRecord, StockMembershipEntry, WriteCounts, @@ -137,6 +144,58 @@ class InMemorySectorRadarRepository: and record.status.value == "available" ) + def load_sector_counts( + self, + trade_date: date, + stock_codes: Sequence[str], + sector_type: SectorType, + ) -> Sequence[SectorCountEntry]: + """Aggregate claimed-stock counts per sector on one snapshot date.""" + + wanted = set(stock_codes) + counts: dict[tuple[str, str], int] = {} + for record in self.memberships.values(): + if ( + record.trade_date == trade_date + and record.sector_type is sector_type + and record.status.value == "available" + and record.stock_code in wanted + ): + key = (record.sector_code, record.sector_name) + counts[key] = counts.get(key, 0) + 1 + return tuple( + SectorCountEntry( + sector_code=sector_code, + sector_name=sector_name, + stock_count=stock_count, + ) + for (sector_code, sector_name), stock_count in sorted( + counts.items(), + key=lambda item: (-item[1], item[0][0]), + ) + ) + + def load_sector_member_codes( + self, + trade_date: date, + sector_code: str, + sector_type: SectorType, + ) -> Sequence[str]: + """Return available member stock codes of one sector on one date.""" + + return tuple( + record.stock_code + for record in sorted( + self.memberships.values(), + key=lambda item: item.stock_code or "", + ) + if record.trade_date == trade_date + and record.sector_type is sector_type + and record.sector_code == sector_code + and record.status.value == "available" + and record.stock_code is not None + ) + def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: """Insert normalized fact revisions idempotently.""" diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py index 6508fd1..7b949cb 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py @@ -30,6 +30,7 @@ from ..domain.persistence import ( PublicationSourceGroup, PublicationSourceRecord, RankingRecord, + SectorCountEntry, StockFactRecord, StockMembershipEntry, WriteCounts, @@ -310,6 +311,60 @@ class PostgresSectorRadarRepository: for row in rows ) + def load_sector_counts( + self, + trade_date: date, + stock_codes: Sequence[str], + sector_type: SectorType, + ) -> Sequence[SectorCountEntry]: + """Aggregate how many of the given stocks each sector claimed on one date.""" + + if not stock_codes: + return () + with self._connection() as connection: + rows = connection.execute( + """ + SELECT sector_code, sector_name, COUNT(*) AS stock_count + FROM sector_radar_membership + WHERE trade_date = %s AND sector_type = %s + AND membership_status = 'available' + AND stock_code = ANY(%s) + GROUP BY sector_code, sector_name + ORDER BY stock_count DESC, sector_code + """, + (trade_date, sector_type.value, list(stock_codes)), + ).fetchall() + return tuple( + SectorCountEntry( + sector_code=row[0], + sector_name=row[1], + stock_count=int(row[2]), + ) + for row in rows + ) + + def load_sector_member_codes( + self, + trade_date: date, + sector_code: str, + sector_type: SectorType, + ) -> Sequence[str]: + """Load every stock one sector claimed on exactly one snapshot date.""" + + with self._connection() as connection: + rows = connection.execute( + """ + SELECT stock_code + FROM sector_radar_membership + WHERE trade_date = %s AND sector_type = %s + AND sector_code = %s + AND membership_status = 'available' + ORDER BY stock_code + """, + (trade_date, sector_type.value, sector_code), + ).fetchall() + return tuple(str(row[0]) for row in rows) + def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: """COPY normalized stock facts while preserving contributing source ids.""" diff --git a/zhixing-server/src/zhixing_server/modules/selection/application/run.py b/zhixing-server/src/zhixing_server/modules/selection/application/run.py index 6add5b3..76c50a0 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/application/run.py +++ b/zhixing-server/src/zhixing_server/modules/selection/application/run.py @@ -7,7 +7,7 @@ import time from collections import Counter from collections.abc import Callable, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass +from dataclasses import dataclass, replace from datetime import date from typing import Protocol, cast @@ -29,6 +29,9 @@ from ..domain.runs import ( SelectionRunItem, SelectionRunStatus, SelectionRunStore, + SelectionSectorCount, + SelectionSectorMembership, + SelectionSectorReader, SelectionStock, SelectionUniverseReader, ) @@ -60,6 +63,16 @@ class PreparedSelectionRun: source: SelectionExecutionSource +@dataclass(frozen=True, slots=True) +class SelectionSectorAggregates: + """A run summary plus its selected stocks' sector membership counts.""" + + run: SelectionRun | None + snapshot_trade_date: date | None + sector_type: str + sectors: tuple[SelectionSectorCount, ...] + + class RunZhixingB1: """Prepare, execute, and query persisted selection strategy batches. @@ -75,6 +88,7 @@ class RunZhixingB1: evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None, pattern_case_loader: PatternCaseLibraryLoader | None = None, pattern_scorer: PatternScorer | None = None, + sector_reader: SelectionSectorReader | None = None, *, pattern_scoring_enabled: bool = False, max_workers: int = 4, @@ -96,6 +110,7 @@ class RunZhixingB1: self.evaluators.update(evaluators) self.pattern_case_loader = pattern_case_loader self.pattern_scorer = pattern_scorer + self.sector_reader = sector_reader self.pattern_scoring_enabled = pattern_scoring_enabled self.max_workers = max_workers self.batch_size = batch_size @@ -211,8 +226,7 @@ class RunZhixingB1: ) history_rows += batch_history_rows batch_missing_turnover = sum( - not _turnover_present(history, target_trade_date) - for history in histories + not _turnover_present(history, target_trade_date) for history in histories ) logger.info( "selection_read_batch_summary strategy=%s target_trade_date=%s " @@ -581,9 +595,22 @@ class RunZhixingB1: *, query: SelectionResultQuery | None = None, ) -> SelectionRun | None: - """Read one persisted run for polling.""" + """Read one persisted run for polling with optional sector filtering.""" - return self.store.get_run(run_id, query=query) + effective = query or SelectionResultQuery() + if not effective.sector: + return self.store.get_run(run_id, query=query) + identity = self.store.get_run_identity(run_id) + if identity is None: + return None + member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector) + if member_codes is None: + return self.store.get_run(run_id, query=query) + return self.store.get_run( + run_id, + query=replace(effective, sector=None), + sector_stock_codes=member_codes, + ) def get_latest( self, @@ -594,7 +621,73 @@ class RunZhixingB1: ) -> SelectionRun | None: """Read the current result by date or the latest result for a strategy.""" - return self.store.get_latest_run(strategy, target_trade_date, query=query) + effective = query or SelectionResultQuery() + if not effective.sector: + return self.store.get_latest_run(strategy, target_trade_date, query=query) + identity = self.store.get_latest_run_identity(strategy, target_trade_date) + if identity is None: + return None + member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector) + if member_codes is None: + return self.store.get_latest_run(strategy, target_trade_date, query=query) + return self.store.get_latest_run( + strategy, + identity.target_trade_date, + query=replace(effective, sector=None), + sector_stock_codes=member_codes, + ) + + def _sector_member_codes( + self, + target_trade_date: date, + sector_code: str, + ) -> tuple[str, ...] | None: + """Resolve one sector's members, or None when the port is absent.""" + + if self.sector_reader is None: + return None + return self.sector_reader.sector_member_codes(target_trade_date, sector_code) + + def list_sector_counts( + self, + strategy: StrategyName, + target_trade_date: date | None = None, + *, + sector_type: str = "concept", + ) -> SelectionSectorAggregates | None: + """Aggregate the current run's selected stocks by point-in-time sector.""" + + run = self.store.get_latest_run(strategy, target_trade_date) + if run is None: + return None + selected_codes = [ + item.ts_code + for item in run.items + if item.status == "selected" and item.signal_count > 0 + ] + membership = self._sector_membership(selected_codes, run.target_trade_date, sector_type) + return SelectionSectorAggregates( + run=run, + snapshot_trade_date=membership.snapshot_trade_date, + sector_type=sector_type, + sectors=membership.sector_counts, + ) + + def _sector_membership( + self, + stock_codes: Sequence[str], + target_trade_date: date, + sector_type: str, + ) -> SelectionSectorMembership: + """Read sector counts for a stock set, tolerating a missing port.""" + + if self.sector_reader is None or not stock_codes: + return SelectionSectorMembership(snapshot_trade_date=None, sector_counts=()) + return self.sector_reader.sector_counts( + stock_codes, + target_trade_date, + sector_type=sector_type, + ) def _to_item( @@ -676,4 +769,5 @@ __all__ = [ "RunZhixingB1", "SelectionRerunRequired", "SelectionRunInProgress", + "SelectionSectorAggregates", ] diff --git a/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py b/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py index 1102910..d9878d2 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py +++ b/zhixing-server/src/zhixing_server/modules/selection/domain/runs.py @@ -31,6 +31,52 @@ class SelectionResultQuery: search: str | None = None category: SelectionSignalCategoryFilter | None = None sort: SelectionResultSort = "code" + sector: str | None = None + + +@dataclass(frozen=True, slots=True) +class SelectionSectorCount: + """One sector and the number of this run's selected stocks it contains.""" + + sector_code: str + sector_name: str + stock_count: int + + +@dataclass(frozen=True, slots=True) +class SelectionSectorMembership: + """Point-in-time sector membership aggregates for a set of stock codes.""" + + snapshot_trade_date: date | None + sector_counts: tuple[SelectionSectorCount, ...] = () + + +@dataclass(frozen=True, slots=True) +class SelectionRunIdentity: + """Minimal run locator used to resolve date-dependent filters.""" + + run_id: str + target_trade_date: date + + +class SelectionSectorReader(Protocol): + """Port to the sector-radar context's point-in-time membership reads.""" + + def sector_counts( + self, + stock_codes: Sequence[str], + target_trade_date: date, + *, + sector_type: str = "concept", + ) -> SelectionSectorMembership: ... + + def sector_member_codes( + self, + target_trade_date: date, + sector_code: str, + *, + sector_type: str = "concept", + ) -> tuple[str, ...]: ... @dataclass(frozen=True, slots=True) @@ -139,6 +185,7 @@ class SelectionRunStore(Protocol): run_id: str, *, query: SelectionResultQuery | None = None, + sector_stock_codes: Sequence[str] | None = None, ) -> SelectionRun | None: ... def get_latest_run( @@ -147,8 +194,17 @@ class SelectionRunStore(Protocol): target_trade_date: date | None = None, *, query: SelectionResultQuery | None = None, + sector_stock_codes: Sequence[str] | None = None, ) -> SelectionRun | None: ... + def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None: ... + + def get_latest_run_identity( + self, + strategy: SelectionStrategyName, + target_trade_date: date | None = None, + ) -> SelectionRunIdentity | None: ... + class SelectionUniverseReader(Protocol): """Read a qualified market-data source snapshot for one strategy run.""" diff --git a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py index 6496933..c099aaa 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py +++ b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py @@ -33,6 +33,7 @@ from ..domain.runs import ( SelectionResultQuery, SelectionRun, SelectionRunError, + SelectionRunIdentity, SelectionRunInProgress, SelectionRunItem, SelectionRunStatus, @@ -46,9 +47,7 @@ _SELECTION_SIGNAL_ORDER: tuple[SelectionSignalCategory, ...] = ( *ZHIXING_B1_SIGNAL_ORDER, *GOLD_BRICK_SIGNAL_ORDER, ) -_SIGNAL_PRIORITY = { - category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER) -} +_SIGNAL_PRIORITY = {category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)} _CATEGORY_PREFIXES = { "pullback": "zhixing_b1_pullback_", "oversold": "zhixing_b1_oversold_", @@ -335,21 +334,41 @@ class PostgresSelectionRunRepository(SelectionRunStore): run_id: str, *, query: SelectionResultQuery | None = None, + sector_stock_codes: Sequence[str] | None = None, ) -> SelectionRun | None: """Read one run with filtered, stock-paged signals and item failures.""" try: with self._connection() as connection: - return self._load_run(connection, run_id, query or SelectionResultQuery()) + return self._load_run( + connection, + run_id, + query or SelectionResultQuery(), + sector_stock_codes=sector_stock_codes, + ) except psycopg.Error as exc: raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc + def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None: + """Read only a run's locator so date-dependent filters resolve first.""" + + try: + with self._connection() as connection: + row = connection.execute( + "SELECT id, target_trade_date FROM selection_run WHERE id = %s", + (run_id,), + ).fetchone() + except psycopg.Error as exc: + raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc + return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None + def get_latest_run( self, strategy: SelectionStrategyName, target_trade_date: date | None = None, *, query: SelectionResultQuery | None = None, + sector_stock_codes: Sequence[str] | None = None, ) -> SelectionRun | None: """Read the current run for a date or the latest date for a strategy.""" @@ -377,18 +396,59 @@ class PostgresSelectionRunRepository(SelectionRunStore): (strategy, target_trade_date), ).fetchone() return ( - self._load_run(connection, str(row[0]), query or SelectionResultQuery()) + self._load_run( + connection, + str(row[0]), + query or SelectionResultQuery(), + sector_stock_codes=sector_stock_codes, + ) if row else None ) except psycopg.Error as exc: raise SelectionRunStoreError("failed to load latest selection run") from exc + def get_latest_run_identity( + self, + strategy: SelectionStrategyName, + target_trade_date: date | None = None, + ) -> SelectionRunIdentity | None: + """Read only the current run's locator so sector filters resolve first.""" + + try: + with self._connection() as connection: + if target_trade_date is None: + row = connection.execute( + """ + SELECT id, target_trade_date + FROM selection_run + WHERE strategy = %s + ORDER BY target_trade_date DESC, created_at DESC, id DESC + LIMIT 1 + """, + (strategy,), + ).fetchone() + else: + row = connection.execute( + """ + SELECT id, target_trade_date + FROM selection_run + WHERE strategy = %s AND target_trade_date = %s + LIMIT 1 + """, + (strategy, target_trade_date), + ).fetchone() + except psycopg.Error as exc: + raise SelectionRunStoreError("failed to load latest selection run") from exc + return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None + @staticmethod def _load_run( connection: Any, run_id: str, query: SelectionResultQuery, + *, + sector_stock_codes: Sequence[str] | None = None, ) -> SelectionRun | None: row = connection.execute( """ @@ -417,7 +477,9 @@ class PostgresSelectionRunRepository(SelectionRunStore): """, (run_id,), ).fetchall() - stock_filter, stock_parameters = _stock_filter(query, run_id) + stock_filter, stock_parameters = _stock_filter( + query, run_id, sector_stock_codes=sector_stock_codes + ) stock_total_row = connection.execute( f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}", tuple(stock_parameters), @@ -568,12 +630,19 @@ def _signal_category(value: str) -> SelectionSignalCategory: return GoldBrickCategory(value) -def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]: +def _stock_filter( + query: SelectionResultQuery, + run_id: str, + *, + sector_stock_codes: Sequence[str] | None = None, +) -> tuple[str, list[object]]: """Build the signal predicate used to select distinct matching stocks. A category narrows which stocks qualify for the page. Once a stock qualifies, the repository loads every signal for that stock so callers - can present all independently persisted categories together. + can present all independently persisted categories together. Resolved + sector membership codes arrive from the sector-radar port, so the SQL + stays inside the selection context. """ clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"] @@ -592,6 +661,12 @@ def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[o ")" ) parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%") + if sector_stock_codes is not None: + if not sector_stock_codes: + clauses.append("FALSE") + else: + clauses.append("item.ts_code = ANY(%s)") + parameters.append(list(sector_stock_codes)) return " AND ".join(clauses), parameters diff --git a/zhixing-server/src/zhixing_server/modules/selection/infrastructure/sector_membership.py b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/sector_membership.py new file mode 100644 index 0000000..86e31da --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/selection/infrastructure/sector_membership.py @@ -0,0 +1,84 @@ +"""Bridge the selection sector-filter port to the sector-radar read model. + +The selection context owns no sector-membership storage. This adapter keeps +the port contract local to selection while delegating point-in-time reads to +the sector-radar application service in the composition root. +""" + +from __future__ import annotations + +from collections.abc import Sequence +from datetime import date + +from zhixing_server.modules.sector_radar.application.read import ReadSectorRadar +from zhixing_server.modules.sector_radar.domain.models import SectorType +from zhixing_server.modules.selection.domain.runs import ( + SelectionSectorCount, + SelectionSectorMembership, +) + +SELECTION_SECTOR_TYPES: dict[str, SectorType] = { + "concept": SectorType.CONCEPT, + "industry": SectorType.INDUSTRY, +} + + +def _sector_type(value: str) -> SectorType: + """Map the public sector-type vocabulary onto the radar domain enum.""" + + try: + return SELECTION_SECTOR_TYPES[value] + except KeyError: + raise ValueError(f"unsupported sector type: {value}") from None + + +class SectorRadarSelectionReader: + """Resolve selection sector aggregates through the sector-radar context.""" + + def __init__(self, reader: ReadSectorRadar) -> None: + self._reader = reader + + def sector_counts( + self, + stock_codes: Sequence[str], + target_trade_date: date, + *, + sector_type: str = "concept", + ) -> SelectionSectorMembership: + """Return per-sector stock counts for one run's selected stocks.""" + + snapshot = self._reader.sector_counts( + stock_codes, + target_trade_date, + sector_type=_sector_type(sector_type), + ) + return SelectionSectorMembership( + snapshot_trade_date=snapshot.trade_date, + sector_counts=tuple( + SelectionSectorCount( + sector_code=entry.sector_code, + sector_name=entry.sector_name, + stock_count=entry.stock_count, + ) + for entry in snapshot.counts + ), + ) + + def sector_member_codes( + self, + target_trade_date: date, + sector_code: str, + *, + sector_type: str = "concept", + ) -> tuple[str, ...]: + """Return one sector's member stock codes on the aligned snapshot.""" + + snapshot = self._reader.sector_member_codes( + target_trade_date, + sector_code, + sector_type=_sector_type(sector_type), + ) + return snapshot.stock_codes + + +__all__ = ["SELECTION_SECTOR_TYPES", "SectorRadarSelectionReader"] diff --git a/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py b/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py index 2c8c7ad..9409bb1 100644 --- a/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py +++ b/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py @@ -9,6 +9,7 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, s from pydantic import BaseModel, Field from zhixing_server.bootstrap.config import Settings, get_settings +from zhixing_server.modules.sector_radar.presentation.http import get_sector_radar_reader from zhixing_server.modules.selection.application.chart import ( GetSelectionChart, SelectionChart, @@ -46,6 +47,9 @@ from zhixing_server.modules.selection.infrastructure.postgres_reader import ( from zhixing_server.modules.selection.infrastructure.postgres_runs import ( PostgresSelectionRunRepository, ) +from zhixing_server.modules.selection.infrastructure.sector_membership import ( + SectorRadarSelectionReader, +) selection_router = APIRouter() _SELECTION_POOL_CACHE_LOCK = threading.Lock() @@ -212,6 +216,32 @@ class SelectionResultsResponse(BaseModel): stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks) +class SelectionSectorCountResponse(BaseModel): + """One sector and the number of this run's selected stocks it contains.""" + + sector_code: str + sector_name: str + stock_count: int = Field(ge=0) + + +def _empty_sectors() -> list[SelectionSectorCountResponse]: + """Create a typed default sector-aggregate list.""" + + return [] + + +class SelectionSectorsResponse(BaseModel): + """The current run's selected stocks aggregated by point-in-time sector.""" + + strategy: StrategyValue + target_trade_date: date | None + run_id: str | None + status: SelectionStatusValue + snapshot_trade_date: date | None + sector_type: Literal["concept", "industry"] + sectors: list[SelectionSectorCountResponse] = Field(default_factory=_empty_sectors) + + def get_selection_service( settings: Annotated[Settings, Depends(get_settings)], ) -> RunZhixingB1: @@ -221,12 +251,14 @@ def get_selection_service( reader = PostgresMarketDataReader(settings, pool=pool) pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool) store = PostgresSelectionRunRepository(settings.database_url, pool=pool) + sector_reader = SectorRadarSelectionReader(get_sector_radar_reader(settings)) return RunZhixingB1( reader, store, evaluators={"gold_brick": EvaluateGoldBrick(reader)}, pattern_case_loader=pattern_case_loader, pattern_scorer=ZhixingB1PatternScorer(), + sector_reader=sector_reader, pattern_scoring_enabled=settings.selection_pattern_scoring_enabled, max_workers=settings.selection_max_workers, batch_size=settings.selection_batch_size, @@ -338,15 +370,14 @@ def get_selection_run( page: Annotated[int, Query(ge=1)] = 1, page_size: Annotated[int, Query(ge=1, le=100)] = 10, search: Annotated[str | None, Query(max_length=100)] = None, - category: Literal[ - "pullback", "oversold", "original", "resonance" - ] | None = None, + category: Literal["pullback", "oversold", "original", "resonance"] | None = None, sort: Literal["code", "score_desc", "score_asc"] = "code", + sector: Annotated[str | None, Query(max_length=60)] = None, ) -> SelectionResultsResponse: """Return one run for asynchronous polling.""" try: - query = _result_query(page, page_size, search, category, sort) + query = _result_query(page, page_size, search, category, sort, sector) run = service.get_run(run_id, query=query) except SelectionRunStoreError as exc: raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc @@ -355,6 +386,51 @@ def get_selection_run( return _run_response(run, query=query) +@selection_router.get("/sectors", response_model=SelectionSectorsResponse) +def get_selection_sectors( + service: Annotated[RunZhixingB1, Depends(get_selection_service)], + strategy: StrategyValue = "zhixing_b1", + target_trade_date: date | None = None, + sector_type: Literal["concept", "industry"] = "concept", +) -> SelectionSectorsResponse: + """Aggregate the current run's selected stocks by point-in-time sector.""" + + try: + aggregates = service.list_sector_counts( + strategy, + target_trade_date, + sector_type=sector_type, + ) + except SelectionRunStoreError as exc: + raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc + if aggregates is None or aggregates.run is None: + return SelectionSectorsResponse( + strategy=strategy, + target_trade_date=target_trade_date, + run_id=None, + status="no_data", + snapshot_trade_date=None, + sector_type=sector_type, + ) + run = aggregates.run + return SelectionSectorsResponse( + strategy=run.strategy, + target_trade_date=run.target_trade_date, + run_id=run.id, + status=run.status, + snapshot_trade_date=aggregates.snapshot_trade_date, + sector_type=sector_type, + sectors=[ + SelectionSectorCountResponse( + sector_code=count.sector_code, + sector_name=count.sector_name, + stock_count=count.stock_count, + ) + for count in aggregates.sectors + ], + ) + + @selection_router.get("/results", response_model=SelectionResultsResponse) def get_selection_results( service: Annotated[RunZhixingB1, Depends(get_selection_service)], @@ -363,15 +439,14 @@ def get_selection_results( page: Annotated[int, Query(ge=1)] = 1, page_size: Annotated[int, Query(ge=1, le=100)] = 10, search: Annotated[str | None, Query(max_length=100)] = None, - category: Literal[ - "pullback", "oversold", "original", "resonance" - ] | None = None, + category: Literal["pullback", "oversold", "original", "resonance"] | None = None, sort: Literal["code", "score_desc", "score_asc"] = "code", + sector: Annotated[str | None, Query(max_length=60)] = None, ) -> SelectionResultsResponse: """Return the current persisted result for a strategy and optional date.""" try: - query = _result_query(page, page_size, search, category, sort) + query = _result_query(page, page_size, search, category, sort, sector) run = service.get_latest(strategy, target_trade_date, query=query) except SelectionRunStoreError as exc: raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc @@ -524,16 +599,19 @@ def _result_query( search: str | None, category: Literal["pullback", "oversold", "original", "resonance"] | None, sort: Literal["code", "score_desc", "score_asc"], + sector: str | None = None, ) -> SelectionResultQuery: """Normalize HTTP query values before handing them to the selection port.""" normalized_search = search.strip() if search else None + normalized_sector = sector.strip() if sector else None return SelectionResultQuery( page=page, page_size=page_size, search=normalized_search or None, category=category, sort=sort, + sector=normalized_sector or None, ) @@ -551,6 +629,7 @@ __all__ = [ "SelectionResultsResponse", "SelectionRunAcceptedResponse", "SelectionRunRequest", + "SelectionSectorsResponse", "SelectionStockResponse", "get_selection_chart_service", "get_selection_service", diff --git a/zhixing-server/tests/test_selection_http.py b/zhixing-server/tests/test_selection_http.py index 60ab841..2294d29 100644 --- a/zhixing-server/tests/test_selection_http.py +++ b/zhixing-server/tests/test_selection_http.py @@ -15,7 +15,10 @@ from zhixing_server.modules.selection.application.chart import ( SelectionChartNotFound, SelectionChartPoint, ) -from zhixing_server.modules.selection.application.run import PreparedSelectionRun +from zhixing_server.modules.selection.application.run import ( + PreparedSelectionRun, + SelectionSectorAggregates, +) from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category from zhixing_server.modules.selection.domain.pattern_scoring import ( PATTERN_SCORING_VERSION, @@ -30,6 +33,7 @@ from zhixing_server.modules.selection.domain.runs import ( SelectionResultQuery, SelectionRun, SelectionRunInProgress, + SelectionSectorCount, SelectionStock, ) from zhixing_server.modules.selection.infrastructure.postgres_reader import ( @@ -45,11 +49,17 @@ TARGET = date(2026, 8, 8) class FakeSelectionService: - def __init__(self, run: SelectionRun | None = None) -> None: + def __init__( + self, + run: SelectionRun | None = None, + sectors: tuple[SelectionSectorCount, ...] = (), + ) -> None: self.run = run + self.sectors = sectors self.executed = False self.mode = "ok" self.last_query: SelectionResultQuery | None = None + self.sector_type_requested: str | None = None def prepare( self, @@ -103,6 +113,23 @@ class FakeSelectionService: return None return self.run + def list_sector_counts( + self, + strategy: str, + target_trade_date: date | None = None, + *, + sector_type: str = "concept", + ) -> SelectionSectorAggregates | None: + self.sector_type_requested = sector_type + if self.run is None: + return None + return SelectionSectorAggregates( + run=self.run, + snapshot_trade_date=self.run.target_trade_date, + sector_type=sector_type, + sectors=self.sectors, + ) + class FakeChartService: """Return or fail one deterministic chart response.""" @@ -434,6 +461,83 @@ def test_query_rejects_invalid_page_size() -> None: assert response.status_code == 422 +def test_query_forwards_sector_filter() -> None: + service = FakeSelectionService(_run("run-http", "success")) + + response = _client(service).get( + "/api/v1/selection/results", + params={"strategy": "zhixing_b1", "sector": " BK0475.DC "}, + ) + + assert response.status_code == 200 + assert service.last_query is not None + assert service.last_query.sector == "BK0475.DC" + + +def test_sectors_returns_aggregated_counts_desc() -> None: + run = _run("run-http", "success") + service = FakeSelectionService( + run, + sectors=( + SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=3), + SelectionSectorCount(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2), + ), + ) + + response = _client(service).get("/api/v1/selection/sectors") + + assert response.status_code == 200 + assert response.json() == { + "strategy": "zhixing_b1", + "target_trade_date": "2026-08-08", + "run_id": "run-http", + "status": "success", + "snapshot_trade_date": "2026-08-08", + "sector_type": "concept", + "sectors": [ + {"sector_code": "BK0001.DC", "sector_name": "机器人", "stock_count": 3}, + {"sector_code": "BK0003.DC", "sector_name": "数字经济", "stock_count": 2}, + ], + } + + +def test_sectors_forwards_sector_type() -> None: + service = FakeSelectionService(_run("run-http", "success")) + + response = _client(service).get( + "/api/v1/selection/sectors", + params={"sector_type": "industry"}, + ) + + assert response.status_code == 200 + assert service.sector_type_requested == "industry" + assert response.json()["sector_type"] == "industry" + + +def test_sectors_without_run_returns_no_data() -> None: + response = _client(FakeSelectionService()).get("/api/v1/selection/sectors") + + assert response.status_code == 200 + assert response.json() == { + "strategy": "zhixing_b1", + "target_trade_date": None, + "run_id": None, + "status": "no_data", + "snapshot_trade_date": None, + "sector_type": "concept", + "sectors": [], + } + + +def test_sectors_rejects_unknown_sector_type() -> None: + response = _client(FakeSelectionService()).get( + "/api/v1/selection/sectors", + params={"sector_type": "macro"}, + ) + + assert response.status_code == 422 + + def test_chart_returns_bounded_qfq_contract() -> None: chart_service = FakeChartService() diff --git a/zhixing-server/tests/unit/sector_radar/test_read.py b/zhixing-server/tests/unit/sector_radar/test_read.py index c789afc..cdd7fd9 100644 --- a/zhixing-server/tests/unit/sector_radar/test_read.py +++ b/zhixing-server/tests/unit/sector_radar/test_read.py @@ -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, " ") diff --git a/zhixing-server/tests/unit/selection/test_run.py b/zhixing-server/tests/unit/selection/test_run.py index fdf0c90..c946648 100644 --- a/zhixing-server/tests/unit/selection/test_run.py +++ b/zhixing-server/tests/unit/selection/test_run.py @@ -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, ) diff --git a/zhixing-server/tests/unit/selection/test_sector_filter.py b/zhixing-server/tests/unit/selection/test_sector_filter.py new file mode 100644 index 0000000..7dccd33 --- /dev/null +++ b/zhixing-server/tests/unit/selection/test_sector_filter.py @@ -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")]