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
@@ -0,0 +1,182 @@
"""Golden and invariant tests for versioned B1 FastDTW scoring."""
from __future__ import annotations
import json
from datetime import date, timedelta
from pathlib import Path
from typing import cast
import pandas as pd
import pytest
from zhixing_server.modules.selection.domain.models import SelectionBar, StockHistory
from zhixing_server.modules.selection.domain.pattern_scoring import (
PATTERN_FASTDTW_RADIUS,
PATTERN_SCORING_VERSION,
ZHIXING_B1_PATTERN_CASES,
PatternCase,
PatternCaseLibraryError,
PatternFeatures,
PatternScore,
PatternScoreBreakdown,
PatternScoringError,
ZhixingB1PatternScorer,
build_pattern_case,
)
FIXTURES = Path(__file__).parents[2] / "fixtures" / "selection" / "zhixing_b1" / "pattern_scoring"
def _history(case_id: str, ts_code: str, name: str) -> StockHistory:
frame = pd.read_csv(FIXTURES / f"{case_id}.csv")
bars = tuple(
SelectionBar(
trade_date=date.fromisoformat(str(row.date)),
open=float(str(row.open)),
high=float(str(row.high)),
low=float(str(row.low)),
close=float(str(row.close)),
volume=float(str(row.volume)),
)
for row in frame.itertuples(index=False)
)
return StockHistory(ts_code=ts_code, name=name, bars=bars)
def _cases() -> tuple[PatternCase, ...]:
return tuple(
build_pattern_case(
definition,
_history(definition.id, definition.ts_code, definition.name),
)
for definition in ZHIXING_B1_PATTERN_CASES
)
def _golden(name: str) -> dict[str, object]:
payload = cast(dict[str, object], json.loads((FIXTURES / "golden.json").read_text()))
return cast(dict[str, object], payload[name])
def _assert_golden(score: PatternScore, expected: dict[str, object]) -> None:
assert score.status == expected["status"]
assert score.value == expected["value"]
assert score.case is not None
assert score.case.id == expected["case_id"]
assert score.breakdown is not None
assert score.breakdown.as_dict() == expected["breakdown"]
def test_fastdtw_v1_self_match_golden_is_finite_and_deterministic() -> None:
cases = _cases()
scorer = ZhixingB1PatternScorer()
first = scorer.score(cases[0].history, cases)
second = scorer.score(cases[0].history, cases)
assert PATTERN_SCORING_VERSION == "zhixing_b1_pattern_fastdtw_v1"
assert PATTERN_FASTDTW_RADIUS == 1
assert first == second
_assert_golden(first, _golden("self_match"))
assert cases[0].features.trend_structure["short_vs_bullbear"] is None
def test_fastdtw_v1_time_warped_curve_golden() -> None:
cases = _cases()
base = cases[0].history.bars
delayed = base[:1] * 3 + base[:-3]
bars = tuple(
SelectionBar(
trade_date=base[index].trade_date,
open=delayed[index].open,
high=delayed[index].high,
low=delayed[index].low,
close=delayed[index].close,
volume=delayed[index].volume,
)
for index in range(25)
)
result = ZhixingB1PatternScorer().score(
StockHistory(ts_code="TEST.SZ", name="time warped", bars=bars),
cases,
)
_assert_golden(result, _golden("time_warped"))
def test_below_threshold_golden_remains_a_successful_computation() -> None:
bars = tuple(
SelectionBar(
trade_date=date(2026, 1, 1) + timedelta(days=index),
open=100.0 if index % 2 == 0 else 1.0,
high=110.0,
low=0.9,
close=1.0 if index % 2 == 0 else 100.0,
volume=1.0 if index < 13 else 1_000_000.0,
)
for index in range(25)
)
result = ZhixingB1PatternScorer().score(
StockHistory(ts_code="TEST.SZ", name="below", bars=bars),
_cases(),
)
_assert_golden(result, _golden("below_threshold"))
def test_case_library_rejects_partial_or_short_input() -> None:
cases = _cases()
with pytest.raises(PatternScoringError, match="incomplete or out of order"):
ZhixingB1PatternScorer().score(cases[0].history, cases[:-1])
definition = ZHIXING_B1_PATTERN_CASES[0]
short = _history(definition.id, definition.ts_code, definition.name)
with pytest.raises(PatternCaseLibraryError, match="requires 25 complete rows"):
build_pattern_case(
definition,
StockHistory(short.ts_code, short.name, short.bars[:-1]),
)
def test_fastdtw_failure_is_not_replaced_by_simple_dtw(monkeypatch: pytest.MonkeyPatch) -> None:
import zhixing_server.modules.selection.domain.pattern_scoring as scoring
def fail(*_args: object, **_kwargs: object) -> tuple[float, list[tuple[int, int]]]:
raise RuntimeError("fastdtw unavailable")
monkeypatch.setattr(scoring, "_fastdtw", lambda: fail)
with pytest.raises(RuntimeError, match="fastdtw unavailable"):
ZhixingB1PatternScorer().score(_cases()[0].history, _cases())
def test_failed_score_requires_a_safe_reason() -> None:
with pytest.raises(ValueError, match="requires a safe reason"):
PatternScore(status="failed")
assert PatternScore.failed(" ").reason == "pattern scoring failed"
def test_threshold_is_inclusive_and_equal_scores_keep_first_case(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import zhixing_server.modules.selection.domain.pattern_scoring as scoring
tied = PatternScoreBreakdown(60.0, 60.0, 60.0, 60.0)
def tied_match(
_candidate: PatternFeatures,
_case: PatternFeatures,
) -> PatternScoreBreakdown:
return tied
monkeypatch.setattr(scoring, "_match", tied_match)
result = ZhixingB1PatternScorer().score(_cases()[0].history, _cases())
assert result.status == "matched"
assert result.value == 60.0
assert result.case == ZHIXING_B1_PATTERN_CASES[0]
@@ -2,17 +2,19 @@
from collections.abc import Generator
from contextlib import contextmanager
from datetime import date
from datetime import date, timedelta
from decimal import Decimal
from typing import cast
import psycopg
import pytest
from zhixing_server.modules.selection.domain.pattern_scoring import ZHIXING_B1_PATTERN_CASES
from zhixing_server.modules.selection.domain.runs import SelectionStock
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
PostgresMarketDataReader,
PostgresPatternCaseLibraryLoader,
SelectionMarketDataNotReady,
)
@@ -237,3 +239,34 @@ def test_reader_rejects_date_without_eligible_market_batch(monkeypatch: pytest.M
"zhixing_b1",
date(2026, 8, 8),
)
def test_pattern_case_loader_reads_one_complete_exclusive_qfq_library() -> None:
rows: list[tuple[object, ...]] = []
for definition in ZHIXING_B1_PATTERN_CASES:
for offset in range(definition.lookback_days, 0, -1):
rows.append(
(
definition.id,
definition.ts_code,
definition.breakout_date - timedelta(days=offset),
"10",
"11",
"9",
str(10 + offset / 100),
str(1000 + offset),
)
)
connection = FakeConnection(rows)
pool = Pool(connection)
owner = SelectionPostgresPool("postgresql://test", max_connections=2, pool=pool)
cases = PostgresPatternCaseLibraryLoader("postgresql://test", pool=owner).load()
assert tuple(case.definition for case in cases) == ZHIXING_B1_PATTERN_CASES
assert all(len(case.history.bars) == 25 for case in cases)
assert all(case.history.bars[-1].trade_date < case.definition.breakout_date for case in cases)
assert "bar.trade_date < definition.breakout_date" in cast(str, connection.query)
assert "bar.source_adj = 'qfq'" in cast(str, connection.query)
assert connection.parameters is not None
assert connection.parameters[0] == [definition.id for definition in ZHIXING_B1_PATTERN_CASES]
@@ -8,6 +8,12 @@ import pytest
from psycopg.types.json import Jsonb
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
from zhixing_server.modules.selection.domain.pattern_scoring import (
PATTERN_SCORING_VERSION,
ZHIXING_B1_PATTERN_CASES,
PatternScore,
PatternScoreBreakdown,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionExecutionSource,
SelectionRerunRequired,
@@ -222,6 +228,14 @@ def test_record_items_uses_one_delete_and_two_batch_upserts(
name="平安银行",
status="selected",
signal_count=2,
pattern_score=PatternScore(
status="matched",
value=86.4,
threshold=60.0,
version=PATTERN_SCORING_VERSION,
case=ZHIXING_B1_PATTERN_CASES[0],
breakdown=PatternScoreBreakdown(71.2, 83.0, 88.0, 90.1),
),
signals=(first, second),
),
SelectionRunItem(
@@ -238,6 +252,28 @@ def test_record_items_uses_one_delete_and_two_batch_upserts(
assert delete_parameters == ("run-1", ["000001.SZ", "600000.SH"])
assert len(connection.executemany_calls) == 2
assert "INSERT INTO selection_run_item" in connection.executemany_calls[0][0]
item_parameters = connection.executemany_calls[0][1]
assert item_parameters[0][6:13] == (
"matched",
86.4,
60.0,
PATTERN_SCORING_VERSION,
"case_001",
"华纳药厂",
date(2025, 5, 12),
)
assert isinstance(item_parameters[0][13], Jsonb)
assert item_parameters[1][6:] == (
"not_executed",
None,
None,
None,
None,
None,
None,
None,
None,
)
assert "INSERT INTO selection_signal" in connection.executemany_calls[1][0]
signal_parameters = connection.executemany_calls[1][1]
assert len(signal_parameters) == 2
@@ -281,11 +317,35 @@ class LoadConnection:
None,
)
)
if "FROM selection_run_item" in query:
return LoadResult(rows=[("000001.SZ", "平安银行", "selected", 2, None)])
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query:
if "FROM selection_run_item\n" in query:
return LoadResult(
rows=[
(
"000001.SZ",
"平安银行",
"selected",
2,
None,
"matched",
Decimal("86.40"),
Decimal("60.00"),
PATTERN_SCORING_VERSION,
"case_001",
"华纳药厂",
date(2025, 5, 12),
{
"trend_structure": 71.2,
"kdj_state": 83.0,
"volume_pattern": 88.0,
"price_shape": 90.1,
},
None,
)
]
)
if "SELECT COUNT(*) FROM selection_run_item AS item" in query:
return LoadResult(row=(2,))
if "SELECT DISTINCT ts_code" in query:
if "SELECT item.ts_code" in query:
return LoadResult(rows=[("000001.SZ",)])
return LoadResult(
rows=[
@@ -315,7 +375,7 @@ class EmptyStockPageConnection(LoadConnection):
"""Return a non-zero filtered total with no stocks on the requested page."""
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
if "SELECT DISTINCT ts_code" in query:
if "SELECT item.ts_code" in query:
self.statements.append((query, parameters))
return LoadResult(rows=[])
return super().execute(query, parameters)
@@ -355,6 +415,7 @@ def test_get_run_pages_stocks_and_loads_all_signals_for_category_matches(
page_size=1,
search="100%",
category="pullback",
sort="score_desc",
),
)
@@ -364,19 +425,21 @@ def test_get_run_pages_stocks_and_loads_all_signals_for_category_matches(
ZHIXING_B1_SIGNAL_ORDER[-1],
]
assert run.stocks_total == 2
assert run.items[0].pattern_score.status == "matched"
assert run.items[0].pattern_score.value == 86.4
count_query, count_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query
if "SELECT COUNT(*) FROM selection_run_item AS item" in query
)
assert "name ILIKE %s ESCAPE" in count_query
assert count_parameters == ("run-1", "%100\\%%", "%100\\%%", "zhixing_b1_pullback_%")
stock_page_query, page_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "SELECT DISTINCT ts_code" in query
if "SELECT item.ts_code" in query
)
assert "ORDER BY ts_code" in stock_page_query
assert "ORDER BY item.score_value DESC NULLS LAST, item.ts_code ASC" in stock_page_query
assert page_parameters[-2:] == (1, 0)
signal_query, signal_parameters = next(
(query, parameters)
@@ -409,8 +472,30 @@ def test_get_run_does_not_load_signals_for_an_empty_stock_page(
stock_page_query, stock_page_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "SELECT DISTINCT ts_code" in query
if "SELECT item.ts_code" in query
)
assert "ORDER BY ts_code" in stock_page_query
assert "ORDER BY item.ts_code ASC" in stock_page_query
assert stock_page_parameters[-2:] == (1, 2)
assert not any("ts_code = ANY(%s)" in query for query, _ in connection.statements)
def test_get_run_sorts_scores_ascending_with_nulls_last_and_code_tiebreak(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = LoadConnection()
def connect(database_url: str) -> LoadConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
run = PostgresSelectionRunRepository("postgresql://test").get_run(
"run-1",
query=SelectionResultQuery(sort="score_asc"),
)
assert run is not None
stock_page_query = next(
query for query, _ in connection.statements if "SELECT item.ts_code" in query
)
assert "ORDER BY item.score_value ASC NULLS LAST, item.ts_code ASC" in stock_page_query
@@ -2,6 +2,7 @@
import threading
import time
from collections.abc import Sequence
from datetime import date
from decimal import Decimal
from typing import Literal
@@ -15,6 +16,14 @@ from zhixing_server.modules.selection.domain.models import (
SelectionSignal,
StockHistory,
)
from zhixing_server.modules.selection.domain.pattern_scoring import (
PATTERN_SCORE_THRESHOLD,
PATTERN_SCORING_VERSION,
ZHIXING_B1_PATTERN_CASES,
PatternCase,
PatternScore,
PatternScoreBreakdown,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionExecutionSource,
SelectionResultQuery,
@@ -198,6 +207,37 @@ class ConcurrentHistoryEvaluator:
return SelectionEvaluation(history.ts_code, target_trade_date, "no_signal")
class FakePatternCaseLoader:
def __init__(self, *, error: Exception | None = None) -> None:
self.calls = 0
self.error = error
def load(self) -> tuple[PatternCase, ...]:
self.calls += 1
if self.error is not None:
raise self.error
return ()
class FakePatternScorer:
def __init__(self, *, error: Exception | None = None) -> None:
self.calls: list[str] = []
self.error = error
def score(self, history: StockHistory, cases: Sequence[PatternCase]) -> PatternScore:
self.calls.append(history.ts_code)
if self.error is not None:
raise self.error
return PatternScore(
status="matched",
value=88.0,
threshold=PATTERN_SCORE_THRESHOLD,
version=PATTERN_SCORING_VERSION,
case=ZHIXING_B1_PATTERN_CASES[0],
breakdown=PatternScoreBreakdown(80.0, 85.0, 90.0, 88.0),
)
def _source() -> SelectionExecutionSource:
return SelectionExecutionSource(
market_sync_batch_id="market-run-1",
@@ -364,3 +404,149 @@ def test_execute_marks_batch_write_failure_as_failed() -> None:
assert store.finished[0:2] == ("run-1", "failed")
assert store.finished[2]["error_type"] == "batch_error"
assert store.finished[2]["failed_count"] == 1
def test_execute_loads_cases_once_and_scores_only_selected_stocks() -> None:
source = _source()
reader = BatchReader(source)
store = FakeStore()
loader = FakePatternCaseLoader()
scorer = FakePatternScorer()
evaluator = FakeEvaluator(
{
"000001.SZ": SelectionEvaluation(
"000001.SZ",
TARGET,
"selected",
signals=(_signal("000001.SZ", "zhixing_b1_original_b1"),),
),
"600000.SH": SelectionEvaluation("600000.SH", TARGET, "no_signal"),
}
)
service = RunZhixingB1(
reader,
store,
evaluator,
loader,
scorer,
pattern_scoring_enabled=True,
batch_size=1,
)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert loader.calls == 1
assert scorer.calls == ["000001.SZ"]
assert [item.pattern_score.status for item in store.items] == ["matched", "not_executed"]
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "success")
assert store.finished[2]["failed_count"] == 0
def test_execute_isolates_pattern_scoring_failure_from_selection_status() -> None:
source = _source()
store = FakeStore()
loader = FakePatternCaseLoader()
scorer = FakePatternScorer(error=RuntimeError("FastDTW unavailable"))
evaluator = FakeEvaluator(
{
"000001.SZ": SelectionEvaluation(
"000001.SZ",
TARGET,
"selected",
signals=(_signal("000001.SZ", "zhixing_b1_original_b1"),),
),
"600000.SH": SelectionEvaluation("600000.SH", TARGET, "no_signal"),
}
)
service = RunZhixingB1(
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_scoring_enabled=True,
)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert store.items[0].status == "selected"
assert store.items[0].pattern_score == PatternScore.failed("FastDTW unavailable")
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "success")
assert store.finished[2]["failed_count"] == 0
def test_execute_skips_pattern_dependencies_when_feature_flag_is_disabled() -> None:
source = _source()
store = FakeStore()
loader = FakePatternCaseLoader(error=AssertionError("loader must not run"))
scorer = FakePatternScorer(error=AssertionError("scorer must not run"))
evaluator = FakeEvaluator(
{
"000001.SZ": SelectionEvaluation(
"000001.SZ",
TARGET,
"selected",
signals=(_signal("000001.SZ", "zhixing_b1_original_b1"),),
),
"600000.SH": SelectionEvaluation("600000.SH", TARGET, "no_signal"),
}
)
service = RunZhixingB1(
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_scoring_enabled=False,
)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert loader.calls == 0
assert scorer.calls == []
assert [item.pattern_score.status for item in store.items] == [
"not_executed",
"not_executed",
]
assert store.items[0].signal_count == 1
assert store.finished is not None
assert store.finished[2]["failed_count"] == 0
def test_execute_marks_scores_failed_when_case_library_is_unavailable() -> None:
source = _source()
store = FakeStore()
loader = FakePatternCaseLoader(error=RuntimeError("case_011 requires 25 qfq rows"))
scorer = FakePatternScorer()
evaluator = FakeEvaluator(
{
"000001.SZ": SelectionEvaluation(
"000001.SZ",
TARGET,
"selected",
signals=(_signal("000001.SZ", "zhixing_b1_original_b1"),),
),
"600000.SH": SelectionEvaluation("600000.SH", TARGET, "no_signal"),
}
)
service = RunZhixingB1(
BatchReader(source),
store,
evaluator,
loader,
scorer,
pattern_scoring_enabled=True,
)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert loader.calls == 1
assert scorer.calls == []
assert store.items[0].status == "selected"
assert store.items[0].pattern_score.status == "failed"
assert store.items[0].signals[0].category.value == "zhixing_b1_original_b1"
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "success")
assert store.finished[2]["failed_count"] == 0