Merge branch 'develop'
Deploy Production / deploy (push) Successful in 28s

This commit is contained in:
yuxuanhui
2026-09-05 19:49:02 +08:00
26 changed files with 1442 additions and 77 deletions
@@ -0,0 +1 @@
{"_example": "Fill with {\"file\": \"<path>\", \"reason\": \"<why>\"}. Put spec/research files only — no code paths. Run `python3 .trellis/scripts/get_context.py --mode packages` to list available specs. Delete this line once real entries are added."}
@@ -0,0 +1 @@
{"_example": "Fill with {\"file\": \"<path>\", \"reason\": \"<why>\"}. Put spec/research files only — no code paths. Run `python3 .trellis/scripts/get_context.py --mode packages` to list available specs. Delete this line once real entries are added."}
@@ -0,0 +1,40 @@
# 选股页面迭代:布局重构与板块筛选
## 需求
1. 选股页面从左右布局改为「上搜索栏 + 左列表 / 右详情」三段式布局。
2. 搜索模块增加板块筛选:下拉选项为当次策略结果的板块聚合,选项旁展示数量,按数量倒序。
## 关键决策
- 板块口径:概念板块(sector_type=concept),与详情面板「板块」标签一致;接口保留 sector_type 参数可扩展行业。
- 单选下拉;排序按 stock_count 倒序(后端保证),名称升序 tie-break。
- 板块数据不在 selection 表中,按 ADR 0001 通过端口委托 sector_radar 读服务(不跨上下文 join SQL)。
## 实现
后端(zhixing-server):
- sector_radar:`domain/persistence.py` 新增 `SectorCountEntry` + 2 个协议方法;`infrastructure/postgres.py` / `infrastructure/memory.py` 实现 `load_sector_counts` / `load_sector_member_codes`;`application/read.py` 新增 `sector_counts` / `sector_member_codes`(先解析 last-good publication)。
- selection:`domain/runs.py` 新增 `SelectionSectorReader` 端口、`SelectionSectorCount` / `SelectionSectorMembership` / `SelectionRunIdentity`,`SelectionResultQuery.sector`;`application/run.py` 新增 `list_sector_counts`,`get_run`/`get_latest` 经 identity → 成员代码 → `sector_stock_codes` 过滤;`infrastructure/postgres_runs.py` `_stock_filter` 支持 `ts_code = ANY(...)` / FALSE;`infrastructure/sector_membership.py` 桥接适配器;`presentation/http.py` 新增 `GET /sectors`,results/runs 增加 `sector` 参数。
前端(zhixing-web):
- `selection.types.ts` / `selection.api.ts` / `selection.query.ts`:`SelectionSectors`、`getSelectionResultSectors`、`useSelectionResultSectors`,query key 加 sector,结果失效同时失效 sectors。
- `route-tree.tsx`:selection 路由 search 增加 `sector`。
- `selection-results-page.tsx`:结果查询带 sector,页面调用 sectors hook 并下传 workbench;切策略重置 sector。
- `selection-results-workbench.tsx`:布局重构为上搜索栏 + 下方 320px 列表/详情两栏;板块下拉(全部板块 + 聚合选项带数量徽标);聚合加载后自动清掉失效 sector。
## 测试
- 后端:`test_read.py` +8、`test_sector_filter.py` 新建 9 个、`test_selection_http.py` +5。
- 前端:页面测试新增「filters by sector and shows per-sector counts」。
## 验收
- [x] 页面布局为上搜索栏 + 左列表右详情(移动端纵向堆叠 搜索→列表→详情)
- [x] 板块下拉选项为当次策略结果的概念板块聚合,选项旁展示数量,按数量倒序
- [x] 选择板块后结果列表服务端过滤,`筛选结果 N 只` 反映叠加计数
- [x] 后端 ruff/pyright/pytest 与前端 lint/typecheck/test/build 门禁通过(遗留项均为基线原有)
## 遗留(均为基线原有,非本次引入)
- 后端 pyright 16 个错误(chart.py/gold_brick.py pandas 相关、test_run.py 旧 fake 类型);前端 4 个执行状态抽屉测试失败;前端 format:check 有基线未格式化文件(.playwright-cli、sector-radar 部分文件、signal-detail-panel.tsx)。
@@ -0,0 +1,26 @@
{
"id": "selection-layout-sector-filter",
"name": "selection-layout-sector-filter",
"title": "选股页面迭代:布局重构与板块筛选",
"description": "",
"status": "planning",
"dev_type": null,
"scope": null,
"package": null,
"priority": "P2",
"creator": "yuxuanhui",
"assignee": "yuxuanhui",
"createdAt": "2026-09-05",
"completedAt": null,
"branch": null,
"base_branch": "main",
"worktree_path": null,
"commit": null,
"pr_url": null,
"subtasks": [],
"children": [],
"parent": null,
"relatedFiles": [],
"notes": "",
"meta": {}
}
@@ -2,7 +2,7 @@
from __future__ import annotations
from collections.abc import Iterator
from collections.abc import Iterator, Sequence
from dataclasses import dataclass
from datetime import date
from enum import StrEnum
@@ -21,7 +21,11 @@ from ..domain.models import (
RankSide,
SectorType,
)
from ..domain.persistence import SectorRadarRepository, StockMembershipEntry
from ..domain.persistence import (
SectorCountEntry,
SectorRadarRepository,
StockMembershipEntry,
)
from ..domain.ranking import select_percentile_side, select_rank_change_side
ReadStatus = Literal["success", "no_data"]
@@ -148,6 +152,26 @@ class StockSectorMembership:
concept_total: int
@dataclass(frozen=True, slots=True)
class SectorCountsSnapshot:
"""Sector counts for a bounded stock set on one last-good snapshot date."""
status: ReadStatus
sector_type: SectorType
trade_date: date | None
counts: tuple[SectorCountEntry, ...] = ()
@dataclass(frozen=True, slots=True)
class SectorMembersSnapshot:
"""One sector's member stock codes on one last-good snapshot date."""
status: ReadStatus
sector_type: SectorType
trade_date: date | None
stock_codes: tuple[str, ...] = ()
_METRIC_DEFINITIONS = {
MetricKind.AMOUNT: RadarMetricDefinition(
metric_kind=MetricKind.AMOUNT,
@@ -290,6 +314,78 @@ class ReadSectorRadar:
concept_total=len(concepts),
)
def sector_counts(
self,
stock_codes: Sequence[str],
target_trade_date: date,
*,
sector_type: SectorType = SectorType.CONCEPT,
) -> SectorCountsSnapshot:
"""Count claimed stocks per sector from the strict last-good snapshot."""
if not stock_codes:
return SectorCountsSnapshot(
status="no_data",
sector_type=sector_type,
trade_date=None,
counts=(),
)
publication = self.repository.get_last_good_publication(target_trade_date)
if publication is None:
return SectorCountsSnapshot(
status="no_data",
sector_type=sector_type,
trade_date=None,
counts=(),
)
counts = tuple(
self.repository.load_sector_counts(
publication.target_trade_date,
stock_codes,
sector_type,
)
)
return SectorCountsSnapshot(
status="success",
sector_type=sector_type,
trade_date=publication.target_trade_date,
counts=counts,
)
def sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
*,
sector_type: SectorType = SectorType.CONCEPT,
) -> SectorMembersSnapshot:
"""Return one sector's members from the strict last-good snapshot."""
normalized_code = sector_code.strip()
if not normalized_code:
raise ValueError("sector_code must not be empty")
publication = self.repository.get_last_good_publication(target_trade_date)
if publication is None:
return SectorMembersSnapshot(
status="no_data",
sector_type=sector_type,
trade_date=None,
stock_codes=(),
)
codes = tuple(
self.repository.load_sector_member_codes(
publication.target_trade_date,
normalized_code,
sector_type,
)
)
return SectorMembersSnapshot(
status="success",
sector_type=sector_type,
trade_date=publication.target_trade_date,
stock_codes=codes,
)
def _sector_refs(entries: Iterator[StockMembershipEntry]) -> tuple[SectorRef, ...]:
"""Project membership entries into ordered public sector references."""
@@ -309,6 +405,8 @@ __all__ = [
"RadarView",
"RankingPage",
"ReadSectorRadar",
"SectorCountsSnapshot",
"SectorMembersSnapshot",
"SectorRef",
"StockSectorMembership",
"StockSectorQuery",
@@ -81,6 +81,23 @@ class StockMembershipEntry:
raise ValueError("membership sector identity fields must not be empty")
@dataclass(frozen=True, slots=True)
class SectorCountEntry:
"""One sector and the number of stocks it claimed on one snapshot date."""
sector_code: str
sector_name: str
stock_count: int
def __post_init__(self) -> None:
"""Reject empty sector identity and non-positive counts."""
if not self.sector_code.strip() or not self.sector_name.strip():
raise ValueError("sector count identity fields must not be empty")
if self.stock_count < 1:
raise ValueError("sector stock_count must be positive")
@dataclass(frozen=True, slots=True)
class StockFactRecord:
"""One normalized stock fact revision with all contributing raw snapshots."""
@@ -217,6 +234,20 @@ class SectorRadarRepository(Protocol):
self, trade_date: date, stock_code: str
) -> Sequence[StockMembershipEntry]: ...
def load_sector_counts(
self,
trade_date: date,
stock_codes: Sequence[str],
sector_type: SectorType,
) -> Sequence[SectorCountEntry]: ...
def load_sector_member_codes(
self,
trade_date: date,
sector_code: str,
sector_type: SectorType,
) -> Sequence[str]: ...
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
@@ -7,13 +7,20 @@ from contextlib import contextmanager
from dataclasses import replace
from datetime import date, datetime
from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate
from ..domain.models import (
PublicationStatus,
RadarPublication,
RankedMetric,
SectorDailyAggregate,
SectorType,
)
from ..domain.persistence import (
DailyAggregateRecord,
MembershipRecord,
PublicationSourceGroup,
PublicationSourceRecord,
RankingRecord,
SectorCountEntry,
StockFactRecord,
StockMembershipEntry,
WriteCounts,
@@ -137,6 +144,58 @@ class InMemorySectorRadarRepository:
and record.status.value == "available"
)
def load_sector_counts(
self,
trade_date: date,
stock_codes: Sequence[str],
sector_type: SectorType,
) -> Sequence[SectorCountEntry]:
"""Aggregate claimed-stock counts per sector on one snapshot date."""
wanted = set(stock_codes)
counts: dict[tuple[str, str], int] = {}
for record in self.memberships.values():
if (
record.trade_date == trade_date
and record.sector_type is sector_type
and record.status.value == "available"
and record.stock_code in wanted
):
key = (record.sector_code, record.sector_name)
counts[key] = counts.get(key, 0) + 1
return tuple(
SectorCountEntry(
sector_code=sector_code,
sector_name=sector_name,
stock_count=stock_count,
)
for (sector_code, sector_name), stock_count in sorted(
counts.items(),
key=lambda item: (-item[1], item[0][0]),
)
)
def load_sector_member_codes(
self,
trade_date: date,
sector_code: str,
sector_type: SectorType,
) -> Sequence[str]:
"""Return available member stock codes of one sector on one date."""
return tuple(
record.stock_code
for record in sorted(
self.memberships.values(),
key=lambda item: item.stock_code or "",
)
if record.trade_date == trade_date
and record.sector_type is sector_type
and record.sector_code == sector_code
and record.status.value == "available"
and record.stock_code is not None
)
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
"""Insert normalized fact revisions idempotently."""
@@ -30,6 +30,7 @@ from ..domain.persistence import (
PublicationSourceGroup,
PublicationSourceRecord,
RankingRecord,
SectorCountEntry,
StockFactRecord,
StockMembershipEntry,
WriteCounts,
@@ -310,6 +311,60 @@ class PostgresSectorRadarRepository:
for row in rows
)
def load_sector_counts(
self,
trade_date: date,
stock_codes: Sequence[str],
sector_type: SectorType,
) -> Sequence[SectorCountEntry]:
"""Aggregate how many of the given stocks each sector claimed on one date."""
if not stock_codes:
return ()
with self._connection() as connection:
rows = connection.execute(
"""
SELECT sector_code, sector_name, COUNT(*) AS stock_count
FROM sector_radar_membership
WHERE trade_date = %s AND sector_type = %s
AND membership_status = 'available'
AND stock_code = ANY(%s)
GROUP BY sector_code, sector_name
ORDER BY stock_count DESC, sector_code
""",
(trade_date, sector_type.value, list(stock_codes)),
).fetchall()
return tuple(
SectorCountEntry(
sector_code=row[0],
sector_name=row[1],
stock_count=int(row[2]),
)
for row in rows
)
def load_sector_member_codes(
self,
trade_date: date,
sector_code: str,
sector_type: SectorType,
) -> Sequence[str]:
"""Load every stock one sector claimed on exactly one snapshot date."""
with self._connection() as connection:
rows = connection.execute(
"""
SELECT stock_code
FROM sector_radar_membership
WHERE trade_date = %s AND sector_type = %s
AND sector_code = %s
AND membership_status = 'available'
ORDER BY stock_code
""",
(trade_date, sector_type.value, sector_code),
).fetchall()
return tuple(str(row[0]) for row in rows)
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
"""COPY normalized stock facts while preserving contributing source ids."""
@@ -7,7 +7,7 @@ import time
from collections import Counter
from collections.abc import Callable, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from dataclasses import dataclass, replace
from datetime import date
from typing import Protocol, cast
@@ -29,6 +29,9 @@ from ..domain.runs import (
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
SelectionSectorCount,
SelectionSectorMembership,
SelectionSectorReader,
SelectionStock,
SelectionUniverseReader,
)
@@ -60,6 +63,16 @@ class PreparedSelectionRun:
source: SelectionExecutionSource
@dataclass(frozen=True, slots=True)
class SelectionSectorAggregates:
"""A run summary plus its selected stocks' sector membership counts."""
run: SelectionRun | None
snapshot_trade_date: date | None
sector_type: str
sectors: tuple[SelectionSectorCount, ...]
class RunZhixingB1:
"""Prepare, execute, and query persisted selection strategy batches.
@@ -75,6 +88,7 @@ class RunZhixingB1:
evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None,
pattern_case_loader: PatternCaseLibraryLoader | None = None,
pattern_scorer: PatternScorer | None = None,
sector_reader: SelectionSectorReader | None = None,
*,
pattern_scoring_enabled: bool = False,
max_workers: int = 4,
@@ -96,6 +110,7 @@ class RunZhixingB1:
self.evaluators.update(evaluators)
self.pattern_case_loader = pattern_case_loader
self.pattern_scorer = pattern_scorer
self.sector_reader = sector_reader
self.pattern_scoring_enabled = pattern_scoring_enabled
self.max_workers = max_workers
self.batch_size = batch_size
@@ -211,8 +226,7 @@ class RunZhixingB1:
)
history_rows += batch_history_rows
batch_missing_turnover = sum(
not _turnover_present(history, target_trade_date)
for history in histories
not _turnover_present(history, target_trade_date) for history in histories
)
logger.info(
"selection_read_batch_summary strategy=%s target_trade_date=%s "
@@ -581,9 +595,22 @@ class RunZhixingB1:
*,
query: SelectionResultQuery | None = None,
) -> SelectionRun | None:
"""Read one persisted run for polling."""
"""Read one persisted run for polling with optional sector filtering."""
effective = query or SelectionResultQuery()
if not effective.sector:
return self.store.get_run(run_id, query=query)
identity = self.store.get_run_identity(run_id)
if identity is None:
return None
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
if member_codes is None:
return self.store.get_run(run_id, query=query)
return self.store.get_run(
run_id,
query=replace(effective, sector=None),
sector_stock_codes=member_codes,
)
def get_latest(
self,
@@ -594,7 +621,73 @@ class RunZhixingB1:
) -> SelectionRun | None:
"""Read the current result by date or the latest result for a strategy."""
effective = query or SelectionResultQuery()
if not effective.sector:
return self.store.get_latest_run(strategy, target_trade_date, query=query)
identity = self.store.get_latest_run_identity(strategy, target_trade_date)
if identity is None:
return None
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
if member_codes is None:
return self.store.get_latest_run(strategy, target_trade_date, query=query)
return self.store.get_latest_run(
strategy,
identity.target_trade_date,
query=replace(effective, sector=None),
sector_stock_codes=member_codes,
)
def _sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
) -> tuple[str, ...] | None:
"""Resolve one sector's members, or None when the port is absent."""
if self.sector_reader is None:
return None
return self.sector_reader.sector_member_codes(target_trade_date, sector_code)
def list_sector_counts(
self,
strategy: StrategyName,
target_trade_date: date | None = None,
*,
sector_type: str = "concept",
) -> SelectionSectorAggregates | None:
"""Aggregate the current run's selected stocks by point-in-time sector."""
run = self.store.get_latest_run(strategy, target_trade_date)
if run is None:
return None
selected_codes = [
item.ts_code
for item in run.items
if item.status == "selected" and item.signal_count > 0
]
membership = self._sector_membership(selected_codes, run.target_trade_date, sector_type)
return SelectionSectorAggregates(
run=run,
snapshot_trade_date=membership.snapshot_trade_date,
sector_type=sector_type,
sectors=membership.sector_counts,
)
def _sector_membership(
self,
stock_codes: Sequence[str],
target_trade_date: date,
sector_type: str,
) -> SelectionSectorMembership:
"""Read sector counts for a stock set, tolerating a missing port."""
if self.sector_reader is None or not stock_codes:
return SelectionSectorMembership(snapshot_trade_date=None, sector_counts=())
return self.sector_reader.sector_counts(
stock_codes,
target_trade_date,
sector_type=sector_type,
)
def _to_item(
@@ -676,4 +769,5 @@ __all__ = [
"RunZhixingB1",
"SelectionRerunRequired",
"SelectionRunInProgress",
"SelectionSectorAggregates",
]
@@ -19,9 +19,7 @@ from .zhixing_b1 import compute_signal_masks, prepare_zhixing_b1_indicators
GOLD_BRICK_MINIMUM_HISTORY = 200
GOLD_BRICK_TURNOVER_RATE_THRESHOLD = 0.99
GOLD_BRICK_SIGNAL_ORDER: tuple[GoldBrickCategory, ...] = (
GoldBrickCategory.RESONANCE,
)
GOLD_BRICK_SIGNAL_ORDER: tuple[GoldBrickCategory, ...] = (GoldBrickCategory.RESONANCE,)
def _safe_ratio(numerator: pd.Series, denominator: pd.Series) -> pd.Series:
@@ -90,9 +88,7 @@ def prepare_gold_brick_indicators(frame: pd.DataFrame, code: str) -> pd.DataFram
)
multiple_volume_bonus = pd.Series(
np.where(
(close > open_price)
& (close > previous_close)
& (volume > previous_volume * 1.8),
(close > open_price) & (close > previous_close) & (volume > previous_volume * 1.8),
multiple_volume_coefficient,
1.0,
),
@@ -115,13 +111,9 @@ def prepare_gold_brick_indicators(frame: pd.DataFrame, code: str) -> pd.DataFram
)
result["j_momentum"] = j_momentum
result["rsi_momentum"] = rsi_momentum
result["yellow_column"] = (
momentum_sum.div(2).mul(shadow_coefficient).mul(multiple_volume_bonus)
)
result["yellow_column"] = momentum_sum.div(2).mul(shadow_coefficient).mul(multiple_volume_bonus)
x_condition = (
(close > open_price)
& (close > previous_close)
& (momentum_sum > previous_momentum_sum)
(close > open_price) & (close > previous_close) & (momentum_sum > previous_momentum_sum)
)
result["x_momentum"] = (
momentum_sum.sub(previous_momentum_sum)
@@ -163,9 +155,8 @@ def prepare_gold_brick_indicators(frame: pd.DataFrame, code: str) -> pd.DataFram
high - close,
high - upper_shadow_floor,
)
result["upper_shadow_condition"] = (
((close >= open_price) | (close > previous_close))
& (result["upper_shadow_strength"] > 0.618)
result["upper_shadow_condition"] = ((close >= open_price) | (close > previous_close)) & (
result["upper_shadow_strength"] > 0.618
)
long = result["long_oscillator"]
@@ -31,6 +31,52 @@ class SelectionResultQuery:
search: str | None = None
category: SelectionSignalCategoryFilter | None = None
sort: SelectionResultSort = "code"
sector: str | None = None
@dataclass(frozen=True, slots=True)
class SelectionSectorCount:
"""One sector and the number of this run's selected stocks it contains."""
sector_code: str
sector_name: str
stock_count: int
@dataclass(frozen=True, slots=True)
class SelectionSectorMembership:
"""Point-in-time sector membership aggregates for a set of stock codes."""
snapshot_trade_date: date | None
sector_counts: tuple[SelectionSectorCount, ...] = ()
@dataclass(frozen=True, slots=True)
class SelectionRunIdentity:
"""Minimal run locator used to resolve date-dependent filters."""
run_id: str
target_trade_date: date
class SelectionSectorReader(Protocol):
"""Port to the sector-radar context's point-in-time membership reads."""
def sector_counts(
self,
stock_codes: Sequence[str],
target_trade_date: date,
*,
sector_type: str = "concept",
) -> SelectionSectorMembership: ...
def sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
*,
sector_type: str = "concept",
) -> tuple[str, ...]: ...
@dataclass(frozen=True, slots=True)
@@ -139,6 +185,7 @@ class SelectionRunStore(Protocol):
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None: ...
def get_latest_run(
@@ -147,8 +194,17 @@ class SelectionRunStore(Protocol):
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None: ...
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None: ...
def get_latest_run_identity(
self,
strategy: SelectionStrategyName,
target_trade_date: date | None = None,
) -> SelectionRunIdentity | None: ...
class SelectionUniverseReader(Protocol):
"""Read a qualified market-data source snapshot for one strategy run."""
@@ -33,6 +33,7 @@ from ..domain.runs import (
SelectionResultQuery,
SelectionRun,
SelectionRunError,
SelectionRunIdentity,
SelectionRunInProgress,
SelectionRunItem,
SelectionRunStatus,
@@ -46,9 +47,7 @@ _SELECTION_SIGNAL_ORDER: tuple[SelectionSignalCategory, ...] = (
*ZHIXING_B1_SIGNAL_ORDER,
*GOLD_BRICK_SIGNAL_ORDER,
)
_SIGNAL_PRIORITY = {
category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)
}
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)}
_CATEGORY_PREFIXES = {
"pullback": "zhixing_b1_pullback_",
"oversold": "zhixing_b1_oversold_",
@@ -335,21 +334,41 @@ class PostgresSelectionRunRepository(SelectionRunStore):
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
"""Read one run with filtered, stock-paged signals and item failures."""
try:
with self._connection() as connection:
return self._load_run(connection, run_id, query or SelectionResultQuery())
return self._load_run(
connection,
run_id,
query or SelectionResultQuery(),
sector_stock_codes=sector_stock_codes,
)
except psycopg.Error as exc:
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
"""Read only a run's locator so date-dependent filters resolve first."""
try:
with self._connection() as connection:
row = connection.execute(
"SELECT id, target_trade_date FROM selection_run WHERE id = %s",
(run_id,),
).fetchone()
except psycopg.Error as exc:
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
def get_latest_run(
self,
strategy: SelectionStrategyName,
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
"""Read the current run for a date or the latest date for a strategy."""
@@ -377,18 +396,59 @@ class PostgresSelectionRunRepository(SelectionRunStore):
(strategy, target_trade_date),
).fetchone()
return (
self._load_run(connection, str(row[0]), query or SelectionResultQuery())
self._load_run(
connection,
str(row[0]),
query or SelectionResultQuery(),
sector_stock_codes=sector_stock_codes,
)
if row
else None
)
except psycopg.Error as exc:
raise SelectionRunStoreError("failed to load latest selection run") from exc
def get_latest_run_identity(
self,
strategy: SelectionStrategyName,
target_trade_date: date | None = None,
) -> SelectionRunIdentity | None:
"""Read only the current run's locator so sector filters resolve first."""
try:
with self._connection() as connection:
if target_trade_date is None:
row = connection.execute(
"""
SELECT id, target_trade_date
FROM selection_run
WHERE strategy = %s
ORDER BY target_trade_date DESC, created_at DESC, id DESC
LIMIT 1
""",
(strategy,),
).fetchone()
else:
row = connection.execute(
"""
SELECT id, target_trade_date
FROM selection_run
WHERE strategy = %s AND target_trade_date = %s
LIMIT 1
""",
(strategy, target_trade_date),
).fetchone()
except psycopg.Error as exc:
raise SelectionRunStoreError("failed to load latest selection run") from exc
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
@staticmethod
def _load_run(
connection: Any,
run_id: str,
query: SelectionResultQuery,
*,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
row = connection.execute(
"""
@@ -417,7 +477,9 @@ class PostgresSelectionRunRepository(SelectionRunStore):
""",
(run_id,),
).fetchall()
stock_filter, stock_parameters = _stock_filter(query, run_id)
stock_filter, stock_parameters = _stock_filter(
query, run_id, sector_stock_codes=sector_stock_codes
)
stock_total_row = connection.execute(
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
tuple(stock_parameters),
@@ -568,12 +630,19 @@ def _signal_category(value: str) -> SelectionSignalCategory:
return GoldBrickCategory(value)
def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]:
def _stock_filter(
query: SelectionResultQuery,
run_id: str,
*,
sector_stock_codes: Sequence[str] | None = None,
) -> tuple[str, list[object]]:
"""Build the signal predicate used to select distinct matching stocks.
A category narrows which stocks qualify for the page. Once a stock
qualifies, the repository loads every signal for that stock so callers
can present all independently persisted categories together.
can present all independently persisted categories together. Resolved
sector membership codes arrive from the sector-radar port, so the SQL
stays inside the selection context.
"""
clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"]
@@ -592,6 +661,12 @@ def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[o
")"
)
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
if sector_stock_codes is not None:
if not sector_stock_codes:
clauses.append("FALSE")
else:
clauses.append("item.ts_code = ANY(%s)")
parameters.append(list(sector_stock_codes))
return " AND ".join(clauses), parameters
@@ -0,0 +1,84 @@
"""Bridge the selection sector-filter port to the sector-radar read model.
The selection context owns no sector-membership storage. This adapter keeps
the port contract local to selection while delegating point-in-time reads to
the sector-radar application service in the composition root.
"""
from __future__ import annotations
from collections.abc import Sequence
from datetime import date
from zhixing_server.modules.sector_radar.application.read import ReadSectorRadar
from zhixing_server.modules.sector_radar.domain.models import SectorType
from zhixing_server.modules.selection.domain.runs import (
SelectionSectorCount,
SelectionSectorMembership,
)
SELECTION_SECTOR_TYPES: dict[str, SectorType] = {
"concept": SectorType.CONCEPT,
"industry": SectorType.INDUSTRY,
}
def _sector_type(value: str) -> SectorType:
"""Map the public sector-type vocabulary onto the radar domain enum."""
try:
return SELECTION_SECTOR_TYPES[value]
except KeyError:
raise ValueError(f"unsupported sector type: {value}") from None
class SectorRadarSelectionReader:
"""Resolve selection sector aggregates through the sector-radar context."""
def __init__(self, reader: ReadSectorRadar) -> None:
self._reader = reader
def sector_counts(
self,
stock_codes: Sequence[str],
target_trade_date: date,
*,
sector_type: str = "concept",
) -> SelectionSectorMembership:
"""Return per-sector stock counts for one run's selected stocks."""
snapshot = self._reader.sector_counts(
stock_codes,
target_trade_date,
sector_type=_sector_type(sector_type),
)
return SelectionSectorMembership(
snapshot_trade_date=snapshot.trade_date,
sector_counts=tuple(
SelectionSectorCount(
sector_code=entry.sector_code,
sector_name=entry.sector_name,
stock_count=entry.stock_count,
)
for entry in snapshot.counts
),
)
def sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
*,
sector_type: str = "concept",
) -> tuple[str, ...]:
"""Return one sector's member stock codes on the aligned snapshot."""
snapshot = self._reader.sector_member_codes(
target_trade_date,
sector_code,
sector_type=_sector_type(sector_type),
)
return snapshot.stock_codes
__all__ = ["SELECTION_SECTOR_TYPES", "SectorRadarSelectionReader"]
@@ -9,6 +9,7 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, s
from pydantic import BaseModel, Field
from zhixing_server.bootstrap.config import Settings, get_settings
from zhixing_server.modules.sector_radar.presentation.http import get_sector_radar_reader
from zhixing_server.modules.selection.application.chart import (
GetSelectionChart,
SelectionChart,
@@ -46,6 +47,9 @@ from zhixing_server.modules.selection.infrastructure.postgres_reader import (
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
PostgresSelectionRunRepository,
)
from zhixing_server.modules.selection.infrastructure.sector_membership import (
SectorRadarSelectionReader,
)
selection_router = APIRouter()
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
@@ -212,6 +216,32 @@ class SelectionResultsResponse(BaseModel):
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
class SelectionSectorCountResponse(BaseModel):
"""One sector and the number of this run's selected stocks it contains."""
sector_code: str
sector_name: str
stock_count: int = Field(ge=0)
def _empty_sectors() -> list[SelectionSectorCountResponse]:
"""Create a typed default sector-aggregate list."""
return []
class SelectionSectorsResponse(BaseModel):
"""The current run's selected stocks aggregated by point-in-time sector."""
strategy: StrategyValue
target_trade_date: date | None
run_id: str | None
status: SelectionStatusValue
snapshot_trade_date: date | None
sector_type: Literal["concept", "industry"]
sectors: list[SelectionSectorCountResponse] = Field(default_factory=_empty_sectors)
def get_selection_service(
settings: Annotated[Settings, Depends(get_settings)],
) -> RunZhixingB1:
@@ -221,12 +251,14 @@ def get_selection_service(
reader = PostgresMarketDataReader(settings, pool=pool)
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
sector_reader = SectorRadarSelectionReader(get_sector_radar_reader(settings))
return RunZhixingB1(
reader,
store,
evaluators={"gold_brick": EvaluateGoldBrick(reader)},
pattern_case_loader=pattern_case_loader,
pattern_scorer=ZhixingB1PatternScorer(),
sector_reader=sector_reader,
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
max_workers=settings.selection_max_workers,
batch_size=settings.selection_batch_size,
@@ -338,15 +370,14 @@ def get_selection_run(
page: Annotated[int, Query(ge=1)] = 1,
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
search: Annotated[str | None, Query(max_length=100)] = None,
category: Literal[
"pullback", "oversold", "original", "resonance"
] | None = None,
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
sort: Literal["code", "score_desc", "score_asc"] = "code",
sector: Annotated[str | None, Query(max_length=60)] = None,
) -> SelectionResultsResponse:
"""Return one run for asynchronous polling."""
try:
query = _result_query(page, page_size, search, category, sort)
query = _result_query(page, page_size, search, category, sort, sector)
run = service.get_run(run_id, query=query)
except SelectionRunStoreError as exc:
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
@@ -355,6 +386,51 @@ def get_selection_run(
return _run_response(run, query=query)
@selection_router.get("/sectors", response_model=SelectionSectorsResponse)
def get_selection_sectors(
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
strategy: StrategyValue = "zhixing_b1",
target_trade_date: date | None = None,
sector_type: Literal["concept", "industry"] = "concept",
) -> SelectionSectorsResponse:
"""Aggregate the current run's selected stocks by point-in-time sector."""
try:
aggregates = service.list_sector_counts(
strategy,
target_trade_date,
sector_type=sector_type,
)
except SelectionRunStoreError as exc:
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
if aggregates is None or aggregates.run is None:
return SelectionSectorsResponse(
strategy=strategy,
target_trade_date=target_trade_date,
run_id=None,
status="no_data",
snapshot_trade_date=None,
sector_type=sector_type,
)
run = aggregates.run
return SelectionSectorsResponse(
strategy=run.strategy,
target_trade_date=run.target_trade_date,
run_id=run.id,
status=run.status,
snapshot_trade_date=aggregates.snapshot_trade_date,
sector_type=sector_type,
sectors=[
SelectionSectorCountResponse(
sector_code=count.sector_code,
sector_name=count.sector_name,
stock_count=count.stock_count,
)
for count in aggregates.sectors
],
)
@selection_router.get("/results", response_model=SelectionResultsResponse)
def get_selection_results(
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
@@ -363,15 +439,14 @@ def get_selection_results(
page: Annotated[int, Query(ge=1)] = 1,
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
search: Annotated[str | None, Query(max_length=100)] = None,
category: Literal[
"pullback", "oversold", "original", "resonance"
] | None = None,
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
sort: Literal["code", "score_desc", "score_asc"] = "code",
sector: Annotated[str | None, Query(max_length=60)] = None,
) -> SelectionResultsResponse:
"""Return the current persisted result for a strategy and optional date."""
try:
query = _result_query(page, page_size, search, category, sort)
query = _result_query(page, page_size, search, category, sort, sector)
run = service.get_latest(strategy, target_trade_date, query=query)
except SelectionRunStoreError as exc:
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
@@ -524,16 +599,19 @@ def _result_query(
search: str | None,
category: Literal["pullback", "oversold", "original", "resonance"] | None,
sort: Literal["code", "score_desc", "score_asc"],
sector: str | None = None,
) -> SelectionResultQuery:
"""Normalize HTTP query values before handing them to the selection port."""
normalized_search = search.strip() if search else None
normalized_sector = sector.strip() if sector else None
return SelectionResultQuery(
page=page,
page_size=page_size,
search=normalized_search or None,
category=category,
sort=sort,
sector=normalized_sector or None,
)
@@ -551,6 +629,7 @@ __all__ = [
"SelectionResultsResponse",
"SelectionRunAcceptedResponse",
"SelectionRunRequest",
"SelectionSectorsResponse",
"SelectionStockResponse",
"get_selection_chart_service",
"get_selection_service",
+106 -2
View File
@@ -15,7 +15,10 @@ from zhixing_server.modules.selection.application.chart import (
SelectionChartNotFound,
SelectionChartPoint,
)
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
from zhixing_server.modules.selection.application.run import (
PreparedSelectionRun,
SelectionSectorAggregates,
)
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
from zhixing_server.modules.selection.domain.pattern_scoring import (
PATTERN_SCORING_VERSION,
@@ -30,6 +33,7 @@ from zhixing_server.modules.selection.domain.runs import (
SelectionResultQuery,
SelectionRun,
SelectionRunInProgress,
SelectionSectorCount,
SelectionStock,
)
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
@@ -45,11 +49,17 @@ TARGET = date(2026, 8, 8)
class FakeSelectionService:
def __init__(self, run: SelectionRun | None = None) -> None:
def __init__(
self,
run: SelectionRun | None = None,
sectors: tuple[SelectionSectorCount, ...] = (),
) -> None:
self.run = run
self.sectors = sectors
self.executed = False
self.mode = "ok"
self.last_query: SelectionResultQuery | None = None
self.sector_type_requested: str | None = None
def prepare(
self,
@@ -103,6 +113,23 @@ class FakeSelectionService:
return None
return self.run
def list_sector_counts(
self,
strategy: str,
target_trade_date: date | None = None,
*,
sector_type: str = "concept",
) -> SelectionSectorAggregates | None:
self.sector_type_requested = sector_type
if self.run is None:
return None
return SelectionSectorAggregates(
run=self.run,
snapshot_trade_date=self.run.target_trade_date,
sector_type=sector_type,
sectors=self.sectors,
)
class FakeChartService:
"""Return or fail one deterministic chart response."""
@@ -434,6 +461,83 @@ def test_query_rejects_invalid_page_size() -> None:
assert response.status_code == 422
def test_query_forwards_sector_filter() -> None:
service = FakeSelectionService(_run("run-http", "success"))
response = _client(service).get(
"/api/v1/selection/results",
params={"strategy": "zhixing_b1", "sector": " BK0475.DC "},
)
assert response.status_code == 200
assert service.last_query is not None
assert service.last_query.sector == "BK0475.DC"
def test_sectors_returns_aggregated_counts_desc() -> None:
run = _run("run-http", "success")
service = FakeSelectionService(
run,
sectors=(
SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=3),
SelectionSectorCount(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2),
),
)
response = _client(service).get("/api/v1/selection/sectors")
assert response.status_code == 200
assert response.json() == {
"strategy": "zhixing_b1",
"target_trade_date": "2026-08-08",
"run_id": "run-http",
"status": "success",
"snapshot_trade_date": "2026-08-08",
"sector_type": "concept",
"sectors": [
{"sector_code": "BK0001.DC", "sector_name": "机器人", "stock_count": 3},
{"sector_code": "BK0003.DC", "sector_name": "数字经济", "stock_count": 2},
],
}
def test_sectors_forwards_sector_type() -> None:
service = FakeSelectionService(_run("run-http", "success"))
response = _client(service).get(
"/api/v1/selection/sectors",
params={"sector_type": "industry"},
)
assert response.status_code == 200
assert service.sector_type_requested == "industry"
assert response.json()["sector_type"] == "industry"
def test_sectors_without_run_returns_no_data() -> None:
response = _client(FakeSelectionService()).get("/api/v1/selection/sectors")
assert response.status_code == 200
assert response.json() == {
"strategy": "zhixing_b1",
"target_trade_date": None,
"run_id": None,
"status": "no_data",
"snapshot_trade_date": None,
"sector_type": "concept",
"sectors": [],
}
def test_sectors_rejects_unknown_sector_type() -> None:
response = _client(FakeSelectionService()).get(
"/api/v1/selection/sectors",
params={"sector_type": "macro"},
)
assert response.status_code == 422
def test_chart_returns_bounded_qfq_contract() -> None:
chart_service = FakeChartService()
@@ -27,6 +27,7 @@ from zhixing_server.modules.sector_radar.domain.models import (
from zhixing_server.modules.sector_radar.domain.persistence import (
MembershipRecord,
RankingRecord,
SectorCountEntry,
)
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
from zhixing_server.modules.sector_radar.infrastructure.memory import (
@@ -315,3 +316,89 @@ def test_stock_membership_query_rejects_invalid_values() -> None:
StockSectorQuery(trade_date=TARGET_DATE, ts_code="x" * 13)
with pytest.raises(ValueError):
StockSectorQuery(trade_date=TARGET_DATE, ts_code="000001.SZ", concept_limit=0)
def test_sector_counts_aggregates_concepts_and_orders_by_count_then_code() -> None:
repository = _membership_repository()
repository.save_memberships(
(
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
_membership_record("000003.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0003.DC", "数字经济"),
)
)
reader = ReadSectorRadar(repository)
snapshot = reader.sector_counts(["000001.SZ", "000002.SZ", "000003.SZ"], TARGET_DATE)
assert snapshot.status == "success"
assert snapshot.trade_date == TARGET_DATE
assert snapshot.sector_type is SectorType.CONCEPT
assert snapshot.counts == (
SectorCountEntry(sector_code="BK0001.DC", sector_name="机器人", stock_count=3),
SectorCountEntry(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2),
SectorCountEntry(sector_code="BK0002.DC", sector_name="人工智能", stock_count=1),
)
def test_sector_counts_supports_industry_type() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_counts(
["000001.SZ", "000002.SZ"],
TARGET_DATE,
sector_type=SectorType.INDUSTRY,
)
assert snapshot.status == "success"
assert snapshot.sector_type is SectorType.INDUSTRY
assert snapshot.counts == (
SectorCountEntry(sector_code="BK0901.DC", sector_name="银行", stock_count=1),
SectorCountEntry(sector_code="BK0902.DC", sector_name="房地产", stock_count=1),
)
def test_sector_counts_without_publication_is_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
snapshot = reader.sector_counts(["000001.SZ"], TARGET_DATE)
assert snapshot.status == "no_data"
assert snapshot.trade_date is None
assert snapshot.counts == ()
def test_sector_counts_with_empty_stock_set_skips_publication_lookup() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_counts([], TARGET_DATE)
assert snapshot.status == "no_data"
assert snapshot.counts == ()
def test_sector_member_codes_returns_available_members() -> None:
reader = ReadSectorRadar(_membership_repository())
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
assert snapshot.status == "success"
assert snapshot.trade_date == TARGET_DATE
assert snapshot.stock_codes == ("000001.SZ",)
def test_sector_member_codes_without_publication_is_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
assert snapshot.status == "no_data"
assert snapshot.trade_date is None
assert snapshot.stock_codes == ()
def test_sector_member_codes_rejects_blank_sector_code() -> None:
reader = ReadSectorRadar(_membership_repository())
with pytest.raises(ValueError):
reader.sector_member_codes(TARGET_DATE, " ")
@@ -136,7 +136,13 @@ class FakeStore:
kwargs["error_message"] = error_message
self.finished = (run_id, status, kwargs)
def get_run(self, run_id: str, *, query: SelectionResultQuery | None = None):
def get_run(
self,
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
):
return None
def get_latest_run(
@@ -145,9 +151,16 @@ class FakeStore:
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
):
return None
def get_run_identity(self, run_id: str):
return None
def get_latest_run_identity(self, strategy: str, target_trade_date: date | None = None):
return None
class BatchStore(FakeStore):
def __init__(self) -> None:
@@ -427,8 +440,8 @@ def test_execute_loads_cases_once_and_scores_only_selected_stocks() -> None:
reader,
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
batch_size=1,
)
@@ -463,8 +476,8 @@ def test_execute_isolates_pattern_scoring_failure_from_selection_status() -> Non
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
)
@@ -497,8 +510,8 @@ def test_execute_skips_pattern_dependencies_when_feature_flag_is_disabled() -> N
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=False,
)
@@ -535,8 +548,8 @@ def test_execute_marks_scores_failed_when_case_library_is_unavailable() -> None:
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_case_loader=loader,
pattern_scorer=scorer,
pattern_scoring_enabled=True,
)
@@ -0,0 +1,264 @@
"""Sector-filter behavior for selection result reads and sector aggregates."""
from collections.abc import Sequence
from dataclasses import replace
from datetime import date
from decimal import Decimal
from zhixing_server.modules.selection.application.run import (
RunZhixingB1,
SelectionSectorAggregates,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionExecutionSource,
SelectionResultQuery,
SelectionRun,
SelectionRunIdentity,
SelectionRunItem,
SelectionRunStatus,
SelectionSectorCount,
SelectionSectorMembership,
)
TARGET = date(2026, 8, 8)
class FakeReader:
def load_execution_source(
self, strategy: str, target_trade_date: date
) -> SelectionExecutionSource:
raise AssertionError("sector reads must not load execution sources")
def load_history(self, ts_code: str, target_trade_date: date):
raise AssertionError("sector reads must not load histories")
class FakeSectorReader:
def __init__(
self,
counts: tuple[SelectionSectorCount, ...] = (),
member_codes: tuple[str, ...] = (),
) -> None:
self.counts = counts
self.member_codes = member_codes
self.count_calls: list[tuple[tuple[str, ...], date, str]] = []
self.member_calls: list[tuple[date, str, str]] = []
def sector_counts(
self,
stock_codes: Sequence[str],
target_trade_date: date,
*,
sector_type: str = "concept",
) -> SelectionSectorMembership:
self.count_calls.append((tuple(stock_codes), target_trade_date, sector_type))
return SelectionSectorMembership(
snapshot_trade_date=target_trade_date,
sector_counts=self.counts,
)
def sector_member_codes(
self,
target_trade_date: date,
sector_code: str,
*,
sector_type: str = "concept",
) -> tuple[str, ...]:
self.member_calls.append((target_trade_date, sector_code, sector_type))
return self.member_codes
class FakeStore:
def __init__(self, run: SelectionRun | None) -> None:
self.run = run
self.sector_codes_seen: Sequence[str] | None = None
self.queries: list[SelectionResultQuery] = []
def prepare_run(self, *args: object, **kwargs: object):
raise AssertionError("sector reads must not prepare runs")
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
raise AssertionError("sector reads must not record items")
def finish_run(self, *args: object, **kwargs: object) -> None:
raise AssertionError("sector reads must not finish runs")
def get_run(
self,
run_id: str,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
self.queries.append(query or SelectionResultQuery())
self.sector_codes_seen = sector_stock_codes
if self.run is not None and self.run.id == run_id:
return self.run
return None
def get_latest_run(
self,
strategy: str,
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
self.queries.append(query or SelectionResultQuery())
self.sector_codes_seen = sector_stock_codes
return self.run
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
if self.run is None or self.run.id != run_id:
return None
return SelectionRunIdentity(
run_id=self.run.id,
target_trade_date=self.run.target_trade_date,
)
def get_latest_run_identity(
self,
strategy: str,
target_trade_date: date | None = None,
) -> SelectionRunIdentity | None:
if self.run is None:
return None
return SelectionRunIdentity(
run_id=self.run.id,
target_trade_date=self.run.target_trade_date,
)
def _run(status: SelectionRunStatus = "success") -> SelectionRun:
return SelectionRun(
id="run-1",
strategy="zhixing_b1",
target_trade_date=TARGET,
market_sync_batch_id="market-run-1",
status=status,
target_count=3,
eligible_count=3,
evaluated_count=3,
selected_stock_count=2,
signal_count=2,
failed_count=0,
coverage=Decimal(1),
items=(
SelectionRunItem(ts_code="000001.SZ", name="A", status="selected", signal_count=1),
SelectionRunItem(ts_code="000002.SZ", name="B", status="selected", signal_count=1),
SelectionRunItem(ts_code="000003.SZ", name="C", status="no_signal", signal_count=0),
),
)
def test_list_sector_counts_aggregates_only_selected_stocks() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(
counts=(SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=2),)
)
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1")
assert isinstance(aggregates, SelectionSectorAggregates)
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "concept")]
assert aggregates.snapshot_trade_date == TARGET
assert aggregates.sector_type == "concept"
assert aggregates.sectors == sector_reader.counts
assert aggregates.run is store.run
def test_list_sector_counts_without_run_returns_none() -> None:
service = RunZhixingB1(FakeReader(), FakeStore(None), sector_reader=FakeSectorReader())
assert service.list_sector_counts("zhixing_b1") is None
def test_list_sector_counts_without_selected_stocks_skips_reader() -> None:
run = replace(_run(), items=())
sector_reader = FakeSectorReader()
service = RunZhixingB1(FakeReader(), FakeStore(run), sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1")
assert sector_reader.count_calls == []
assert aggregates is not None
assert aggregates.snapshot_trade_date is None
assert aggregates.sectors == ()
def test_get_latest_resolves_sector_filter_against_run_snapshot() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(member_codes=("000001.SZ", "000002.SZ"))
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
run = service.get_latest(
"zhixing_b1",
query=SelectionResultQuery(page=2, page_size=20, sector="BK0001.DC"),
)
assert run is store.run
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
assert store.sector_codes_seen == ("000001.SZ", "000002.SZ")
assert store.queries[-1].sector is None
assert store.queries[-1].page == 2
def test_get_latest_with_empty_sector_members_returns_empty_codes_page() -> None:
store = FakeStore(_run())
service = RunZhixingB1(
FakeReader(),
store,
sector_reader=FakeSectorReader(member_codes=()),
)
service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK9999.DC"))
assert store.sector_codes_seen == ()
def test_get_latest_without_sector_reader_ignores_the_filter() -> None:
store = FakeStore(_run())
service = RunZhixingB1(FakeReader(), store)
run = service.get_latest(
"zhixing_b1",
query=SelectionResultQuery(sector="BK0001.DC"),
)
assert run is store.run
assert store.sector_codes_seen is None
assert store.queries[-1].sector == "BK0001.DC"
def test_get_run_resolves_sector_filter_against_run_snapshot() -> None:
store = FakeStore(_run())
sector_reader = FakeSectorReader(member_codes=("000001.SZ",))
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
run = service.get_run("run-1", query=SelectionResultQuery(sector="BK0001.DC"))
assert run is store.run
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
assert store.sector_codes_seen == ("000001.SZ",)
def test_get_latest_without_run_and_sector_filter_returns_none() -> None:
service = RunZhixingB1(
FakeReader(),
FakeStore(None),
sector_reader=FakeSectorReader(member_codes=()),
)
assert service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK0001.DC")) is None
def test_list_sector_counts_accepts_industry_type() -> None:
sector_reader = FakeSectorReader()
service = RunZhixingB1(FakeReader(), FakeStore(_run()), sector_reader=sector_reader)
aggregates = service.list_sector_counts("zhixing_b1", sector_type="industry")
assert aggregates is not None
assert aggregates.sector_type == "industry"
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "industry")]
@@ -7,6 +7,7 @@ import {
type SelectionResultsQuery,
type SelectionRunAccepted,
type SelectionRunRequest,
type SelectionSectors,
type SelectionStrategy,
} from "./selection.types"
@@ -63,9 +64,27 @@ function buildSelectionQueryParams(query: SelectionResultsQuery) {
if (query.search) params.set("search", query.search)
if (query.category) params.set("category", query.category)
if (query.sort) params.set("sort", query.sort)
if (query.sector) params.set("sector", query.sector)
return params
}
export function getSelectionResultSectors(
strategy: SelectionStrategy,
targetTradeDate?: string,
sectorType: "concept" | "industry" = "concept",
signal?: AbortSignal,
) {
const params = new URLSearchParams({
strategy,
sector_type: sectorType,
})
if (targetTradeDate) params.set("target_trade_date", targetTradeDate)
return requestJson<SelectionSectors>(
`/api/v1/selection/sectors?${params.toString()}`,
{ signal },
)
}
export function triggerSelectionRun(
request: SelectionRunRequest,
signal?: AbortSignal,
@@ -69,6 +69,7 @@ describe("selection query hooks", () => {
"平安",
"pullback",
"score_desc",
"",
],
)
expect(
@@ -10,6 +10,7 @@ import { useEffect } from "react"
import {
getSelectionChart,
getSelectionResultSectors,
getSelectionResults,
getSelectionRun,
triggerSelectionRun,
@@ -19,6 +20,7 @@ import {
type SelectionResults,
type SelectionResultsQuery,
type SelectionRunRequest,
type SelectionSectors,
type SelectionStrategy,
} from "./selection.types"
@@ -55,6 +57,7 @@ export const selectionResultsQueryKey = (
query.search ?? "",
query.category ?? "all",
query.sort ?? "code",
query.sector ?? "",
] as const
export const selectionRunQueryKey = (
@@ -69,8 +72,15 @@ export const selectionRunQueryKey = (
query.search ?? "",
query.category ?? "all",
query.sort ?? "code",
query.sector ?? "",
] as const
export const selectionSectorsQueryKey = (
strategy: SelectionStrategy,
targetTradeDate?: string,
sectorType: "concept" | "industry" = "concept",
) => ["selection", "sectors", strategy, targetTradeDate, sectorType] as const
export function flattenSelectionResults(
data: InfiniteData<SelectionResults> | undefined,
): SelectionResults | undefined {
@@ -120,6 +130,18 @@ export function useSelectionChart(
})
}
export function useSelectionResultSectors(
strategy: SelectionStrategy,
targetTradeDate?: string,
sectorType: "concept" | "industry" = "concept",
) {
return useQuery<SelectionSectors, Error>({
queryFn: ({ signal }) =>
getSelectionResultSectors(strategy, targetTradeDate, sectorType, signal),
queryKey: selectionSectorsQueryKey(strategy, targetTradeDate, sectorType),
})
}
export function useSelectionRun(
runId: string | null,
resultQuery: SelectionResultsListQuery = defaultListQuery,
@@ -173,9 +195,14 @@ export function invalidateSelectionResults(
queryClient: QueryClient,
strategy: SelectionStrategy,
) {
return queryClient.invalidateQueries({
return Promise.all([
queryClient.invalidateQueries({
queryKey: ["selection", "results", strategy],
})
}),
queryClient.invalidateQueries({
queryKey: ["selection", "sectors", strategy],
}),
])
}
function nextSelectionPageParam(
@@ -25,6 +25,27 @@ export interface SelectionResultsQuery {
search?: string
category?: Exclude<SelectionCategoryFilter, "all">
sort?: SelectionSort
sector?: string
}
export type SelectionSectorType = "concept" | "industry"
export const selectionSectorTypes = ["concept", "industry"] as const
export interface SelectionSectorAggregate {
sector_code: string
sector_name: string
stock_count: number
}
export interface SelectionSectors {
strategy: SelectionStrategy
target_trade_date: string | null
run_id: string | null
status: SelectionRunStatus
snapshot_trade_date: string | null
sector_type: SelectionSectorType
sectors: SelectionSectorAggregate[]
}
export type SelectionRunStatus =
@@ -1,4 +1,4 @@
import { useMemo, useRef, useState } from "react"
import { useEffect, useMemo, useRef, useState } from "react"
import { useNavigate, useSearch } from "@tanstack/react-router"
import { Input } from "@/shared/ui/input"
@@ -11,7 +11,11 @@ import {
SelectValue,
} from "@/shared/ui/select"
import type { SelectionResults, SelectionSort } from "../api/selection.types"
import type {
SelectionResults,
SelectionSectorAggregate,
SelectionSort,
} from "../api/selection.types"
import { SignalDetailPanel } from "./signal-detail-panel"
import { SignalRecordList } from "./signal-record-list"
import {
@@ -30,12 +34,15 @@ const SCORE_SORT_OPTIONS: ReadonlyArray<{
{ label: "评分从低到高", value: "score_asc" },
]
const ALL_SECTORS_VALUE = "all"
interface SelectionResultsWorkbenchProps {
hasNextPage: boolean
isFetchNextPageError: boolean
isFetchingNextPage: boolean
onLoadMore: () => void | Promise<unknown>
result: SelectionResults
sectorOptions: ReadonlyArray<SelectionSectorAggregate>
}
export function SelectionResultsWorkbench({
@@ -44,12 +51,14 @@ export function SelectionResultsWorkbench({
isFetchingNextPage,
onLoadMore,
result,
sectorOptions,
}: SelectionResultsWorkbenchProps) {
const search = useSearch({ from: "/_workspace/selection" })
const navigate = useNavigate({ from: "/selection" })
const query = search.search ?? ""
const category = search.category ?? "all"
const sort = search.sort ?? "code"
const sector = search.sector
const [selectedKey, setSelectedKey] = useState<string | null>(null)
const loadMoreRequestPending = useRef(false)
@@ -68,11 +77,32 @@ export function SelectionResultsWorkbench({
const selectedStock =
visibleStocks.find((stock) => getStockKey(stock) === selectedKey) ??
visibleStocks[0]
const sectorSelectItems = useMemo(
() => [
{ label: "全部板块", value: ALL_SECTORS_VALUE },
...sectorOptions.map((option) => ({
label: option.sector_name,
value: option.sector_code,
})),
],
[sectorOptions],
)
// The aggregate list is the source of truth for valid sectors; drop a stale
// value left over from another strategy or trade date.
useEffect(() => {
if (!sector) return
if (sectorOptions.some((option) => option.sector_code === sector)) return
void navigate({
search: (previous) => ({ ...previous, sector: undefined }),
})
}, [navigate, sector, sectorOptions])
function updateSearch(next: {
search?: string
category?: SignalCategoryFilter
sort?: SelectionSort
sector?: string
}) {
void navigate({ search: (previous) => ({ ...previous, ...next }) })
}
@@ -89,6 +119,10 @@ export function SelectionResultsWorkbench({
updateSearch({ sort: value })
}
function handleSectorChange(value: string) {
updateSearch({ sector: value === ALL_SECTORS_VALUE ? undefined : value })
}
function requestLoadMore() {
if (
loadMoreRequestPending.current ||
@@ -105,9 +139,9 @@ export function SelectionResultsWorkbench({
}
return (
<div className="flex min-h-0 min-w-0 flex-1 flex-col gap-2 overflow-hidden md:grid md:grid-cols-[320px_minmax(0,1fr)]">
<section className="flex min-h-0 min-w-0 flex-1 flex-col overflow-hidden rounded-md border border-border/80 bg-card">
<div className="flex shrink-0 flex-wrap items-center gap-1.5 border-b border-border/60 p-2.5">
<div className="flex min-h-0 min-w-0 flex-1 flex-col gap-2 overflow-hidden">
<section className="shrink-0 rounded-md border border-border/80 bg-card">
<div className="flex flex-wrap items-center gap-1.5 p-2.5">
<Input
aria-label="搜索命中股票"
className="h-11 min-w-0 flex-1 bg-background text-sm sm:max-w-[260px] md:h-8"
@@ -116,6 +150,38 @@ export function SelectionResultsWorkbench({
type="search"
value={query}
/>
<Select
items={sectorSelectItems}
onValueChange={(value) => {
if (typeof value === "string") handleSectorChange(value)
}}
value={sector ?? ALL_SECTORS_VALUE}
>
<SelectTrigger
aria-label="筛选板块"
className="h-11 w-full bg-background text-sm sm:w-auto sm:min-w-36 md:h-8"
>
<SelectValue placeholder="筛选板块" />
</SelectTrigger>
<SelectContent>
<SelectGroup>
<SelectItem value={ALL_SECTORS_VALUE}>全部板块</SelectItem>
{sectorOptions.map((option) => (
<SelectItem
key={option.sector_code}
value={option.sector_code}
>
<span className="flex w-full items-center justify-between gap-3">
<span>{option.sector_name}</span>
<span className="tabular-nums text-muted-foreground">
{option.stock_count}
</span>
</span>
</SelectItem>
))}
</SelectGroup>
</SelectContent>
</Select>
<Select
items={signalCategoryOptions}
onValueChange={(value) => {
@@ -170,7 +236,10 @@ export function SelectionResultsWorkbench({
筛选结果 {stocksTotal} 只
</span>
</div>
</section>
<div className="flex min-h-0 min-w-0 flex-1 flex-col gap-2 overflow-hidden md:grid md:grid-cols-[320px_minmax(0,1fr)]">
<section className="flex min-h-0 min-w-0 flex-1 flex-col overflow-hidden rounded-md border border-border/80 bg-card">
<div className="flex min-h-0 min-w-0 flex-1 flex-col overflow-hidden">
<SignalRecordList
isFetchNextPageError={isFetchNextPageError}
@@ -189,5 +258,6 @@ export function SelectionResultsWorkbench({
<SignalDetailPanel stock={selectedStock} />
</div>
</div>
)
}
@@ -16,6 +16,7 @@ import { SelectionResultsPage } from "./selection-results-page"
const routerNavigate = vi.hoisted(() => vi.fn())
const useSelectionResults = vi.fn()
const useSelectionResultSectors = vi.fn()
const useSelectionChart = vi.fn()
const useStockSectorMembership = vi.fn()
const useSelectionRun = vi.fn()
@@ -31,6 +32,8 @@ vi.mock("@/features/selection/api/selection.query", async () => {
return {
...actual,
useSelectionResults: (...args: unknown[]) => useSelectionResults(...args),
useSelectionResultSectors: (...args: unknown[]) =>
useSelectionResultSectors(...args),
useSelectionChart: (...args: unknown[]) => useSelectionChart(...args),
useSelectionRun: (...args: unknown[]) => useSelectionRun(...args),
useTriggerSelectionRun: () => useTriggerSelectionRun(),
@@ -169,6 +172,11 @@ describe("SelectionResultsPage", () => {
isPending: false,
})
useSelectionResults.mockReturnValue(selectionQueryResult(selectedResult))
useSelectionResultSectors.mockReturnValue({
data: undefined,
isError: false,
isPending: false,
})
useSelectionRun.mockReturnValue(emptySelectionQuery())
useTriggerSelectionRun.mockReturnValue({
isError: false,
@@ -535,6 +543,45 @@ describe("SelectionResultsPage", () => {
})
})
it("filters by sector and shows per-sector counts", async () => {
useSelectionResultSectors.mockReturnValue({
data: {
strategy: "zhixing_b1",
target_trade_date: "2026-08-08",
run_id: "run-1",
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 },
],
},
isError: false,
isPending: false,
})
render(<SelectionResultsPage />)
const sectorTrigger = screen.getByRole("combobox", { name: "筛选板块" })
expect(sectorTrigger).toHaveTextContent("全部板块")
fireEvent.click(sectorTrigger)
const robotOption = await screen.findByRole("option", {
name: /机器人/,
})
expect(robotOption).toHaveTextContent("3")
fireEvent.pointerDown(robotOption, { pointerType: "mouse" })
fireEvent.click(robotOption)
expect(routerNavigate).toHaveBeenCalled()
const lastCall =
routerNavigate.mock.calls[routerNavigate.mock.calls.length - 1]
const searchUpdate = lastCall[0].search as (previous: {
sector?: string
}) => Record<string, unknown>
expect(searchUpdate({})).toEqual({ sector: "BK0001.DC" })
})
it("requests database-backed score sorting", async () => {
render(<SelectionResultsPage />)
@@ -5,6 +5,7 @@ import { useRef, useState, type RefObject } from "react"
import { PageLayout } from "@/app/layout/page-layout"
import {
flattenSelectionResults,
useSelectionResultSectors,
useSelectionResults,
useSelectionRun,
useTriggerSelectionRun,
@@ -14,6 +15,8 @@ import {
selectionStrategies,
type SelectionResults,
type SelectionResultsQuery,
type SelectionSectorAggregate,
type SelectionSectors,
type SelectionStrategy,
} from "@/features/selection/api/selection.types"
import { Button } from "@/shared/ui/button"
@@ -71,6 +74,7 @@ export function SelectionResultsPage() {
pageSize: selectionResultPageSize,
...(search.search ? { search: search.search } : {}),
...(search.category !== "all" ? { category: search.category } : {}),
...(search.sector ? { sector: search.sector } : {}),
sort: search.sort,
}
@@ -79,6 +83,10 @@ export function SelectionResultsPage() {
targetTradeDate || undefined,
resultQuery,
)
const sectorsQuery = useSelectionResultSectors(
strategy,
targetTradeDate || undefined,
)
const resultsSnapshot = flattenSelectionResults(results.data)
const persistedRunningRunId =
resultsSnapshot?.status === "running" ? resultsSnapshot.run_id : null
@@ -156,7 +164,11 @@ export function SelectionResultsPage() {
setRerunDialogOpen(false)
trigger.reset()
void navigate({
search: (previous) => ({ ...previous, category: "all" }),
search: (previous) => ({
...previous,
category: "all",
sector: undefined,
}),
})
}
@@ -212,6 +224,7 @@ export function SelectionResultsPage() {
isFetchingNextPage={Boolean(listQuery.isFetchingNextPage)}
onLoadMore={() => listQuery.fetchNextPage()}
result={displayedResult}
sectors={sectorsQuery.data}
/>
) : null}
</div>
@@ -376,8 +389,11 @@ interface ResultStateProps {
isFetchingNextPage: boolean
onLoadMore: () => void | Promise<unknown>
result: SelectionResults
sectors: SelectionSectors | undefined
}
const EMPTY_SECTOR_OPTIONS: ReadonlyArray<SelectionSectorAggregate> = []
function parseLocalTradeDate(value: string): Date | undefined {
const match = /^(\d{4})-(\d{2})-(\d{2})$/.exec(value)
if (!match) return undefined
@@ -511,6 +527,7 @@ function ResultState({
isFetchingNextPage,
onLoadMore,
result,
sectors,
}: ResultStateProps) {
return result.signal_count === 0 ? (
<NoSignalState />
@@ -521,6 +538,7 @@ function ResultState({
isFetchingNextPage={isFetchingNextPage}
onLoadMore={onLoadMore}
result={result}
sectorOptions={sectors?.sectors ?? EMPTY_SECTOR_OPTIONS}
/>
)
}
+5 -1
View File
@@ -124,7 +124,11 @@ const selectionRoute = createRoute({
)
? (rawSort as SelectionSort)
: "code"
return { page, pageSize, search: searchValue, category, sort }
const sector =
typeof search.sector === "string" && search.sector.trim()
? search.sector.trim().slice(0, 60)
: undefined
return { page, pageSize, search: searchValue, category, sort, sector }
},
component: SelectionResultsPage,
})