refactor(selection): 统一信号与股票处理逻辑,重构相关数据结构与接口,更新前端展示以支持股票信息
This commit is contained in:
@@ -79,7 +79,7 @@ class SelectionRun:
|
||||
finished_at: datetime | None = None
|
||||
items: tuple[SelectionRunItem, ...] = field(default_factory=tuple)
|
||||
signals: tuple[SelectionSignal, ...] = field(default_factory=tuple)
|
||||
signals_total: int | None = None
|
||||
stocks_total: int | None = None
|
||||
|
||||
|
||||
class SelectionRunError(RuntimeError):
|
||||
|
||||
+46
-21
@@ -284,7 +284,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read one run with filtered, paged signals and item failures."""
|
||||
"""Read one run with filtered, stock-paged signals and item failures."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
@@ -361,27 +361,47 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchall()
|
||||
signal_filter, signal_parameters = _signal_filter(query, run_id)
|
||||
signal_total_row = connection.execute(
|
||||
f"SELECT COUNT(*) FROM selection_signal WHERE {signal_filter}",
|
||||
tuple(signal_parameters),
|
||||
stock_filter, stock_parameters = _stock_filter(query, run_id)
|
||||
stock_total_row = connection.execute(
|
||||
f"SELECT COUNT(DISTINCT ts_code) FROM selection_signal WHERE {stock_filter}",
|
||||
tuple(stock_parameters),
|
||||
).fetchone()
|
||||
signal_total = int(signal_total_row[0] or 0) if signal_total_row else 0
|
||||
stock_total = int(stock_total_row[0] or 0) if stock_total_row else 0
|
||||
offset = (query.page - 1) * query.page_size
|
||||
signal_rows = connection.execute(
|
||||
f"""
|
||||
SELECT
|
||||
ts_code, name, target_trade_date, strategy, category, close, details
|
||||
FROM selection_signal
|
||||
WHERE {signal_filter}
|
||||
ORDER BY ts_code, {_SIGNAL_ORDER_SQL}
|
||||
LIMIT %s OFFSET %s
|
||||
""",
|
||||
tuple((*signal_parameters, query.page_size, offset)),
|
||||
).fetchall()
|
||||
stock_rows = cast(
|
||||
list[tuple[object, ...]],
|
||||
connection.execute(
|
||||
f"""
|
||||
SELECT DISTINCT ts_code
|
||||
FROM selection_signal
|
||||
WHERE {stock_filter}
|
||||
ORDER BY ts_code
|
||||
LIMIT %s OFFSET %s
|
||||
""",
|
||||
tuple((*stock_parameters, query.page_size, offset)),
|
||||
).fetchall(),
|
||||
)
|
||||
stock_codes = [str(value[0]) for value in stock_rows]
|
||||
signal_rows = (
|
||||
cast(
|
||||
list[tuple[object, ...]],
|
||||
connection.execute(
|
||||
f"""
|
||||
SELECT
|
||||
ts_code, name, target_trade_date, strategy, category, close, details
|
||||
FROM selection_signal
|
||||
WHERE run_id = %s AND ts_code = ANY(%s)
|
||||
ORDER BY ts_code, {_SIGNAL_ORDER_SQL}
|
||||
""",
|
||||
(run_id, stock_codes),
|
||||
).fetchall(),
|
||||
)
|
||||
if stock_codes
|
||||
else []
|
||||
)
|
||||
signals = tuple(
|
||||
sorted(
|
||||
(_signal_from_row(cast(tuple[object, ...], value)) for value in signal_rows),
|
||||
(_signal_from_row(value) for value in signal_rows),
|
||||
key=lambda signal: (
|
||||
signal.ts_code,
|
||||
_SIGNAL_PRIORITY.get(signal.category, len(_SIGNAL_PRIORITY)),
|
||||
@@ -430,7 +450,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
finished_at=cast(datetime | None, row[15]),
|
||||
items=items,
|
||||
signals=signals,
|
||||
signals_total=signal_total,
|
||||
stocks_total=stock_total,
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
@@ -481,8 +501,13 @@ def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
|
||||
)
|
||||
|
||||
|
||||
def _signal_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]:
|
||||
"""Build the parameterized WHERE clause shared by count and page reads."""
|
||||
def _stock_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]:
|
||||
"""Build the signal predicate used to select distinct matching stocks.
|
||||
|
||||
A category narrows which stocks qualify for the page. Once a stock
|
||||
qualifies, the repository loads every signal for that stock so callers
|
||||
can present all independently persisted categories together.
|
||||
"""
|
||||
|
||||
clauses = ["run_id = %s"]
|
||||
parameters: list[object] = [run_id]
|
||||
|
||||
@@ -12,6 +12,7 @@ from zhixing_server.bootstrap.config import Settings, get_settings
|
||||
from zhixing_server.modules.selection.application.run import (
|
||||
RunZhixingB1,
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.models import SelectionSignal
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionRerunRequired,
|
||||
SelectionResultQuery,
|
||||
@@ -93,8 +94,25 @@ def _empty_signals() -> list[SelectionSignalResponse]:
|
||||
return []
|
||||
|
||||
|
||||
class SelectionStockResponse(BaseModel):
|
||||
"""One selected stock with all independently persisted signals."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
target_trade_date: date
|
||||
strategy: StrategyValue
|
||||
close: float
|
||||
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
|
||||
|
||||
|
||||
def _empty_stocks() -> list[SelectionStockResponse]:
|
||||
"""Create a typed default stock-result list."""
|
||||
|
||||
return []
|
||||
|
||||
|
||||
class SelectionResultsResponse(BaseModel):
|
||||
"""Batch summary and one filtered page of signals consumed by the Web feature."""
|
||||
"""Batch summary and one filtered page of selected stocks."""
|
||||
|
||||
strategy: StrategyValue
|
||||
target_trade_date: date | None
|
||||
@@ -114,9 +132,9 @@ class SelectionResultsResponse(BaseModel):
|
||||
finished_at: datetime | None = None
|
||||
page: int = Field(default=1, ge=1)
|
||||
page_size: int = Field(default=10, ge=1, le=100)
|
||||
signals_total: int = Field(default=0, ge=0)
|
||||
stocks_total: int = Field(default=0, ge=0)
|
||||
failures: list[SelectionFailureResponse] = Field(default_factory=_empty_failures)
|
||||
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
|
||||
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
|
||||
|
||||
|
||||
def get_selection_service(
|
||||
@@ -247,7 +265,7 @@ def get_selection_results(
|
||||
coverage=0,
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
signals_total=0,
|
||||
stocks_total=0,
|
||||
)
|
||||
return _run_response(run, query=query)
|
||||
|
||||
@@ -255,6 +273,10 @@ def get_selection_results(
|
||||
def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> SelectionResultsResponse:
|
||||
"""Translate a domain run without exposing storage-specific fields."""
|
||||
|
||||
signals_by_stock: dict[str, list[SelectionSignalResponse]] = {}
|
||||
for signal in run.signals:
|
||||
signals_by_stock.setdefault(signal.ts_code, []).append(_signal_response(signal))
|
||||
|
||||
return SelectionResultsResponse(
|
||||
strategy=run.strategy,
|
||||
target_trade_date=run.target_trade_date,
|
||||
@@ -274,7 +296,9 @@ def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> Selectio
|
||||
finished_at=run.finished_at,
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
signals_total=run.signals_total if run.signals_total is not None else run.signal_count,
|
||||
stocks_total=(
|
||||
run.stocks_total if run.stocks_total is not None else run.selected_stock_count
|
||||
),
|
||||
failures=[
|
||||
SelectionFailureResponse(
|
||||
ts_code=item.ts_code,
|
||||
@@ -285,21 +309,34 @@ def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> Selectio
|
||||
for item in run.items
|
||||
if item.status in {"insufficient_history", "missing_target_bar", "data_error"}
|
||||
],
|
||||
signals=[
|
||||
SelectionSignalResponse(
|
||||
ts_code=signal.ts_code,
|
||||
name=signal.name,
|
||||
target_trade_date=signal.target_trade_date,
|
||||
strategy=signal.strategy,
|
||||
category=signal.category.value,
|
||||
close=signal.close,
|
||||
details=dict(signal.details),
|
||||
stocks=[
|
||||
SelectionStockResponse(
|
||||
ts_code=signals[0].ts_code,
|
||||
name=signals[0].name,
|
||||
target_trade_date=signals[0].target_trade_date,
|
||||
strategy=signals[0].strategy,
|
||||
close=signals[0].close,
|
||||
signals=signals,
|
||||
)
|
||||
for signal in run.signals
|
||||
for signals in signals_by_stock.values()
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _signal_response(signal: SelectionSignal) -> SelectionSignalResponse:
|
||||
"""Map a domain signal while preserving its category-specific details."""
|
||||
|
||||
return SelectionSignalResponse(
|
||||
ts_code=signal.ts_code,
|
||||
name=signal.name,
|
||||
target_trade_date=signal.target_trade_date,
|
||||
strategy=signal.strategy,
|
||||
category=signal.category.value,
|
||||
close=signal.close,
|
||||
details=dict(signal.details),
|
||||
)
|
||||
|
||||
|
||||
def _result_query(
|
||||
page: int,
|
||||
page_size: int,
|
||||
@@ -330,6 +367,7 @@ __all__ = [
|
||||
"SelectionResultsResponse",
|
||||
"SelectionRunAcceptedResponse",
|
||||
"SelectionRunRequest",
|
||||
"SelectionStockResponse",
|
||||
"get_selection_service",
|
||||
"get_selection_postgres_pool",
|
||||
"selection_router",
|
||||
|
||||
@@ -91,7 +91,7 @@ class FakeSelectionService:
|
||||
|
||||
|
||||
def _run(run_id: str, status: str) -> SelectionRun:
|
||||
signal = SelectionSignal(
|
||||
original_signal = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
@@ -100,6 +100,15 @@ def _run(run_id: str, status: str) -> SelectionRun:
|
||||
close=10.5,
|
||||
details={"j": 12.0},
|
||||
)
|
||||
pullback_signal = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZhixingB1Category.PULLBACK_WHITE,
|
||||
close=10.5,
|
||||
details={"j": 13.0, "rsi": 20.0},
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.runs import SelectionRunItem
|
||||
|
||||
return SelectionRun(
|
||||
@@ -112,7 +121,7 @@ def _run(run_id: str, status: str) -> SelectionRun:
|
||||
eligible_count=1,
|
||||
evaluated_count=1,
|
||||
selected_stock_count=1,
|
||||
signal_count=1,
|
||||
signal_count=2,
|
||||
failed_count=0,
|
||||
coverage=Decimal("1"),
|
||||
items=(
|
||||
@@ -120,11 +129,12 @@ def _run(run_id: str, status: str) -> SelectionRun:
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
status="selected",
|
||||
signal_count=1,
|
||||
signals=(signal,),
|
||||
signal_count=2,
|
||||
signals=(original_signal, pullback_signal),
|
||||
),
|
||||
),
|
||||
signals=(signal,),
|
||||
signals=(original_signal, pullback_signal),
|
||||
stocks_total=1,
|
||||
)
|
||||
|
||||
|
||||
@@ -209,7 +219,7 @@ def test_query_returns_no_data_without_fabricating_a_result() -> None:
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["status"] == "no_data"
|
||||
assert response.json()["signals"] == []
|
||||
assert response.json()["stocks"] == []
|
||||
|
||||
|
||||
def test_query_returns_persisted_signal_details() -> None:
|
||||
@@ -220,12 +230,22 @@ def test_query_returns_persisted_signal_details() -> None:
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["run_id"] == "run-http"
|
||||
assert body["signal_count"] == 1
|
||||
assert body["signal_count"] == 2
|
||||
assert body["page"] == 1
|
||||
assert body["page_size"] == 10
|
||||
assert body["signals_total"] == 1
|
||||
assert body["signals"][0]["category"] == "zhixing_b1_original_b1"
|
||||
assert body["signals"][0]["details"] == {"j": 12.0}
|
||||
assert body["stocks_total"] == 1
|
||||
assert "signals_total" not in body
|
||||
assert "signals" not in body
|
||||
assert len(body["stocks"]) == 1
|
||||
assert body["stocks"][0]["ts_code"] == "000001.SZ"
|
||||
assert [signal["category"] for signal in body["stocks"][0]["signals"]] == [
|
||||
"zhixing_b1_original_b1",
|
||||
"zhixing_b1_pullback_white",
|
||||
]
|
||||
assert [signal["details"] for signal in body["stocks"][0]["signals"]] == [
|
||||
{"j": 12.0},
|
||||
{"j": 13.0, "rsi": 20.0},
|
||||
]
|
||||
|
||||
|
||||
def test_query_forwards_pagination_and_filters() -> None:
|
||||
@@ -269,4 +289,4 @@ def test_run_polling_returns_the_persisted_terminal_result() -> None:
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["status"] == "success"
|
||||
assert response.json()["signals"][0]["category"] == "zhixing_b1_original_b1"
|
||||
assert response.json()["stocks"][0]["signals"][0]["category"] == ("zhixing_b1_original_b1")
|
||||
|
||||
@@ -283,8 +283,10 @@ class LoadConnection:
|
||||
)
|
||||
if "FROM selection_run_item" in query:
|
||||
return LoadResult(rows=[("000001.SZ", "平安银行", "selected", 2, None)])
|
||||
if "COUNT(*) FROM selection_signal" in query:
|
||||
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query:
|
||||
return LoadResult(row=(2,))
|
||||
if "SELECT DISTINCT ts_code" in query:
|
||||
return LoadResult(rows=[("000001.SZ",)])
|
||||
return LoadResult(
|
||||
rows=[
|
||||
(
|
||||
@@ -301,7 +303,7 @@ class LoadConnection:
|
||||
"平安银行",
|
||||
TARGET,
|
||||
"zhixing_b1",
|
||||
ZHIXING_B1_SIGNAL_ORDER[0].value,
|
||||
ZhixingB1Category.ORIGINAL_B1.value,
|
||||
Decimal("10.5"),
|
||||
{},
|
||||
),
|
||||
@@ -309,6 +311,16 @@ class LoadConnection:
|
||||
)
|
||||
|
||||
|
||||
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:
|
||||
self.statements.append((query, parameters))
|
||||
return LoadResult(rows=[])
|
||||
return super().execute(query, parameters)
|
||||
|
||||
|
||||
class LoadResult:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -326,7 +338,9 @@ class LoadResult:
|
||||
return self.rows
|
||||
|
||||
|
||||
def test_get_run_orders_signals_by_formula_priority(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_get_run_pages_stocks_and_loads_all_signals_for_category_matches(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
connection = LoadConnection()
|
||||
|
||||
def connect(database_url: str) -> LoadConnection:
|
||||
@@ -337,7 +351,7 @@ def test_get_run_orders_signals_by_formula_priority(monkeypatch: pytest.MonkeyPa
|
||||
run = PostgresSelectionRunRepository("postgresql://test").get_run(
|
||||
"run-1",
|
||||
query=SelectionResultQuery(
|
||||
page=2,
|
||||
page=1,
|
||||
page_size=1,
|
||||
search="100%",
|
||||
category="pullback",
|
||||
@@ -346,20 +360,57 @@ def test_get_run_orders_signals_by_formula_priority(monkeypatch: pytest.MonkeyPa
|
||||
|
||||
assert run is not None
|
||||
assert [signal.category for signal in run.signals] == [
|
||||
ZHIXING_B1_SIGNAL_ORDER[0],
|
||||
ZhixingB1Category.ORIGINAL_B1,
|
||||
ZHIXING_B1_SIGNAL_ORDER[-1],
|
||||
]
|
||||
assert run.stocks_total == 2
|
||||
count_query, count_parameters = next(
|
||||
(query, parameters)
|
||||
for query, parameters in connection.statements
|
||||
if "COUNT(*) FROM selection_signal" in query
|
||||
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query
|
||||
)
|
||||
assert "name ILIKE %s ESCAPE" in count_query
|
||||
assert count_parameters == ("run-1", "%100\\%%", "%100\\%%", "zhixing_b1_pullback_%")
|
||||
page_query, page_parameters = next(
|
||||
stock_page_query, page_parameters = next(
|
||||
(query, parameters)
|
||||
for query, parameters in connection.statements
|
||||
if "LIMIT %s OFFSET %s" in query
|
||||
if "SELECT DISTINCT ts_code" in query
|
||||
)
|
||||
assert "ORDER BY ts_code, CASE category" in page_query
|
||||
assert page_parameters[-2:] == (1, 1)
|
||||
assert "ORDER BY ts_code" in stock_page_query
|
||||
assert page_parameters[-2:] == (1, 0)
|
||||
signal_query, signal_parameters = next(
|
||||
(query, parameters)
|
||||
for query, parameters in connection.statements
|
||||
if "ts_code = ANY(%s)" in query and "SELECT\n" in query
|
||||
)
|
||||
assert "ORDER BY ts_code, CASE category" in signal_query
|
||||
assert "category LIKE" not in signal_query
|
||||
assert signal_parameters == ("run-1", ["000001.SZ"])
|
||||
|
||||
|
||||
def test_get_run_does_not_load_signals_for_an_empty_stock_page(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
connection = EmptyStockPageConnection()
|
||||
|
||||
def connect(database_url: str) -> EmptyStockPageConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
run = PostgresSelectionRunRepository("postgresql://test").get_run(
|
||||
"run-1",
|
||||
query=SelectionResultQuery(page=3, page_size=1),
|
||||
)
|
||||
|
||||
assert run is not None
|
||||
assert run.stocks_total == 2
|
||||
assert run.signals == ()
|
||||
stock_page_query, stock_page_parameters = next(
|
||||
(query, parameters)
|
||||
for query, parameters in connection.statements
|
||||
if "SELECT DISTINCT ts_code" in query
|
||||
)
|
||||
assert "ORDER BY ts_code" 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)
|
||||
|
||||
Reference in New Issue
Block a user