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