feat libs,sdk,trend

This commit is contained in:
2026-09-07 14:04:26 +08:00
parent d37f9edefc
commit e320de3241
30 changed files with 515 additions and 507 deletions

View File

@@ -1,75 +1,89 @@
import sqlite3
import ast
import tempfile
import unittest
from contextlib import closing
from dataclasses import asdict
from dataclasses import asdict, fields
from datetime import datetime
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import Mock
from libs.order import OrderBook as ActiveOrders
from libs.orderbook import OrderBook
from sdk.models import DealItem
from sdk.models import Assets, DealItem, OrderItem, PositionItem
from sdk.portfolio import PortfolioMixin
class DealModelTests(unittest.TestCase):
class ApiModelTests(unittest.TestCase):
def setUp(self):
self.raw = {
'm_strOrderSysID': 'sys-123',
'm_strInstrumentID': '600000',
'm_strExchangeID': 'SH',
'm_strInstrumentName': 'Test stock',
'm_nOffsetFlag': '24',
'm_nOrderStatus': '56',
'm_nVolumeTotal': '0',
'm_nVolumeTraded': '200',
'm_nOrderTime': '101530',
'm_strInsertDate': '20260907',
'm_strInsertTime': '10:15:30',
'm_strRemark': 'zt-t-sell-local1|zt',
'm_dPrice': '15.8',
'm_dTradePrice': '15.6',
'm_dTradeAmount': '3120',
}
source = Path(__file__).resolve().parents[2] / 'api' / 'qmt_rest_new.py'
names = {'format_assets', 'format_holding', 'format_orders', 'format_deals'}
nodes = [n for n in ast.parse(source.read_text(encoding='utf-8')).body
if isinstance(n, ast.FunctionDef) and n.name in names]
ns = {'HTTPError': RuntimeError}
exec(compile(ast.Module(body=nodes, type_ignores=[]), str(source), 'exec'), ns)
attrs = {n.attr: '' if n.attr.startswith('m_str') else 0
for node in nodes for n in ast.walk(node)
if isinstance(n, ast.Attribute) and n.attr.startswith('m_')}
attrs.update(m_strInstrumentID='600000', m_strExchangeID='SH',
m_strOrderSysID='sys1', m_strRemark='trend-BUY-1|trend',
m_nOffsetFlag=23, m_nOrderStatus=56, m_nVolume=100,
m_nVolumeTraded=100, m_nVolumeTotalOriginal=100,
m_dPrice=10.0, m_dTradeAmount=1000.0, m_dBalance=2000.0,
m_dAvailable=1000.0, m_strInsertDate='20260907',
m_strInsertTime='100000', m_strTradeDate='20260907', m_strTradeTime='100000')
obj = SimpleNamespace(**attrs)
self.assets = ns['format_assets']([obj])
self.positions = ns['format_holding']([obj])
self.orders = ns['format_orders']([obj])
self.deals = ns['format_deals']([obj])
self.client = PortfolioMixin()
self.client._get_json = {
'/api/portfolio/assets': self.assets, '/api/portfolio/positions': self.positions,
'/api/portfolio/order': self.orders, '/api/portfolio/deal': self.deals,
'/api/portfolio': {'assets': self.assets, 'positions': self.positions, 'orders': self.orders},
}.__getitem__
def test_all_api_fields_are_parsed(self):
self.assertEqual(asdict(DealItem.from_trade_detail(self.raw)), {
'sys_order_id': 'sys-123', 'local_order_id': 'zt-t-sell-local1',
'code': '600000.SH', 'instrument_id': '600000', 'exchange_id': 'SH',
'name': 'Test stock', 'offset_flag': '24', 'side': 'SELL', 'status': '56',
'remaining_volume': 0, 'traded_volume': 200, 'order_time': 101530,
'insert_date': '20260907', 'insert_time': '10:15:30',
'remark': 'zt-t-sell-local1|zt', 'price': 15.8,
'trade_price': 15.6, 'trade_amount': 3120.0,
})
def test_models_exactly_match_api_keys_and_values(self):
for model, row in ((Assets, self.assets), (PositionItem, self.positions['600000.SH']),
(OrderItem, self.orders[0]), (DealItem, self.deals[0])):
self.assertEqual({field.name for field in fields(model)}, set(row))
self.assertEqual(asdict(model(**row)), row)
def test_api_deals_response_returns_deal_items(self):
client = PortfolioMixin()
for response in ({'deals': [self.raw]}, [self.raw]):
with self.subTest(response_type=type(response).__name__):
client._get_json = lambda path: response
deals = client.deals()
self.assertIsInstance(deals[0], DealItem)
self.assertEqual(deals[0].sys_order_id, 'sys-123')
self.assertEqual(deals[0].traded_volume, 200)
def test_all_endpoints(self):
self.assertEqual(asdict(self.client.assets()), self.assets)
codes, positions = self.client.positions()
self.assertEqual(codes, ['600000.SH'])
self.assertEqual(asdict(positions[0]), self.positions[codes[0]])
self.assertEqual(asdict(self.client.orders()[0]), self.orders[0])
self.assertEqual(asdict(self.client.deals()[0]), self.deals[0])
portfolio = self.client.portfolio()
self.assertEqual(asdict(portfolio.positions[codes[0]]), self.positions[codes[0]])
self.assertEqual(asdict(portfolio.orders[0]), self.orders[0])
def test_model_matches_sql_columns_and_restart(self):
deal = DealItem.from_trade_detail(self.raw)
def test_derived_properties_and_order_cache(self):
order = self.client.orders()[0]
self.assertEqual(order.side, 'BUY')
self.assertEqual(order.local_order_id, 'trend-BUY-1')
self.assertEqual(order.created_at, datetime(2026, 9, 7, 10))
order.order_status = 50
order.insert_date = '20000101'
client = Mock()
book = ActiveOrders('trend')
book.refresh(client, [order])
client.cancel_by_id.assert_called_once_with('sys1')
self.assertTrue(book.busy('600000.SH', 'BUY'))
def test_storage_and_price_fallback(self):
deal = self.client.deals()[0]
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / 'state.db'
book = OrderBook(path)
book.sync_deals([deal])
with closing(sqlite3.connect(path)) as db:
columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')}
self.assertEqual(columns, {'id', *asdict(deal)})
loaded = OrderBook(path).deals['sys-123']
expected = asdict(deal)
expected['insert_date'] = '2026-09-07'
self.assertEqual({key: value for key, value in loaded.items() if key != 'id'}, expected)
def test_missing_amount_uses_execution_price_not_order_price(self):
deal = DealItem.from_trade_detail({**self.raw, 'm_dTradeAmount': '0'})
row = OrderBook.deal_record(deal)
self.assertEqual(row['trade_amount'], 3120)
deal.trade_price = 0
loaded = OrderBook(path).deals['sys1']
self.assertEqual(loaded['volume'], deal.volume)
self.assertEqual(loaded['trade_date'], '2026-09-07')
deal.trade_amount = 0
self.assertEqual(OrderBook.deal_record(deal)['trade_amount'], 1000)
deal.price = 0
with self.assertRaises(ValueError):
OrderBook.deal_record(deal)

View File

@@ -2,6 +2,7 @@ import sqlite3
import tempfile
import unittest
from contextlib import closing
from dataclasses import asdict, fields
from pathlib import Path
from datetime import datetime
from unittest.mock import patch
@@ -23,18 +24,12 @@ class OrderBookTests(unittest.TestCase):
def deal(self, kind, sys_order_id, qty, price, date='2026-09-01'):
prefix = {'base': 'zt-base-', 'sell': 'zt-t-sell-', 'buy': 'zt-t-buy-'}[kind]
return DealItem.from_trade_detail({
'm_strOrderSysID': sys_order_id,
'm_strInstrumentID': '600000', 'm_strExchangeID': 'SH',
'm_strInstrumentName': 'Test stock',
'm_nOffsetFlag': '24' if kind == 'sell' else '23',
'm_nOrderStatus': '56', 'm_nVolumeTotal': '0',
'm_nVolumeTraded': str(qty), 'm_nOrderTime': '100000',
'm_strInsertDate': date, 'm_strInsertTime': '10:00:00',
'm_strRemark': prefix + 'order1|zt',
'm_dPrice': str(price + 1), 'm_dTradePrice': str(price),
'm_dTradeAmount': str(qty * price),
})
return DealItem(
order_sys_id=sys_order_id, stock_code='600000.SH',
offset_flag=24 if kind == 'sell' else 23,
volume=qty, price=price, trade_amount=qty * price,
trade_date=date, trade_time='10:00:00', remark=prefix + 'order1|zt',
)
def test_partial_fills_restart_dedup_and_daily_cycle(self):
state = TState(self.path)
@@ -83,10 +78,12 @@ class OrderBookTests(unittest.TestCase):
insert.assert_called_once()
self.assertEqual(len(insert.call_args.args[1]), 2)
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
self.assertEqual(book.deals['d1']['insert_date'], '2026-09-01')
self.assertEqual(book.deals['d2']['traded_volume'], 60)
self.assertEqual(book.deals['d1']['trade_date'], '2026-09-01')
self.assertEqual(book.deals['d2']['volume'], 60)
self.assertEqual(book.deals['d2']['order_local_id'], 'zt-base-order1')
book = OrderBook(self.path)
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
self.assertEqual(book.deals['d1']['order_local_id'], 'zt-base-order1')
with patch.object(book, '_insert_deals') as insert:
book.sync_deals([first, second])
book.sync_deals([])
@@ -97,25 +94,25 @@ class OrderBookTests(unittest.TestCase):
book = OrderBook(self.path)
first = self.deal('base', 'd1', 100, 10)
invalid = self.deal('base', 'd2', 100, 10)
invalid.side = 'INVALID'
invalid.offset_flag = -1
with self.assertRaises(sqlite3.IntegrityError):
book.sync_deals([first, invalid])
self.assertEqual(book.deals, {})
self.assertEqual(book.deals_sys_ids, set())
self.assertEqual(OrderBook(self.path).deals, {})
invalid.side = 'BUY'
invalid.offset_flag = 23
book.sync_deals([first, invalid])
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
def test_load_refreshes_all_caches(self):
book = OrderBook(self.path)
writer = OrderBook(self.path)
writer.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, [])
writer.sync_positions([PositionItem(stock_code='600000.SH', volume=100)])
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
book.load()
self.assertEqual(book.positions['600000.SH']['status'], READY)
self.assertEqual(book.positions['600000.SH']['volume'], 100)
self.assertEqual(book.deals_sys_ids, {'d1'})
self.assertEqual(book.deals['d1']['local_order_id'], 'zt-base-order1')
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
def test_first_start_after_full_sale_keeps_buyback_quantity(self):
state = TState(self.path)
@@ -148,12 +145,17 @@ class OrderBookTests(unittest.TestCase):
tables = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='table'")}
self.assertEqual(tables, {'positions', 'deals', 'sqlite_sequence'})
columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')}
self.assertFalse({'kind', 'confirmed_date', 'deal_ids', 'deal_id', 'order_id', 'qty', 'filled_qty', 'filled_cost', 'amount', 'trade_date', 'trade_time'} & columns)
self.assertTrue({'sys_order_id', 'local_order_id'} <= columns)
self.assertEqual(columns, {'id', 'order_local_id', *(field.name for field in fields(DealItem))})
indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")}
self.assertTrue({'idx_positions_code', 'idx_positions_base_order_id',
'idx_positions_added_order_id', 'idx_deals_local_order_id',
'idx_deals_code_date', 'idx_deals_date_time'} <= indexes)
self.assertTrue({'idx_positions_stock_code', 'idx_deals_order_sys_id',
'idx_deals_order_ref', 'idx_deals_stock_code_date', 'idx_deals_date_time'} <= indexes)
for index, expected in (
('idx_deals_order_sys_id', ['order_sys_id']),
('idx_deals_order_ref', ['order_local_id']),
('idx_deals_stock_code_date', ['stock_code']),
('idx_deals_date_time', ['trade_date']),
):
self.assertEqual([row[2] for row in db.execute(f'PRAGMA index_info({index})')], expected)
self.assertNotIn('kind', state.deals[0])
self.assertEqual(TState(self.path).deals, state.deals)
state.deals.append(dict(state.deals[0]))
@@ -168,29 +170,23 @@ class OrderBookTests(unittest.TestCase):
def test_position_columns_defaults_indexes_and_stable_id(self):
store = OrderBook(self.path)
with closing(sqlite3.connect(self.path)) as db:
columns = [row[1] for row in db.execute('PRAGMA table_info(positions)')]
self.assertEqual(columns, [
'id', 'code', 'base_order_id', 'base_qty', 'base_cost',
'added_order_id', 'added_num', 'added_qty', 'added_cost', 'status',
])
columns = {row[1] for row in db.execute('PRAGMA table_info(positions)')}
self.assertEqual(columns, {'id', *(field.name for field in fields(PositionItem))})
indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')}
self.assertEqual(indexes, {
'idx_positions_code', 'idx_positions_base_order_id', 'idx_positions_added_order_id',
})
store.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, [])
position = store.positions['600000.SH']
first_id = position['id']
self.assertGreater(first_id, 0)
self.assertEqual(position['base_order_id'], '')
self.assertEqual(position['base_qty'], 0)
self.assertEqual(position['added_cost'], 0.0)
position.update(base_order_id='base1', base_qty=200, base_cost=10.5,
added_order_id='add1', added_num=1, added_qty=100, added_cost=9.0,
status='ACTIVE')
store.save({'600000.SH': position}, [])
self.assertEqual(store.positions['600000.SH'], position)
store.save({}, [])
store.save({'600001.SH': {'code': '600001.SH', 'status': READY}}, [])
self.assertEqual(indexes, {'idx_positions_stock_code'})
position = PositionItem(stock_code='600000.SH', volume=100, open_price=10,
stock_name='stock', can_use_volume=100, float_profit=-2.5)
store.sync_positions([position])
saved = store.positions[position.stock_code]
first_id = saved['id']
self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, asdict(position))
position.volume = 200
store.sync_positions([position])
self.assertEqual(store.positions[position.stock_code]['id'], first_id)
self.assertEqual(store.positions[position.stock_code]['volume'], 200)
store.sync_positions([])
self.assertEqual(store.positions, {})
store.sync_positions([PositionItem(stock_code='600001.SH')])
self.assertGreater(store.positions['600001.SH']['id'], first_id)
def test_base_split_fills_and_snapshot_do_not_double_count(self):
@@ -207,12 +203,12 @@ class OrderBookTests(unittest.TestCase):
state = TState(self.path)
first = self.deal('base', 'd1', 100, 10, '20260901')
other = self.deal('base', 'd2', 100, 10)
other.local_order_id = 'trend-base-order'
other.remark = 'trend-base-order'
state.reconcile([], [first, other])
first.insert_date = '2026-09-01'
first.trade_date = '2026-09-01'
state.reconcile([], [first])
self.assertEqual(len(state.deals), 1)
self.assertEqual(state.deals[0]['insert_date'], '2026-09-01')
self.assertEqual(state.deals[0]['trade_date'], '2026-09-01')
if __name__ == '__main__':

View File

@@ -0,0 +1,88 @@
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
from sdk import Assets, 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 = boot._collector_snapshot
self.addCleanup(setattr, boot, '_collector_snapshot', old_snapshot)
boot._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()
boot._cache_portfolio('account', Assets(available=100), [])
assets = Assets(available=200)
positions = [PositionItem(stock_code='600000.SH', volume=100)]
boot._cache_portfolio('account', assets, positions)
collector.submit_trend_data()
push.assert_called_once_with('account', assets, positions)
uploaded = push.call_args.args
uploaded[1].available = 0
uploaded[2].clear()
self.assertEqual(boot.get_collector_snapshot()[1].available, 200)
self.assertEqual(len(boot.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=[])
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(boot.get_collector_snapshot(), ('account', assets, []))
push.assert_not_called()
run.executor.submit.assert_called_once_with(boot.manage_positions, run, {}, [], True, 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()