This commit is contained in:
@@ -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 __future__ import annotations
|
||||||
|
|
||||||
from collections.abc import Iterator
|
from collections.abc import Iterator, Sequence
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
@@ -21,7 +21,11 @@ from ..domain.models import (
|
|||||||
RankSide,
|
RankSide,
|
||||||
SectorType,
|
SectorType,
|
||||||
)
|
)
|
||||||
from ..domain.persistence import SectorRadarRepository, StockMembershipEntry
|
from ..domain.persistence import (
|
||||||
|
SectorCountEntry,
|
||||||
|
SectorRadarRepository,
|
||||||
|
StockMembershipEntry,
|
||||||
|
)
|
||||||
from ..domain.ranking import select_percentile_side, select_rank_change_side
|
from ..domain.ranking import select_percentile_side, select_rank_change_side
|
||||||
|
|
||||||
ReadStatus = Literal["success", "no_data"]
|
ReadStatus = Literal["success", "no_data"]
|
||||||
@@ -148,6 +152,26 @@ class StockSectorMembership:
|
|||||||
concept_total: int
|
concept_total: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SectorCountsSnapshot:
|
||||||
|
"""Sector counts for a bounded stock set on one last-good snapshot date."""
|
||||||
|
|
||||||
|
status: ReadStatus
|
||||||
|
sector_type: SectorType
|
||||||
|
trade_date: date | None
|
||||||
|
counts: tuple[SectorCountEntry, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SectorMembersSnapshot:
|
||||||
|
"""One sector's member stock codes on one last-good snapshot date."""
|
||||||
|
|
||||||
|
status: ReadStatus
|
||||||
|
sector_type: SectorType
|
||||||
|
trade_date: date | None
|
||||||
|
stock_codes: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
_METRIC_DEFINITIONS = {
|
_METRIC_DEFINITIONS = {
|
||||||
MetricKind.AMOUNT: RadarMetricDefinition(
|
MetricKind.AMOUNT: RadarMetricDefinition(
|
||||||
metric_kind=MetricKind.AMOUNT,
|
metric_kind=MetricKind.AMOUNT,
|
||||||
@@ -290,6 +314,78 @@ class ReadSectorRadar:
|
|||||||
concept_total=len(concepts),
|
concept_total=len(concepts),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def sector_counts(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
*,
|
||||||
|
sector_type: SectorType = SectorType.CONCEPT,
|
||||||
|
) -> SectorCountsSnapshot:
|
||||||
|
"""Count claimed stocks per sector from the strict last-good snapshot."""
|
||||||
|
|
||||||
|
if not stock_codes:
|
||||||
|
return SectorCountsSnapshot(
|
||||||
|
status="no_data",
|
||||||
|
sector_type=sector_type,
|
||||||
|
trade_date=None,
|
||||||
|
counts=(),
|
||||||
|
)
|
||||||
|
publication = self.repository.get_last_good_publication(target_trade_date)
|
||||||
|
if publication is None:
|
||||||
|
return SectorCountsSnapshot(
|
||||||
|
status="no_data",
|
||||||
|
sector_type=sector_type,
|
||||||
|
trade_date=None,
|
||||||
|
counts=(),
|
||||||
|
)
|
||||||
|
counts = tuple(
|
||||||
|
self.repository.load_sector_counts(
|
||||||
|
publication.target_trade_date,
|
||||||
|
stock_codes,
|
||||||
|
sector_type,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return SectorCountsSnapshot(
|
||||||
|
status="success",
|
||||||
|
sector_type=sector_type,
|
||||||
|
trade_date=publication.target_trade_date,
|
||||||
|
counts=counts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
*,
|
||||||
|
sector_type: SectorType = SectorType.CONCEPT,
|
||||||
|
) -> SectorMembersSnapshot:
|
||||||
|
"""Return one sector's members from the strict last-good snapshot."""
|
||||||
|
|
||||||
|
normalized_code = sector_code.strip()
|
||||||
|
if not normalized_code:
|
||||||
|
raise ValueError("sector_code must not be empty")
|
||||||
|
publication = self.repository.get_last_good_publication(target_trade_date)
|
||||||
|
if publication is None:
|
||||||
|
return SectorMembersSnapshot(
|
||||||
|
status="no_data",
|
||||||
|
sector_type=sector_type,
|
||||||
|
trade_date=None,
|
||||||
|
stock_codes=(),
|
||||||
|
)
|
||||||
|
codes = tuple(
|
||||||
|
self.repository.load_sector_member_codes(
|
||||||
|
publication.target_trade_date,
|
||||||
|
normalized_code,
|
||||||
|
sector_type,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return SectorMembersSnapshot(
|
||||||
|
status="success",
|
||||||
|
sector_type=sector_type,
|
||||||
|
trade_date=publication.target_trade_date,
|
||||||
|
stock_codes=codes,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _sector_refs(entries: Iterator[StockMembershipEntry]) -> tuple[SectorRef, ...]:
|
def _sector_refs(entries: Iterator[StockMembershipEntry]) -> tuple[SectorRef, ...]:
|
||||||
"""Project membership entries into ordered public sector references."""
|
"""Project membership entries into ordered public sector references."""
|
||||||
@@ -309,6 +405,8 @@ __all__ = [
|
|||||||
"RadarView",
|
"RadarView",
|
||||||
"RankingPage",
|
"RankingPage",
|
||||||
"ReadSectorRadar",
|
"ReadSectorRadar",
|
||||||
|
"SectorCountsSnapshot",
|
||||||
|
"SectorMembersSnapshot",
|
||||||
"SectorRef",
|
"SectorRef",
|
||||||
"StockSectorMembership",
|
"StockSectorMembership",
|
||||||
"StockSectorQuery",
|
"StockSectorQuery",
|
||||||
|
|||||||
@@ -81,6 +81,23 @@ class StockMembershipEntry:
|
|||||||
raise ValueError("membership sector identity fields must not be empty")
|
raise ValueError("membership sector identity fields must not be empty")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SectorCountEntry:
|
||||||
|
"""One sector and the number of stocks it claimed on one snapshot date."""
|
||||||
|
|
||||||
|
sector_code: str
|
||||||
|
sector_name: str
|
||||||
|
stock_count: int
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
"""Reject empty sector identity and non-positive counts."""
|
||||||
|
|
||||||
|
if not self.sector_code.strip() or not self.sector_name.strip():
|
||||||
|
raise ValueError("sector count identity fields must not be empty")
|
||||||
|
if self.stock_count < 1:
|
||||||
|
raise ValueError("sector stock_count must be positive")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class StockFactRecord:
|
class StockFactRecord:
|
||||||
"""One normalized stock fact revision with all contributing raw snapshots."""
|
"""One normalized stock fact revision with all contributing raw snapshots."""
|
||||||
@@ -217,6 +234,20 @@ class SectorRadarRepository(Protocol):
|
|||||||
self, trade_date: date, stock_code: str
|
self, trade_date: date, stock_code: str
|
||||||
) -> Sequence[StockMembershipEntry]: ...
|
) -> Sequence[StockMembershipEntry]: ...
|
||||||
|
|
||||||
|
def load_sector_counts(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[SectorCountEntry]: ...
|
||||||
|
|
||||||
|
def load_sector_member_codes(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[str]: ...
|
||||||
|
|
||||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
|
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
|
||||||
|
|
||||||
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
||||||
|
|||||||
@@ -7,13 +7,20 @@ from contextlib import contextmanager
|
|||||||
from dataclasses import replace
|
from dataclasses import replace
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
|
|
||||||
from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate
|
from ..domain.models import (
|
||||||
|
PublicationStatus,
|
||||||
|
RadarPublication,
|
||||||
|
RankedMetric,
|
||||||
|
SectorDailyAggregate,
|
||||||
|
SectorType,
|
||||||
|
)
|
||||||
from ..domain.persistence import (
|
from ..domain.persistence import (
|
||||||
DailyAggregateRecord,
|
DailyAggregateRecord,
|
||||||
MembershipRecord,
|
MembershipRecord,
|
||||||
PublicationSourceGroup,
|
PublicationSourceGroup,
|
||||||
PublicationSourceRecord,
|
PublicationSourceRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
StockFactRecord,
|
StockFactRecord,
|
||||||
StockMembershipEntry,
|
StockMembershipEntry,
|
||||||
WriteCounts,
|
WriteCounts,
|
||||||
@@ -137,6 +144,58 @@ class InMemorySectorRadarRepository:
|
|||||||
and record.status.value == "available"
|
and record.status.value == "available"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def load_sector_counts(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[SectorCountEntry]:
|
||||||
|
"""Aggregate claimed-stock counts per sector on one snapshot date."""
|
||||||
|
|
||||||
|
wanted = set(stock_codes)
|
||||||
|
counts: dict[tuple[str, str], int] = {}
|
||||||
|
for record in self.memberships.values():
|
||||||
|
if (
|
||||||
|
record.trade_date == trade_date
|
||||||
|
and record.sector_type is sector_type
|
||||||
|
and record.status.value == "available"
|
||||||
|
and record.stock_code in wanted
|
||||||
|
):
|
||||||
|
key = (record.sector_code, record.sector_name)
|
||||||
|
counts[key] = counts.get(key, 0) + 1
|
||||||
|
return tuple(
|
||||||
|
SectorCountEntry(
|
||||||
|
sector_code=sector_code,
|
||||||
|
sector_name=sector_name,
|
||||||
|
stock_count=stock_count,
|
||||||
|
)
|
||||||
|
for (sector_code, sector_name), stock_count in sorted(
|
||||||
|
counts.items(),
|
||||||
|
key=lambda item: (-item[1], item[0][0]),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_sector_member_codes(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[str]:
|
||||||
|
"""Return available member stock codes of one sector on one date."""
|
||||||
|
|
||||||
|
return tuple(
|
||||||
|
record.stock_code
|
||||||
|
for record in sorted(
|
||||||
|
self.memberships.values(),
|
||||||
|
key=lambda item: item.stock_code or "",
|
||||||
|
)
|
||||||
|
if record.trade_date == trade_date
|
||||||
|
and record.sector_type is sector_type
|
||||||
|
and record.sector_code == sector_code
|
||||||
|
and record.status.value == "available"
|
||||||
|
and record.stock_code is not None
|
||||||
|
)
|
||||||
|
|
||||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
||||||
"""Insert normalized fact revisions idempotently."""
|
"""Insert normalized fact revisions idempotently."""
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ from ..domain.persistence import (
|
|||||||
PublicationSourceGroup,
|
PublicationSourceGroup,
|
||||||
PublicationSourceRecord,
|
PublicationSourceRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
StockFactRecord,
|
StockFactRecord,
|
||||||
StockMembershipEntry,
|
StockMembershipEntry,
|
||||||
WriteCounts,
|
WriteCounts,
|
||||||
@@ -310,6 +311,60 @@ class PostgresSectorRadarRepository:
|
|||||||
for row in rows
|
for row in rows
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def load_sector_counts(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[SectorCountEntry]:
|
||||||
|
"""Aggregate how many of the given stocks each sector claimed on one date."""
|
||||||
|
|
||||||
|
if not stock_codes:
|
||||||
|
return ()
|
||||||
|
with self._connection() as connection:
|
||||||
|
rows = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT sector_code, sector_name, COUNT(*) AS stock_count
|
||||||
|
FROM sector_radar_membership
|
||||||
|
WHERE trade_date = %s AND sector_type = %s
|
||||||
|
AND membership_status = 'available'
|
||||||
|
AND stock_code = ANY(%s)
|
||||||
|
GROUP BY sector_code, sector_name
|
||||||
|
ORDER BY stock_count DESC, sector_code
|
||||||
|
""",
|
||||||
|
(trade_date, sector_type.value, list(stock_codes)),
|
||||||
|
).fetchall()
|
||||||
|
return tuple(
|
||||||
|
SectorCountEntry(
|
||||||
|
sector_code=row[0],
|
||||||
|
sector_name=row[1],
|
||||||
|
stock_count=int(row[2]),
|
||||||
|
)
|
||||||
|
for row in rows
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_sector_member_codes(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
sector_type: SectorType,
|
||||||
|
) -> Sequence[str]:
|
||||||
|
"""Load every stock one sector claimed on exactly one snapshot date."""
|
||||||
|
|
||||||
|
with self._connection() as connection:
|
||||||
|
rows = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT stock_code
|
||||||
|
FROM sector_radar_membership
|
||||||
|
WHERE trade_date = %s AND sector_type = %s
|
||||||
|
AND sector_code = %s
|
||||||
|
AND membership_status = 'available'
|
||||||
|
ORDER BY stock_code
|
||||||
|
""",
|
||||||
|
(trade_date, sector_type.value, sector_code),
|
||||||
|
).fetchall()
|
||||||
|
return tuple(str(row[0]) for row in rows)
|
||||||
|
|
||||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
||||||
"""COPY normalized stock facts while preserving contributing source ids."""
|
"""COPY normalized stock facts while preserving contributing source ids."""
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import time
|
|||||||
from collections import Counter
|
from collections import Counter
|
||||||
from collections.abc import Callable, Mapping, Sequence
|
from collections.abc import Callable, Mapping, Sequence
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, replace
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from typing import Protocol, cast
|
from typing import Protocol, cast
|
||||||
|
|
||||||
@@ -29,6 +29,9 @@ from ..domain.runs import (
|
|||||||
SelectionRunItem,
|
SelectionRunItem,
|
||||||
SelectionRunStatus,
|
SelectionRunStatus,
|
||||||
SelectionRunStore,
|
SelectionRunStore,
|
||||||
|
SelectionSectorCount,
|
||||||
|
SelectionSectorMembership,
|
||||||
|
SelectionSectorReader,
|
||||||
SelectionStock,
|
SelectionStock,
|
||||||
SelectionUniverseReader,
|
SelectionUniverseReader,
|
||||||
)
|
)
|
||||||
@@ -60,6 +63,16 @@ class PreparedSelectionRun:
|
|||||||
source: SelectionExecutionSource
|
source: SelectionExecutionSource
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SelectionSectorAggregates:
|
||||||
|
"""A run summary plus its selected stocks' sector membership counts."""
|
||||||
|
|
||||||
|
run: SelectionRun | None
|
||||||
|
snapshot_trade_date: date | None
|
||||||
|
sector_type: str
|
||||||
|
sectors: tuple[SelectionSectorCount, ...]
|
||||||
|
|
||||||
|
|
||||||
class RunZhixingB1:
|
class RunZhixingB1:
|
||||||
"""Prepare, execute, and query persisted selection strategy batches.
|
"""Prepare, execute, and query persisted selection strategy batches.
|
||||||
|
|
||||||
@@ -75,6 +88,7 @@ class RunZhixingB1:
|
|||||||
evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None,
|
evaluators: Mapping[StrategyName, SelectionEvaluator] | None = None,
|
||||||
pattern_case_loader: PatternCaseLibraryLoader | None = None,
|
pattern_case_loader: PatternCaseLibraryLoader | None = None,
|
||||||
pattern_scorer: PatternScorer | None = None,
|
pattern_scorer: PatternScorer | None = None,
|
||||||
|
sector_reader: SelectionSectorReader | None = None,
|
||||||
*,
|
*,
|
||||||
pattern_scoring_enabled: bool = False,
|
pattern_scoring_enabled: bool = False,
|
||||||
max_workers: int = 4,
|
max_workers: int = 4,
|
||||||
@@ -96,6 +110,7 @@ class RunZhixingB1:
|
|||||||
self.evaluators.update(evaluators)
|
self.evaluators.update(evaluators)
|
||||||
self.pattern_case_loader = pattern_case_loader
|
self.pattern_case_loader = pattern_case_loader
|
||||||
self.pattern_scorer = pattern_scorer
|
self.pattern_scorer = pattern_scorer
|
||||||
|
self.sector_reader = sector_reader
|
||||||
self.pattern_scoring_enabled = pattern_scoring_enabled
|
self.pattern_scoring_enabled = pattern_scoring_enabled
|
||||||
self.max_workers = max_workers
|
self.max_workers = max_workers
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
@@ -211,8 +226,7 @@ class RunZhixingB1:
|
|||||||
)
|
)
|
||||||
history_rows += batch_history_rows
|
history_rows += batch_history_rows
|
||||||
batch_missing_turnover = sum(
|
batch_missing_turnover = sum(
|
||||||
not _turnover_present(history, target_trade_date)
|
not _turnover_present(history, target_trade_date) for history in histories
|
||||||
for history in histories
|
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"selection_read_batch_summary strategy=%s target_trade_date=%s "
|
"selection_read_batch_summary strategy=%s target_trade_date=%s "
|
||||||
@@ -581,9 +595,22 @@ class RunZhixingB1:
|
|||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read one persisted run for polling."""
|
"""Read one persisted run for polling with optional sector filtering."""
|
||||||
|
|
||||||
return self.store.get_run(run_id, query=query)
|
effective = query or SelectionResultQuery()
|
||||||
|
if not effective.sector:
|
||||||
|
return self.store.get_run(run_id, query=query)
|
||||||
|
identity = self.store.get_run_identity(run_id)
|
||||||
|
if identity is None:
|
||||||
|
return None
|
||||||
|
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||||
|
if member_codes is None:
|
||||||
|
return self.store.get_run(run_id, query=query)
|
||||||
|
return self.store.get_run(
|
||||||
|
run_id,
|
||||||
|
query=replace(effective, sector=None),
|
||||||
|
sector_stock_codes=member_codes,
|
||||||
|
)
|
||||||
|
|
||||||
def get_latest(
|
def get_latest(
|
||||||
self,
|
self,
|
||||||
@@ -594,7 +621,73 @@ class RunZhixingB1:
|
|||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read the current result by date or the latest result for a strategy."""
|
"""Read the current result by date or the latest result for a strategy."""
|
||||||
|
|
||||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
effective = query or SelectionResultQuery()
|
||||||
|
if not effective.sector:
|
||||||
|
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||||
|
identity = self.store.get_latest_run_identity(strategy, target_trade_date)
|
||||||
|
if identity is None:
|
||||||
|
return None
|
||||||
|
member_codes = self._sector_member_codes(identity.target_trade_date, effective.sector)
|
||||||
|
if member_codes is None:
|
||||||
|
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||||
|
return self.store.get_latest_run(
|
||||||
|
strategy,
|
||||||
|
identity.target_trade_date,
|
||||||
|
query=replace(effective, sector=None),
|
||||||
|
sector_stock_codes=member_codes,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
) -> tuple[str, ...] | None:
|
||||||
|
"""Resolve one sector's members, or None when the port is absent."""
|
||||||
|
|
||||||
|
if self.sector_reader is None:
|
||||||
|
return None
|
||||||
|
return self.sector_reader.sector_member_codes(target_trade_date, sector_code)
|
||||||
|
|
||||||
|
def list_sector_counts(
|
||||||
|
self,
|
||||||
|
strategy: StrategyName,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorAggregates | None:
|
||||||
|
"""Aggregate the current run's selected stocks by point-in-time sector."""
|
||||||
|
|
||||||
|
run = self.store.get_latest_run(strategy, target_trade_date)
|
||||||
|
if run is None:
|
||||||
|
return None
|
||||||
|
selected_codes = [
|
||||||
|
item.ts_code
|
||||||
|
for item in run.items
|
||||||
|
if item.status == "selected" and item.signal_count > 0
|
||||||
|
]
|
||||||
|
membership = self._sector_membership(selected_codes, run.target_trade_date, sector_type)
|
||||||
|
return SelectionSectorAggregates(
|
||||||
|
run=run,
|
||||||
|
snapshot_trade_date=membership.snapshot_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
sectors=membership.sector_counts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _sector_membership(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_type: str,
|
||||||
|
) -> SelectionSectorMembership:
|
||||||
|
"""Read sector counts for a stock set, tolerating a missing port."""
|
||||||
|
|
||||||
|
if self.sector_reader is None or not stock_codes:
|
||||||
|
return SelectionSectorMembership(snapshot_trade_date=None, sector_counts=())
|
||||||
|
return self.sector_reader.sector_counts(
|
||||||
|
stock_codes,
|
||||||
|
target_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _to_item(
|
def _to_item(
|
||||||
@@ -676,4 +769,5 @@ __all__ = [
|
|||||||
"RunZhixingB1",
|
"RunZhixingB1",
|
||||||
"SelectionRerunRequired",
|
"SelectionRerunRequired",
|
||||||
"SelectionRunInProgress",
|
"SelectionRunInProgress",
|
||||||
|
"SelectionSectorAggregates",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -19,9 +19,7 @@ from .zhixing_b1 import compute_signal_masks, prepare_zhixing_b1_indicators
|
|||||||
|
|
||||||
GOLD_BRICK_MINIMUM_HISTORY = 200
|
GOLD_BRICK_MINIMUM_HISTORY = 200
|
||||||
GOLD_BRICK_TURNOVER_RATE_THRESHOLD = 0.99
|
GOLD_BRICK_TURNOVER_RATE_THRESHOLD = 0.99
|
||||||
GOLD_BRICK_SIGNAL_ORDER: tuple[GoldBrickCategory, ...] = (
|
GOLD_BRICK_SIGNAL_ORDER: tuple[GoldBrickCategory, ...] = (GoldBrickCategory.RESONANCE,)
|
||||||
GoldBrickCategory.RESONANCE,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _safe_ratio(numerator: pd.Series, denominator: pd.Series) -> pd.Series:
|
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(
|
multiple_volume_bonus = pd.Series(
|
||||||
np.where(
|
np.where(
|
||||||
(close > open_price)
|
(close > open_price) & (close > previous_close) & (volume > previous_volume * 1.8),
|
||||||
& (close > previous_close)
|
|
||||||
& (volume > previous_volume * 1.8),
|
|
||||||
multiple_volume_coefficient,
|
multiple_volume_coefficient,
|
||||||
1.0,
|
1.0,
|
||||||
),
|
),
|
||||||
@@ -115,13 +111,9 @@ def prepare_gold_brick_indicators(frame: pd.DataFrame, code: str) -> pd.DataFram
|
|||||||
)
|
)
|
||||||
result["j_momentum"] = j_momentum
|
result["j_momentum"] = j_momentum
|
||||||
result["rsi_momentum"] = rsi_momentum
|
result["rsi_momentum"] = rsi_momentum
|
||||||
result["yellow_column"] = (
|
result["yellow_column"] = momentum_sum.div(2).mul(shadow_coefficient).mul(multiple_volume_bonus)
|
||||||
momentum_sum.div(2).mul(shadow_coefficient).mul(multiple_volume_bonus)
|
|
||||||
)
|
|
||||||
x_condition = (
|
x_condition = (
|
||||||
(close > open_price)
|
(close > open_price) & (close > previous_close) & (momentum_sum > previous_momentum_sum)
|
||||||
& (close > previous_close)
|
|
||||||
& (momentum_sum > previous_momentum_sum)
|
|
||||||
)
|
)
|
||||||
result["x_momentum"] = (
|
result["x_momentum"] = (
|
||||||
momentum_sum.sub(previous_momentum_sum)
|
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 - close,
|
||||||
high - upper_shadow_floor,
|
high - upper_shadow_floor,
|
||||||
)
|
)
|
||||||
result["upper_shadow_condition"] = (
|
result["upper_shadow_condition"] = ((close >= open_price) | (close > previous_close)) & (
|
||||||
((close >= open_price) | (close > previous_close))
|
result["upper_shadow_strength"] > 0.618
|
||||||
& (result["upper_shadow_strength"] > 0.618)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
long = result["long_oscillator"]
|
long = result["long_oscillator"]
|
||||||
|
|||||||
@@ -31,6 +31,52 @@ class SelectionResultQuery:
|
|||||||
search: str | None = None
|
search: str | None = None
|
||||||
category: SelectionSignalCategoryFilter | None = None
|
category: SelectionSignalCategoryFilter | None = None
|
||||||
sort: SelectionResultSort = "code"
|
sort: SelectionResultSort = "code"
|
||||||
|
sector: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SelectionSectorCount:
|
||||||
|
"""One sector and the number of this run's selected stocks it contains."""
|
||||||
|
|
||||||
|
sector_code: str
|
||||||
|
sector_name: str
|
||||||
|
stock_count: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SelectionSectorMembership:
|
||||||
|
"""Point-in-time sector membership aggregates for a set of stock codes."""
|
||||||
|
|
||||||
|
snapshot_trade_date: date | None
|
||||||
|
sector_counts: tuple[SelectionSectorCount, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SelectionRunIdentity:
|
||||||
|
"""Minimal run locator used to resolve date-dependent filters."""
|
||||||
|
|
||||||
|
run_id: str
|
||||||
|
target_trade_date: date
|
||||||
|
|
||||||
|
|
||||||
|
class SelectionSectorReader(Protocol):
|
||||||
|
"""Port to the sector-radar context's point-in-time membership reads."""
|
||||||
|
|
||||||
|
def sector_counts(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorMembership: ...
|
||||||
|
|
||||||
|
def sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> tuple[str, ...]: ...
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
@@ -139,6 +185,7 @@ class SelectionRunStore(Protocol):
|
|||||||
run_id: str,
|
run_id: str,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None: ...
|
) -> SelectionRun | None: ...
|
||||||
|
|
||||||
def get_latest_run(
|
def get_latest_run(
|
||||||
@@ -147,8 +194,17 @@ class SelectionRunStore(Protocol):
|
|||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None: ...
|
) -> SelectionRun | None: ...
|
||||||
|
|
||||||
|
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None: ...
|
||||||
|
|
||||||
|
def get_latest_run_identity(
|
||||||
|
self,
|
||||||
|
strategy: SelectionStrategyName,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
) -> SelectionRunIdentity | None: ...
|
||||||
|
|
||||||
|
|
||||||
class SelectionUniverseReader(Protocol):
|
class SelectionUniverseReader(Protocol):
|
||||||
"""Read a qualified market-data source snapshot for one strategy run."""
|
"""Read a qualified market-data source snapshot for one strategy run."""
|
||||||
|
|||||||
+83
-8
@@ -33,6 +33,7 @@ from ..domain.runs import (
|
|||||||
SelectionResultQuery,
|
SelectionResultQuery,
|
||||||
SelectionRun,
|
SelectionRun,
|
||||||
SelectionRunError,
|
SelectionRunError,
|
||||||
|
SelectionRunIdentity,
|
||||||
SelectionRunInProgress,
|
SelectionRunInProgress,
|
||||||
SelectionRunItem,
|
SelectionRunItem,
|
||||||
SelectionRunStatus,
|
SelectionRunStatus,
|
||||||
@@ -46,9 +47,7 @@ _SELECTION_SIGNAL_ORDER: tuple[SelectionSignalCategory, ...] = (
|
|||||||
*ZHIXING_B1_SIGNAL_ORDER,
|
*ZHIXING_B1_SIGNAL_ORDER,
|
||||||
*GOLD_BRICK_SIGNAL_ORDER,
|
*GOLD_BRICK_SIGNAL_ORDER,
|
||||||
)
|
)
|
||||||
_SIGNAL_PRIORITY = {
|
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)}
|
||||||
category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)
|
|
||||||
}
|
|
||||||
_CATEGORY_PREFIXES = {
|
_CATEGORY_PREFIXES = {
|
||||||
"pullback": "zhixing_b1_pullback_",
|
"pullback": "zhixing_b1_pullback_",
|
||||||
"oversold": "zhixing_b1_oversold_",
|
"oversold": "zhixing_b1_oversold_",
|
||||||
@@ -335,21 +334,41 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
run_id: str,
|
run_id: str,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read one run with filtered, stock-paged signals and item failures."""
|
"""Read one run with filtered, stock-paged signals and item failures."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with self._connection() as connection:
|
with self._connection() as connection:
|
||||||
return self._load_run(connection, run_id, query or SelectionResultQuery())
|
return self._load_run(
|
||||||
|
connection,
|
||||||
|
run_id,
|
||||||
|
query or SelectionResultQuery(),
|
||||||
|
sector_stock_codes=sector_stock_codes,
|
||||||
|
)
|
||||||
except psycopg.Error as exc:
|
except psycopg.Error as exc:
|
||||||
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
|
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
|
||||||
|
|
||||||
|
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
|
||||||
|
"""Read only a run's locator so date-dependent filters resolve first."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._connection() as connection:
|
||||||
|
row = connection.execute(
|
||||||
|
"SELECT id, target_trade_date FROM selection_run WHERE id = %s",
|
||||||
|
(run_id,),
|
||||||
|
).fetchone()
|
||||||
|
except psycopg.Error as exc:
|
||||||
|
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
|
||||||
|
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
|
||||||
|
|
||||||
def get_latest_run(
|
def get_latest_run(
|
||||||
self,
|
self,
|
||||||
strategy: SelectionStrategyName,
|
strategy: SelectionStrategyName,
|
||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
"""Read the current run for a date or the latest date for a strategy."""
|
"""Read the current run for a date or the latest date for a strategy."""
|
||||||
|
|
||||||
@@ -377,18 +396,59 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
(strategy, target_trade_date),
|
(strategy, target_trade_date),
|
||||||
).fetchone()
|
).fetchone()
|
||||||
return (
|
return (
|
||||||
self._load_run(connection, str(row[0]), query or SelectionResultQuery())
|
self._load_run(
|
||||||
|
connection,
|
||||||
|
str(row[0]),
|
||||||
|
query or SelectionResultQuery(),
|
||||||
|
sector_stock_codes=sector_stock_codes,
|
||||||
|
)
|
||||||
if row
|
if row
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
except psycopg.Error as exc:
|
except psycopg.Error as exc:
|
||||||
raise SelectionRunStoreError("failed to load latest selection run") from exc
|
raise SelectionRunStoreError("failed to load latest selection run") from exc
|
||||||
|
|
||||||
|
def get_latest_run_identity(
|
||||||
|
self,
|
||||||
|
strategy: SelectionStrategyName,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
) -> SelectionRunIdentity | None:
|
||||||
|
"""Read only the current run's locator so sector filters resolve first."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._connection() as connection:
|
||||||
|
if target_trade_date is None:
|
||||||
|
row = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, target_trade_date
|
||||||
|
FROM selection_run
|
||||||
|
WHERE strategy = %s
|
||||||
|
ORDER BY target_trade_date DESC, created_at DESC, id DESC
|
||||||
|
LIMIT 1
|
||||||
|
""",
|
||||||
|
(strategy,),
|
||||||
|
).fetchone()
|
||||||
|
else:
|
||||||
|
row = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, target_trade_date
|
||||||
|
FROM selection_run
|
||||||
|
WHERE strategy = %s AND target_trade_date = %s
|
||||||
|
LIMIT 1
|
||||||
|
""",
|
||||||
|
(strategy, target_trade_date),
|
||||||
|
).fetchone()
|
||||||
|
except psycopg.Error as exc:
|
||||||
|
raise SelectionRunStoreError("failed to load latest selection run") from exc
|
||||||
|
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _load_run(
|
def _load_run(
|
||||||
connection: Any,
|
connection: Any,
|
||||||
run_id: str,
|
run_id: str,
|
||||||
query: SelectionResultQuery,
|
query: SelectionResultQuery,
|
||||||
|
*,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
) -> SelectionRun | None:
|
) -> SelectionRun | None:
|
||||||
row = connection.execute(
|
row = connection.execute(
|
||||||
"""
|
"""
|
||||||
@@ -417,7 +477,9 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
|||||||
""",
|
""",
|
||||||
(run_id,),
|
(run_id,),
|
||||||
).fetchall()
|
).fetchall()
|
||||||
stock_filter, stock_parameters = _stock_filter(query, run_id)
|
stock_filter, stock_parameters = _stock_filter(
|
||||||
|
query, run_id, sector_stock_codes=sector_stock_codes
|
||||||
|
)
|
||||||
stock_total_row = connection.execute(
|
stock_total_row = connection.execute(
|
||||||
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
|
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
|
||||||
tuple(stock_parameters),
|
tuple(stock_parameters),
|
||||||
@@ -568,12 +630,19 @@ def _signal_category(value: str) -> SelectionSignalCategory:
|
|||||||
return GoldBrickCategory(value)
|
return GoldBrickCategory(value)
|
||||||
|
|
||||||
|
|
||||||
def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]:
|
def _stock_filter(
|
||||||
|
query: SelectionResultQuery,
|
||||||
|
run_id: str,
|
||||||
|
*,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
|
) -> tuple[str, list[object]]:
|
||||||
"""Build the signal predicate used to select distinct matching stocks.
|
"""Build the signal predicate used to select distinct matching stocks.
|
||||||
|
|
||||||
A category narrows which stocks qualify for the page. Once a stock
|
A category narrows which stocks qualify for the page. Once a stock
|
||||||
qualifies, the repository loads every signal for that stock so callers
|
qualifies, the repository loads every signal for that stock so callers
|
||||||
can present all independently persisted categories together.
|
can present all independently persisted categories together. Resolved
|
||||||
|
sector membership codes arrive from the sector-radar port, so the SQL
|
||||||
|
stays inside the selection context.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"]
|
clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"]
|
||||||
@@ -592,6 +661,12 @@ def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[o
|
|||||||
")"
|
")"
|
||||||
)
|
)
|
||||||
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
|
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
|
||||||
|
if sector_stock_codes is not None:
|
||||||
|
if not sector_stock_codes:
|
||||||
|
clauses.append("FALSE")
|
||||||
|
else:
|
||||||
|
clauses.append("item.ts_code = ANY(%s)")
|
||||||
|
parameters.append(list(sector_stock_codes))
|
||||||
return " AND ".join(clauses), parameters
|
return " AND ".join(clauses), parameters
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+84
@@ -0,0 +1,84 @@
|
|||||||
|
"""Bridge the selection sector-filter port to the sector-radar read model.
|
||||||
|
|
||||||
|
The selection context owns no sector-membership storage. This adapter keeps
|
||||||
|
the port contract local to selection while delegating point-in-time reads to
|
||||||
|
the sector-radar application service in the composition root.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from datetime import date
|
||||||
|
|
||||||
|
from zhixing_server.modules.sector_radar.application.read import ReadSectorRadar
|
||||||
|
from zhixing_server.modules.sector_radar.domain.models import SectorType
|
||||||
|
from zhixing_server.modules.selection.domain.runs import (
|
||||||
|
SelectionSectorCount,
|
||||||
|
SelectionSectorMembership,
|
||||||
|
)
|
||||||
|
|
||||||
|
SELECTION_SECTOR_TYPES: dict[str, SectorType] = {
|
||||||
|
"concept": SectorType.CONCEPT,
|
||||||
|
"industry": SectorType.INDUSTRY,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _sector_type(value: str) -> SectorType:
|
||||||
|
"""Map the public sector-type vocabulary onto the radar domain enum."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
return SELECTION_SECTOR_TYPES[value]
|
||||||
|
except KeyError:
|
||||||
|
raise ValueError(f"unsupported sector type: {value}") from None
|
||||||
|
|
||||||
|
|
||||||
|
class SectorRadarSelectionReader:
|
||||||
|
"""Resolve selection sector aggregates through the sector-radar context."""
|
||||||
|
|
||||||
|
def __init__(self, reader: ReadSectorRadar) -> None:
|
||||||
|
self._reader = reader
|
||||||
|
|
||||||
|
def sector_counts(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorMembership:
|
||||||
|
"""Return per-sector stock counts for one run's selected stocks."""
|
||||||
|
|
||||||
|
snapshot = self._reader.sector_counts(
|
||||||
|
stock_codes,
|
||||||
|
target_trade_date,
|
||||||
|
sector_type=_sector_type(sector_type),
|
||||||
|
)
|
||||||
|
return SelectionSectorMembership(
|
||||||
|
snapshot_trade_date=snapshot.trade_date,
|
||||||
|
sector_counts=tuple(
|
||||||
|
SelectionSectorCount(
|
||||||
|
sector_code=entry.sector_code,
|
||||||
|
sector_name=entry.sector_name,
|
||||||
|
stock_count=entry.stock_count,
|
||||||
|
)
|
||||||
|
for entry in snapshot.counts
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
"""Return one sector's member stock codes on the aligned snapshot."""
|
||||||
|
|
||||||
|
snapshot = self._reader.sector_member_codes(
|
||||||
|
target_trade_date,
|
||||||
|
sector_code,
|
||||||
|
sector_type=_sector_type(sector_type),
|
||||||
|
)
|
||||||
|
return snapshot.stock_codes
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["SELECTION_SECTOR_TYPES", "SectorRadarSelectionReader"]
|
||||||
@@ -9,6 +9,7 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, s
|
|||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from zhixing_server.bootstrap.config import Settings, get_settings
|
from zhixing_server.bootstrap.config import Settings, get_settings
|
||||||
|
from zhixing_server.modules.sector_radar.presentation.http import get_sector_radar_reader
|
||||||
from zhixing_server.modules.selection.application.chart import (
|
from zhixing_server.modules.selection.application.chart import (
|
||||||
GetSelectionChart,
|
GetSelectionChart,
|
||||||
SelectionChart,
|
SelectionChart,
|
||||||
@@ -46,6 +47,9 @@ from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
|||||||
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
||||||
PostgresSelectionRunRepository,
|
PostgresSelectionRunRepository,
|
||||||
)
|
)
|
||||||
|
from zhixing_server.modules.selection.infrastructure.sector_membership import (
|
||||||
|
SectorRadarSelectionReader,
|
||||||
|
)
|
||||||
|
|
||||||
selection_router = APIRouter()
|
selection_router = APIRouter()
|
||||||
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
|
||||||
@@ -212,6 +216,32 @@ class SelectionResultsResponse(BaseModel):
|
|||||||
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
|
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
|
||||||
|
|
||||||
|
|
||||||
|
class SelectionSectorCountResponse(BaseModel):
|
||||||
|
"""One sector and the number of this run's selected stocks it contains."""
|
||||||
|
|
||||||
|
sector_code: str
|
||||||
|
sector_name: str
|
||||||
|
stock_count: int = Field(ge=0)
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_sectors() -> list[SelectionSectorCountResponse]:
|
||||||
|
"""Create a typed default sector-aggregate list."""
|
||||||
|
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
class SelectionSectorsResponse(BaseModel):
|
||||||
|
"""The current run's selected stocks aggregated by point-in-time sector."""
|
||||||
|
|
||||||
|
strategy: StrategyValue
|
||||||
|
target_trade_date: date | None
|
||||||
|
run_id: str | None
|
||||||
|
status: SelectionStatusValue
|
||||||
|
snapshot_trade_date: date | None
|
||||||
|
sector_type: Literal["concept", "industry"]
|
||||||
|
sectors: list[SelectionSectorCountResponse] = Field(default_factory=_empty_sectors)
|
||||||
|
|
||||||
|
|
||||||
def get_selection_service(
|
def get_selection_service(
|
||||||
settings: Annotated[Settings, Depends(get_settings)],
|
settings: Annotated[Settings, Depends(get_settings)],
|
||||||
) -> RunZhixingB1:
|
) -> RunZhixingB1:
|
||||||
@@ -221,12 +251,14 @@ def get_selection_service(
|
|||||||
reader = PostgresMarketDataReader(settings, pool=pool)
|
reader = PostgresMarketDataReader(settings, pool=pool)
|
||||||
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
|
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
|
||||||
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
|
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
|
||||||
|
sector_reader = SectorRadarSelectionReader(get_sector_radar_reader(settings))
|
||||||
return RunZhixingB1(
|
return RunZhixingB1(
|
||||||
reader,
|
reader,
|
||||||
store,
|
store,
|
||||||
evaluators={"gold_brick": EvaluateGoldBrick(reader)},
|
evaluators={"gold_brick": EvaluateGoldBrick(reader)},
|
||||||
pattern_case_loader=pattern_case_loader,
|
pattern_case_loader=pattern_case_loader,
|
||||||
pattern_scorer=ZhixingB1PatternScorer(),
|
pattern_scorer=ZhixingB1PatternScorer(),
|
||||||
|
sector_reader=sector_reader,
|
||||||
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
|
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
|
||||||
max_workers=settings.selection_max_workers,
|
max_workers=settings.selection_max_workers,
|
||||||
batch_size=settings.selection_batch_size,
|
batch_size=settings.selection_batch_size,
|
||||||
@@ -338,15 +370,14 @@ def get_selection_run(
|
|||||||
page: Annotated[int, Query(ge=1)] = 1,
|
page: Annotated[int, Query(ge=1)] = 1,
|
||||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||||
category: Literal[
|
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
|
||||||
"pullback", "oversold", "original", "resonance"
|
|
||||||
] | None = None,
|
|
||||||
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
||||||
|
sector: Annotated[str | None, Query(max_length=60)] = None,
|
||||||
) -> SelectionResultsResponse:
|
) -> SelectionResultsResponse:
|
||||||
"""Return one run for asynchronous polling."""
|
"""Return one run for asynchronous polling."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
query = _result_query(page, page_size, search, category, sort)
|
query = _result_query(page, page_size, search, category, sort, sector)
|
||||||
run = service.get_run(run_id, query=query)
|
run = service.get_run(run_id, query=query)
|
||||||
except SelectionRunStoreError as exc:
|
except SelectionRunStoreError as exc:
|
||||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||||
@@ -355,6 +386,51 @@ def get_selection_run(
|
|||||||
return _run_response(run, query=query)
|
return _run_response(run, query=query)
|
||||||
|
|
||||||
|
|
||||||
|
@selection_router.get("/sectors", response_model=SelectionSectorsResponse)
|
||||||
|
def get_selection_sectors(
|
||||||
|
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||||
|
strategy: StrategyValue = "zhixing_b1",
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
sector_type: Literal["concept", "industry"] = "concept",
|
||||||
|
) -> SelectionSectorsResponse:
|
||||||
|
"""Aggregate the current run's selected stocks by point-in-time sector."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
aggregates = service.list_sector_counts(
|
||||||
|
strategy,
|
||||||
|
target_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
)
|
||||||
|
except SelectionRunStoreError as exc:
|
||||||
|
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||||
|
if aggregates is None or aggregates.run is None:
|
||||||
|
return SelectionSectorsResponse(
|
||||||
|
strategy=strategy,
|
||||||
|
target_trade_date=target_trade_date,
|
||||||
|
run_id=None,
|
||||||
|
status="no_data",
|
||||||
|
snapshot_trade_date=None,
|
||||||
|
sector_type=sector_type,
|
||||||
|
)
|
||||||
|
run = aggregates.run
|
||||||
|
return SelectionSectorsResponse(
|
||||||
|
strategy=run.strategy,
|
||||||
|
target_trade_date=run.target_trade_date,
|
||||||
|
run_id=run.id,
|
||||||
|
status=run.status,
|
||||||
|
snapshot_trade_date=aggregates.snapshot_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
sectors=[
|
||||||
|
SelectionSectorCountResponse(
|
||||||
|
sector_code=count.sector_code,
|
||||||
|
sector_name=count.sector_name,
|
||||||
|
stock_count=count.stock_count,
|
||||||
|
)
|
||||||
|
for count in aggregates.sectors
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@selection_router.get("/results", response_model=SelectionResultsResponse)
|
@selection_router.get("/results", response_model=SelectionResultsResponse)
|
||||||
def get_selection_results(
|
def get_selection_results(
|
||||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||||
@@ -363,15 +439,14 @@ def get_selection_results(
|
|||||||
page: Annotated[int, Query(ge=1)] = 1,
|
page: Annotated[int, Query(ge=1)] = 1,
|
||||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||||
category: Literal[
|
category: Literal["pullback", "oversold", "original", "resonance"] | None = None,
|
||||||
"pullback", "oversold", "original", "resonance"
|
|
||||||
] | None = None,
|
|
||||||
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
sort: Literal["code", "score_desc", "score_asc"] = "code",
|
||||||
|
sector: Annotated[str | None, Query(max_length=60)] = None,
|
||||||
) -> SelectionResultsResponse:
|
) -> SelectionResultsResponse:
|
||||||
"""Return the current persisted result for a strategy and optional date."""
|
"""Return the current persisted result for a strategy and optional date."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
query = _result_query(page, page_size, search, category, sort)
|
query = _result_query(page, page_size, search, category, sort, sector)
|
||||||
run = service.get_latest(strategy, target_trade_date, query=query)
|
run = service.get_latest(strategy, target_trade_date, query=query)
|
||||||
except SelectionRunStoreError as exc:
|
except SelectionRunStoreError as exc:
|
||||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||||
@@ -524,16 +599,19 @@ def _result_query(
|
|||||||
search: str | None,
|
search: str | None,
|
||||||
category: Literal["pullback", "oversold", "original", "resonance"] | None,
|
category: Literal["pullback", "oversold", "original", "resonance"] | None,
|
||||||
sort: Literal["code", "score_desc", "score_asc"],
|
sort: Literal["code", "score_desc", "score_asc"],
|
||||||
|
sector: str | None = None,
|
||||||
) -> SelectionResultQuery:
|
) -> SelectionResultQuery:
|
||||||
"""Normalize HTTP query values before handing them to the selection port."""
|
"""Normalize HTTP query values before handing them to the selection port."""
|
||||||
|
|
||||||
normalized_search = search.strip() if search else None
|
normalized_search = search.strip() if search else None
|
||||||
|
normalized_sector = sector.strip() if sector else None
|
||||||
return SelectionResultQuery(
|
return SelectionResultQuery(
|
||||||
page=page,
|
page=page,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
search=normalized_search or None,
|
search=normalized_search or None,
|
||||||
category=category,
|
category=category,
|
||||||
sort=sort,
|
sort=sort,
|
||||||
|
sector=normalized_sector or None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -551,6 +629,7 @@ __all__ = [
|
|||||||
"SelectionResultsResponse",
|
"SelectionResultsResponse",
|
||||||
"SelectionRunAcceptedResponse",
|
"SelectionRunAcceptedResponse",
|
||||||
"SelectionRunRequest",
|
"SelectionRunRequest",
|
||||||
|
"SelectionSectorsResponse",
|
||||||
"SelectionStockResponse",
|
"SelectionStockResponse",
|
||||||
"get_selection_chart_service",
|
"get_selection_chart_service",
|
||||||
"get_selection_service",
|
"get_selection_service",
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ from zhixing_server.modules.selection.application.chart import (
|
|||||||
SelectionChartNotFound,
|
SelectionChartNotFound,
|
||||||
SelectionChartPoint,
|
SelectionChartPoint,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
|
from zhixing_server.modules.selection.application.run import (
|
||||||
|
PreparedSelectionRun,
|
||||||
|
SelectionSectorAggregates,
|
||||||
|
)
|
||||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||||
from zhixing_server.modules.selection.domain.pattern_scoring import (
|
from zhixing_server.modules.selection.domain.pattern_scoring import (
|
||||||
PATTERN_SCORING_VERSION,
|
PATTERN_SCORING_VERSION,
|
||||||
@@ -30,6 +33,7 @@ from zhixing_server.modules.selection.domain.runs import (
|
|||||||
SelectionResultQuery,
|
SelectionResultQuery,
|
||||||
SelectionRun,
|
SelectionRun,
|
||||||
SelectionRunInProgress,
|
SelectionRunInProgress,
|
||||||
|
SelectionSectorCount,
|
||||||
SelectionStock,
|
SelectionStock,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||||
@@ -45,11 +49,17 @@ TARGET = date(2026, 8, 8)
|
|||||||
|
|
||||||
|
|
||||||
class FakeSelectionService:
|
class FakeSelectionService:
|
||||||
def __init__(self, run: SelectionRun | None = None) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
run: SelectionRun | None = None,
|
||||||
|
sectors: tuple[SelectionSectorCount, ...] = (),
|
||||||
|
) -> None:
|
||||||
self.run = run
|
self.run = run
|
||||||
|
self.sectors = sectors
|
||||||
self.executed = False
|
self.executed = False
|
||||||
self.mode = "ok"
|
self.mode = "ok"
|
||||||
self.last_query: SelectionResultQuery | None = None
|
self.last_query: SelectionResultQuery | None = None
|
||||||
|
self.sector_type_requested: str | None = None
|
||||||
|
|
||||||
def prepare(
|
def prepare(
|
||||||
self,
|
self,
|
||||||
@@ -103,6 +113,23 @@ class FakeSelectionService:
|
|||||||
return None
|
return None
|
||||||
return self.run
|
return self.run
|
||||||
|
|
||||||
|
def list_sector_counts(
|
||||||
|
self,
|
||||||
|
strategy: str,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorAggregates | None:
|
||||||
|
self.sector_type_requested = sector_type
|
||||||
|
if self.run is None:
|
||||||
|
return None
|
||||||
|
return SelectionSectorAggregates(
|
||||||
|
run=self.run,
|
||||||
|
snapshot_trade_date=self.run.target_trade_date,
|
||||||
|
sector_type=sector_type,
|
||||||
|
sectors=self.sectors,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class FakeChartService:
|
class FakeChartService:
|
||||||
"""Return or fail one deterministic chart response."""
|
"""Return or fail one deterministic chart response."""
|
||||||
@@ -434,6 +461,83 @@ def test_query_rejects_invalid_page_size() -> None:
|
|||||||
assert response.status_code == 422
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
def test_query_forwards_sector_filter() -> None:
|
||||||
|
service = FakeSelectionService(_run("run-http", "success"))
|
||||||
|
|
||||||
|
response = _client(service).get(
|
||||||
|
"/api/v1/selection/results",
|
||||||
|
params={"strategy": "zhixing_b1", "sector": " BK0475.DC "},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert service.last_query is not None
|
||||||
|
assert service.last_query.sector == "BK0475.DC"
|
||||||
|
|
||||||
|
|
||||||
|
def test_sectors_returns_aggregated_counts_desc() -> None:
|
||||||
|
run = _run("run-http", "success")
|
||||||
|
service = FakeSelectionService(
|
||||||
|
run,
|
||||||
|
sectors=(
|
||||||
|
SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=3),
|
||||||
|
SelectionSectorCount(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
response = _client(service).get("/api/v1/selection/sectors")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {
|
||||||
|
"strategy": "zhixing_b1",
|
||||||
|
"target_trade_date": "2026-08-08",
|
||||||
|
"run_id": "run-http",
|
||||||
|
"status": "success",
|
||||||
|
"snapshot_trade_date": "2026-08-08",
|
||||||
|
"sector_type": "concept",
|
||||||
|
"sectors": [
|
||||||
|
{"sector_code": "BK0001.DC", "sector_name": "机器人", "stock_count": 3},
|
||||||
|
{"sector_code": "BK0003.DC", "sector_name": "数字经济", "stock_count": 2},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sectors_forwards_sector_type() -> None:
|
||||||
|
service = FakeSelectionService(_run("run-http", "success"))
|
||||||
|
|
||||||
|
response = _client(service).get(
|
||||||
|
"/api/v1/selection/sectors",
|
||||||
|
params={"sector_type": "industry"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert service.sector_type_requested == "industry"
|
||||||
|
assert response.json()["sector_type"] == "industry"
|
||||||
|
|
||||||
|
|
||||||
|
def test_sectors_without_run_returns_no_data() -> None:
|
||||||
|
response = _client(FakeSelectionService()).get("/api/v1/selection/sectors")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {
|
||||||
|
"strategy": "zhixing_b1",
|
||||||
|
"target_trade_date": None,
|
||||||
|
"run_id": None,
|
||||||
|
"status": "no_data",
|
||||||
|
"snapshot_trade_date": None,
|
||||||
|
"sector_type": "concept",
|
||||||
|
"sectors": [],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sectors_rejects_unknown_sector_type() -> None:
|
||||||
|
response = _client(FakeSelectionService()).get(
|
||||||
|
"/api/v1/selection/sectors",
|
||||||
|
params={"sector_type": "macro"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
def test_chart_returns_bounded_qfq_contract() -> None:
|
def test_chart_returns_bounded_qfq_contract() -> None:
|
||||||
chart_service = FakeChartService()
|
chart_service = FakeChartService()
|
||||||
|
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from zhixing_server.modules.sector_radar.domain.models import (
|
|||||||
from zhixing_server.modules.sector_radar.domain.persistence import (
|
from zhixing_server.modules.sector_radar.domain.persistence import (
|
||||||
MembershipRecord,
|
MembershipRecord,
|
||||||
RankingRecord,
|
RankingRecord,
|
||||||
|
SectorCountEntry,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
||||||
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
||||||
@@ -315,3 +316,89 @@ def test_stock_membership_query_rejects_invalid_values() -> None:
|
|||||||
StockSectorQuery(trade_date=TARGET_DATE, ts_code="x" * 13)
|
StockSectorQuery(trade_date=TARGET_DATE, ts_code="x" * 13)
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
StockSectorQuery(trade_date=TARGET_DATE, ts_code="000001.SZ", concept_limit=0)
|
StockSectorQuery(trade_date=TARGET_DATE, ts_code="000001.SZ", concept_limit=0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_counts_aggregates_concepts_and_orders_by_count_then_code() -> None:
|
||||||
|
repository = _membership_repository()
|
||||||
|
repository.save_memberships(
|
||||||
|
(
|
||||||
|
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
|
||||||
|
_membership_record("000003.SZ", SectorType.CONCEPT, "BK0001.DC", "机器人"),
|
||||||
|
_membership_record("000002.SZ", SectorType.CONCEPT, "BK0003.DC", "数字经济"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
reader = ReadSectorRadar(repository)
|
||||||
|
|
||||||
|
snapshot = reader.sector_counts(["000001.SZ", "000002.SZ", "000003.SZ"], TARGET_DATE)
|
||||||
|
|
||||||
|
assert snapshot.status == "success"
|
||||||
|
assert snapshot.trade_date == TARGET_DATE
|
||||||
|
assert snapshot.sector_type is SectorType.CONCEPT
|
||||||
|
assert snapshot.counts == (
|
||||||
|
SectorCountEntry(sector_code="BK0001.DC", sector_name="机器人", stock_count=3),
|
||||||
|
SectorCountEntry(sector_code="BK0003.DC", sector_name="数字经济", stock_count=2),
|
||||||
|
SectorCountEntry(sector_code="BK0002.DC", sector_name="人工智能", stock_count=1),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_counts_supports_industry_type() -> None:
|
||||||
|
reader = ReadSectorRadar(_membership_repository())
|
||||||
|
|
||||||
|
snapshot = reader.sector_counts(
|
||||||
|
["000001.SZ", "000002.SZ"],
|
||||||
|
TARGET_DATE,
|
||||||
|
sector_type=SectorType.INDUSTRY,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert snapshot.status == "success"
|
||||||
|
assert snapshot.sector_type is SectorType.INDUSTRY
|
||||||
|
assert snapshot.counts == (
|
||||||
|
SectorCountEntry(sector_code="BK0901.DC", sector_name="银行", stock_count=1),
|
||||||
|
SectorCountEntry(sector_code="BK0902.DC", sector_name="房地产", stock_count=1),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_counts_without_publication_is_no_data() -> None:
|
||||||
|
reader = ReadSectorRadar(InMemorySectorRadarRepository())
|
||||||
|
|
||||||
|
snapshot = reader.sector_counts(["000001.SZ"], TARGET_DATE)
|
||||||
|
|
||||||
|
assert snapshot.status == "no_data"
|
||||||
|
assert snapshot.trade_date is None
|
||||||
|
assert snapshot.counts == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_counts_with_empty_stock_set_skips_publication_lookup() -> None:
|
||||||
|
reader = ReadSectorRadar(_membership_repository())
|
||||||
|
|
||||||
|
snapshot = reader.sector_counts([], TARGET_DATE)
|
||||||
|
|
||||||
|
assert snapshot.status == "no_data"
|
||||||
|
assert snapshot.counts == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_member_codes_returns_available_members() -> None:
|
||||||
|
reader = ReadSectorRadar(_membership_repository())
|
||||||
|
|
||||||
|
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
|
||||||
|
|
||||||
|
assert snapshot.status == "success"
|
||||||
|
assert snapshot.trade_date == TARGET_DATE
|
||||||
|
assert snapshot.stock_codes == ("000001.SZ",)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_member_codes_without_publication_is_no_data() -> None:
|
||||||
|
reader = ReadSectorRadar(InMemorySectorRadarRepository())
|
||||||
|
|
||||||
|
snapshot = reader.sector_member_codes(TARGET_DATE, "BK0001.DC")
|
||||||
|
|
||||||
|
assert snapshot.status == "no_data"
|
||||||
|
assert snapshot.trade_date is None
|
||||||
|
assert snapshot.stock_codes == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_sector_member_codes_rejects_blank_sector_code() -> None:
|
||||||
|
reader = ReadSectorRadar(_membership_repository())
|
||||||
|
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
reader.sector_member_codes(TARGET_DATE, " ")
|
||||||
|
|||||||
@@ -136,7 +136,13 @@ class FakeStore:
|
|||||||
kwargs["error_message"] = error_message
|
kwargs["error_message"] = error_message
|
||||||
self.finished = (run_id, status, kwargs)
|
self.finished = (run_id, status, kwargs)
|
||||||
|
|
||||||
def get_run(self, run_id: str, *, query: SelectionResultQuery | None = None):
|
def get_run(
|
||||||
|
self,
|
||||||
|
run_id: str,
|
||||||
|
*,
|
||||||
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_latest_run(
|
def get_latest_run(
|
||||||
@@ -145,9 +151,16 @@ class FakeStore:
|
|||||||
target_trade_date: date | None = None,
|
target_trade_date: date | None = None,
|
||||||
*,
|
*,
|
||||||
query: SelectionResultQuery | None = None,
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def get_run_identity(self, run_id: str):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_latest_run_identity(self, strategy: str, target_trade_date: date | None = None):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class BatchStore(FakeStore):
|
class BatchStore(FakeStore):
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
@@ -427,8 +440,8 @@ def test_execute_loads_cases_once_and_scores_only_selected_stocks() -> None:
|
|||||||
reader,
|
reader,
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
pattern_scoring_enabled=True,
|
||||||
batch_size=1,
|
batch_size=1,
|
||||||
)
|
)
|
||||||
@@ -463,8 +476,8 @@ def test_execute_isolates_pattern_scoring_failure_from_selection_status() -> Non
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
pattern_scoring_enabled=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -497,8 +510,8 @@ def test_execute_skips_pattern_dependencies_when_feature_flag_is_disabled() -> N
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=False,
|
pattern_scoring_enabled=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -535,8 +548,8 @@ def test_execute_marks_scores_failed_when_case_library_is_unavailable() -> None:
|
|||||||
BatchReader(source),
|
BatchReader(source),
|
||||||
store,
|
store,
|
||||||
evaluator,
|
evaluator,
|
||||||
loader,
|
pattern_case_loader=loader,
|
||||||
scorer,
|
pattern_scorer=scorer,
|
||||||
pattern_scoring_enabled=True,
|
pattern_scoring_enabled=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,264 @@
|
|||||||
|
"""Sector-filter behavior for selection result reads and sector aggregates."""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from dataclasses import replace
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from zhixing_server.modules.selection.application.run import (
|
||||||
|
RunZhixingB1,
|
||||||
|
SelectionSectorAggregates,
|
||||||
|
)
|
||||||
|
from zhixing_server.modules.selection.domain.runs import (
|
||||||
|
SelectionExecutionSource,
|
||||||
|
SelectionResultQuery,
|
||||||
|
SelectionRun,
|
||||||
|
SelectionRunIdentity,
|
||||||
|
SelectionRunItem,
|
||||||
|
SelectionRunStatus,
|
||||||
|
SelectionSectorCount,
|
||||||
|
SelectionSectorMembership,
|
||||||
|
)
|
||||||
|
|
||||||
|
TARGET = date(2026, 8, 8)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeReader:
|
||||||
|
def load_execution_source(
|
||||||
|
self, strategy: str, target_trade_date: date
|
||||||
|
) -> SelectionExecutionSource:
|
||||||
|
raise AssertionError("sector reads must not load execution sources")
|
||||||
|
|
||||||
|
def load_history(self, ts_code: str, target_trade_date: date):
|
||||||
|
raise AssertionError("sector reads must not load histories")
|
||||||
|
|
||||||
|
|
||||||
|
class FakeSectorReader:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
counts: tuple[SelectionSectorCount, ...] = (),
|
||||||
|
member_codes: tuple[str, ...] = (),
|
||||||
|
) -> None:
|
||||||
|
self.counts = counts
|
||||||
|
self.member_codes = member_codes
|
||||||
|
self.count_calls: list[tuple[tuple[str, ...], date, str]] = []
|
||||||
|
self.member_calls: list[tuple[date, str, str]] = []
|
||||||
|
|
||||||
|
def sector_counts(
|
||||||
|
self,
|
||||||
|
stock_codes: Sequence[str],
|
||||||
|
target_trade_date: date,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> SelectionSectorMembership:
|
||||||
|
self.count_calls.append((tuple(stock_codes), target_trade_date, sector_type))
|
||||||
|
return SelectionSectorMembership(
|
||||||
|
snapshot_trade_date=target_trade_date,
|
||||||
|
sector_counts=self.counts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def sector_member_codes(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
sector_code: str,
|
||||||
|
*,
|
||||||
|
sector_type: str = "concept",
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
self.member_calls.append((target_trade_date, sector_code, sector_type))
|
||||||
|
return self.member_codes
|
||||||
|
|
||||||
|
|
||||||
|
class FakeStore:
|
||||||
|
def __init__(self, run: SelectionRun | None) -> None:
|
||||||
|
self.run = run
|
||||||
|
self.sector_codes_seen: Sequence[str] | None = None
|
||||||
|
self.queries: list[SelectionResultQuery] = []
|
||||||
|
|
||||||
|
def prepare_run(self, *args: object, **kwargs: object):
|
||||||
|
raise AssertionError("sector reads must not prepare runs")
|
||||||
|
|
||||||
|
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||||
|
raise AssertionError("sector reads must not record items")
|
||||||
|
|
||||||
|
def finish_run(self, *args: object, **kwargs: object) -> None:
|
||||||
|
raise AssertionError("sector reads must not finish runs")
|
||||||
|
|
||||||
|
def get_run(
|
||||||
|
self,
|
||||||
|
run_id: str,
|
||||||
|
*,
|
||||||
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
|
) -> SelectionRun | None:
|
||||||
|
self.queries.append(query or SelectionResultQuery())
|
||||||
|
self.sector_codes_seen = sector_stock_codes
|
||||||
|
if self.run is not None and self.run.id == run_id:
|
||||||
|
return self.run
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_latest_run(
|
||||||
|
self,
|
||||||
|
strategy: str,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
*,
|
||||||
|
query: SelectionResultQuery | None = None,
|
||||||
|
sector_stock_codes: Sequence[str] | None = None,
|
||||||
|
) -> SelectionRun | None:
|
||||||
|
self.queries.append(query or SelectionResultQuery())
|
||||||
|
self.sector_codes_seen = sector_stock_codes
|
||||||
|
return self.run
|
||||||
|
|
||||||
|
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
|
||||||
|
if self.run is None or self.run.id != run_id:
|
||||||
|
return None
|
||||||
|
return SelectionRunIdentity(
|
||||||
|
run_id=self.run.id,
|
||||||
|
target_trade_date=self.run.target_trade_date,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_latest_run_identity(
|
||||||
|
self,
|
||||||
|
strategy: str,
|
||||||
|
target_trade_date: date | None = None,
|
||||||
|
) -> SelectionRunIdentity | None:
|
||||||
|
if self.run is None:
|
||||||
|
return None
|
||||||
|
return SelectionRunIdentity(
|
||||||
|
run_id=self.run.id,
|
||||||
|
target_trade_date=self.run.target_trade_date,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _run(status: SelectionRunStatus = "success") -> SelectionRun:
|
||||||
|
return SelectionRun(
|
||||||
|
id="run-1",
|
||||||
|
strategy="zhixing_b1",
|
||||||
|
target_trade_date=TARGET,
|
||||||
|
market_sync_batch_id="market-run-1",
|
||||||
|
status=status,
|
||||||
|
target_count=3,
|
||||||
|
eligible_count=3,
|
||||||
|
evaluated_count=3,
|
||||||
|
selected_stock_count=2,
|
||||||
|
signal_count=2,
|
||||||
|
failed_count=0,
|
||||||
|
coverage=Decimal(1),
|
||||||
|
items=(
|
||||||
|
SelectionRunItem(ts_code="000001.SZ", name="A", status="selected", signal_count=1),
|
||||||
|
SelectionRunItem(ts_code="000002.SZ", name="B", status="selected", signal_count=1),
|
||||||
|
SelectionRunItem(ts_code="000003.SZ", name="C", status="no_signal", signal_count=0),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_sector_counts_aggregates_only_selected_stocks() -> None:
|
||||||
|
store = FakeStore(_run())
|
||||||
|
sector_reader = FakeSectorReader(
|
||||||
|
counts=(SelectionSectorCount(sector_code="BK0001.DC", sector_name="机器人", stock_count=2),)
|
||||||
|
)
|
||||||
|
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
|
||||||
|
|
||||||
|
aggregates = service.list_sector_counts("zhixing_b1")
|
||||||
|
|
||||||
|
assert isinstance(aggregates, SelectionSectorAggregates)
|
||||||
|
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "concept")]
|
||||||
|
assert aggregates.snapshot_trade_date == TARGET
|
||||||
|
assert aggregates.sector_type == "concept"
|
||||||
|
assert aggregates.sectors == sector_reader.counts
|
||||||
|
assert aggregates.run is store.run
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_sector_counts_without_run_returns_none() -> None:
|
||||||
|
service = RunZhixingB1(FakeReader(), FakeStore(None), sector_reader=FakeSectorReader())
|
||||||
|
|
||||||
|
assert service.list_sector_counts("zhixing_b1") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_sector_counts_without_selected_stocks_skips_reader() -> None:
|
||||||
|
run = replace(_run(), items=())
|
||||||
|
sector_reader = FakeSectorReader()
|
||||||
|
service = RunZhixingB1(FakeReader(), FakeStore(run), sector_reader=sector_reader)
|
||||||
|
|
||||||
|
aggregates = service.list_sector_counts("zhixing_b1")
|
||||||
|
|
||||||
|
assert sector_reader.count_calls == []
|
||||||
|
assert aggregates is not None
|
||||||
|
assert aggregates.snapshot_trade_date is None
|
||||||
|
assert aggregates.sectors == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_latest_resolves_sector_filter_against_run_snapshot() -> None:
|
||||||
|
store = FakeStore(_run())
|
||||||
|
sector_reader = FakeSectorReader(member_codes=("000001.SZ", "000002.SZ"))
|
||||||
|
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
|
||||||
|
|
||||||
|
run = service.get_latest(
|
||||||
|
"zhixing_b1",
|
||||||
|
query=SelectionResultQuery(page=2, page_size=20, sector="BK0001.DC"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert run is store.run
|
||||||
|
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
|
||||||
|
assert store.sector_codes_seen == ("000001.SZ", "000002.SZ")
|
||||||
|
assert store.queries[-1].sector is None
|
||||||
|
assert store.queries[-1].page == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_latest_with_empty_sector_members_returns_empty_codes_page() -> None:
|
||||||
|
store = FakeStore(_run())
|
||||||
|
service = RunZhixingB1(
|
||||||
|
FakeReader(),
|
||||||
|
store,
|
||||||
|
sector_reader=FakeSectorReader(member_codes=()),
|
||||||
|
)
|
||||||
|
|
||||||
|
service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK9999.DC"))
|
||||||
|
|
||||||
|
assert store.sector_codes_seen == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_latest_without_sector_reader_ignores_the_filter() -> None:
|
||||||
|
store = FakeStore(_run())
|
||||||
|
service = RunZhixingB1(FakeReader(), store)
|
||||||
|
|
||||||
|
run = service.get_latest(
|
||||||
|
"zhixing_b1",
|
||||||
|
query=SelectionResultQuery(sector="BK0001.DC"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert run is store.run
|
||||||
|
assert store.sector_codes_seen is None
|
||||||
|
assert store.queries[-1].sector == "BK0001.DC"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_run_resolves_sector_filter_against_run_snapshot() -> None:
|
||||||
|
store = FakeStore(_run())
|
||||||
|
sector_reader = FakeSectorReader(member_codes=("000001.SZ",))
|
||||||
|
service = RunZhixingB1(FakeReader(), store, sector_reader=sector_reader)
|
||||||
|
|
||||||
|
run = service.get_run("run-1", query=SelectionResultQuery(sector="BK0001.DC"))
|
||||||
|
|
||||||
|
assert run is store.run
|
||||||
|
assert sector_reader.member_calls == [(TARGET, "BK0001.DC", "concept")]
|
||||||
|
assert store.sector_codes_seen == ("000001.SZ",)
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_latest_without_run_and_sector_filter_returns_none() -> None:
|
||||||
|
service = RunZhixingB1(
|
||||||
|
FakeReader(),
|
||||||
|
FakeStore(None),
|
||||||
|
sector_reader=FakeSectorReader(member_codes=()),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert service.get_latest("zhixing_b1", query=SelectionResultQuery(sector="BK0001.DC")) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_sector_counts_accepts_industry_type() -> None:
|
||||||
|
sector_reader = FakeSectorReader()
|
||||||
|
service = RunZhixingB1(FakeReader(), FakeStore(_run()), sector_reader=sector_reader)
|
||||||
|
|
||||||
|
aggregates = service.list_sector_counts("zhixing_b1", sector_type="industry")
|
||||||
|
|
||||||
|
assert aggregates is not None
|
||||||
|
assert aggregates.sector_type == "industry"
|
||||||
|
assert sector_reader.count_calls == [(("000001.SZ", "000002.SZ"), TARGET, "industry")]
|
||||||
@@ -7,6 +7,7 @@ import {
|
|||||||
type SelectionResultsQuery,
|
type SelectionResultsQuery,
|
||||||
type SelectionRunAccepted,
|
type SelectionRunAccepted,
|
||||||
type SelectionRunRequest,
|
type SelectionRunRequest,
|
||||||
|
type SelectionSectors,
|
||||||
type SelectionStrategy,
|
type SelectionStrategy,
|
||||||
} from "./selection.types"
|
} from "./selection.types"
|
||||||
|
|
||||||
@@ -63,9 +64,27 @@ function buildSelectionQueryParams(query: SelectionResultsQuery) {
|
|||||||
if (query.search) params.set("search", query.search)
|
if (query.search) params.set("search", query.search)
|
||||||
if (query.category) params.set("category", query.category)
|
if (query.category) params.set("category", query.category)
|
||||||
if (query.sort) params.set("sort", query.sort)
|
if (query.sort) params.set("sort", query.sort)
|
||||||
|
if (query.sector) params.set("sector", query.sector)
|
||||||
return params
|
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(
|
export function triggerSelectionRun(
|
||||||
request: SelectionRunRequest,
|
request: SelectionRunRequest,
|
||||||
signal?: AbortSignal,
|
signal?: AbortSignal,
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ describe("selection query hooks", () => {
|
|||||||
"平安",
|
"平安",
|
||||||
"pullback",
|
"pullback",
|
||||||
"score_desc",
|
"score_desc",
|
||||||
|
"",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
expect(
|
expect(
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import { useEffect } from "react"
|
|||||||
|
|
||||||
import {
|
import {
|
||||||
getSelectionChart,
|
getSelectionChart,
|
||||||
|
getSelectionResultSectors,
|
||||||
getSelectionResults,
|
getSelectionResults,
|
||||||
getSelectionRun,
|
getSelectionRun,
|
||||||
triggerSelectionRun,
|
triggerSelectionRun,
|
||||||
@@ -19,6 +20,7 @@ import {
|
|||||||
type SelectionResults,
|
type SelectionResults,
|
||||||
type SelectionResultsQuery,
|
type SelectionResultsQuery,
|
||||||
type SelectionRunRequest,
|
type SelectionRunRequest,
|
||||||
|
type SelectionSectors,
|
||||||
type SelectionStrategy,
|
type SelectionStrategy,
|
||||||
} from "./selection.types"
|
} from "./selection.types"
|
||||||
|
|
||||||
@@ -55,6 +57,7 @@ export const selectionResultsQueryKey = (
|
|||||||
query.search ?? "",
|
query.search ?? "",
|
||||||
query.category ?? "all",
|
query.category ?? "all",
|
||||||
query.sort ?? "code",
|
query.sort ?? "code",
|
||||||
|
query.sector ?? "",
|
||||||
] as const
|
] as const
|
||||||
|
|
||||||
export const selectionRunQueryKey = (
|
export const selectionRunQueryKey = (
|
||||||
@@ -69,8 +72,15 @@ export const selectionRunQueryKey = (
|
|||||||
query.search ?? "",
|
query.search ?? "",
|
||||||
query.category ?? "all",
|
query.category ?? "all",
|
||||||
query.sort ?? "code",
|
query.sort ?? "code",
|
||||||
|
query.sector ?? "",
|
||||||
] as const
|
] as const
|
||||||
|
|
||||||
|
export const selectionSectorsQueryKey = (
|
||||||
|
strategy: SelectionStrategy,
|
||||||
|
targetTradeDate?: string,
|
||||||
|
sectorType: "concept" | "industry" = "concept",
|
||||||
|
) => ["selection", "sectors", strategy, targetTradeDate, sectorType] as const
|
||||||
|
|
||||||
export function flattenSelectionResults(
|
export function flattenSelectionResults(
|
||||||
data: InfiniteData<SelectionResults> | undefined,
|
data: InfiniteData<SelectionResults> | undefined,
|
||||||
): 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(
|
export function useSelectionRun(
|
||||||
runId: string | null,
|
runId: string | null,
|
||||||
resultQuery: SelectionResultsListQuery = defaultListQuery,
|
resultQuery: SelectionResultsListQuery = defaultListQuery,
|
||||||
@@ -173,9 +195,14 @@ export function invalidateSelectionResults(
|
|||||||
queryClient: QueryClient,
|
queryClient: QueryClient,
|
||||||
strategy: SelectionStrategy,
|
strategy: SelectionStrategy,
|
||||||
) {
|
) {
|
||||||
return queryClient.invalidateQueries({
|
return Promise.all([
|
||||||
queryKey: ["selection", "results", strategy],
|
queryClient.invalidateQueries({
|
||||||
})
|
queryKey: ["selection", "results", strategy],
|
||||||
|
}),
|
||||||
|
queryClient.invalidateQueries({
|
||||||
|
queryKey: ["selection", "sectors", strategy],
|
||||||
|
}),
|
||||||
|
])
|
||||||
}
|
}
|
||||||
|
|
||||||
function nextSelectionPageParam(
|
function nextSelectionPageParam(
|
||||||
|
|||||||
@@ -25,6 +25,27 @@ export interface SelectionResultsQuery {
|
|||||||
search?: string
|
search?: string
|
||||||
category?: Exclude<SelectionCategoryFilter, "all">
|
category?: Exclude<SelectionCategoryFilter, "all">
|
||||||
sort?: SelectionSort
|
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 =
|
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 { useNavigate, useSearch } from "@tanstack/react-router"
|
||||||
|
|
||||||
import { Input } from "@/shared/ui/input"
|
import { Input } from "@/shared/ui/input"
|
||||||
@@ -11,7 +11,11 @@ import {
|
|||||||
SelectValue,
|
SelectValue,
|
||||||
} from "@/shared/ui/select"
|
} 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 { SignalDetailPanel } from "./signal-detail-panel"
|
||||||
import { SignalRecordList } from "./signal-record-list"
|
import { SignalRecordList } from "./signal-record-list"
|
||||||
import {
|
import {
|
||||||
@@ -30,12 +34,15 @@ const SCORE_SORT_OPTIONS: ReadonlyArray<{
|
|||||||
{ label: "评分从低到高", value: "score_asc" },
|
{ label: "评分从低到高", value: "score_asc" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
const ALL_SECTORS_VALUE = "all"
|
||||||
|
|
||||||
interface SelectionResultsWorkbenchProps {
|
interface SelectionResultsWorkbenchProps {
|
||||||
hasNextPage: boolean
|
hasNextPage: boolean
|
||||||
isFetchNextPageError: boolean
|
isFetchNextPageError: boolean
|
||||||
isFetchingNextPage: boolean
|
isFetchingNextPage: boolean
|
||||||
onLoadMore: () => void | Promise<unknown>
|
onLoadMore: () => void | Promise<unknown>
|
||||||
result: SelectionResults
|
result: SelectionResults
|
||||||
|
sectorOptions: ReadonlyArray<SelectionSectorAggregate>
|
||||||
}
|
}
|
||||||
|
|
||||||
export function SelectionResultsWorkbench({
|
export function SelectionResultsWorkbench({
|
||||||
@@ -44,12 +51,14 @@ export function SelectionResultsWorkbench({
|
|||||||
isFetchingNextPage,
|
isFetchingNextPage,
|
||||||
onLoadMore,
|
onLoadMore,
|
||||||
result,
|
result,
|
||||||
|
sectorOptions,
|
||||||
}: SelectionResultsWorkbenchProps) {
|
}: SelectionResultsWorkbenchProps) {
|
||||||
const search = useSearch({ from: "/_workspace/selection" })
|
const search = useSearch({ from: "/_workspace/selection" })
|
||||||
const navigate = useNavigate({ from: "/selection" })
|
const navigate = useNavigate({ from: "/selection" })
|
||||||
const query = search.search ?? ""
|
const query = search.search ?? ""
|
||||||
const category = search.category ?? "all"
|
const category = search.category ?? "all"
|
||||||
const sort = search.sort ?? "code"
|
const sort = search.sort ?? "code"
|
||||||
|
const sector = search.sector
|
||||||
const [selectedKey, setSelectedKey] = useState<string | null>(null)
|
const [selectedKey, setSelectedKey] = useState<string | null>(null)
|
||||||
const loadMoreRequestPending = useRef(false)
|
const loadMoreRequestPending = useRef(false)
|
||||||
|
|
||||||
@@ -68,11 +77,32 @@ export function SelectionResultsWorkbench({
|
|||||||
const selectedStock =
|
const selectedStock =
|
||||||
visibleStocks.find((stock) => getStockKey(stock) === selectedKey) ??
|
visibleStocks.find((stock) => getStockKey(stock) === selectedKey) ??
|
||||||
visibleStocks[0]
|
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: {
|
function updateSearch(next: {
|
||||||
search?: string
|
search?: string
|
||||||
category?: SignalCategoryFilter
|
category?: SignalCategoryFilter
|
||||||
sort?: SelectionSort
|
sort?: SelectionSort
|
||||||
|
sector?: string
|
||||||
}) {
|
}) {
|
||||||
void navigate({ search: (previous) => ({ ...previous, ...next }) })
|
void navigate({ search: (previous) => ({ ...previous, ...next }) })
|
||||||
}
|
}
|
||||||
@@ -89,6 +119,10 @@ export function SelectionResultsWorkbench({
|
|||||||
updateSearch({ sort: value })
|
updateSearch({ sort: value })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function handleSectorChange(value: string) {
|
||||||
|
updateSearch({ sector: value === ALL_SECTORS_VALUE ? undefined : value })
|
||||||
|
}
|
||||||
|
|
||||||
function requestLoadMore() {
|
function requestLoadMore() {
|
||||||
if (
|
if (
|
||||||
loadMoreRequestPending.current ||
|
loadMoreRequestPending.current ||
|
||||||
@@ -105,9 +139,9 @@ export function SelectionResultsWorkbench({
|
|||||||
}
|
}
|
||||||
|
|
||||||
return (
|
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)]">
|
<div className="flex min-h-0 min-w-0 flex-1 flex-col gap-2 overflow-hidden">
|
||||||
<section className="flex min-h-0 min-w-0 flex-1 flex-col overflow-hidden rounded-md border border-border/80 bg-card">
|
<section className="shrink-0 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 flex-wrap items-center gap-1.5 p-2.5">
|
||||||
<Input
|
<Input
|
||||||
aria-label="搜索命中股票"
|
aria-label="搜索命中股票"
|
||||||
className="h-11 min-w-0 flex-1 bg-background text-sm sm:max-w-[260px] md:h-8"
|
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"
|
type="search"
|
||||||
value={query}
|
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
|
<Select
|
||||||
items={signalCategoryOptions}
|
items={signalCategoryOptions}
|
||||||
onValueChange={(value) => {
|
onValueChange={(value) => {
|
||||||
@@ -170,24 +236,28 @@ export function SelectionResultsWorkbench({
|
|||||||
筛选结果 {stocksTotal} 只
|
筛选结果 {stocksTotal} 只
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div className="flex min-h-0 min-w-0 flex-1 flex-col overflow-hidden">
|
|
||||||
<SignalRecordList
|
|
||||||
isFetchNextPageError={isFetchNextPageError}
|
|
||||||
isFetchingNextPage={isFetchingNextPage}
|
|
||||||
onLoadMore={requestLoadMore}
|
|
||||||
onRetryLoadMore={() => {
|
|
||||||
loadMoreRequestPending.current = false
|
|
||||||
onLoadMore()
|
|
||||||
}}
|
|
||||||
onSelect={(stock) => setSelectedKey(getStockKey(stock))}
|
|
||||||
selectedKey={selectedStock ? getStockKey(selectedStock) : null}
|
|
||||||
stocks={visibleStocks}
|
|
||||||
/>
|
|
||||||
</div>
|
|
||||||
</section>
|
</section>
|
||||||
|
|
||||||
<SignalDetailPanel stock={selectedStock} />
|
<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}
|
||||||
|
isFetchingNextPage={isFetchingNextPage}
|
||||||
|
onLoadMore={requestLoadMore}
|
||||||
|
onRetryLoadMore={() => {
|
||||||
|
loadMoreRequestPending.current = false
|
||||||
|
onLoadMore()
|
||||||
|
}}
|
||||||
|
onSelect={(stock) => setSelectedKey(getStockKey(stock))}
|
||||||
|
selectedKey={selectedStock ? getStockKey(selectedStock) : null}
|
||||||
|
stocks={visibleStocks}
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
</section>
|
||||||
|
|
||||||
|
<SignalDetailPanel stock={selectedStock} />
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import { SelectionResultsPage } from "./selection-results-page"
|
|||||||
|
|
||||||
const routerNavigate = vi.hoisted(() => vi.fn())
|
const routerNavigate = vi.hoisted(() => vi.fn())
|
||||||
const useSelectionResults = vi.fn()
|
const useSelectionResults = vi.fn()
|
||||||
|
const useSelectionResultSectors = vi.fn()
|
||||||
const useSelectionChart = vi.fn()
|
const useSelectionChart = vi.fn()
|
||||||
const useStockSectorMembership = vi.fn()
|
const useStockSectorMembership = vi.fn()
|
||||||
const useSelectionRun = vi.fn()
|
const useSelectionRun = vi.fn()
|
||||||
@@ -31,6 +32,8 @@ vi.mock("@/features/selection/api/selection.query", async () => {
|
|||||||
return {
|
return {
|
||||||
...actual,
|
...actual,
|
||||||
useSelectionResults: (...args: unknown[]) => useSelectionResults(...args),
|
useSelectionResults: (...args: unknown[]) => useSelectionResults(...args),
|
||||||
|
useSelectionResultSectors: (...args: unknown[]) =>
|
||||||
|
useSelectionResultSectors(...args),
|
||||||
useSelectionChart: (...args: unknown[]) => useSelectionChart(...args),
|
useSelectionChart: (...args: unknown[]) => useSelectionChart(...args),
|
||||||
useSelectionRun: (...args: unknown[]) => useSelectionRun(...args),
|
useSelectionRun: (...args: unknown[]) => useSelectionRun(...args),
|
||||||
useTriggerSelectionRun: () => useTriggerSelectionRun(),
|
useTriggerSelectionRun: () => useTriggerSelectionRun(),
|
||||||
@@ -169,6 +172,11 @@ describe("SelectionResultsPage", () => {
|
|||||||
isPending: false,
|
isPending: false,
|
||||||
})
|
})
|
||||||
useSelectionResults.mockReturnValue(selectionQueryResult(selectedResult))
|
useSelectionResults.mockReturnValue(selectionQueryResult(selectedResult))
|
||||||
|
useSelectionResultSectors.mockReturnValue({
|
||||||
|
data: undefined,
|
||||||
|
isError: false,
|
||||||
|
isPending: false,
|
||||||
|
})
|
||||||
useSelectionRun.mockReturnValue(emptySelectionQuery())
|
useSelectionRun.mockReturnValue(emptySelectionQuery())
|
||||||
useTriggerSelectionRun.mockReturnValue({
|
useTriggerSelectionRun.mockReturnValue({
|
||||||
isError: false,
|
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 () => {
|
it("requests database-backed score sorting", async () => {
|
||||||
render(<SelectionResultsPage />)
|
render(<SelectionResultsPage />)
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import { useRef, useState, type RefObject } from "react"
|
|||||||
import { PageLayout } from "@/app/layout/page-layout"
|
import { PageLayout } from "@/app/layout/page-layout"
|
||||||
import {
|
import {
|
||||||
flattenSelectionResults,
|
flattenSelectionResults,
|
||||||
|
useSelectionResultSectors,
|
||||||
useSelectionResults,
|
useSelectionResults,
|
||||||
useSelectionRun,
|
useSelectionRun,
|
||||||
useTriggerSelectionRun,
|
useTriggerSelectionRun,
|
||||||
@@ -14,6 +15,8 @@ import {
|
|||||||
selectionStrategies,
|
selectionStrategies,
|
||||||
type SelectionResults,
|
type SelectionResults,
|
||||||
type SelectionResultsQuery,
|
type SelectionResultsQuery,
|
||||||
|
type SelectionSectorAggregate,
|
||||||
|
type SelectionSectors,
|
||||||
type SelectionStrategy,
|
type SelectionStrategy,
|
||||||
} from "@/features/selection/api/selection.types"
|
} from "@/features/selection/api/selection.types"
|
||||||
import { Button } from "@/shared/ui/button"
|
import { Button } from "@/shared/ui/button"
|
||||||
@@ -71,6 +74,7 @@ export function SelectionResultsPage() {
|
|||||||
pageSize: selectionResultPageSize,
|
pageSize: selectionResultPageSize,
|
||||||
...(search.search ? { search: search.search } : {}),
|
...(search.search ? { search: search.search } : {}),
|
||||||
...(search.category !== "all" ? { category: search.category } : {}),
|
...(search.category !== "all" ? { category: search.category } : {}),
|
||||||
|
...(search.sector ? { sector: search.sector } : {}),
|
||||||
sort: search.sort,
|
sort: search.sort,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -79,6 +83,10 @@ export function SelectionResultsPage() {
|
|||||||
targetTradeDate || undefined,
|
targetTradeDate || undefined,
|
||||||
resultQuery,
|
resultQuery,
|
||||||
)
|
)
|
||||||
|
const sectorsQuery = useSelectionResultSectors(
|
||||||
|
strategy,
|
||||||
|
targetTradeDate || undefined,
|
||||||
|
)
|
||||||
const resultsSnapshot = flattenSelectionResults(results.data)
|
const resultsSnapshot = flattenSelectionResults(results.data)
|
||||||
const persistedRunningRunId =
|
const persistedRunningRunId =
|
||||||
resultsSnapshot?.status === "running" ? resultsSnapshot.run_id : null
|
resultsSnapshot?.status === "running" ? resultsSnapshot.run_id : null
|
||||||
@@ -156,7 +164,11 @@ export function SelectionResultsPage() {
|
|||||||
setRerunDialogOpen(false)
|
setRerunDialogOpen(false)
|
||||||
trigger.reset()
|
trigger.reset()
|
||||||
void navigate({
|
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)}
|
isFetchingNextPage={Boolean(listQuery.isFetchingNextPage)}
|
||||||
onLoadMore={() => listQuery.fetchNextPage()}
|
onLoadMore={() => listQuery.fetchNextPage()}
|
||||||
result={displayedResult}
|
result={displayedResult}
|
||||||
|
sectors={sectorsQuery.data}
|
||||||
/>
|
/>
|
||||||
) : null}
|
) : null}
|
||||||
</div>
|
</div>
|
||||||
@@ -376,8 +389,11 @@ interface ResultStateProps {
|
|||||||
isFetchingNextPage: boolean
|
isFetchingNextPage: boolean
|
||||||
onLoadMore: () => void | Promise<unknown>
|
onLoadMore: () => void | Promise<unknown>
|
||||||
result: SelectionResults
|
result: SelectionResults
|
||||||
|
sectors: SelectionSectors | undefined
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const EMPTY_SECTOR_OPTIONS: ReadonlyArray<SelectionSectorAggregate> = []
|
||||||
|
|
||||||
function parseLocalTradeDate(value: string): Date | undefined {
|
function parseLocalTradeDate(value: string): Date | undefined {
|
||||||
const match = /^(\d{4})-(\d{2})-(\d{2})$/.exec(value)
|
const match = /^(\d{4})-(\d{2})-(\d{2})$/.exec(value)
|
||||||
if (!match) return undefined
|
if (!match) return undefined
|
||||||
@@ -511,6 +527,7 @@ function ResultState({
|
|||||||
isFetchingNextPage,
|
isFetchingNextPage,
|
||||||
onLoadMore,
|
onLoadMore,
|
||||||
result,
|
result,
|
||||||
|
sectors,
|
||||||
}: ResultStateProps) {
|
}: ResultStateProps) {
|
||||||
return result.signal_count === 0 ? (
|
return result.signal_count === 0 ? (
|
||||||
<NoSignalState />
|
<NoSignalState />
|
||||||
@@ -521,6 +538,7 @@ function ResultState({
|
|||||||
isFetchingNextPage={isFetchingNextPage}
|
isFetchingNextPage={isFetchingNextPage}
|
||||||
onLoadMore={onLoadMore}
|
onLoadMore={onLoadMore}
|
||||||
result={result}
|
result={result}
|
||||||
|
sectorOptions={sectors?.sectors ?? EMPTY_SECTOR_OPTIONS}
|
||||||
/>
|
/>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -124,7 +124,11 @@ const selectionRoute = createRoute({
|
|||||||
)
|
)
|
||||||
? (rawSort as SelectionSort)
|
? (rawSort as SelectionSort)
|
||||||
: "code"
|
: "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,
|
component: SelectionResultsPage,
|
||||||
})
|
})
|
||||||
|
|||||||
Reference in New Issue
Block a user