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