Files
big-qmt/py-client/tests/test_trend_collector.py
2026-09-11 14:38:24 +08:00

107 lines
5.2 KiB
Python

import importlib
import io
import logging
import unittest
from concurrent.futures import Future
from contextlib import ExitStack, redirect_stdout
from types import SimpleNamespace
from unittest.mock import Mock, patch
from libs import collector, snapshot
from sdk import Assets, DealItem, PositionItem
from strategy.trend import boot
class TrendCollectorTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
with patch('logging.FileHandler', return_value=logging.NullHandler()):
cls.app = importlib.import_module('main')
def setUp(self):
old_snapshot = snapshot._collector_snapshot
self.addCleanup(setattr, snapshot, '_collector_snapshot', old_snapshot)
snapshot._collector_snapshot = None
def test_submission_reads_latest_cache_and_skips_empty(self):
with patch.object(collector, 'collector_push') as push:
collector.submit_trend_data()
push.assert_not_called()
snapshot.cache_portfolio('account', Assets(available=100), [], [])
assets = Assets(available=200)
positions = [PositionItem(stock_code='600000.SH', volume=100)]
deals = [DealItem(stock_code='600000.SH', volume=100)]
snapshot.cache_portfolio('account', assets, positions, deals)
collector.submit_trend_data()
push.assert_called_once_with('account', assets, positions, deals)
uploaded = push.call_args.args
uploaded[1].available = 0
uploaded[2].clear()
uploaded[3][0].volume = 0
self.assertEqual(snapshot.get_collector_snapshot()[3][0].volume, 100)
self.assertEqual(snapshot.get_collector_snapshot()[1].available, 200)
self.assertEqual(len(snapshot.get_collector_snapshot()[2]), 1)
def test_run_once_caches_portfolio_without_submitting_data(self):
completed = Future()
completed.set_result(None)
run = SimpleNamespace(
client=Mock(), orders=Mock(), executor=Mock(),
account_cfg=SimpleNamespace(account_id='account', min_cash_ratio=0.1),
)
assets = Assets(available=100, total=1000)
run.client.portfolio.return_value = SimpleNamespace(assets=assets, positions={}, orders=[])
deals = [DealItem(stock_code='600000.SH', volume=100)]
run.client.deals.return_value = deals
run.client.full_tick.return_value = {}
run.executor.submit.return_value = completed
with patch.object(boot, 'trading_time', return_value=True), \
patch.object(boot, 'market_allow_open', return_value=True), \
patch.object(collector, 'collector_push') as push, redirect_stdout(io.StringIO()):
boot.RunOnce(run, [])
self.assertEqual(snapshot.get_collector_snapshot(), ('account', assets, [], deals))
run.client.deals.assert_called_once_with()
push.assert_not_called()
run.executor.submit.assert_called_once_with(boot.manage_positions, run, {}, [], True, 100)
def test_submission_serializes_deals(self):
snapshot.cache_portfolio(
'account', Assets(available=100), [],
[DealItem(stock_code='600000.SH', volume=100)],
)
with patch.object(collector.httpx, 'post') as post:
collector.submit_trend_data()
payload = post.call_args.kwargs['json']
self.assertEqual(payload['account_id'], 'account')
self.assertEqual(payload['deals'][0]['stock_code'], '600000.SH')
self.assertEqual(payload['deals'][0]['volume'], 100)
def test_main_registers_five_minute_collector_job(self):
for strategy in ('trend', 'zt'):
with self.subTest(strategy=strategy), ExitStack() as stack:
scheduler = Mock(running=True)
stack.enter_context(patch.object(self.app, 'BackgroundScheduler', return_value=scheduler))
stack.enter_context(patch.object(self.app, 'require_windows', return_value=True))
stack.enter_context(patch.object(self.app, 'check_single_instance', return_value=True))
stack.enter_context(patch.object(self.app, 'wait_for_qmt_api'))
stack.enter_context(patch.object(self.app.config, 'load'))
stack.enter_context(patch.object(self.app.config, 'global_config', SimpleNamespace(api_host='unused')))
stack.enter_context(patch.object(self.app.config, 'account_config', SimpleNamespace(strategy=strategy)))
stack.enter_context(patch.dict(self.app.STRATEGIES, {
strategy: SimpleNamespace(start_strategy=Mock()),
}))
self.assertEqual(self.app.main(), 0)
jobs = [call for call in scheduler.add_job.call_args_list
if call.kwargs.get('id') == 'trend_collector']
self.assertEqual(len(jobs), 1)
if jobs:
self.assertIs(jobs[0].args[0], collector.submit_trend_data)
self.assertEqual(jobs[0].kwargs['trigger'], 'interval')
self.assertEqual(jobs[0].kwargs['minutes'], 5)
scheduler.start.assert_called_once()
scheduler.shutdown.assert_called_once_with(wait=True)
if __name__ == '__main__':
unittest.main()