feat(selection): 集成 B1 FastDTW 图形评分
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
"""HTTP contracts for triggering and querying persisted selection runs."""
|
||||
|
||||
from dataclasses import replace
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
@@ -11,6 +12,12 @@ from zhixing_server.bootstrap.app import create_app
|
||||
from zhixing_server.bootstrap.config import Settings
|
||||
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
|
||||
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,
|
||||
@@ -130,6 +137,14 @@ def _run(run_id: str, status: str) -> SelectionRun:
|
||||
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=(original_signal, pullback_signal),
|
||||
),
|
||||
),
|
||||
@@ -238,6 +253,24 @@ def test_query_returns_persisted_signal_details() -> None:
|
||||
assert "signals" not in body
|
||||
assert len(body["stocks"]) == 1
|
||||
assert body["stocks"][0]["ts_code"] == "000001.SZ"
|
||||
assert body["stocks"][0]["score"] == {
|
||||
"status": "matched",
|
||||
"value": 86.4,
|
||||
"threshold": 60.0,
|
||||
"version": PATTERN_SCORING_VERSION,
|
||||
"case": {
|
||||
"id": "case_001",
|
||||
"name": "华纳药厂",
|
||||
"breakout_date": "2025-05-12",
|
||||
},
|
||||
"breakdown": {
|
||||
"trend_structure": 71.2,
|
||||
"kdj_state": 83.0,
|
||||
"volume_pattern": 88.0,
|
||||
"price_shape": 90.1,
|
||||
},
|
||||
"reason": None,
|
||||
}
|
||||
assert [signal["category"] for signal in body["stocks"][0]["signals"]] == [
|
||||
"zhixing_b1_original_b1",
|
||||
"zhixing_b1_pullback_white",
|
||||
@@ -248,6 +281,59 @@ def test_query_returns_persisted_signal_details() -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("pattern_score", "expected_score"),
|
||||
[
|
||||
(
|
||||
PatternScore(
|
||||
status="below_threshold",
|
||||
value=42.5,
|
||||
threshold=60.0,
|
||||
version=PATTERN_SCORING_VERSION,
|
||||
case=ZHIXING_B1_PATTERN_CASES[0],
|
||||
breakdown=PatternScoreBreakdown(40.0, 42.0, 43.0, 44.0),
|
||||
),
|
||||
{
|
||||
"status": "below_threshold",
|
||||
"value": None,
|
||||
"threshold": 60.0,
|
||||
"version": PATTERN_SCORING_VERSION,
|
||||
"case": None,
|
||||
"breakdown": None,
|
||||
"reason": "未匹配到评分阈值以上案例",
|
||||
},
|
||||
),
|
||||
(
|
||||
PatternScore.failed("FastDTW unavailable"),
|
||||
{
|
||||
"status": "failed",
|
||||
"value": None,
|
||||
"threshold": None,
|
||||
"version": None,
|
||||
"case": None,
|
||||
"breakdown": None,
|
||||
"reason": "FastDTW unavailable",
|
||||
},
|
||||
),
|
||||
(PatternScore(), None),
|
||||
],
|
||||
)
|
||||
def test_query_preserves_signals_for_every_pattern_score_state(
|
||||
pattern_score: PatternScore,
|
||||
expected_score: dict[str, object] | None,
|
||||
) -> None:
|
||||
run = _run("run-http", "success")
|
||||
run = replace(run, items=(replace(run.items[0], pattern_score=pattern_score),))
|
||||
|
||||
response = _client(FakeSelectionService(run)).get("/api/v1/selection/results")
|
||||
|
||||
assert response.status_code == 200
|
||||
stock = response.json()["stocks"][0]
|
||||
assert stock["score"] == expected_score
|
||||
assert len(stock["signals"]) == 2
|
||||
assert response.json()["failures"] == []
|
||||
|
||||
|
||||
def test_query_forwards_pagination_and_filters() -> None:
|
||||
service = FakeSelectionService(_run("run-http", "success"))
|
||||
|
||||
@@ -259,6 +345,7 @@ def test_query_forwards_pagination_and_filters() -> None:
|
||||
"page_size": 5,
|
||||
"search": " 平安银行 ",
|
||||
"category": "original",
|
||||
"sort": "score_desc",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -268,6 +355,7 @@ def test_query_forwards_pagination_and_filters() -> None:
|
||||
page_size=5,
|
||||
search="平安银行",
|
||||
category="original",
|
||||
sort="score_desc",
|
||||
)
|
||||
assert response.json()["page"] == 2
|
||||
assert response.json()["page_size"] == 5
|
||||
@@ -282,6 +370,18 @@ def test_query_rejects_invalid_page_size() -> None:
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
def test_query_forwards_score_ascending_sort() -> None:
|
||||
service = FakeSelectionService(_run("run-http", "success"))
|
||||
|
||||
response = _client(service).get(
|
||||
"/api/v1/selection/results",
|
||||
params={"strategy": "zhixing_b1", "sort": "score_asc"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert service.last_query == SelectionResultQuery(sort="score_asc")
|
||||
|
||||
|
||||
def test_run_polling_returns_the_persisted_terminal_result() -> None:
|
||||
response = _client(FakeSelectionService(_run("run-http", "success"))).get(
|
||||
"/api/v1/selection/runs/run-http"
|
||||
|
||||
Reference in New Issue
Block a user