216 lines
9.4 KiB
Python
216 lines
9.4 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
from datetime import datetime, timedelta
|
|
from tempfile import TemporaryDirectory
|
|
from types import SimpleNamespace
|
|
from unittest.mock import patch
|
|
|
|
from libs.grid_take_profit import GridState, GridTrailingTracker
|
|
from sdk import Assets, OrderItem, PositionItem, Tick
|
|
from strategy.trend.order import OrderBook, PlaceOrderRequest
|
|
from strategy.trend.positions import LOSS_TIERS, handle_loss, manage_positions
|
|
from strategy.trend.boot import RunOnce
|
|
from strategy.trend.state import STATUS_OK, STATUS_UNKNOWN, State, StateItem
|
|
|
|
|
|
class FakeClient:
|
|
def __init__(self):
|
|
self.orders = []
|
|
|
|
def passorder_latest_tagged(self, op, code, volume, strategy_name, order_id):
|
|
self.orders.append((op, code, volume, strategy_name, order_id))
|
|
return {"status": "success", "order_ref": f"broker-{len(self.orders)}"}
|
|
|
|
|
|
class FakeOrderClient:
|
|
def __init__(self, orders):
|
|
self.orders = orders
|
|
self.canceled = []
|
|
|
|
def trade_detail_data(self, _datatype):
|
|
return self.orders
|
|
|
|
def cancel_by_id(self, order_id):
|
|
self.canceled.append(order_id)
|
|
|
|
|
|
class TrendTests(unittest.TestCase):
|
|
def test_grid_states_and_account_isolation(self):
|
|
tracker = GridTrailingTracker(1)
|
|
self.assertEqual(tracker.observe("A:code", 2.1).state, GridState.ARMED)
|
|
self.assertEqual(tracker.observe("A:code", 3.1).state, GridState.RAISED)
|
|
self.assertEqual(tracker.observe("A:code", 2.9).state, GridState.RETREAT)
|
|
self.assertEqual(tracker.observe("B:code", 2.9).state, GridState.ARMED)
|
|
tracker.retain([])
|
|
self.assertEqual(tracker.observe("A:code", 2.9).state, GridState.ARMED)
|
|
|
|
def test_order_book_locks_duplicate_order(self):
|
|
client = FakeClient()
|
|
book = OrderBook()
|
|
request = PlaceOrderRequest(client, 23, "000001.SZ", 100, "local", "morning")
|
|
self.assertTrue(book.place(request))
|
|
self.assertTrue(book.busy("000001.SZ", "BUY"))
|
|
|
|
def test_refresh_tracks_active_and_completed_and_cancels_expired(self):
|
|
old = datetime.now() - timedelta(seconds=20)
|
|
orders = [
|
|
OrderItem("active", "A", "BUY", "", "49", old, 100),
|
|
OrderItem("completed", "B", "SELL", "", "56", old, 100),
|
|
OrderItem("canceled", "C", "BUY", "", "54", old, 100),
|
|
OrderItem("failed", "D", "BUY", "", "57", old, 100),
|
|
]
|
|
client = FakeOrderClient(orders)
|
|
book = OrderBook(cancel_timeout_sec=10)
|
|
|
|
book.refresh(client)
|
|
|
|
self.assertEqual({item.id for item in book.data}, {"completed"})
|
|
self.assertEqual(client.canceled, ["active"])
|
|
|
|
def test_position_dataclasses_execute_without_type_error(self):
|
|
with TemporaryDirectory() as directory:
|
|
state = State.for_strategy(directory, "trend", "A")
|
|
position = PositionItem(
|
|
stock_code="000001.SZ", volume=100, can_use_volume=100,
|
|
open_price=10, market_value=1000,
|
|
)
|
|
state.sync_positions([position])
|
|
runtime = SimpleNamespace(
|
|
client=FakeClient(), state=state, orders=OrderBook(),
|
|
open_watch=SimpleNamespace(forget=lambda _code: None),
|
|
add_watch=SimpleNamespace(triggered=lambda *_args: False, forget=lambda _code: None),
|
|
profit_tracker=GridTrailingTracker(1),
|
|
account_cfg=SimpleNamespace(
|
|
account_id="A", excluded_codes=[], grid_step_pct=1,
|
|
enable_loss_add_position=False, buy_value=5000,
|
|
strategy="trend",
|
|
),
|
|
)
|
|
manage_positions(runtime, {"000001.SZ": Tick(last_price=10.1)}, [position], True, 5000)
|
|
|
|
def test_loss_tier_boundary_does_not_overflow(self):
|
|
self.assertEqual(len(LOSS_TIERS), 2)
|
|
with TemporaryDirectory() as directory:
|
|
state = State.for_strategy(directory, "trend", "A")
|
|
position = PositionItem(stock_code="A", volume=100, open_price=10, market_value=1000)
|
|
state.sync_positions([position])
|
|
item = state.get("A")
|
|
item.added_num = len(LOSS_TIERS)
|
|
state.set(item)
|
|
runtime = SimpleNamespace(
|
|
state=state, account_cfg=SimpleNamespace(buy_value=5000, strategy="trend"),
|
|
add_watch=SimpleNamespace(triggered=lambda *_args: True), orders=OrderBook(),
|
|
client=FakeClient(),
|
|
)
|
|
decision = handle_loss(runtime, position, Tick(last_price=5), -60, 5000)
|
|
self.assertFalse(decision.submitted)
|
|
|
|
def test_loss_tiers_zero_and_one(self):
|
|
with TemporaryDirectory() as directory:
|
|
state = State.for_strategy(directory, "trend", "A")
|
|
position = PositionItem(stock_code="A", volume=100, open_price=10, market_value=1000)
|
|
state.sync_positions([position])
|
|
runtime = SimpleNamespace(
|
|
state=state, account_cfg=SimpleNamespace(buy_value=5000, strategy="trend"),
|
|
add_watch=SimpleNamespace(triggered=lambda *_args: False),
|
|
orders=OrderBook(), client=FakeClient(),
|
|
)
|
|
first = handle_loss(runtime, position, Tick(last_price=7), -30, 5000)
|
|
self.assertIn("等待", first.message)
|
|
item = state.get("A")
|
|
item.added_num = 1
|
|
state.set(item)
|
|
before_second_tier = handle_loss(runtime, position, Tick(last_price=6), -40, 5000)
|
|
self.assertEqual(before_second_tier.message, "")
|
|
second = handle_loss(runtime, position, Tick(last_price=5), -50, 5000)
|
|
self.assertIn("等待", second.message)
|
|
|
|
def test_reconcile_split_orders_complete_only_when_all_are_status_56(self):
|
|
with TemporaryDirectory() as directory:
|
|
state = State.for_strategy(directory, "trend", "A")
|
|
position = PositionItem(stock_code="A", volume=100, open_price=10)
|
|
state.set(StateItem("A", base_order_id="local-1", base_status="ING"))
|
|
completed = OrderItem("1", "A", "BUY", "", "56", None, 50, "local-1")
|
|
processing = OrderItem("2", "A", "BUY", "", "50", None, 50, "local-1")
|
|
|
|
state.reconcile([position], [completed, processing])
|
|
self.assertEqual(state.get("A").base_status, "ING")
|
|
|
|
state.reconcile(
|
|
[position],
|
|
[completed, OrderItem("2", "A", "BUY", "", "56", None, 50, "local-1")],
|
|
)
|
|
self.assertEqual(state.get("A").base_status, STATUS_OK)
|
|
|
|
def test_low_cash_still_runs_position_management(self):
|
|
client = SimpleNamespace(
|
|
assets=lambda: Assets(total=10000, available=10),
|
|
positions=lambda: (["A"], [PositionItem(stock_code="A", volume=100, open_price=10)]),
|
|
trade_detail_data=lambda _datatype: [],
|
|
deals=lambda: [],
|
|
full_tick=lambda _codes: {"A": Tick(last_price=11)},
|
|
)
|
|
runtime = SimpleNamespace(
|
|
client=client,
|
|
account_cfg=SimpleNamespace(min_cash_ratio=0.1),
|
|
global_cfg=SimpleNamespace(api_host="http://example"),
|
|
orders=SimpleNamespace(refresh=lambda _client: None),
|
|
state=SimpleNamespace(
|
|
codes=["A"],
|
|
unresolved_codes=[],
|
|
reconcile=lambda *_args: None,
|
|
),
|
|
)
|
|
with (
|
|
patch("strategy.trend.boot.trading_time", return_value=True),
|
|
patch("strategy.trend.boot.market_allow_open", return_value=True),
|
|
patch("strategy.trend.boot.open_signal") as open_mock,
|
|
patch("strategy.trend.boot.manage_positions") as manage_mock,
|
|
):
|
|
RunOnce(runtime, [])
|
|
open_mock.assert_not_called()
|
|
manage_mock.assert_called_once()
|
|
|
|
def test_unknown_order_without_position_blocks_reopen(self):
|
|
with TemporaryDirectory() as directory:
|
|
state = State.for_strategy(directory, "trend", "A")
|
|
state.set(StateItem(
|
|
"A",
|
|
base_order_id="missing-order",
|
|
base_qty=100,
|
|
base_status=STATUS_UNKNOWN,
|
|
))
|
|
state.save()
|
|
client = SimpleNamespace(
|
|
assets=lambda: Assets(total=10000, available=5000),
|
|
positions=lambda: ([], []),
|
|
trade_detail_data=lambda _datatype: [],
|
|
deals=lambda: [],
|
|
full_tick=lambda _codes: {"A": Tick(last_price=10)},
|
|
)
|
|
runtime = SimpleNamespace(
|
|
client=client,
|
|
account_cfg=SimpleNamespace(min_cash_ratio=0.1),
|
|
global_cfg=SimpleNamespace(api_host="http://example"),
|
|
orders=OrderBook(),
|
|
state=state,
|
|
open_watch=SimpleNamespace(forget=lambda _code: None),
|
|
add_watch=SimpleNamespace(forget=lambda _code: None),
|
|
)
|
|
signal = SimpleNamespace(code="A", signal_key="morning")
|
|
with (
|
|
patch("strategy.trend.boot.trading_time", return_value=True),
|
|
patch("strategy.trend.boot.market_allow_open", return_value=True),
|
|
patch("strategy.trend.boot.open_signal") as open_mock,
|
|
patch("strategy.trend.boot.manage_positions"),
|
|
):
|
|
RunOnce(runtime, [signal])
|
|
|
|
open_mock.assert_not_called()
|
|
self.assertTrue(state.has_unresolved_order("A"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|