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