Files
big-qmt/py-client/tests/test_trend.py

360 lines
15 KiB
Python

from __future__ import annotations
import unittest
from concurrent.futures import ThreadPoolExecutor
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 APIError, Assets, OrderItem, Portfolio, PositionItem, Tick
from strategy.trend.order import OrderBook, PlaceOrderRequest
from strategy.trend.open import do_open
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 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 FailedOrderClient:
def passorder_latest_tagged(self, *_args):
raise APIError(502, "QMT did not return a valid order reference")
class TrendTests(unittest.TestCase):
def test_trend_order_id_format(self):
self.assertRegex(OrderBook.new_order_id(), r"^trend-[0-9a-f]{24}$")
def test_open_records_only_pending_order(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
forgotten = []
runtime = SimpleNamespace(
client=FakeClient(),
orders=OrderBook(),
state=state,
open_watch=SimpleNamespace(forget=forgotten.append),
)
do_open(runtime, "000001.SZ", 100, "morning", 12.345)
item = state.get("000001.SZ")
self.assertRegex(item.base_order_id, r"^trend-[0-9a-f]{24}$")
self.assertEqual(item.base_qty, 0)
self.assertEqual(item.base_cost, 0)
self.assertEqual(item.base_status, "ING")
self.assertEqual(forgotten, ["000001.SZ"])
def test_pending_base_order_survives_position_delay(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
state.set(StateItem(
"000001.SZ",
base_order_id="trend-12345678",
base_status="ING",
))
pending = OrderItem(
"1", "000001.SZ", "BUY", "", "50", None, 100,
local_order_id="trend-12345678",
)
state.reconcile([], [pending])
self.assertEqual(
state.get("000001.SZ").base_order_id,
"trend-12345678",
)
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_order_api_error_returns_false_with_traceback(self):
book = OrderBook()
request = PlaceOrderRequest(FailedOrderClient(), 23, "000001.SZ", 100, "local", "morning")
with self.assertLogs(level="ERROR") as captured:
self.assertFalse(book.place(request))
output = "\n".join(captured.output)
self.assertIn("HTTP状态=502", output)
self.assertIn("Traceback", output)
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, orders)
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_position_log_contains_code_name_profit_and_loss_actions(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = PositionItem(
stock_code="000001.SZ", stock_name="平安银行", volume=100,
can_use_volume=100, open_price=10, market_value=1000,
)
state.sync_positions([position])
runtime = SimpleNamespace(
client=FakeClient(), state=state, orders=OrderBook(),
add_watch=SimpleNamespace(triggered=lambda *_args: False),
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",
),
)
with self.assertLogs(level="INFO") as captured:
manage_positions(
runtime,
{"000001.SZ": Tick(last_price=10.1)},
[position],
True,
5000,
)
output = "\n".join(captured.output)
self.assertIn("代码=000001.SZ", output)
self.assertIn("名称=平安银行", output)
self.assertIn("止盈=未触发", output)
self.assertIn("补仓=未启用", output)
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",
traded_volume=50, trade_price=10.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",
traded_volume=50, trade_price=10.3,
),
],
)
item = state.get("A")
self.assertEqual(item.base_status, STATUS_OK)
self.assertEqual(item.base_qty, 100)
self.assertEqual(item.base_cost, 10.2)
canceled = OrderItem("2", "A", "BUY", "", "54", None, 50, "local-1")
item = state.get("A")
item.base_status = "ING"
state.set(item)
state.reconcile([position], [completed, canceled])
self.assertEqual(state.get("A").base_status, "")
def test_reconcile_records_filled_add_order(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = PositionItem(stock_code="A", volume=200, open_price=10)
state.set(StateItem(
"A",
base_qty=100,
base_cost=10,
base_status=STATUS_OK,
added_order_id="add-1",
added_status="ING",
))
completed = OrderItem(
"1", "A", "BUY", "", "56", None, 100, "add-1",
traded_volume=100, trade_amount=950,
)
state.reconcile([position], [completed])
item = state.get("A")
self.assertEqual(item.added_status, STATUS_OK)
self.assertEqual(item.added_num, 1)
self.assertEqual(item.added_qty, 100)
self.assertEqual(item.added_cost, 9.5)
def test_low_cash_still_runs_position_management(self):
client = SimpleNamespace(
portfolio=lambda: Portfolio(
assets=Assets(total=10000, available=10),
positions={"A": PositionItem(stock_code="A", volume=100, open_price=10)},
orders=[],
),
full_tick=lambda _codes: {"A": Tick(last_price=11)},
)
with ThreadPoolExecutor(max_workers=2) as executor:
runtime = SimpleNamespace(
client=client,
account_cfg=SimpleNamespace(min_cash_ratio=0.1),
global_cfg=SimpleNamespace(api_host="http://example"),
orders=SimpleNamespace(refresh=lambda _client, _orders: None, data=[]),
state=SimpleNamespace(
codes=["A"],
reconcile=lambda *_args: None,
),
executor=executor,
)
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_state_without_broker_order_allows_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="ING",
))
state.save()
client = SimpleNamespace(
portfolio=lambda: Portfolio(
assets=Assets(total=10000, available=5000),
positions={},
orders=[],
),
full_tick=lambda _codes: {"A": Tick(last_price=10)},
)
signal = SimpleNamespace(code="A", signal_key="morning")
with ThreadPoolExecutor(max_workers=2) as executor:
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),
executor=executor,
)
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_called_once()
self.assertEqual(state.codes, [])
if __name__ == "__main__":
unittest.main()