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
@@ -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",
@@ -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: ...
@@ -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."""
@@ -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."""
@@ -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."""
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."""
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",
]
@@ -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."""
@@ -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
@@ -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"]
@@ -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",
+106 -2
View File
@@ -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()
@@ -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")]