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:
@@ -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."""
|
||||
|
||||
+83
-8
@@ -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
|
||||
|
||||
|
||||
|
||||
+84
@@ -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",
|
||||
|
||||
@@ -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")]
|
||||
Reference in New Issue
Block a user