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 __future__ import annotations
|
||||||
|
|
||||||
from collections.abc import Iterator
|
from collections.abc import Iterator, Sequence
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
@@ -21,7 +21,11 @@ from ..domain.models import (
|
|||||||
RankSide,
|
RankSide,
|
||||||
SectorType,
|
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
|
from ..domain.ranking import select_percentile_side, select_rank_change_side
|
||||||
|
|
||||||
ReadStatus = Literal["success", "no_data"]
|
ReadStatus = Literal["success", "no_data"]
|
||||||
@@ -148,6 +152,26 @@ class StockSectorMembership:
|
|||||||
concept_total: int
|
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 = {
|
_METRIC_DEFINITIONS = {
|
||||||
MetricKind.AMOUNT: RadarMetricDefinition(
|
MetricKind.AMOUNT: RadarMetricDefinition(
|
||||||
metric_kind=MetricKind.AMOUNT,
|
metric_kind=MetricKind.AMOUNT,
|
||||||
@@ -290,6 +314,78 @@ class ReadSectorRadar:
|
|||||||
concept_total=len(concepts),
|
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, ...]:
|
def _sector_refs(entries: Iterator[StockMembershipEntry]) -> tuple[SectorRef, ...]:
|
||||||
"""Project membership entries into ordered public sector references."""
|
"""Project membership entries into ordered public sector references."""
|
||||||
@@ -309,6 +405,8 @@ __all__ = [
|
|||||||
"RadarView",
|
"RadarView",
|
||||||
"RankingPage",
|
"RankingPage",
|
||||||
"ReadSectorRadar",
|
"ReadSectorRadar",
|
||||||
|
"SectorCountsSnapshot",
|
||||||
|
"SectorMembersSnapshot",
|
||||||
"SectorRef",
|
"SectorRef",
|
||||||
"StockSectorMembership",
|
"StockSectorMembership",
|
||||||
"StockSectorQuery",
|
"StockSectorQuery",
|
||||||
|
|||||||
@@ -81,6 +81,23 @@ class StockMembershipEntry:
|
|||||||
raise ValueError("membership sector identity fields must not be empty")
|
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)
|
@dataclass(frozen=True, slots=True)
|
||||||
class StockFactRecord:
|
class StockFactRecord:
|
||||||
"""One normalized stock fact revision with all contributing raw snapshots."""
|
"""One normalized stock fact revision with all contributing raw snapshots."""
|
||||||
@@ -217,6 +234,20 @@ class SectorRadarRepository(Protocol):
|
|||||||
self, trade_date: date, stock_code: str
|
self, trade_date: date, stock_code: str
|
||||||
) -> Sequence[StockMembershipEntry]: ...
|
) -> 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_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
|
||||||
|
|
||||||
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
||||||
|
|||||||
@@ -7,13 +7,20 @@ from contextlib import contextmanager
|
|||||||
from dataclasses import replace
|
from dataclasses import replace
|
||||||
from datetime import date, datetime
|
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 (
|
from ..domain.persistence import (
|
||||||
DailyAggregateRecord,
|
DailyAggregateRecord,
|
||||||
MembershipRecord,
|
MembershipRecord,
|
||||||
PublicationSourceGroup,
|
PublicationSourceGroup,
|
||||||
PublicationSourceRecord,
|
PublicationSourceRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
StockFactRecord,
|
StockFactRecord,
|
||||||
StockMembershipEntry,
|
StockMembershipEntry,
|
||||||
WriteCounts,
|
WriteCounts,
|
||||||
@@ -137,6 +144,58 @@ class InMemorySectorRadarRepository:
|
|||||||
and record.status.value == "available"
|
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:
|
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
||||||
"""Insert normalized fact revisions idempotently."""
|
"""Insert normalized fact revisions idempotently."""
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ from ..domain.persistence import (
|
|||||||
PublicationSourceGroup,
|
PublicationSourceGroup,
|
||||||
PublicationSourceRecord,
|
PublicationSourceRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
StockFactRecord,
|
StockFactRecord,
|
||||||
StockMembershipEntry,
|
StockMembershipEntry,
|
||||||
WriteCounts,
|
WriteCounts,
|
||||||
@@ -310,6 +311,60 @@ class PostgresSectorRadarRepository:
|
|||||||
for row in rows
|
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:
|
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
||||||
"""COPY normalized stock facts while preserving contributing source ids."""
|
"""COPY normalized stock facts while preserving contributing source ids."""
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import time
|
|||||||
from collections import Counter
|
from collections import Counter
|
||||||
from collections.abc import Callable, Mapping, Sequence
|
from collections.abc import Callable, Mapping, Sequence
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, replace
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from typing import Protocol, cast
|
from typing import Protocol, cast
|
||||||
|
|
||||||
@@ -29,6 +29,9 @@ from ..domain.runs import (
|
|||||||
SelectionRunItem,
|
SelectionRunItem,
|
||||||
SelectionRunStatus,
|
SelectionRunStatus,
|
||||||
SelectionRunStore,
|
SelectionRunStore,
|
||||||
|
SelectionSectorCount,
|
||||||
|
SelectionSectorMembership,
|
||||||
|
SelectionSectorReader,
|
||||||
SelectionStock,
|
SelectionStock,
|
||||||
SelectionUniverseReader,
|
SelectionUniverseReader,
|
||||||
)
|
)
|
||||||
@@ -60,6 +63,16 @@ class PreparedSelectionRun:
|
|||||||
source: SelectionExecutionSource
|
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:
|
class RunZhixingB1:
|
||||||
"""Prepare, execute, and query persisted selection strategy batches.
|
"""Prepare, execute, and query persisted selection strategy batches.
|
||||||
|
|
||||||
@@ -75,6 +88,7 @@ class RunZhixingB1:
|
|||||||
evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None,
|
evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None,
|
||||||
pattern_case_loader: PatternCaseLibraryLoader | None = None,
|
pattern_case_loader: PatternCaseLibraryLoader | None = None,
|
||||||
pattern_scorer: PatternScorer | None = None,
|
pattern_scorer: PatternScorer | None = None,
|
||||||
|
sector_reader: SelectionSectorReader | None = None,
|
||||||
*,
|
*,
|
||||||
pattern_scoring_enabled: bool = False,
|
pattern_scoring_enabled: bool = False,
|
||||||
max_workers: int = 4,
|
max_workers: int = 4,
|
||||||
@@ -96,6 +110,7 @@ class RunZhixingB1:
|
|||||||
self.evaluators.update(evaluators)
|
self.evaluators.update(evaluators)
|
||||||
self.pattern_case_loader = pattern_case_loader
|
self.pattern_case_loader = pattern_case_loader
|
||||||
self.pattern_scorer = pattern_scorer
|
self.pattern_scorer = pattern_scorer
|
||||||
|
self.sector_reader = sector_reader
|
||||||
self.pattern_scoring_enabled = pattern_scoring_enabled
|
self.pattern_scoring_enabled = pattern_scoring_enabled
|
||||||
self.max_workers = max_workers
|
self.max_workers = max_workers
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
@@ -211,8 +226,7 @@ class RunZhixingB1:
|
|||||||
)
|
)
|
||||||
history_rows += batch_history_rows
|
history_rows += batch_history_rows
|
||||||
batch_missing_turnover = sum(
|
batch_missing_turnover = sum(
|
||||||
not _turnover_present(history, target_trade_date)
|
not _turnover_present(history, target_trade_date) for history in histories
|
||||||
for history in histories
|
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"selection_read_batch_summary strategy=%s target_trade_date=%s "
|
"selection_read_batch_summary strategy=%s target_trade_date=%s "
|
||||||
@@ -581,9 +595,22 @@ class RunZhixingB1:
|
|||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read one persisted run for polling."""
|
"""Read one persisted run for polling with optional sector filtering."""
|
||||||
|
|
||||||
return self.store.get_run(run_id, query=query)
|
effective = query or SelectionResultQuery()
|
||||||
|
if not effective.sector:
|
||||||
|
return self.store.get_run(run_id, query=query)
|
||||||
|
identity = self.store.get_run_identity(run_id)
|
||||||
|
if identity is None:
|
||||||
|
return None
|
||||||
|
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||||
|
if member_codes is None:
|
||||||
|
return self.store.get_run(run_id, query=query)
|
||||||
|
return self.store.get_run(
|
||||||
|
run_id,
|
||||||
|
query=replace(effective, sector=None),
|
||||||
|
sector_stock_codes=member_codes,
|
||||||
|
)
|
||||||
|
|
||||||
def get_latest(
|
def get_latest(
|
||||||
self,
|
self,
|
||||||
@@ -594,7 +621,73 @@ class RunZhixingB1:
|
|||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read the current result by date or the latest result for a strategy."""
|
"""Read the current result by date or the latest result for a strategy."""
|
||||||
|
|
||||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
effective = query or SelectionResultQuery()
|
||||||
|
if not effective.sector:
|
||||||
|
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||||
|
identity = self.store.get_latest_run_identity(strategy, target_trade_date)
|
||||||
|
if identity is None:
|
||||||
|
return None
|
||||||
|
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||||
|
if member_codes is None:
|
||||||
|
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||||
|
return self.store.get_latest_run(
|
||||||
|
strategy,
|
||||||
|
identity.target_trade_date,
|
||||||
|
query=replace(effective, sector=None),
|
||||||
|
sector_stock_codes=member_codes,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
) -> tuple[str, ...] | None:
|
||||||
|
"""Resolve one sector's members, or None when the port is absent."""
|
||||||
|
|
||||||
|
if self.sector_reader is None:
|
||||||
|
return None
|
||||||
|
return self.sector_reader.sector_member_codes(target_trade_date, sector_code)
|
||||||
|
|
||||||
|
def list_sector_counts(
|
||||||
|
self,
|
||||||
|
strategy: StrategyName,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorAggregates | None:
|
||||||
|
"""Aggregate the current run's selected stocks by point-in-time sector."""
|
||||||
|
|
||||||
|
run = self.store.get_latest_run(strategy, target_trade_date)
|
||||||
|
if run is None:
|
||||||
|
return None
|
||||||
|
selected_codes = [
|
||||||
|
item.ts_code
|
||||||
|
for item in run.items
|
||||||
|
if item.status == "selected" and item.signal_count > 0
|
||||||
|
]
|
||||||
|
membership = self._sector_membership(selected_codes, run.target_trade_date, sector_type)
|
||||||
|
return SelectionSectorAggregates(
|
||||||
|
run=run,
|
||||||
|
snapshot_trade_date=membership.snapshot_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
sectors=membership.sector_counts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _sector_membership(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_type: str,
|
||||||
|
) -> SelectionSectorMembership:
|
||||||
|
"""Read sector counts for a stock set, tolerating a missing port."""
|
||||||
|
|
||||||
|
if self.sector_reader is None or not stock_codes:
|
||||||
|
return SelectionSectorMembership(snapshot_trade_date=None, sector_counts=())
|
||||||
|
return self.sector_reader.sector_counts(
|
||||||
|
stock_codes,
|
||||||
|
target_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _to_item(
|
def _to_item(
|
||||||
@@ -676,4 +769,5 @@ __all__ = [
|
|||||||
"RunZhixingB1",
|
"RunZhixingB1",
|
||||||
"SelectionRerunRequired",
|
"SelectionRerunRequired",
|
||||||
"SelectionRunInProgress",
|
"SelectionRunInProgress",
|
||||||
|
"SelectionSectorAggregates",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -31,6 +31,52 @@ class SelectionResultQuery:
|
|||||||
search: str | None = None
|
search: str | None = None
|
||||||
category: SelectionSignalCategoryFilter | None = None
|
category: SelectionSignalCategoryFilter | None = None
|
||||||
sort: SelectionResultSort = "code"
|
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)
|
@dataclass(frozen=True, slots=True)
|
||||||
@@ -139,6 +185,7 @@ class SelectionRunStore(Protocol):
|
|||||||
run_id: str,
|
run_id: str,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None: ...
|
) -> SelectionRun | None: ...
|
||||||
|
|
||||||
def get_latest_run(
|
def get_latest_run(
|
||||||
@@ -147,8 +194,17 @@ class SelectionRunStore(Protocol):
|
|||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | 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):
|
class SelectionUniverseReader(Protocol):
|
||||||
"""Read a qualified market-data source snapshot for one strategy run."""
|
"""Read a qualified market-data source snapshot for one strategy run."""
|
||||||
|
|||||||
+83
-8
@@ -33,6 +33,7 @@ from ..domain.runs import (
|
|||||||
SelectionResultQuery,
|
SelectionResultQuery,
|
||||||
SelectionRun,
|
SelectionRun,
|
||||||
SelectionRunError,
|
SelectionRunError,
|
||||||
|
SelectionRunIdentity,
|
||||||
SelectionRunInProgress,
|
SelectionRunInProgress,
|
||||||
SelectionRunItem,
|
SelectionRunItem,
|
||||||
SelectionRunStatus,
|
SelectionRunStatus,
|
||||||
@@ -46,9 +47,7 @@ _SELECTION_SIGNAL_ORDER: tuple[SelectionSignalCategory, ...] = (
|
|||||||
*ZHIXING_B1_SIGNAL_ORDER,
|
*ZHIXING_B1_SIGNAL_ORDER,
|
||||||
*GOLD_BRICK_SIGNAL_ORDER,
|
*GOLD_BRICK_SIGNAL_ORDER,
|
||||||
)
|
)
|
||||||
_SIGNAL_PRIORITY = {
|
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)}
|
||||||
category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)
|
|
||||||
}
|
|
||||||
_CATEGORY_PREFIXES = {
|
_CATEGORY_PREFIXES = {
|
||||||
"pullback": "zhixing_b1_pullback_",
|
"pullback": "zhixing_b1_pullback_",
|
||||||
"oversold": "zhixing_b1_oversold_",
|
"oversold": "zhixing_b1_oversold_",
|
||||||
@@ -335,21 +334,41 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
run_id: str,
|
run_id: str,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read one run with filtered, stock-paged signals and item failures."""
|
"""Read one run with filtered, stock-paged signals and item failures."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with self._connection() as connection:
|
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:
|
except psycopg.Error as exc:
|
||||||
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from 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(
|
def get_latest_run(
|
||||||
self,
|
self,
|
||||||
strategy: SelectionStrategyName,
|
strategy: SelectionStrategyName,
|
||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read the current run for a date or the latest date for a strategy."""
|
"""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),
|
(strategy, target_trade_date),
|
||||||
).fetchone()
|
).fetchone()
|
||||||
return (
|
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
|
if row
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
except psycopg.Error as exc:
|
except psycopg.Error as exc:
|
||||||
raise SelectionRunStoreError("failed to load latest selection run") from 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
|
@staticmethod
|
||||||
def _load_run(
|
def _load_run(
|
||||||
connection: Any,
|
connection: Any,
|
||||||
run_id: str,
|
run_id: str,
|
||||||
query: SelectionResultQuery,
|
query: SelectionResultQuery,
|
||||||
|
*,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
row = connection.execute(
|
row = connection.execute(
|
||||||
"""
|
"""
|
||||||
@@ -417,7 +477,9 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
""",
|
""",
|
||||||
(run_id,),
|
(run_id,),
|
||||||
).fetchall()
|
).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(
|
stock_total_row = connection.execute(
|
||||||
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
|
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
|
||||||
tuple(stock_parameters),
|
tuple(stock_parameters),
|
||||||
@@ -568,12 +630,19 @@ def _signal_category(value: str) -> SelectionSignalCategory:
|
|||||||
return GoldBrickCategory(value)
|
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.
|
"""Build the signal predicate used to select distinct matching stocks.
|
||||||
|
|
||||||
A category narrows which stocks qualify for the page. Once a stock
|
A category narrows which stocks qualify for the page. Once a stock
|
||||||
qualifies, the repository loads every signal for that stock so callers
|
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"]
|
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]}%")
|
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
|
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 pydantic import BaseModel, Field
|
||||||
|
|
||||||
from zhixing_server.bootstrap.config import Settings, get_settings
|
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 (
|
from zhixing_server.modules.selection.application.chart import (
|
||||||
GetSelectionChart,
|
GetSelectionChart,
|
||||||
SelectionChart,
|
SelectionChart,
|
||||||
@@ -46,6 +47,9 @@ from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
|||||||
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
||||||
PostgresSelectionRunRepository,
|
PostgresSelectionRunRepository,
|
||||||
)
|
)
|
||||||
|
from zhixing_server.modules.selection.infrastructure.sector_membership import (
|
||||||
|
SectorRadarSelectionReader,
|
||||||
|
)
|
||||||
|
|
||||||
selection_router = APIRouter()
|
selection_router = APIRouter()
|
||||||
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
||||||
@@ -212,6 +216,32 @@ class SelectionResultsResponse(BaseModel):
|
|||||||
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
|
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(
|
def get_selection_service(
|
||||||
settings: Annotated[Settings, Depends(get_settings)],
|
settings: Annotated[Settings, Depends(get_settings)],
|
||||||
) -> RunZhixingB1:
|
) -> RunZhixingB1:
|
||||||
@@ -221,12 +251,14 @@ def get_selection_service(
|
|||||||
reader = PostgresMarketDataReader(settings, pool=pool)
|
reader = PostgresMarketDataReader(settings, pool=pool)
|
||||||
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
|
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
|
||||||
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
|
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
|
||||||
|
sector_reader = SectorRadarSelectionReader(get_sector_radar_reader(settings))
|
||||||
return RunZhixingB1(
|
return RunZhixingB1(
|
||||||
reader,
|
reader,
|
||||||
store,
|
store,
|
||||||
evaluators={"gold_brick": EvaluateGoldBrick(reader)},
|
evaluators={"gold_brick": EvaluateGoldBrick(reader)},
|
||||||
pattern_case_loader=pattern_case_loader,
|
pattern_case_loader=pattern_case_loader,
|
||||||
pattern_scorer=ZhixingB1PatternScorer(),
|
pattern_scorer=ZhixingB1PatternScorer(),
|
||||||
|
sector_reader=sector_reader,
|
||||||
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
|
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
|
||||||
max_workers=settings.selection_max_workers,
|
max_workers=settings.selection_max_workers,
|
||||||
batch_size=settings.selection_batch_size,
|
batch_size=settings.selection_batch_size,
|
||||||
@@ -338,15 +370,14 @@ def get_selection_run(
|
|||||||
page: Annotated[int, Query(ge=1)] = 1,
|
page: Annotated[int, Query(ge=1)] = 1,
|
||||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||||
category: Literal[
|
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
|
||||||
"pullback", "oversold", "original", "resonance"
|
|
||||||
] | None = None,
|
|
||||||
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
||||||
|
sector: Annotated[str | None, Query(max_length=60)] = None,
|
||||||
) -> SelectionResultsResponse:
|
) -> SelectionResultsResponse:
|
||||||
"""Return one run for asynchronous polling."""
|
"""Return one run for asynchronous polling."""
|
||||||
|
|
||||||
try:
|
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)
|
run = service.get_run(run_id, query=query)
|
||||||
except SelectionRunStoreError as exc:
|
except SelectionRunStoreError as exc:
|
||||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from 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)
|
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)
|
@selection_router.get("/results", response_model=SelectionResultsResponse)
|
||||||
def get_selection_results(
|
def get_selection_results(
|
||||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||||
@@ -363,15 +439,14 @@ def get_selection_results(
|
|||||||
page: Annotated[int, Query(ge=1)] = 1,
|
page: Annotated[int, Query(ge=1)] = 1,
|
||||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||||
category: Literal[
|
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
|
||||||
"pullback", "oversold", "original", "resonance"
|
|
||||||
] | None = None,
|
|
||||||
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
||||||
|
sector: Annotated[str | None, Query(max_length=60)] = None,
|
||||||
) -> SelectionResultsResponse:
|
) -> SelectionResultsResponse:
|
||||||
"""Return the current persisted result for a strategy and optional date."""
|
"""Return the current persisted result for a strategy and optional date."""
|
||||||
|
|
||||||
try:
|
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)
|
run = service.get_latest(strategy, target_trade_date, query=query)
|
||||||
except SelectionRunStoreError as exc:
|
except SelectionRunStoreError as exc:
|
||||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||||
@@ -524,16 +599,19 @@ def _result_query(
|
|||||||
search: str | None,
|
search: str | None,
|
||||||
category: Literal["pullback", "oversold", "original", "resonance"] | None,
|
category: Literal["pullback", "oversold", "original", "resonance"] | None,
|
||||||
sort: Literal["code", "score_desc", "score_asc"],
|
sort: Literal["code", "score_desc", "score_asc"],
|
||||||
|
sector: str | None = None,
|
||||||
) -> SelectionResultQuery:
|
) -> SelectionResultQuery:
|
||||||
"""Normalize HTTP query values before handing them to the selection port."""
|
"""Normalize HTTP query values before handing them to the selection port."""
|
||||||
|
|
||||||
normalized_search = search.strip() if search else None
|
normalized_search = search.strip() if search else None
|
||||||
|
normalized_sector = sector.strip() if sector else None
|
||||||
return SelectionResultQuery(
|
return SelectionResultQuery(
|
||||||
page=page,
|
page=page,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
search=normalized_search or None,
|
search=normalized_search or None,
|
||||||
category=category,
|
category=category,
|
||||||
sort=sort,
|
sort=sort,
|
||||||
|
sector=normalized_sector or None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -551,6 +629,7 @@ __all__ = [
|
|||||||
"SelectionResultsResponse",
|
"SelectionResultsResponse",
|
||||||
"SelectionRunAcceptedResponse",
|
"SelectionRunAcceptedResponse",
|
||||||
"SelectionRunRequest",
|
"SelectionRunRequest",
|
||||||
|
"SelectionSectorsResponse",
|
||||||
"SelectionStockResponse",
|
"SelectionStockResponse",
|
||||||
"get_selection_chart_service",
|
"get_selection_chart_service",
|
||||||
"get_selection_service",
|
"get_selection_service",
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ from zhixing_server.modules.selection.application.chart import (
|
|||||||
SelectionChartNotFound,
|
SelectionChartNotFound,
|
||||||
SelectionChartPoint,
|
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.models import SelectionSignal, ZhixingB1Category
|
||||||
from zhixing_server.modules.selection.domain.pattern_scoring import (
|
from zhixing_server.modules.selection.domain.pattern_scoring import (
|
||||||
PATTERN_SCORING_VERSION,
|
PATTERN_SCORING_VERSION,
|
||||||
@@ -30,6 +33,7 @@ from zhixing_server.modules.selection.domain.runs import (
|
|||||||
SelectionResultQuery,
|
SelectionResultQuery,
|
||||||
SelectionRun,
|
SelectionRun,
|
||||||
SelectionRunInProgress,
|
SelectionRunInProgress,
|
||||||
|
SelectionSectorCount,
|
||||||
SelectionStock,
|
SelectionStock,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||||
@@ -45,11 +49,17 @@ TARGET = date(2026, 8, 8)
|
|||||||
|
|
||||||
|
|
||||||
class FakeSelectionService:
|
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.run = run
|
||||||
|
self.sectors = sectors
|
||||||
self.executed = False
|
self.executed = False
|
||||||
self.mode = "ok"
|
self.mode = "ok"
|
||||||
self.last_query: SelectionResultQuery | None = None
|
self.last_query: SelectionResultQuery | None = None
|
||||||
|
self.sector_type_requested: str | None = None
|
||||||
|
|
||||||
def prepare(
|
def prepare(
|
||||||
self,
|
self,
|
||||||
@@ -103,6 +113,23 @@ class FakeSelectionService:
|
|||||||
return None
|
return None
|
||||||
return self.run
|
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:
|
class FakeChartService:
|
||||||
"""Return or fail one deterministic chart response."""
|
"""Return or fail one deterministic chart response."""
|
||||||
@@ -434,6 +461,83 @@ def test_query_rejects_invalid_page_size() -> None:
|
|||||||
assert response.status_code == 422
|
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:
|
def test_chart_returns_bounded_qfq_contract() -> None:
|
||||||
chart_service = FakeChartService()
|
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 (
|
from zhixing_server.modules.sector_radar.domain.persistence import (
|
||||||
MembershipRecord,
|
MembershipRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
||||||
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
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)
|
StockSectorQuery(trade_date=TARGET_DATE, ts_code="x" * 13)
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
StockSectorQuery(trade_date=TARGET_DATE, ts_code="000001.SZ", concept_limit=0)
|
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
|
kwargs["error_message"] = error_message
|
||||||
self.finished = (run_id, status, kwargs)
|
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
|
return None
|
||||||
|
|
||||||
def get_latest_run(
|
def get_latest_run(
|
||||||
@@ -145,9 +151,16 @@ class FakeStore:
|
|||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
):
|
):
|
||||||
return 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):
|
class BatchStore(FakeStore):
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
@@ -427,8 +440,8 @@ def test_execute_loads_cases_once_and_scores_only_selected_stocks() -> None:
|
|||||||
reader,
|
reader,
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
pattern_scoring_enabled=True,
|
||||||
batch_size=1,
|
batch_size=1,
|
||||||
)
|
)
|
||||||
@@ -463,8 +476,8 @@ def test_execute_isolates_pattern_scoring_failure_from_selection_status() -> Non
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
pattern_scoring_enabled=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -497,8 +510,8 @@ def test_execute_skips_pattern_dependencies_when_feature_flag_is_disabled() -> N
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=False,
|
pattern_scoring_enabled=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -535,8 +548,8 @@ def test_execute_marks_scores_failed_when_case_library_is_unavailable() -> None:
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
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