359 lines
15 KiB
Python
359 lines
15 KiB
Python
"""ETF 开仓与持仓:白名单/入场门槛/反弹确认、主出口、副出口、百分比补仓。"""
|
||
|
||
from datetime import date, datetime, timedelta
|
||
import unittest
|
||
from unittest.mock import Mock, patch
|
||
|
||
from config import AccountConfig, EtfConfig, EtfDefaults, EtfSymbolConfig, GlobalConfig
|
||
from libs.order import OrderBook
|
||
from libs.runtime import Runtime
|
||
from libs.signal import SignalItem
|
||
from libs import watch
|
||
from libs.watch import DipWatch
|
||
from sdk import Assets, PositionItem, Tick
|
||
from strategy.etf import open as etf_open
|
||
from strategy.etf import positions as etf_positions
|
||
|
||
CODE = "510300.SH"
|
||
OTHER = "159915.SZ"
|
||
# 固定"当前时刻",与 tick 的时间戳保持同一交易日且不过期。
|
||
NOW = datetime(2026, 9, 16, 10, 0, 0)
|
||
|
||
|
||
class FrozenDateTime(datetime):
|
||
"""冻结 ``datetime.now()``,其余行为与标准库一致。"""
|
||
|
||
@classmethod
|
||
def now(cls, tz=None):
|
||
return NOW if tz is None else NOW.astimezone(tz)
|
||
|
||
|
||
def stamp(now: datetime) -> str:
|
||
return now.strftime("%Y%m%d %H:%M:%S")
|
||
|
||
|
||
def tick(price: float, now: datetime | None = None) -> Tick:
|
||
now = now or NOW
|
||
return Tick(last_price=price, last_close=price, raw={"timetag": stamp(now)})
|
||
|
||
|
||
def symbol(**overrides) -> EtfSymbolConfig:
|
||
base = dict(is_t0=False, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7)
|
||
base.update(overrides)
|
||
return EtfSymbolConfig(**base)
|
||
|
||
|
||
def etf_config(**symbol_overrides) -> EtfConfig:
|
||
return EtfConfig(
|
||
defaults=EtfDefaults(),
|
||
symbols={CODE: symbol(**symbol_overrides)},
|
||
)
|
||
|
||
|
||
def position(volume: int, cost: float, can_use: int | None = None, name: str = "") -> PositionItem:
|
||
return PositionItem(
|
||
stock_code=CODE,
|
||
stock_name=name,
|
||
volume=volume,
|
||
open_price=cost,
|
||
can_use_volume=volume if can_use is None else can_use,
|
||
yesterday_volume=volume if can_use is None else can_use,
|
||
last_price=cost,
|
||
)
|
||
|
||
|
||
def signal(entry: float = 9.65, price: float = 10.0, code: str = CODE) -> SignalItem:
|
||
return SignalItem(
|
||
signal_key="etf",
|
||
code=code,
|
||
last_close=price,
|
||
tech_indicator={
|
||
"etf_entry": entry,
|
||
"etf_price": price,
|
||
"etf_grid": 1.0,
|
||
"etf_add_price": price * 0.97,
|
||
"etf_ma60": 10.0,
|
||
},
|
||
)
|
||
|
||
|
||
class ETFTradeBase(unittest.TestCase):
|
||
def setUp(self):
|
||
self.client = Mock()
|
||
self.client.passorder.return_value = {"status": "success"}
|
||
# Runtime.__post_init__ 会拉一次服务端初始化数据,测试里不发真实请求。
|
||
patch("libs.runtime.get_json", side_effect=OSError("offline")).start()
|
||
self.run = Runtime(
|
||
client=self.client,
|
||
global_cfg=GlobalConfig(api_host="http://api.test"),
|
||
account_cfg=AccountConfig(account_id="acct", strategy="etf", min_cash_ratio=0.0),
|
||
etf_cfg=etf_config(),
|
||
orders=OrderBook(),
|
||
open_watch=DipWatch(expire_seconds=600, rebound_threshold=0.5),
|
||
add_watch=DipWatch(expire_seconds=600, rebound_threshold=0.5),
|
||
)
|
||
for module in (etf_open, etf_positions, watch):
|
||
patch.object(module, "trading_time", return_value=True, create=True).start()
|
||
patch.object(module, "datetime", FrozenDateTime).start()
|
||
patch.dict(etf_positions._progress, {}, clear=True).start()
|
||
patch.dict(etf_positions._trackers, {}, clear=True).start()
|
||
self.addCleanup(patch.stopall)
|
||
|
||
def last_order(self) -> dict:
|
||
self.assertTrue(self.client.passorder.called, "未提交任何委托")
|
||
return self.client.passorder.call_args.kwargs
|
||
|
||
|
||
class OpenSignalTests(ETFTradeBase):
|
||
def test_whitelist_outside_config_is_skipped(self):
|
||
item = signal(code=OTHER)
|
||
etf_open.open_signal(self.run, {OTHER: tick(9.4)}, [item])
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_price_above_entry_no_observation_no_order(self):
|
||
etf_open.open_signal(self.run, {CODE: tick(9.9)}, [signal(entry=9.65)])
|
||
self.client.passorder.assert_not_called()
|
||
self.assertEqual(self.run.open_watch.data, {})
|
||
|
||
def test_seesaw_below_entry_requires_rebound_confirmation(self):
|
||
item = signal(entry=9.65)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item]) # 进入入场区,记低点
|
||
etf_open.open_signal(self.run, {CODE: tick(9.39)}, [item]) # 刷新低点
|
||
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item]) # 反弹不足 0.5%
|
||
self.client.passorder.assert_not_called()
|
||
etf_open.open_signal(self.run, {CODE: tick(9.42)}, [item]) # (9.42-9.39)/9.39 = 0.32% 仍不足
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_rebound_places_base_limit_order_at_anchor(self):
|
||
item = signal(entry=9.65)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item])
|
||
etf_open.open_signal(self.run, {CODE: tick(9.45)}, [item]) # 反弹 0.53% 确认
|
||
request = self.last_order()
|
||
self.assertEqual(request["op_type"], 23)
|
||
self.assertEqual(request["volume"], 1000)
|
||
self.assertEqual(request["price"], 9.45)
|
||
self.assertEqual(request["pr_type"], 11)
|
||
self.assertEqual(request["strategy_name"], "etf")
|
||
self.assertEqual(self.run.open_watch.data, {})
|
||
|
||
def test_leaving_entry_band_forgets_the_watch(self):
|
||
item = signal(entry=9.65)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item])
|
||
etf_open.open_signal(self.run, {CODE: tick(9.8)}, [item])
|
||
self.assertEqual(self.run.open_watch.data, {})
|
||
|
||
def test_missing_entry_indicator_is_skipped(self):
|
||
item = signal()
|
||
item.tech_indicator.clear()
|
||
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item])
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_stale_tick_is_skipped(self):
|
||
old = datetime(2026, 9, 16, 9, 50, 0)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.4, now=old)}, [signal(entry=9.65)])
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_previous_day_tick_is_skipped(self):
|
||
yesterday = datetime(2026, 9, 15, 14, 0, 0)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.4, now=yesterday)}, [signal(entry=9.65)])
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_insufficient_budget_cancels_the_anchor(self):
|
||
self.run.client.assets.return_value = Assets(total=1000.0, available=100.0)
|
||
item = signal(entry=9.65)
|
||
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item])
|
||
etf_open.open_signal(self.run, {CODE: tick(9.45)}, [item])
|
||
self.client.passorder.assert_not_called()
|
||
self.assertEqual(self.run.open_watch.data, {})
|
||
|
||
|
||
class MainExitTests(ETFTradeBase):
|
||
def test_profit_above_target_clears_the_whole_grid(self):
|
||
etf_positions.manage_positions(
|
||
self.run, {CODE: tick(10.2)}, [position(2000, 10.0)], True, 100000.0
|
||
)
|
||
request = self.last_order()
|
||
self.assertEqual(request["op_type"], 24)
|
||
self.assertEqual(request["volume"], 2000)
|
||
self.assertEqual(request["strategy_name"], "etf")
|
||
self.assertEqual(request["price"], 10.2)
|
||
self.assertEqual(request["pr_type"], 11)
|
||
|
||
def test_profit_below_target_does_not_sell(self):
|
||
etf_positions.manage_positions(
|
||
self.run, {CODE: tick(10.05)}, [position(2000, 10.0)], True, 100000.0
|
||
)
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_t_plus_1_position_bought_today_is_not_sellable(self):
|
||
held = position(1000, 10.0, can_use=0)
|
||
held.yesterday_volume = 0
|
||
etf_positions.manage_positions(
|
||
self.run, {CODE: tick(10.5)}, [held], True, 100000.0
|
||
)
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_t0_symbol_sells_on_the_same_day(self):
|
||
self.run.etf_cfg = etf_config(is_t0=True)
|
||
held = position(1000, 10.0, can_use=1000)
|
||
held.yesterday_volume = 0
|
||
etf_positions.manage_positions(
|
||
self.run, {CODE: tick(10.5)}, [held], True, 100000.0
|
||
)
|
||
self.assertEqual(self.last_order()["volume"], 1000)
|
||
|
||
|
||
class LevelExitTests(ETFTradeBase):
|
||
"""副出口:主出口在盈亏率 ≥1% 时会先吃掉整仓,因此这里直接喂盈亏率验证峰值回撤。
|
||
|
||
inner_step = 0.7、inner_grids = 2:只有峰值抬到第 2 格后的回撤才允许卖出。
|
||
"""
|
||
|
||
def observe(self, series, cost: float = 11.6):
|
||
held = position(1000, cost)
|
||
symbol = self.run.etf_cfg.symbols[CODE]
|
||
level = etf_positions.position_level(self.run, held)
|
||
decisions = []
|
||
for pnl_rate in series:
|
||
price = cost * (1 + pnl_rate / 100)
|
||
decisions.append(
|
||
etf_positions.handle_level_exit(
|
||
self.run, symbol, held, tick(price), pnl_rate, level
|
||
)
|
||
)
|
||
return decisions
|
||
|
||
def test_peak_retreat_sells_only_that_level(self):
|
||
first, second, third = self.observe([0.5, 1.5, 1.2])
|
||
self.assertFalse(first.submitted) # 首次观察,只建基准
|
||
self.assertFalse(second.submitted) # 峰值抬到第 2 格
|
||
self.assertTrue(third.submitted) # 回撤到第 1 格
|
||
request = self.last_order()
|
||
self.assertEqual(request["op_type"], 24)
|
||
self.assertEqual(request["volume"], 1000)
|
||
self.assertEqual(request["strategy_name"], "etf")
|
||
|
||
def test_retreat_below_inner_grids_is_held(self):
|
||
first, second = self.observe([0.5, 0.1])
|
||
self.assertFalse(first.submitted)
|
||
self.assertFalse(second.submitted) # 峰值只有 0 格
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_peak_is_kept_when_the_order_is_rejected(self):
|
||
self.client.passorder.return_value = {"status": "rejected"}
|
||
self.run.orders.place = Mock(return_value=False)
|
||
first, second, third = self.observe([0.5, 1.5, 1.2])
|
||
self.assertFalse(third.submitted)
|
||
# 下单失败必须保留峰值:下一轮同样能再次触发。
|
||
fourth = self.observe([1.2])[0]
|
||
self.assertFalse(fourth.submitted)
|
||
|
||
def test_manage_positions_runs_the_secondary_exit(self):
|
||
"""成本 11.6、现价 11.7/11.65 的盈亏率都低于 1%,主出口不参与。"""
|
||
held = position(1000, 11.6)
|
||
etf_positions.manage_positions(self.run, {CODE: tick(11.7)}, [held], True, 0.0)
|
||
etf_positions.manage_positions(self.run, {CODE: tick(11.65)}, [held], True, 0.0)
|
||
self.client.passorder.assert_not_called() # 峰值未达 2 格
|
||
|
||
|
||
class AddTests(ETFTradeBase):
|
||
"""补仓规则单测:直接调 handle_add,避免其它出口的委托锁干扰。"""
|
||
|
||
def held(self, volume: int = 3000, cost: float = 10.0) -> PositionItem:
|
||
return position(volume, cost)
|
||
|
||
def add(self, price: float, volume: int = 3000, cost: float = 10.0,
|
||
budget: float = 100000.0, level: int | None = None,
|
||
symbol_overrides: dict | None = None, first_low: float | None = None):
|
||
if symbol_overrides:
|
||
self.run.etf_cfg = etf_config(**symbol_overrides)
|
||
held = self.held(volume, cost)
|
||
if first_low is not None:
|
||
# 先造出一个观察低点,再由本次调用验证反弹确认。
|
||
self.run.add_watch.triggered("补仓", CODE, first_low)
|
||
return etf_positions.handle_add(
|
||
self.run,
|
||
self.run.etf_cfg.symbols[CODE],
|
||
held,
|
||
tick(price),
|
||
price,
|
||
budget,
|
||
level if level is not None else etf_positions.position_level(self.run, held),
|
||
)
|
||
|
||
def test_add_requires_add_pct_drop(self):
|
||
decision = self.add(9.95)
|
||
self.assertFalse(decision.submitted)
|
||
self.client.passorder.assert_not_called()
|
||
self.assertEqual(self.run.add_watch.data, {})
|
||
|
||
def test_add_waits_for_rebound_before_buying(self):
|
||
# 跌幅 4% ≥ add_pct 3%,但还没反弹确认:只观察,不下单。
|
||
first = self.add(9.60)
|
||
self.assertFalse(first.submitted)
|
||
self.assertIn(CODE, self.run.add_watch.data)
|
||
# 从观察低点 9.60 反弹 0.63%:确认后按现价买一档。
|
||
second = self.add(9.66)
|
||
self.assertTrue(second.submitted)
|
||
request = self.last_order()
|
||
self.assertEqual(request["op_type"], 23)
|
||
self.assertEqual(request["volume"], 1000)
|
||
self.assertEqual(request["price"], 9.66)
|
||
self.assertEqual(request["strategy_name"], "etf")
|
||
self.assertEqual(self.run.add_watch.data, {})
|
||
|
||
def test_add_stops_at_max_adds(self):
|
||
decision = self.add(9.60, volume=10000)
|
||
self.assertFalse(decision.submitted)
|
||
self.assertIn("档", decision.message)
|
||
|
||
def test_add_respects_max_shares(self):
|
||
decision = self.add(9.60, volume=3000, symbol_overrides={"max_shares": 3000})
|
||
self.assertFalse(decision.submitted)
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_add_needs_budget(self):
|
||
# 跌幅 4%、反弹 0.63% 都满足,但预算为 0:不消耗观察状态,也不下单。
|
||
first = self.add(9.60, budget=0.0)
|
||
self.assertFalse(first.submitted)
|
||
second = self.add(9.66, budget=0.0, first_low=9.60)
|
||
self.assertFalse(second.submitted)
|
||
self.assertIn(CODE, self.run.add_watch.data)
|
||
# 资金到位后同一个观察低点仍可确认。
|
||
third = self.add(9.66, budget=100000.0, first_low=None)
|
||
self.assertTrue(third.submitted)
|
||
|
||
def test_add_uses_broker_cost_as_previous_level(self):
|
||
# 上一档 = 券商成本 9.5:跌到 9.16 是 3.58% ≥ 3%,反弹到 9.21 确认。
|
||
decision = self.add(9.21, cost=9.5, first_low=9.16)
|
||
self.assertTrue(decision.submitted)
|
||
self.assertEqual(self.last_order()["price"], 9.21)
|
||
|
||
def test_add_blocked_when_market_disallows(self):
|
||
held = self.held()
|
||
etf_positions.manage_positions(self.run, {CODE: tick(9.6)}, [held], False, 100000.0)
|
||
etf_positions.manage_positions(self.run, {CODE: tick(9.66)}, [held], False, 100000.0)
|
||
self.client.passorder.assert_not_called()
|
||
|
||
def test_add_blocked_by_in_flight_buy(self):
|
||
self.run.orders.busy_cache.set("BUY-" + CODE, True, timeout=180)
|
||
self.add(9.60)
|
||
self.add(9.66)
|
||
self.client.passorder.assert_not_called()
|
||
|
||
|
||
class PositionLevelTests(ETFTradeBase):
|
||
def test_level_is_derived_from_volume(self):
|
||
for volume, expected in ((1000, 1), (2000, 2), (3500, 4), (10000, 10)):
|
||
with self.subTest(volume=volume):
|
||
self.assertEqual(
|
||
etf_positions.position_level(self.run, position(volume, 10.0)), expected
|
||
)
|
||
|
||
def test_unknown_symbol_yields_baseline_level(self):
|
||
self.assertEqual(etf_positions.position_level(self.run, position(1000, 10.0)), 1)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|