from __future__ import annotations import unittest 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, Position, 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, 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 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_position_dataclasses_execute_without_type_error(self): with TemporaryDirectory() as directory: state = State.for_strategy(directory, "trend", "A") position = Position( 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 = Position(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 = Position(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_ing_order_from_deal(self): with TemporaryDirectory() as directory: state = State.for_strategy(directory, "trend", "A") position = Position(stock_code="A", volume=100, open_price=10) state.set(StateItem("A", base_order_id="local-1", base_status="ING")) state.reconcile( [position], [], [{"m_strRemark": "local-1|morning"}], ) 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"], [Position(stock_code="A", volume=100, open_price=10)]), 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(cancel_expired=lambda _client: None), state=SimpleNamespace(codes=["A"]), ) 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() if __name__ == "__main__": unittest.main()