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."""
|
||||
|
||||
return self.store.get_run(run_id, query=query)
|
||||
effective = query or SelectionResultQuery()
|
||||
if not effective.sector:
|
||||
return self.store.get_run(run_id, query=query)
|
||||
identity = self.store.get_run_identity(run_id)
|
||||
if identity is None:
|
||||
return None
|
||||
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||
if member_codes is None:
|
||||
return self.store.get_run(run_id, query=query)
|
||||
return self.store.get_run(
|
||||
run_id,
|
||||
query=replace(effective, sector=None),
|
||||
sector_stock_codes=member_codes,
|
||||
)
|
||||
|
||||
def get_latest(
|
||||
self,
|
||||
@@ -594,7 +621,73 @@ class RunZhixingB1:
|
||||
) -> SelectionRun | None:
|
||||
"""Read the current result by date or the latest result for a strategy."""
|
||||
|
||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||
effective = query or SelectionResultQuery()
|
||||
if not effective.sector:
|
||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||
identity = self.store.get_latest_run_identity(strategy, target_trade_date)
|
||||
if identity is None:
|
||||
return None
|
||||
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||
if member_codes is None:
|
||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||
return self.store.get_latest_run(
|
||||
strategy,
|
||||
identity.target_trade_date,
|
||||
query=replace(effective, sector=None),
|
||||
sector_stock_codes=member_codes,
|
||||
)
|
||||
|
||||
def _sector_member_codes(
|
||||
self,
|
||||
target_trade_date: date,
|
||||
sector_code: str,
|
||||
) -> tuple[str, ...] | None:
|
||||
"""Resolve one sector's members, or None when the port is absent."""
|
||||
|
||||
if self.sector_reader is None:
|
||||
return None
|
||||
return self.sector_reader.sector_member_codes(target_trade_date, sector_code)
|
||||
|
||||
def list_sector_counts(
|
||||
self,
|
||||
strategy: StrategyName,
|
||||
target_trade_date: date | None = None,
|
||||
*,
|
||||
sector_type: str = "concept",
|
||||
) -> SelectionSectorAggregates | None:
|
||||
"""Aggregate the current run's selected stocks by point-in-time sector."""
|
||||
|
||||
run = self.store.get_latest_run(strategy, target_trade_date)
|
||||
if run is None:
|
||||
return None
|
||||
selected_codes = [
|
||||
item.ts_code
|
||||
for item in run.items
|
||||
if item.status == "selected" and item.signal_count > 0
|
||||
]
|
||||
membership = self._sector_membership(selected_codes, run.target_trade_date, sector_type)
|
||||
return SelectionSectorAggregates(
|
||||
run=run,
|
||||
snapshot_trade_date=membership.snapshot_trade_date,
|
||||
sector_type=sector_type,
|
||||
sectors=membership.sector_counts,
|
||||
)
|
||||
|
||||
def _sector_membership(
|
||||
self,
|
||||
stock_codes: Sequence[str],
|
||||
target_trade_date: date,
|
||||
sector_type: str,
|
||||
) -> SelectionSectorMembership:
|
||||
"""Read sector counts for a stock set, tolerating a missing port."""
|
||||
|
||||
if self.sector_reader is None or not stock_codes:
|
||||
return SelectionSectorMembership(snapshot_trade_date=None, sector_counts=())
|
||||
return self.sector_reader.sector_counts(
|
||||
stock_codes,
|
||||
target_trade_date,
|
||||
sector_type=sector_type,
|
||||
)
|
||||
|
||||
|
||||
def _to_item(
|
||||
@@ -676,4 +769,5 @@ __all__ = [
|
||||
"RunZhixingB1",
|
||||
"SelectionRerunRequired",
|
||||
"SelectionRunInProgress",
|
||||
"SelectionSectorAggregates",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user