feat(selection): 迁移知行B1选股策略

This commit is contained in:
yuxuanhui
2026-08-08 22:41:45 +08:00
parent 0c999fb828
commit e9d06df5de
32 changed files with 2423 additions and 0 deletions
@@ -0,0 +1,31 @@
"""Application-level state and port mapping tests."""
from datetime import date
from zhixing_server.modules.selection.application.evaluate import EvaluateZhixingB1
from zhixing_server.modules.selection.domain.models import StockHistory
from zhixing_server.modules.selection.domain.ports import MarketDataReaderError
TARGET = date(2024, 1, 2)
class EmptyReader:
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
return StockHistory(ts_code=ts_code, name="", bars=())
class FailingReader:
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
raise MarketDataReaderError(f"database unavailable for {ts_code}")
def test_evaluate_maps_reader_error_to_data_error() -> None:
result = EvaluateZhixingB1(FailingReader()).execute("000001.SZ", TARGET)
assert result.status == "data_error"
assert result.signals == ()
assert "000001.SZ" in (result.reason or "")
def test_evaluate_distinguishes_missing_target_from_reader_error() -> None:
result = EvaluateZhixingB1(EmptyReader()).execute("000001.SZ", TARGET)
assert result.status == "missing_target_bar"
@@ -0,0 +1,69 @@
"""Boundary tests for TDX-style selection indicators."""
import numpy as np
import pandas as pd
import pytest
from zhixing_server.modules.selection.domain.indicators import (
BARSLAST,
COUNT,
CROSS,
EVERY,
HHVBARS,
MA,
REF,
compute_amplitude_params,
compute_kdj,
compute_rsi,
)
def test_rolling_primitives_use_trading_rows_and_keep_every_warmup() -> None:
values = pd.Series([1.0, 2.0, 3.0, 2.0])
assert MA(values, 3).tolist() == [1.0, 1.5, 2.0, 7 / 3]
assert REF(values, 1).isna().iloc[0]
assert EVERY(pd.Series([True, True, True]), 3).tolist() == [False, False, True]
assert COUNT(pd.Series([True, False, True]), 2).tolist() == [1.0, 1.0, 1.0]
def test_hhvbars_and_barslast_are_stable_for_ties_and_missing_prefix() -> None:
values = pd.Series([1.0, 3.0, 3.0, 2.0])
assert HHVBARS(values, 3).tolist() == [0.0, 0.0, 0.0, 1.0]
bars_last = BARSLAST(pd.Series([False, True, False, True]))
assert np.isnan(bars_last.iloc[0])
assert bars_last.iloc[1:].tolist() == [0.0, 1.0, 0.0]
def test_cross_does_not_match_without_a_previous_complete_row() -> None:
assert CROSS(pd.Series([1.0, 3.0, 2.0]), pd.Series([2.0, 2.0, 2.0])).tolist() == [
False,
True,
False,
]
def test_zero_range_and_zero_rsi_denominator_do_not_create_finite_signals() -> None:
frame = pd.DataFrame(
{
"low": [10.0, 10.0, 10.0],
"high": [10.0, 10.0, 10.0],
"close": [10.0, 10.0, 10.0],
}
)
kdj = compute_kdj(frame, 3)
rsi = compute_rsi(frame["close"], 3)
assert kdj["J"].isna().all()
assert rsi.isna().all()
def test_amplitude_parameters_cover_wide_prefix_and_historical_wide_move() -> None:
close = pd.Series([10.0, 10.0, 11.6, 11.0])
assert compute_amplitude_params("688001", close) == (8.0, 0.9)
assert compute_amplitude_params("000001", close) == (8.0, 0.9)
assert compute_amplitude_params("000001", pd.Series([10.0, 10.1])) == (5.0, 1.0)
def test_invalid_indicator_windows_fail_loudly() -> None:
with pytest.raises(ValueError):
MA(pd.Series([1.0]), 0)
@@ -0,0 +1,84 @@
"""PostgreSQL reader contract tests using a fake connection."""
from datetime import date
from typing import cast
import psycopg
import pytest
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
PostgresMarketDataReader,
)
class FakeConnection:
def __init__(self, rows: list[tuple[object, ...]]) -> None:
self.rows = rows
self.query: str | None = None
self.parameters: tuple[object, ...] | None = None
def __enter__(self) -> "FakeConnection":
return self
def __exit__(self, *args: object) -> None:
return None
def execute(self, query: str, parameters: tuple[object, ...]) -> "FakeResult":
self.query = query
self.parameters = parameters
return FakeResult(self.rows)
class FakeResult:
def __init__(self, rows: list[tuple[object, ...]]) -> None:
self.rows = rows
def fetchall(self) -> list[tuple[object, ...]]:
return self.rows
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(
[
(
"000001.SZ",
"平安银行",
date(2024, 1, 2),
"10",
"11",
"9",
"10.5",
"1000",
None,
None,
),
(
"000001.SZ",
"平安银行",
date(2024, 1, 3),
"10.5",
"11",
"10",
"10.8",
"1200",
"1.2",
"100000",
),
]
)
def connect(database_url: str) -> FakeConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
history = PostgresMarketDataReader("postgresql://test").load_history(
"000001.SZ", date(2024, 1, 3)
)
assert [bar.trade_date for bar in history.bars] == [date(2024, 1, 2), date(2024, 1, 3)]
assert history.daily_basic[date(2024, 1, 2)].turnover_rate is None
assert history.daily_basic[date(2024, 1, 3)].total_mv == 100000.0
assert connection.parameters == ("000001.SZ", date(2024, 1, 3))
assert "source_adj = 'qfq'" in cast(str, connection.query)
assert "trade_date <= %s" in cast(str, connection.query)
@@ -0,0 +1,94 @@
"""Behavior tests for explicit-date Zhixing B1 evaluation."""
from datetime import date, timedelta
import pandas as pd
import pytest
from zhixing_server.modules.selection.domain import zhixing_b1
from zhixing_server.modules.selection.domain.models import SelectionBar, StockHistory
from zhixing_server.modules.selection.domain.zhixing_b1 import (
MINIMUM_HISTORY,
ZHIXING_B1_SIGNAL_ORDER,
ZhixingB1Strategy,
compute_signal_masks,
prepare_zhixing_b1_indicators,
)
def make_history(count: int = MINIMUM_HISTORY, code: str = "000001.SZ") -> StockHistory:
bars = tuple(
SelectionBar(
trade_date=date(2020, 1, 1) + timedelta(days=index),
open=10.0 + index * 0.02,
high=10.2 + index * 0.02,
low=9.9 + index * 0.02,
close=10.1 + index * 0.02,
volume=1000.0 + (index % 7) * 30,
)
for index in range(count)
)
return StockHistory(ts_code=code, name="测试股票", bars=bars)
def test_strategy_has_seven_stable_categories_and_prepared_masks() -> None:
history = make_history()
frame = pd.DataFrame(
{
"open": [bar.open for bar in history.bars],
"high": [bar.high for bar in history.bars],
"low": [bar.low for bar in history.bars],
"close": [bar.close for bar in history.bars],
"volume": [bar.volume for bar in history.bars],
}
)
prepared = prepare_zhixing_b1_indicators(frame, history.ts_code)
masks = compute_signal_masks(prepared)
assert tuple(masks) == ZHIXING_B1_SIGNAL_ORDER
assert all(mask.dtype == bool for mask in masks.values())
assert all(len(mask) == len(history.bars) for mask in masks.values())
def test_strategy_explicit_target_ignores_future_rows() -> None:
history = make_history()
target = history.bars[-1].trade_date
future = SelectionBar(
trade_date=target + timedelta(days=1),
open=1.0,
high=100.0,
low=0.5,
close=99.0,
volume=1_000_000.0,
)
with_future = StockHistory(history.ts_code, history.name, history.bars + (future,))
strategy = ZhixingB1Strategy()
assert strategy.evaluate(with_future, target) == strategy.evaluate(history, target)
def test_strategy_returns_missing_and_warmup_states() -> None:
strategy = ZhixingB1Strategy()
history = make_history(MINIMUM_HISTORY - 1)
target = history.bars[-1].trade_date
assert strategy.evaluate(history, target).status == "insufficient_history"
assert strategy.evaluate(history, target + timedelta(days=1)).status == "missing_target_bar"
def test_strategy_keeps_all_same_day_subsignals_in_priority_order(
monkeypatch: pytest.MonkeyPatch,
) -> None:
history = make_history()
target = history.bars[-1].trade_date
def all_masks(frame: pd.DataFrame) -> dict[zhixing_b1.ZhixingB1Category, pd.Series]:
return {
category: pd.Series(True, index=frame.index) for category in ZHIXING_B1_SIGNAL_ORDER
}
monkeypatch.setattr(zhixing_b1, "compute_signal_masks", all_masks)
result = ZhixingB1Strategy().evaluate(history, target)
assert result.status == "selected"
assert tuple(signal.category for signal in result.signals) == ZHIXING_B1_SIGNAL_ORDER
assert len({signal.identity for signal in result.signals}) == 7