feat(selection): 迁移知行B1选股策略
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user