feat(selection): 集成 B1 FastDTW 图形评分

This commit is contained in:
yuxuanhui
2026-08-31 16:14:16 +08:00
parent 86762c0d9a
commit 6ce291e242
51 changed files with 2917 additions and 41 deletions
@@ -26,6 +26,7 @@ class Settings(BaseSettings):
market_data_advisory_lock_key: int = 7_380_521
selection_max_workers: int = Field(default=4, ge=1)
selection_batch_size: int = Field(default=200, ge=1)
selection_pattern_scoring_enabled: bool = True
model_config = SettingsConfigDict(
env_file=".env",
@@ -2,6 +2,7 @@
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
@@ -22,6 +23,24 @@ from sqlalchemy.dialects.postgresql import JSONB
metadata = MetaData()
_PATTERN_BREAKDOWN_CHECK = """
match_breakdown IS NULL OR (
jsonb_typeof(match_breakdown) = 'object'
AND CASE WHEN jsonb_typeof(match_breakdown -> 'trend_structure') = 'number'
THEN (match_breakdown ->> 'trend_structure')::numeric BETWEEN 0 AND 100
ELSE false END
AND CASE WHEN jsonb_typeof(match_breakdown -> 'kdj_state') = 'number'
THEN (match_breakdown ->> 'kdj_state')::numeric BETWEEN 0 AND 100
ELSE false END
AND CASE WHEN jsonb_typeof(match_breakdown -> 'volume_pattern') = 'number'
THEN (match_breakdown ->> 'volume_pattern')::numeric BETWEEN 0 AND 100
ELSE false END
AND CASE WHEN jsonb_typeof(match_breakdown -> 'price_shape') = 'number'
THEN (match_breakdown ->> 'price_shape')::numeric BETWEEN 0 AND 100
ELSE false END
)
"""
market_stock = Table(
"market_stock",
metadata,
@@ -153,8 +172,51 @@ selection_run_item = Table(
Column("status", String(32), nullable=False),
Column("signal_count", Integer, nullable=False, server_default="0"),
Column("reason", Text),
Column("score_status", String(32), nullable=False, server_default="not_executed"),
Column("score_value", Numeric(5, 2)),
Column("score_threshold", Numeric(5, 2)),
Column("score_version", String(64)),
Column("match_case_id", String(32)),
Column("match_case_name", String(128)),
Column("match_case_breakout_date", Date),
Column("match_breakdown", JSONB),
Column("score_reason", Text),
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
PrimaryKeyConstraint("run_id", "ts_code"),
CheckConstraint(
"score_status IN ('not_executed', 'matched', 'below_threshold', 'failed')",
name="ck_selection_run_item_score_status",
),
CheckConstraint(
"score_value IS NULL OR score_value BETWEEN 0 AND 100",
name="ck_selection_run_item_score_value_range",
),
CheckConstraint(
"score_threshold IS NULL OR score_threshold BETWEEN 0 AND 100",
name="ck_selection_run_item_score_threshold_range",
),
CheckConstraint(
_PATTERN_BREAKDOWN_CHECK,
name="ck_selection_run_item_breakdown_range",
),
CheckConstraint(
"(score_status = 'not_executed' AND score_value IS NULL AND score_threshold IS NULL "
"AND score_version IS NULL AND match_case_id IS NULL AND match_case_name IS NULL "
"AND match_case_breakout_date IS NULL AND match_breakdown IS NULL "
"AND score_reason IS NULL) "
"OR (score_status = 'failed' AND score_value IS NULL AND score_threshold IS NULL "
"AND score_version IS NULL AND match_case_id IS NULL AND match_case_name IS NULL "
"AND match_case_breakout_date IS NULL AND match_breakdown IS NULL "
"AND score_reason IS NOT NULL) "
"OR (score_status IN ('matched', 'below_threshold') AND score_value IS NOT NULL "
"AND score_threshold IS NOT NULL AND score_version IS NOT NULL "
"AND match_case_id IS NOT NULL AND match_case_name IS NOT NULL "
"AND match_case_breakout_date IS NOT NULL AND match_breakdown IS NOT NULL "
"AND score_reason IS NULL "
"AND ((score_status = 'matched' AND score_value >= score_threshold) "
"OR (score_status = 'below_threshold' AND score_value < score_threshold)))",
name="ck_selection_run_item_score_shape",
),
)
selection_signal = Table(
@@ -223,6 +285,12 @@ Index(
selection_run.c.target_trade_date,
)
Index("ix_selection_run_item_status", selection_run_item.c.run_id, selection_run_item.c.status)
Index(
"ix_selection_run_item_score",
selection_run_item.c.run_id,
selection_run_item.c.score_value.desc(),
selection_run_item.c.ts_code,
)
Index(
"ix_selection_signal_strategy_date",
selection_signal.c.strategy,
@@ -11,6 +11,12 @@ from datetime import date
from typing import Literal, Protocol, cast
from ..domain.models import SelectionEvaluation, StockHistory
from ..domain.pattern_scoring import (
PatternCase,
PatternCaseLibraryLoader,
PatternScore,
PatternScorer,
)
from ..domain.runs import (
BatchSelectionRunStore,
BatchSelectionUniverseReader,
@@ -54,7 +60,10 @@ class RunZhixingB1:
reader: SelectionUniverseReader,
store: SelectionRunStore,
evaluator: SelectionEvaluator | None = None,
pattern_case_loader: PatternCaseLibraryLoader | None = None,
pattern_scorer: PatternScorer | None = None,
*,
pattern_scoring_enabled: bool = False,
max_workers: int = 4,
batch_size: int = 200,
) -> None:
@@ -67,6 +76,9 @@ class RunZhixingB1:
self.reader = reader
self.store = store
self.evaluator = evaluator or EvaluateZhixingB1(reader)
self.pattern_case_loader = pattern_case_loader
self.pattern_scorer = pattern_scorer
self.pattern_scoring_enabled = pattern_scoring_enabled
self.max_workers = max_workers
self.batch_size = batch_size
@@ -106,7 +118,9 @@ class RunZhixingB1:
read_seconds = 0.0
evaluate_seconds = 0.0
persist_seconds = 0.0
scoring_seconds = 0.0
try:
pattern_cases, pattern_library_error = self._prepare_pattern_cases(prepared.run.id)
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
for batch_stocks in _chunks(stocks, self.batch_size):
read_started = time.perf_counter()
@@ -120,24 +134,39 @@ class RunZhixingB1:
)
evaluate_started = time.perf_counter()
evaluations = tuple(
executor.map(
self._evaluate_stock,
batch_stocks,
histories,
[prepared.source.target_trade_date] * len(batch_stocks),
)
)
evaluate_seconds += time.perf_counter() - evaluate_started
scoring_started = time.perf_counter()
items = tuple(
_to_item(
stock.ts_code,
stock.name,
evaluation,
)
for stock, evaluation in zip(
batch_stocks,
executor.map(
self._evaluate_stock,
batch_stocks,
histories,
[prepared.source.target_trade_date] * len(batch_stocks),
pattern_score=self._score_stock(
prepared.run.id,
stock,
history,
evaluation,
pattern_cases,
pattern_library_error,
),
)
for stock, history, evaluation in zip(
batch_stocks,
histories,
evaluations,
strict=True,
)
)
evaluate_seconds += time.perf_counter() - evaluate_started
scoring_seconds += time.perf_counter() - scoring_started
evaluated_count += len(items)
selected_stock_count += sum(item.status == "selected" for item in items)
@@ -184,7 +213,7 @@ class RunZhixingB1:
logger.info(
"selection_run_summary run_id=%s stock_count=%d history_rows=%d "
"batch_count=%d worker_count=%d read_seconds=%.3f "
"evaluate_seconds=%.3f persist_seconds=%.3f",
"evaluate_seconds=%.3f scoring_seconds=%.3f persist_seconds=%.3f",
prepared.run.id,
len(stocks),
history_rows,
@@ -192,9 +221,64 @@ class RunZhixingB1:
self.max_workers,
read_seconds,
evaluate_seconds,
scoring_seconds,
persist_seconds,
)
def _prepare_pattern_cases(
self,
run_id: str,
) -> tuple[tuple[PatternCase, ...] | None, str | None]:
"""Load the complete case library once without failing selection."""
if not self.pattern_scoring_enabled:
return None, None
if self.pattern_case_loader is None or self.pattern_scorer is None:
reason = "pattern scoring is enabled but not configured"
logger.error("selection_pattern_library_failed run_id=%s reason=%s", run_id, reason)
return None, reason
try:
return self.pattern_case_loader.load(), None
except Exception as exc: # noqa: BLE001 - scoring enrichment must not fail selection
reason = _safe_item_error(exc)
logger.warning(
"selection_pattern_library_failed run_id=%s error_type=%s reason=%s",
run_id,
exc.__class__.__name__,
reason,
)
return None, reason
def _score_stock(
self,
run_id: str,
stock: SelectionStock,
history: StockHistory | None,
evaluation: SelectionEvaluation,
cases: tuple[PatternCase, ...] | None,
library_error: str | None,
) -> PatternScore:
"""Score one selected stock once and isolate enrichment failures."""
if not self.pattern_scoring_enabled or evaluation.status != "selected":
return PatternScore()
if library_error is not None:
return PatternScore.failed(library_error)
if history is None or cases is None or self.pattern_scorer is None:
return PatternScore.failed("pattern scoring history or case library is unavailable")
try:
return self.pattern_scorer.score(history, cases)
except Exception as exc: # noqa: BLE001 - one score must not fail the selection run
reason = _safe_item_error(exc)
logger.warning(
"selection_pattern_score_failed run_id=%s ts_code=%s error_type=%s reason=%s",
run_id,
stock.ts_code,
exc.__class__.__name__,
reason,
)
return PatternScore.failed(reason)
def _load_histories(
self,
stocks: Sequence[SelectionStock],
@@ -287,7 +371,13 @@ class RunZhixingB1:
return self.store.get_latest_run(strategy, target_trade_date, query=query)
def _to_item(ts_code: str, name: str, evaluation: SelectionEvaluation) -> SelectionRunItem:
def _to_item(
ts_code: str,
name: str,
evaluation: SelectionEvaluation,
*,
pattern_score: PatternScore | None = None,
) -> SelectionRunItem:
"""Translate a single-stock domain result into a stored item."""
return SelectionRunItem(
@@ -296,6 +386,7 @@ def _to_item(ts_code: str, name: str, evaluation: SelectionEvaluation) -> Select
status=evaluation.status,
signal_count=len(evaluation.signals),
reason=evaluation.reason,
pattern_score=pattern_score or PatternScore(),
signals=evaluation.signals,
)
@@ -0,0 +1,572 @@
"""Versioned Zhixing B1 pattern-similarity scoring.
The module deliberately keeps the algorithm and its ten case definitions in
one bounded-context-owned contract. Infrastructure supplies qfq histories;
the scorer performs no I/O and never falls back to a different DTW algorithm.
"""
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import date
from math import isfinite
from numbers import Real
from typing import Literal, Protocol, cast
import numpy as np
import pandas as pd
from fastdtw import fastdtw # type: ignore[reportMissingTypeStubs]
from .models import SelectionBar, StockHistory
PATTERN_SCORING_VERSION = "zhixing_b1_pattern_fastdtw_v1"
PATTERN_LOOKBACK_DAYS = 25
PATTERN_SCORE_THRESHOLD = 60.0
PATTERN_FASTDTW_RADIUS = 1
PATTERN_WEIGHTS = {
"trend_structure": 0.10,
"kdj_state": 0.20,
"volume_pattern": 0.25,
"price_shape": 0.45,
}
PATTERN_TOLERANCES = {
"trend_ratio": 0.10,
"price_bias": 10.0,
"trend_spread": 10.0,
"j_value": 30.0,
"drawdown": 15.0,
}
PatternScoreStatus = Literal["not_executed", "matched", "below_threshold", "failed"]
@dataclass(frozen=True, slots=True)
class PatternCaseDefinition:
"""A versioned pattern template and its exclusive breakout boundary."""
id: str
name: str
ts_code: str
breakout_date: date
lookback_days: int = PATTERN_LOOKBACK_DAYS
ZHIXING_B1_PATTERN_CASES: tuple[PatternCaseDefinition, ...] = (
PatternCaseDefinition("case_001", "华纳药厂", "688799.SH", date(2025, 5, 12)),
PatternCaseDefinition("case_002", "宁波韵升", "600366.SH", date(2025, 8, 6)),
PatternCaseDefinition("case_003", "微芯生物", "688321.SH", date(2025, 6, 20)),
PatternCaseDefinition("case_004", "方正科技", "600601.SH", date(2025, 7, 23)),
PatternCaseDefinition("case_006", "国轩高科", "002074.SZ", date(2025, 8, 4)),
PatternCaseDefinition("case_007", "野马电池", "605378.SH", date(2025, 8, 1)),
PatternCaseDefinition("case_008", "光电股份", "600184.SH", date(2025, 7, 10)),
PatternCaseDefinition("case_009", "新瀚新材", "301076.SZ", date(2025, 8, 1)),
PatternCaseDefinition("case_010", "昂利康", "002940.SZ", date(2025, 7, 11)),
PatternCaseDefinition("case_011", "航天发展", "000547.SZ", date(2025, 11, 12)),
)
@dataclass(frozen=True, slots=True)
class PatternFeatures:
"""Immutable, finite-or-null features used by the matcher."""
trend_structure: Mapping[str, float | bool | None]
kdj_state: Mapping[str, float | bool | str | None]
volume_pattern: Mapping[str, float | bool | str | int | None]
price_shape: Mapping[str, float | str | int | tuple[float, ...] | None]
@dataclass(frozen=True, slots=True)
class PatternCase:
"""One complete case history with precomputed immutable features."""
definition: PatternCaseDefinition
history: StockHistory
features: PatternFeatures
@dataclass(frozen=True, slots=True)
class PatternScoreBreakdown:
"""Finite 0-100 scores for the four versioned pattern dimensions."""
trend_structure: float
kdj_state: float
volume_pattern: float
price_shape: float
def __post_init__(self) -> None:
"""Reject non-finite or out-of-range values before persistence."""
for name in ("trend_structure", "kdj_state", "volume_pattern", "price_shape"):
_validate_score(getattr(self, name), name)
def as_dict(self) -> dict[str, float]:
"""Return the JSONB/HTTP field names without exposing dataclass internals."""
return {
"trend_structure": self.trend_structure,
"kdj_state": self.kdj_state,
"volume_pattern": self.volume_pattern,
"price_shape": self.price_shape,
}
@dataclass(frozen=True, slots=True)
class PatternScore:
"""One stock-level scoring outcome independent of selection status."""
status: PatternScoreStatus = "not_executed"
value: float | None = None
threshold: float | None = None
version: str | None = None
case: PatternCaseDefinition | None = None
breakdown: PatternScoreBreakdown | None = None
reason: str | None = None
def __post_init__(self) -> None:
"""Enforce complete successful results and value-free failures."""
if self.status in {"matched", "below_threshold"}:
if (
self.value is None
or self.threshold is None
or self.version is None
or self.case is None
or self.breakdown is None
):
raise ValueError("computed pattern score requires complete match context")
_validate_score(self.value, "value")
_validate_score(self.threshold, "threshold")
if (self.value >= self.threshold) != (self.status == "matched"):
raise ValueError("pattern score status must agree with threshold")
elif any(
value is not None
for value in (self.value, self.threshold, self.version, self.case, self.breakdown)
):
raise ValueError("uncomputed pattern score cannot carry match values")
if self.status == "failed" and not (self.reason and self.reason.strip()):
raise ValueError("failed pattern score requires a safe reason")
if self.status == "not_executed" and self.reason is not None:
raise ValueError("not-executed pattern score cannot carry a reason")
@classmethod
def failed(cls, reason: str) -> PatternScore:
"""Create a safe failure without retaining raw exception details."""
normalized = " ".join(reason.split())[:500] or "pattern scoring failed"
return cls(status="failed", reason=normalized)
class PatternCaseLibraryError(RuntimeError):
"""The immutable ten-case library could not be loaded completely."""
class PatternScoringError(RuntimeError):
"""A candidate could not be scored under the versioned algorithm."""
class PatternCaseLibraryLoader(Protocol):
"""Load the complete versioned case library once for a selection run."""
def load(self) -> tuple[PatternCase, ...]: ...
class PatternScorer(Protocol):
"""Score one selected stock against an already prepared case library."""
def score(self, history: StockHistory, cases: Sequence[PatternCase]) -> PatternScore: ...
class PatternFeatureExtractor:
"""Reproduce the legacy 25-row feature formulas with finite outputs."""
def extract(self, history: StockHistory) -> PatternFeatures:
"""Extract features from the latest 25 ascending, complete OHLCV rows.
Raises:
PatternScoringError: If the history does not contain exactly the
required complete window or dates are not strictly ascending.
"""
bars = history.bars[-PATTERN_LOOKBACK_DAYS:]
_validate_window(bars, history.ts_code)
frame = pd.DataFrame(
{
"open": [bar.open for bar in bars],
"high": [bar.high for bar in bars],
"low": [bar.low for bar in bars],
"close": [bar.close for bar in bars],
"volume": [bar.volume for bar in bars],
},
dtype=float,
)
white = frame["close"].ewm(span=10, adjust=False).mean()
white = white.ewm(span=10, adjust=False).mean()
yellow = (
frame["close"].rolling(14, min_periods=14).mean()
+ frame["close"].rolling(28, min_periods=28).mean()
+ frame["close"].rolling(57, min_periods=57).mean()
+ frame["close"].rolling(114, min_periods=114).mean()
) / 4.0
frame["short_term_trend"] = white
frame["bull_bear_line"] = yellow
frame = _legacy_kdj(frame)
return PatternFeatures(
trend_structure=_trend_features(frame),
kdj_state=_kdj_features(frame),
volume_pattern=_volume_features(frame),
price_shape=_price_features(frame),
)
class ZhixingB1PatternScorer:
"""Select the stable best case using working scalar FastDTW radius one."""
def __init__(self, extractor: PatternFeatureExtractor | None = None) -> None:
"""Inject an extractor for deterministic unit tests."""
self.extractor = extractor or PatternFeatureExtractor()
def score(self, history: StockHistory, cases: Sequence[PatternCase]) -> PatternScore:
"""Score one selected history once against all ten ordered cases.
Raises:
PatternScoringError: If the library is incomplete or FastDTW
cannot produce a finite distance. No alternative algorithm is
used when FastDTW fails.
"""
if tuple(case.definition for case in cases) != ZHIXING_B1_PATTERN_CASES:
raise PatternScoringError("pattern case library is incomplete or out of order")
candidate = self.extractor.extract(history)
best: tuple[float, PatternCase, PatternScoreBreakdown] | None = None
for case in cases:
breakdown = _match(candidate, case.features)
value = round(
sum(
breakdown.as_dict()[name] / 100.0 * weight
for name, weight in PATTERN_WEIGHTS.items()
)
* 100.0,
2,
)
_validate_score(value, "value")
if best is None or value > best[0]:
best = (value, case, breakdown)
if best is None:
raise PatternScoringError("pattern case library is empty")
value, case, breakdown = best
return PatternScore(
status="matched" if value >= PATTERN_SCORE_THRESHOLD else "below_threshold",
value=value,
threshold=PATTERN_SCORE_THRESHOLD,
version=PATTERN_SCORING_VERSION,
case=case.definition,
breakdown=breakdown,
)
def build_pattern_case(
definition: PatternCaseDefinition,
history: StockHistory,
extractor: PatternFeatureExtractor | None = None,
) -> PatternCase:
"""Validate and precompute one versioned case for run-wide reuse."""
if history.ts_code != definition.ts_code:
raise PatternCaseLibraryError(f"case {definition.id} code does not match definition")
if len(history.bars) != definition.lookback_days:
raise PatternCaseLibraryError(
f"case {definition.id} requires {definition.lookback_days} complete rows"
)
try:
features = (extractor or PatternFeatureExtractor()).extract(history)
except PatternScoringError as exc:
raise PatternCaseLibraryError(f"case {definition.id} history is invalid") from exc
return PatternCase(definition=definition, history=history, features=features)
def _match(candidate: PatternFeatures, case: PatternFeatures) -> PatternScoreBreakdown:
return PatternScoreBreakdown(
trend_structure=round(_trend_similarity(candidate, case) * 100.0, 2),
kdj_state=round(_kdj_similarity(candidate, case) * 100.0, 2),
volume_pattern=round(_volume_similarity(candidate, case) * 100.0, 2),
price_shape=round(_price_similarity(candidate, case) * 100.0, 2),
)
def _trend_similarity(candidate: PatternFeatures, case: PatternFeatures) -> float:
c, s = candidate.trend_structure, case.trend_structure
values = [
_difference_similarity(c.get("short_vs_bullbear"), s.get("short_vs_bullbear"), 0.10),
_slope_similarity(c.get("short_slope"), s.get("short_slope")),
1.0 if c.get("is_in_bowl") == s.get("is_in_bowl") else 0.2,
_difference_similarity(c.get("price_vs_short_pct"), s.get("price_vs_short_pct"), 10.0),
_difference_similarity(c.get("trend_spread_pct"), s.get("trend_spread_pct"), 10.0),
_difference_similarity(c.get("price_bias_pct"), s.get("price_bias_pct"), 10.0),
]
return float(np.mean(values))
def _kdj_similarity(candidate: PatternFeatures, case: PatternFeatures) -> float:
c, s = candidate.kdj_state, case.kdj_state
values = [
1.0 if c.get("j_position") == s.get("j_position") else 0.4,
_difference_similarity(c.get("j_value"), s.get("j_value"), 30.0),
1.0 if c.get("k_cross_d") == s.get("k_cross_d") else 0.6,
1.0 if c.get("j_rebound") == s.get("j_rebound") else 0.7,
]
return float(np.mean(values))
def _volume_similarity(candidate: PatternFeatures, case: PatternFeatures) -> float:
c, s = candidate.volume_pattern, case.volume_pattern
values = [
_difference_similarity(c.get("avg_volume_ratio"), s.get("avg_volume_ratio"), 1.5),
1.0 if c.get("shrink_then_expand") == s.get("shrink_then_expand") else 0.5,
1.0 if c.get("volume_trend") == s.get("volume_trend") else 0.6,
_difference_similarity(c.get("max_volume_ratio"), s.get("max_volume_ratio"), 3.0),
]
return float(np.mean(values))
def _price_similarity(candidate: PatternFeatures, case: PatternFeatures) -> float:
c, s = candidate.price_shape, case.price_shape
candidate_curve = cast(tuple[float, ...], c.get("normalized_curve"))
case_curve = cast(tuple[float, ...], s.get("normalized_curve"))
distance, _path = _fastdtw()( # radius and scalar metric are versioned behavior
candidate_curve,
case_curve,
radius=PATTERN_FASTDTW_RADIUS,
dist=_scalar_euclidean,
)
if not isfinite(float(distance)):
raise PatternScoringError("FastDTW returned a non-finite distance")
values = [
max(0.0, 1.0 - float(distance) / max(len(candidate_curve), len(case_curve))),
_difference_similarity(c.get("max_drawdown"), s.get("max_drawdown"), 15.0),
_difference_similarity(c.get("breakout_strength"), s.get("breakout_strength"), 5.0),
1.0 if c.get("overall_trend") == s.get("overall_trend") else 0.5,
_difference_similarity(c.get("consolidation_days"), s.get("consolidation_days"), 10.0),
]
return float(np.mean(values))
def _fastdtw() -> Callable[..., tuple[float, list[tuple[int, int]]]]:
"""Give the untyped extension one narrow, checked call signature."""
return cast(Callable[..., tuple[float, list[tuple[int, int]]]], fastdtw)
def _scalar_euclidean(left: float, right: float) -> float:
"""Return Euclidean distance for scalar one-dimensional curve points."""
return abs(float(left) - float(right))
def _difference_similarity(left: object, right: object, tolerance: float) -> float:
left_number = _finite_float(left)
right_number = _finite_float(right)
if left_number is None or right_number is None:
return 0.0
return max(0.0, 1.0 - abs(left_number - right_number) / tolerance)
def _slope_similarity(left: object, right: object) -> float:
left_number = _finite_float(left)
right_number = _finite_float(right)
if left_number is None or right_number is None:
return 0.0
difference = abs(left_number - right_number)
if (left_number > 0) == (right_number > 0):
return max(0.7, 1.0 - difference / 10.0)
return max(0.0, 0.3 - difference / 20.0)
def _trend_features(frame: pd.DataFrame) -> dict[str, float | bool | None]:
latest = frame.iloc[-1]
short = float(latest["short_term_trend"])
bullbear = float(latest["bull_bear_line"])
short_previous = float(frame["short_term_trend"].iloc[-5])
bullbear_previous = float(frame["bull_bear_line"].iloc[-5])
close = float(latest["close"])
average = (short + bullbear) / 2.0
return {
"short_vs_bullbear": _finite_round(short / bullbear if bullbear else 1.0, 4),
"short_slope": _finite_round(
(short / short_previous - 1.0) * 100.0 if short_previous else 0.0,
4,
),
"bullbear_slope": _finite_round(
(bullbear / bullbear_previous - 1.0) * 100.0 if bullbear_previous else 0.0,
4,
),
"price_vs_short_pct": _finite_round((close - short) / short * 100.0 if short else 0.0, 4),
"price_vs_bullbear_pct": _finite_round(
(close - bullbear) / bullbear * 100.0 if bullbear else 0.0,
4,
),
"is_in_bowl": bool(short > close > bullbear),
"trend_spread_pct": _finite_round(
(short - bullbear) / bullbear * 100.0 if bullbear else 0.0,
4,
),
"price_bias_pct": _finite_round((close - average) / average * 100.0 if average else 0.0, 4),
}
def _kdj_features(frame: pd.DataFrame) -> dict[str, float | bool | str | None]:
latest = frame.iloc[-1]
j_values = frame["J"].to_numpy(dtype=float)
recent = j_values[-5:]
j_trend = float(np.polyfit(np.arange(5), recent, 1)[0]) if np.isfinite(recent).all() else 0.0
previous = frame.iloc[-2]
j_value = float(latest["J"]) if pd.notna(latest["J"]) else 50.0
return {
"j_value": _finite_round(j_value, 2),
"j_trend": _finite_round(j_trend, 4),
"j_min_lookback": _finite_round(float(frame["J"].min()), 2),
"k_cross_d": bool(previous["K"] < previous["D"] and latest["K"] > latest["D"]),
"j_position": "低位" if j_value <= 20 else ("高位" if j_value >= 80 else "中位"),
"j_rebound": bool(j_values[-1] > j_values[-3]),
}
def _volume_features(frame: pd.DataFrame) -> dict[str, float | bool | str | int | None]:
volumes = frame["volume"].to_numpy(dtype=float)
recent_average = float(np.mean(volumes[-10:]))
before_average = float(np.mean(volumes[-20:-10]))
average_ratio = recent_average / before_average if before_average > 0 else 1.0
ratios = [
volumes[index] / volumes[index - 1] for index in range(1, 20) if volumes[index - 1] > 0
]
midpoint = len(volumes) // 2
early, late = float(np.mean(volumes[:midpoint])), float(np.mean(volumes[midpoint:]))
shrink_expand = bool(late > early * 1.3 and early < float(np.mean(volumes)) * 0.9)
key_count = sum(
1
for index in range(1, len(frame))
if frame["volume"].iloc[index] > frame["volume"].iloc[index - 1] * 2
and frame["close"].iloc[index] > frame["open"].iloc[index]
)
slope = float(np.polyfit(np.arange(len(volumes)), volumes, 1)[0])
slope_pct = slope / float(np.mean(volumes)) * 100.0 if float(np.mean(volumes)) > 0 else 0.0
trend = (
"持续放量"
if slope_pct > 5
else "持续缩量"
if slope_pct < -5
else "缩量后放量"
if shrink_expand
else "量能平稳"
)
return {
"avg_volume_ratio": _finite_round(average_ratio, 2),
"max_volume_ratio": _finite_round(max(ratios, default=1.0), 2),
"volume_trend": trend,
"key_candles_count": key_count,
"shrink_then_expand": shrink_expand,
}
def _price_features(frame: pd.DataFrame) -> dict[str, float | str | int | tuple[float, ...] | None]:
closes = frame["close"].to_numpy(dtype=float)
minimum, maximum = float(closes.min()), float(closes.max())
normalized = (
tuple(float(value) for value in (closes - minimum) / (maximum - minimum))
if maximum > minimum
else (0.0,) * len(closes)
)
peak = np.maximum.accumulate(closes)
max_drawdown = float(((peak - closes) / peak).max()) * 100.0
breakout = (closes[-1] / closes[-2] - 1.0) * 100.0
returns = np.diff(closes) / closes[:-1]
volatility = float(np.std(returns)) * 100.0
consolidation, current = 0, 0
for index in range(len(frame) - 5):
window = closes[index : index + 5]
if window.max() > 0 and (window.max() - window.min()) / window.max() < 0.05:
current += 1
consolidation = max(consolidation, current)
else:
current = 0
trend = (
"上升"
if closes[-1] > closes[0] * 1.05
else "下降"
if closes[-1] < closes[0] * 0.95
else "震荡"
)
return {
"consolidation_days": consolidation,
"max_drawdown": _finite_round(max_drawdown, 2),
"breakout_strength": _finite_round(breakout, 2),
"normalized_curve": normalized,
"volatility": _finite_round(volatility, 4),
"overall_trend": trend,
}
def _legacy_kdj(frame: pd.DataFrame) -> pd.DataFrame:
low = frame["low"].rolling(window=9, min_periods=1).min()
high = frame["high"].rolling(window=9, min_periods=1).max()
rsv = ((frame["close"] - low) / (high - low + 1e-9) * 100.0).to_numpy(dtype=float)
k = np.empty(len(rsv), dtype=float)
d = np.empty(len(rsv), dtype=float)
k[0] = d[0] = 50.0
for index in range(1, len(rsv)):
k[index] = 2.0 / 3.0 * k[index - 1] + 1.0 / 3.0 * rsv[index]
d[index] = 2.0 / 3.0 * d[index - 1] + 1.0 / 3.0 * k[index]
return frame.assign(K=k, D=d, J=3.0 * k - 2.0 * d)
def _validate_window(bars: Sequence[SelectionBar], ts_code: str) -> None:
if len(bars) != PATTERN_LOOKBACK_DAYS:
raise PatternScoringError(f"{ts_code} requires {PATTERN_LOOKBACK_DAYS} complete rows")
if any(
left.trade_date >= right.trade_date for left, right in zip(bars, bars[1:], strict=False)
):
raise PatternScoringError(f"{ts_code} pattern rows must be strictly ascending")
if any(
value is None
for bar in bars
for value in (bar.open, bar.high, bar.low, bar.close, bar.volume)
):
raise PatternScoringError(f"{ts_code} pattern rows require complete OHLCV")
def _finite_float(value: object) -> float | None:
if isinstance(value, bool) or not isinstance(value, Real):
return None
number = float(value)
return number if isfinite(number) else None
def _finite_round(value: float, digits: int) -> float | None:
return round(float(value), digits) if isfinite(float(value)) else None
def _validate_score(value: float, name: str) -> None:
if not isfinite(value) or value < 0 or value > 100:
raise ValueError(f"{name} must be finite and between 0 and 100")
__all__ = [
"PATTERN_FASTDTW_RADIUS",
"PATTERN_LOOKBACK_DAYS",
"PATTERN_SCORE_THRESHOLD",
"PATTERN_SCORING_VERSION",
"PatternCase",
"PatternCaseDefinition",
"PatternCaseLibraryError",
"PatternCaseLibraryLoader",
"PatternFeatureExtractor",
"PatternFeatures",
"PatternScore",
"PatternScoreBreakdown",
"PatternScorer",
"PatternScoringError",
"ZHIXING_B1_PATTERN_CASES",
"ZhixingB1PatternScorer",
"build_pattern_case",
]
@@ -9,10 +9,12 @@ from decimal import Decimal
from typing import Literal, Protocol
from .models import SelectionEvaluationStatus, SelectionSignal, StockHistory
from .pattern_scoring import PatternScore
SelectionRunStatus = Literal["running", "success", "partial_success", "failed"]
SelectionRunItemStatus = SelectionEvaluationStatus
SelectionSignalCategoryFilter = Literal["pullback", "oversold", "original"]
SelectionResultSort = Literal["code", "score_desc", "score_asc"]
@dataclass(frozen=True, slots=True)
@@ -23,6 +25,7 @@ class SelectionResultQuery:
page_size: int = 10
search: str | None = None
category: SelectionSignalCategoryFilter | None = None
sort: SelectionResultSort = "code"
@dataclass(frozen=True, slots=True)
@@ -54,6 +57,7 @@ class SelectionRunItem:
status: SelectionRunItemStatus
signal_count: int = 0
reason: str | None = None
pattern_score: PatternScore = field(default_factory=PatternScore)
signals: tuple[SelectionSignal, ...] = field(default_factory=tuple)
@@ -12,6 +12,12 @@ import psycopg
from ....bootstrap.config import Settings
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
from ..domain.pattern_scoring import (
ZHIXING_B1_PATTERN_CASES,
PatternCase,
PatternCaseLibraryError,
build_pattern_case,
)
from ..domain.ports import MarketDataReaderError
from ..domain.runs import SelectionExecutionSource, SelectionStock
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
@@ -103,6 +109,38 @@ WHERE stock.is_active = true
ORDER BY stock.ts_code
"""
_PATTERN_CASES_QUERY = """
WITH case_definition AS (
SELECT *
FROM unnest(%s::text[], %s::text[], %s::date[], %s::integer[])
AS definition(case_id, ts_code, breakout_date, lookback_days)
), ranked AS (
SELECT
definition.case_id,
bar.ts_code,
bar.trade_date,
bar.open,
bar.high,
bar.low,
bar.close,
bar.vol,
row_number() OVER (
PARTITION BY definition.case_id
ORDER BY bar.trade_date DESC
) AS recency_rank,
definition.lookback_days
FROM case_definition AS definition
JOIN market_daily_bar AS bar
ON bar.ts_code = definition.ts_code
AND bar.source_adj = 'qfq'
AND bar.trade_date < definition.breakout_date
)
SELECT case_id, ts_code, trade_date, open, high, low, close, vol
FROM ranked
WHERE recency_rank <= lookback_days
ORDER BY case_id ASC, trade_date ASC
"""
def _as_date(value: object) -> date:
"""Convert a PostgreSQL date-like scalar to a date."""
@@ -386,3 +424,106 @@ class PostgresMarketDataReader:
raise
except Exception as exc: # noqa: BLE001 - normalize pool/driver failures
raise SelectionReaderError("selection database operation failed") from exc
class PostgresPatternCaseLibraryLoader:
"""Build the complete versioned FastDTW case library from PostgreSQL qfq bars."""
def __init__(
self,
settings: Settings | str,
*,
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
) -> None:
"""Create a loader sharing the process selection connection pool."""
self.database_url = settings.database_url if isinstance(settings, Settings) else settings
if isinstance(pool, SelectionPostgresPool):
self.pool: SelectionPostgresPool | None = pool
elif pool is not None:
self.pool = SelectionPostgresPool(self.database_url, max_connections=1, pool=pool)
else:
self.pool = None
def load(self) -> tuple[PatternCase, ...]:
"""Load all ten exclusive pre-breakout windows exactly once.
Returns:
Ordered, feature-precomputed cases matching the versioned definitions.
Raises:
PatternCaseLibraryError: If the query fails or any case lacks a
complete finite 25-row qfq window.
"""
definitions = ZHIXING_B1_PATTERN_CASES
parameters = (
[definition.id for definition in definitions],
[definition.ts_code for definition in definitions],
[definition.breakout_date for definition in definitions],
[definition.lookback_days for definition in definitions],
)
try:
with self._connection() as connection:
rows = connection.execute(_PATTERN_CASES_QUERY, parameters).fetchall()
rows_by_case: dict[str, list[tuple[object, ...]]] = {
definition.id: [] for definition in definitions
}
for raw_row in rows:
row = cast(tuple[object, ...], raw_row)
case_id = str(row[0])
if case_id not in rows_by_case:
raise PatternCaseLibraryError(f"unexpected pattern case row: {case_id}")
rows_by_case[case_id].append(row)
cases: list[PatternCase] = []
for definition in definitions:
case_rows = rows_by_case[definition.id]
if len(case_rows) != definition.lookback_days:
raise PatternCaseLibraryError(
f"case {definition.id} requires {definition.lookback_days} qfq rows"
)
bars = tuple(
SelectionBar(
trade_date=_as_date(row[2]),
open=_as_float(row[3]),
high=_as_float(row[4]),
low=_as_float(row[5]),
close=_as_float(row[6]),
volume=_as_float(row[7]),
)
for row in case_rows
)
if any(bar.trade_date >= definition.breakout_date for bar in bars):
raise PatternCaseLibraryError(
f"case {definition.id} contains a non-exclusive breakout row"
)
history = StockHistory(
ts_code=definition.ts_code,
name=definition.name,
bars=bars,
)
cases.append(build_pattern_case(definition, history))
return tuple(cases)
except PatternCaseLibraryError:
raise
except Exception as exc: # noqa: BLE001 - redact database details at the port boundary
raise PatternCaseLibraryError(
"failed to load the complete pattern case library"
) from exc
@contextmanager
def _connection(self) -> Generator[Any, None, None]:
"""Borrow a shared connection without exposing driver failures."""
try:
if self.pool is None:
with psycopg.connect(self.database_url) as connection:
yield connection
else:
with self.pool.connection() as connection:
yield connection
except PatternCaseLibraryError:
raise
except Exception as exc: # noqa: BLE001 - normalize driver/pool errors
raise PatternCaseLibraryError("pattern case database operation failed") from exc
@@ -15,6 +15,11 @@ import psycopg
from psycopg.types.json import Jsonb
from ..domain.models import SelectionSignal, ZhixingB1Category
from ..domain.pattern_scoring import (
ZHIXING_B1_PATTERN_CASES,
PatternScore,
PatternScoreBreakdown,
)
from ..domain.runs import (
SelectionExecutionSource,
SelectionRerunRequired,
@@ -44,17 +49,37 @@ _SIGNAL_ORDER_SQL = (
)
+ f" ELSE {len(ZHIXING_B1_SIGNAL_ORDER)} END"
)
_PATTERN_CASES_BY_ID = {definition.id: definition for definition in ZHIXING_B1_PATTERN_CASES}
_STOCK_ORDER_SQL = {
"code": "item.ts_code ASC",
"score_desc": "item.score_value DESC NULLS LAST, item.ts_code ASC",
"score_asc": "item.score_value ASC NULLS LAST, item.ts_code ASC",
}
_ITEM_UPSERT = """
INSERT INTO selection_run_item
(run_id, ts_code, name, status, signal_count, reason)
VALUES (%s, %s, %s, %s, %s, %s)
(
run_id, ts_code, name, status, signal_count, reason,
score_status, score_value, score_threshold, score_version,
match_case_id, match_case_name, match_case_breakout_date,
match_breakdown, score_reason
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
ON CONFLICT (run_id, ts_code) DO UPDATE SET
name = EXCLUDED.name,
status = EXCLUDED.status,
signal_count = EXCLUDED.signal_count,
reason = EXCLUDED.reason
reason = EXCLUDED.reason,
score_status = EXCLUDED.score_status,
score_value = EXCLUDED.score_value,
score_threshold = EXCLUDED.score_threshold,
score_version = EXCLUDED.score_version,
match_case_id = EXCLUDED.match_case_id,
match_case_name = EXCLUDED.match_case_name,
match_case_breakout_date = EXCLUDED.match_case_breakout_date,
match_breakdown = EXCLUDED.match_breakdown,
score_reason = EXCLUDED.score_reason
"""
_SIGNAL_UPSERT = """
INSERT INTO selection_signal
@@ -200,6 +225,19 @@ class PostgresSelectionRunRepository(SelectionRunStore):
item.status,
item.signal_count,
item.reason,
item.pattern_score.status,
item.pattern_score.value,
item.pattern_score.threshold,
item.pattern_score.version,
item.pattern_score.case.id if item.pattern_score.case else None,
item.pattern_score.case.name if item.pattern_score.case else None,
item.pattern_score.case.breakout_date if item.pattern_score.case else None,
(
Jsonb(item.pattern_score.breakdown.as_dict())
if item.pattern_score.breakdown
else None
),
item.pattern_score.reason,
)
for item in items
)
@@ -354,7 +392,11 @@ class PostgresSelectionRunRepository(SelectionRunStore):
return None
item_rows = connection.execute(
"""
SELECT ts_code, name, status, signal_count, reason
SELECT
ts_code, name, status, signal_count, reason,
score_status, score_value, score_threshold, score_version,
match_case_id, match_case_name, match_case_breakout_date,
match_breakdown, score_reason
FROM selection_run_item
WHERE run_id = %s
ORDER BY ts_code
@@ -363,7 +405,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
).fetchall()
stock_filter, stock_parameters = _stock_filter(query, run_id)
stock_total_row = connection.execute(
f"SELECT COUNT(DISTINCT ts_code) FROM selection_signal WHERE {stock_filter}",
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
tuple(stock_parameters),
).fetchone()
stock_total = int(stock_total_row[0] or 0) if stock_total_row else 0
@@ -372,10 +414,10 @@ class PostgresSelectionRunRepository(SelectionRunStore):
list[tuple[object, ...]],
connection.execute(
f"""
SELECT DISTINCT ts_code
FROM selection_signal
SELECT item.ts_code
FROM selection_run_item AS item
WHERE {stock_filter}
ORDER BY ts_code
ORDER BY {_STOCK_ORDER_SQL[query.sort]}
LIMIT %s OFFSET %s
""",
tuple((*stock_parameters, query.page_size, offset)),
@@ -403,7 +445,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
sorted(
(_signal_from_row(value) for value in signal_rows),
key=lambda signal: (
signal.ts_code,
stock_codes.index(signal.ts_code),
_SIGNAL_PRIORITY.get(signal.category, len(_SIGNAL_PRIORITY)),
),
)
@@ -427,6 +469,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
),
signal_count=int(value[3] or 0),
reason=str(value[4]) if value[4] is not None else None,
pattern_score=_pattern_score_from_row(value[5:14]),
signals=tuple(signals_by_stock.get(str(value[0]), ())),
)
for value in item_rows
@@ -509,18 +552,76 @@ def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[o
can present all independently persisted categories together.
"""
clauses = ["run_id = %s"]
clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"]
parameters: list[object] = [run_id]
if query.search:
pattern = f"%{_escape_like(query.search)}%"
clauses.append("(name ILIKE %s ESCAPE '\\' OR ts_code ILIKE %s ESCAPE '\\')")
clauses.append("(item.name ILIKE %s ESCAPE '\\' OR item.ts_code ILIKE %s ESCAPE '\\')")
parameters.extend((pattern, pattern))
if query.category:
clauses.append("category LIKE %s")
clauses.append(
"EXISTS ("
"SELECT 1 FROM selection_signal AS signal "
"WHERE signal.run_id = item.run_id "
"AND signal.ts_code = item.ts_code "
"AND signal.category LIKE %s"
")"
)
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
return " AND ".join(clauses), parameters
def _pattern_score_from_row(row: Sequence[object]) -> PatternScore:
"""Reconstruct a validated stock-level score from nullable item columns."""
if len(row) < 9:
return PatternScore()
status = str(row[0] or "not_executed")
if status == "not_executed":
return PatternScore()
if status == "failed":
return PatternScore.failed(str(row[8] or "pattern scoring failed"))
if status not in {"matched", "below_threshold"}:
return PatternScore.failed("persisted pattern score status is invalid")
definition = _PATTERN_CASES_BY_ID.get(str(row[4]))
breakdown = _pattern_breakdown(row[7])
if definition is None or breakdown is None:
return PatternScore.failed("persisted pattern score is incomplete")
try:
return PatternScore(
status=cast(Literal["matched", "below_threshold"], status),
value=float(str(row[1])),
threshold=float(str(row[2])),
version=str(row[3]),
case=definition,
breakdown=breakdown,
)
except (TypeError, ValueError):
return PatternScore.failed("persisted pattern score is invalid")
def _pattern_breakdown(value: object) -> PatternScoreBreakdown | None:
"""Parse the four finite JSONB score dimensions."""
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError:
return None
if not isinstance(value, Mapping):
return None
values = cast(Mapping[object, object], value)
try:
return PatternScoreBreakdown(
trend_structure=float(str(values["trend_structure"])),
kdj_state=float(str(values["kdj_state"])),
volume_pattern=float(str(values["volume_pattern"])),
price_shape=float(str(values["price_shape"])),
)
except (KeyError, TypeError, ValueError):
return None
def _escape_like(value: str) -> str:
"""Escape user wildcards before placing text inside a SQL LIKE pattern."""
@@ -13,6 +13,10 @@ from zhixing_server.modules.selection.application.run import (
RunZhixingB1,
)
from zhixing_server.modules.selection.domain.models import SelectionSignal
from zhixing_server.modules.selection.domain.pattern_scoring import (
PatternScore,
ZhixingB1PatternScorer,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionRerunRequired,
SelectionResultQuery,
@@ -23,6 +27,7 @@ from zhixing_server.modules.selection.domain.runs import (
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
PostgresMarketDataReader,
PostgresPatternCaseLibraryLoader,
SelectionMarketDataNotReady,
SelectionReaderError,
)
@@ -82,6 +87,35 @@ class SelectionFailureResponse(BaseModel):
reason: str | None
class SelectionPatternCaseResponse(BaseModel):
"""The best matching versioned case for one computed score."""
id: str
name: str
breakout_date: date
class SelectionPatternBreakdownResponse(BaseModel):
"""The four finite 0-100 similarity dimensions."""
trend_structure: float = Field(ge=0, le=100)
kdj_state: float = Field(ge=0, le=100)
volume_pattern: float = Field(ge=0, le=100)
price_shape: float = Field(ge=0, le=100)
class SelectionPatternScoreResponse(BaseModel):
"""A stock-level enrichment independent of selection evaluation status."""
status: Literal["matched", "below_threshold", "failed"]
value: float | None = Field(default=None, ge=0, le=100)
threshold: float | None = Field(default=None, ge=0, le=100)
version: str | None = None
case: SelectionPatternCaseResponse | None = None
breakdown: SelectionPatternBreakdownResponse | None = None
reason: str | None = None
def _empty_failures() -> list[SelectionFailureResponse]:
"""Create a typed default list for Pydantic's strict checker."""
@@ -102,6 +136,7 @@ class SelectionStockResponse(BaseModel):
target_trade_date: date
strategy: StrategyValue
close: float
score: SelectionPatternScoreResponse | None = None
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
@@ -144,10 +179,14 @@ def get_selection_service(
pool = get_selection_postgres_pool(settings)
reader = PostgresMarketDataReader(settings, pool=pool)
pattern_case_loader = PostgresPatternCaseLibraryLoader(settings, pool=pool)
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
return RunZhixingB1(
reader,
store,
pattern_case_loader=pattern_case_loader,
pattern_scorer=ZhixingB1PatternScorer(),
pattern_scoring_enabled=settings.selection_pattern_scoring_enabled,
max_workers=settings.selection_max_workers,
batch_size=settings.selection_batch_size,
)
@@ -225,11 +264,12 @@ def get_selection_run(
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
search: Annotated[str | None, Query(max_length=100)] = None,
category: Literal["pullback", "oversold", "original"] | None = None,
sort: Literal["code", "score_desc", "score_asc"] = "code",
) -> SelectionResultsResponse:
"""Return one run for asynchronous polling."""
try:
query = _result_query(page, page_size, search, category)
query = _result_query(page, page_size, search, category, sort)
run = service.get_run(run_id, query=query)
except SelectionRunStoreError as exc:
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
@@ -247,11 +287,12 @@ def get_selection_results(
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
search: Annotated[str | None, Query(max_length=100)] = None,
category: Literal["pullback", "oversold", "original"] | None = None,
sort: Literal["code", "score_desc", "score_asc"] = "code",
) -> SelectionResultsResponse:
"""Return the current persisted result for a strategy and optional date."""
try:
query = _result_query(page, page_size, search, category)
query = _result_query(page, page_size, search, category, sort)
run = service.get_latest(strategy, target_trade_date, query=query)
except SelectionRunStoreError as exc:
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
@@ -276,6 +317,7 @@ def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> Selectio
signals_by_stock: dict[str, list[SelectionSignalResponse]] = {}
for signal in run.signals:
signals_by_stock.setdefault(signal.ts_code, []).append(_signal_response(signal))
items_by_stock = {item.ts_code: item for item in run.items}
return SelectionResultsResponse(
strategy=run.strategy,
@@ -316,6 +358,7 @@ def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> Selectio
target_trade_date=signals[0].target_trade_date,
strategy=signals[0].strategy,
close=signals[0].close,
score=_pattern_score_response(items_by_stock[signals[0].ts_code].pattern_score),
signals=signals,
)
for signals in signals_by_stock.values()
@@ -337,11 +380,43 @@ def _signal_response(signal: SelectionSignal) -> SelectionSignalResponse:
)
def _pattern_score_response(score: PatternScore) -> SelectionPatternScoreResponse | None:
"""Hide not-executed scores and expose validated computed/failure states."""
if score.status == "not_executed":
return None
if score.status == "failed":
return SelectionPatternScoreResponse(status="failed", reason=score.reason)
if score.status == "below_threshold":
return SelectionPatternScoreResponse(
status="below_threshold",
threshold=score.threshold,
version=score.version,
reason="未匹配到评分阈值以上案例",
)
if score.case is None or score.breakdown is None:
return SelectionPatternScoreResponse(status="failed", reason="评分结果不完整")
return SelectionPatternScoreResponse(
status=score.status,
value=score.value,
threshold=score.threshold,
version=score.version,
case=SelectionPatternCaseResponse(
id=score.case.id,
name=score.case.name,
breakout_date=score.case.breakout_date,
),
breakdown=SelectionPatternBreakdownResponse(**score.breakdown.as_dict()),
reason=score.reason,
)
def _result_query(
page: int,
page_size: int,
search: str | None,
category: Literal["pullback", "oversold", "original"] | None,
sort: Literal["code", "score_desc", "score_asc"],
) -> SelectionResultQuery:
"""Normalize HTTP query values before handing them to the selection port."""
@@ -351,6 +426,7 @@ def _result_query(
page_size=page_size,
search=normalized_search or None,
category=category,
sort=sort,
)