feat(selection): 集成 B1 FastDTW 图形评分
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
+141
@@ -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
|
||||
|
||||
+113
-12
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user