This commit is contained in:
2026-09-11 10:44:34 +08:00
parent 4e693d773a
commit 6e554e1565
10 changed files with 140 additions and 36 deletions

View File

@@ -19,4 +19,4 @@ signals:
dcm: {url: /a/dcm_signal, timezone: "*", gt_last_price_is_open: false} dcm: {url: /a/dcm_signal, timezone: "*", gt_last_price_is_open: false}
morning: {url: /a/morning_signal, timezone: "9:30-10:30", gt_last_price_is_open: true} morning: {url: /a/morning_signal, timezone: "9:30-10:30", gt_last_price_is_open: true}
tail: {url: /a/tail_signal, timezone: "14:30-14:55", gt_last_price_is_open: false} tail: {url: /a/tail_signal, timezone: "14:30-14:55", gt_last_price_is_open: false}
arbitrage: {url: /a/arbitrage_signal, timezone: "*", gt_last_price_is_open: false} arbitrage: {url: /a/arbitrage_signal, timezone: "13:00-14:50", gt_last_price_is_open: false}

View File

@@ -35,6 +35,7 @@ class OrderBook:
def __init__( def __init__(
self, order_prefix: str, lock_timeout_sec: int = 180, cancel_timeout_sec: float = 30 self, order_prefix: str, lock_timeout_sec: int = 180, cancel_timeout_sec: float = 30
) -> None: ) -> None:
self.order_prefix = order_prefix
self.lock_timeout_sec = max(1, lock_timeout_sec) self.lock_timeout_sec = max(1, lock_timeout_sec)
self.cancel_timeout_sec = timedelta(seconds=cancel_timeout_sec) self.cancel_timeout_sec = timedelta(seconds=cancel_timeout_sec)
self.data: list[OrderItem] = [] self.data: list[OrderItem] = []

View File

@@ -137,6 +137,7 @@ class State:
amount, date, deal.trade_time, deal.remark, deal.close_profit, 0), amount, date, deal.trade_time, deal.remark, deal.close_profit, 0),
) )
self.load_deals() self.load_deals()
def sync_state(self, positions: list[PositionItem]) -> None: def sync_state(self, positions: list[PositionItem]) -> None:
"""同步完整持仓:无状态则插入底仓,已有则保留,清仓则删除。 """同步完整持仓:无状态则插入底仓,已有则保留,清仓则删除。
""" """
@@ -165,20 +166,29 @@ class State:
) )
self.load_state() self.load_state()
def archiving(self, *, base_order_prefix: str = '') -> dict[str, str]: def archiving(self) -> None:
"""将未归档成交累加到当前底仓和加仓;失败证券保留记录供重试。""" """将未归档成交累加到当前底仓和加仓;失败证券保留记录供重试。
errors = {}
归档是把成交反映到持仓状态中,再标记为已处理,不会删除成交记录。
本地委托号以 zt-base- 开头的买入计入底仓,其他买入计入补仓。
不返回结果;失败原因记录到日志,对应成交保留未归档标记供重试。
"""
with closing(self._connect()) as db, db: with closing(self._connect()) as db, db:
# 提前取得数据库写入锁,让持仓更新和成交标记在同一事务内完成。
db.execute('BEGIN IMMEDIATE') db.execute('BEGIN IMMEDIATE')
# 只找尚未处理的成交23、48 是买入24、49 是卖出。
codes = db.execute( codes = db.execute(
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0 AND offset_flag IN (23, 24, 48, 49)' 'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
).fetchall() ).fetchall()
for entry in codes: for entry in codes:
code = entry['stock_code'] code = entry['stock_code']
# 每只股票设一个回滚点;这只处理失败时,不撤销其他股票的结果。
db.execute('SAVEPOINT archive_stock') db.execute('SAVEPOINT archive_stock')
try: try:
current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone() current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone()
# 已有持仓就接着计算;没有记录则从底仓、补仓均为零开始。
state = dict(current) if current else asdict(StateItem(stock_code=code)) state = dict(current) if current else asdict(StateItem(stock_code=code))
# 按成交日期、时间、记录编号依次处理,保证先买后卖等顺序正确。
deals = db.execute( deals = db.execute(
'SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49) ' 'SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49) '
"ORDER BY trade_date, REPLACE(trade_time, ':', ''), id", (code,) "ORDER BY trade_date, REPLACE(trade_time, ':', ''), id", (code,)
@@ -186,34 +196,43 @@ class State:
for deal in deals: for deal in deals:
qty = deal['volume'] qty = deal['volume']
if deal['offset_flag'] in (23, 48): if deal['offset_flag'] in (23, 48):
bucket = 'base' if base_order_prefix and deal['order_local_id'].startswith(base_order_prefix) else 'added' # 按已有的 ZT 底仓委托号约定识别,无需调用方传入规则。
bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added'
total = state[f'{bucket}_qty'] + qty total = state[f'{bucket}_qty'] + qty
# 新均价 =(原数量 × 原均价 + 本次成交金额)÷ 买入后总数量。
state[f'{bucket}_price'] = ( state[f'{bucket}_price'] = (
state[f'{bucket}_qty'] * state[f'{bucket}_price'] + deal['trade_amount'] state[f'{bucket}_qty'] * state[f'{bucket}_price'] + deal['trade_amount']
) / total ) / total
state[f'{bucket}_qty'] = total state[f'{bucket}_qty'] = total
# 保存这一类仓位最近一次买入的委托号和成交时间。
state[f'{bucket}_order_local_id'] = deal['order_local_id'] state[f'{bucket}_order_local_id'] = deal['order_local_id']
state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip() state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip()
else: else:
total = state['base_qty'] + state['added_qty'] total = state['base_qty'] + state['added_qty']
# 卖出不能超过本地记录的总持仓;不一致时留待排查后重试。
if qty > total: if qty > total:
raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}') raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}')
# 卖出优先扣补仓;补仓未卖完时只减数量,保留原均价。
if qty < state['added_qty']: if qty < state['added_qty']:
state['added_qty'] -= qty state['added_qty'] -= qty
else: else:
# 补仓全部卖完后,剩余卖出量扣底仓,并清空补仓信息。
state['base_qty'] = total - qty state['base_qty'] = total - qty
state['added_qty'] = 0 state['added_qty'] = 0
state['added_price'] = 0.0 state['added_price'] = 0.0
state['added_order_local_id'] = state['added_created_at'] = '' state['added_order_local_id'] = state['added_created_at'] = ''
if state['base_qty'] + state['added_qty'] == 0: if state['base_qty'] + state['added_qty'] == 0:
# 中途清仓先重置,后续若又买入,就从零重新累计。
state = asdict(StateItem(stock_code=code)) state = asdict(StateItem(stock_code=code))
if state['base_qty'] + state['added_qty'] == 0: if state['base_qty'] + state['added_qty'] == 0:
# 全部成交处理完仍无持仓,则删除该股票的策略状态。
db.execute('DELETE FROM state WHERE stock_code = ?', (code,)) db.execute('DELETE FROM state WHERE stock_code = ?', (code,))
else: else:
# 更新持仓,保留策略状态及已有记录主键。 # 更新持仓,保留策略状态及已有记录主键。
state['status'] = current['status'] if current else state['status'] state['status'] = current['status'] if current else state['status']
state.pop('id', None) state.pop('id', None)
columns = tuple(state) columns = tuple(state)
# 没有该股票就新增,已有则更新仓位字段,不替换原记录主键。
db.execute( db.execute(
f"INSERT INTO state ({', '.join(columns)}) " f"INSERT INTO state ({', '.join(columns)}) "
f"VALUES ({', '.join(':' + key for key in columns)}) " f"VALUES ({', '.join(':' + key for key in columns)}) "
@@ -221,17 +240,19 @@ class State:
+ ', '.join(f'{key} = excluded.{key}' for key in columns if key != 'stock_code'), + ', '.join(f'{key} = excluded.{key}' for key in columns if key != 'stock_code'),
state, state,
) )
# 持仓处理成功后才标记成交,防止下次重复加仓或重复扣减。
db.execute( db.execute(
'UPDATE deals SET is_arch = 1 WHERE stock_code = ? ' 'UPDATE deals SET is_arch = 1 WHERE stock_code = ? '
'AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49)', (code,) 'AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49)', (code,)
) )
except (ValueError, sqlite3.IntegrityError) as exc: except (ValueError, sqlite3.IntegrityError) as exc:
# 撤销这只股票的全部归档修改,成交仍保持未归档,供下次重试。
db.execute('ROLLBACK TO archive_stock') db.execute('ROLLBACK TO archive_stock')
errors[code] = str(exc)
log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc) log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc)
finally: finally:
# 释放当前股票的回滚点;整个事务在退出外层 with 时提交。
db.execute('RELEASE archive_stock') db.execute('RELEASE archive_stock')
# 数据库提交完成后刷新内存缓存,让策略读到最新持仓和归档标记。
self.load() self.load()
return errors

View File

@@ -137,7 +137,7 @@ def handle_profit(
volume = position.can_use_volume - position.can_use_volume % 100 volume = position.can_use_volume - position.can_use_volume % 100
if volume <= 0: if volume <= 0:
return TradeDecision(False, "无可用整手持仓") return TradeDecision(False, "无可用整手持仓")
order_id = runtime.orders.new_order_id("SELL") order_id = runtime.orders.new_order_id("TREN","SELL")
request = PlaceOrderRequest( request = PlaceOrderRequest(
op=OP_SELL, op=OP_SELL,
code=position.stock_code, code=position.stock_code,
@@ -178,7 +178,7 @@ def handle_loss(
if volume <= 0 or amount > available: if volume <= 0 or amount > available:
return TradeDecision(False, "本轮可用资金不足") return TradeDecision(False, "本轮可用资金不足")
order_id = runtime.orders.new_order_id("BUY") order_id = runtime.orders.new_order_id("TREN","BUY")
request = PlaceOrderRequest( request = PlaceOrderRequest(
op=OP_BUY, op=OP_BUY,
code=position.stock_code, code=position.stock_code,

View File

@@ -2,8 +2,10 @@
import logging as log import logging as log
import time import time
from contextlib import closing
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from tempfile import TemporaryDirectory
import config import config
from libs.calc import trading_time from libs.calc import trading_time
@@ -45,7 +47,8 @@ def StartZT() -> None:
log.exception('[ZT] 本轮失败,下一轮重试') log.exception('[ZT] 本轮失败,下一轮重试')
time.sleep(30 - time.time() % 30) time.sleep(30 - time.time() % 30)
# 收盘后补记最后一轮成交,不再下单。 # 收盘后补记最后一轮成交,不再下单。
sync_account_state(state, list(client.portfolio().positions.values()), client.deals()) deals = client.deals()
sync_account_state(state, list(client.portfolio().positions.values()), deals)
def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None: def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
@@ -85,12 +88,31 @@ def sync_account_state(
) -> None: ) -> None:
"""初次持仓作为底仓;后续只按成交减仓,避免延迟快照删除持仓。""" """初次持仓作为底仓;后续只按成交减仓,避免延迟快照删除持仓。"""
zt_deals = [d for d in deals if d.get_local_order_id.startswith('zt-')] zt_deals = [d for d in deals if d.get_local_order_id.startswith('zt-')]
state.sync_deals(zt_deals)
if initialize: if initialize:
state.sync_state(positions) if state.state or state.deals:
raise ValueError('ZT 初始化仅允许空状态库,禁止覆盖已有持仓和成交')
# 在同目录临时库完成初始化,再原子替换空库,避免中途退出留下半份快照。
with TemporaryDirectory(dir=state.path.parent, prefix='zt-init-') as directory:
initial = State(Path(directory) / 'state.db')
initial.sync_deals(zt_deals)
initial.sync_state(positions)
# 初始快照已包含历史成交;保留成交计算做 T 欠仓,但不再累加持仓。
with closing(initial._connect()) as db, db:
db.execute('UPDATE deals SET is_arch = 1')
initial.path.replace(state.path)
state.load()
return return
errors = state.archiving(base_order_prefix='zt-base-') state.sync_deals(zt_deals)
if errors: # 在外层统一归档:新增成交写入后处理,也重试此前失败的未归档成交。
raise ValueError(f'ZT 成交归档失败:{errors}') state.archiving()
# 从刷新后的成交缓存检查失败记录,避免归档未完成时继续交易。
pending = sorted({d['stock_code'] for d in state.deals.values() if d['is_arch'] == 0})
if pending:
raise ValueError(f'ZT 成交归档未完成:{pending},原因见归档日志')
traded = {d['stock_code'] for d in state.deals.values()} traded = {d['stock_code'] for d in state.deals.values()}
state.sync_state([p for p in positions if p.stock_code not in traded], remove_missing=False) # sync_state 要求完整持仓。保留成交账本中的现有仓位,避免延迟快照删仓;
# 只从账户快照补入没有策略成交历史的股票,避免刚卖完又被旧快照重建。
holdings = [PositionItem(stock_code=code, volume=row['base_qty'] + row['added_qty'])
for code, row in state.state.items()]
holdings.extend(p for p in positions if p.stock_code not in traded and p.stock_code not in state.state)
state.sync_state(holdings)

View File

@@ -42,7 +42,7 @@ def open_signal(run: Runtime, ticks, signals, available: float) -> float:
# 等待价格从观察低点反弹,防止直接接下跌中的“飞刀”。 # 等待价格从观察低点反弹,防止直接接下跌中的“飞刀”。
if not run.open_watch.triggered("ZT 建仓", item.code, price): if not run.open_watch.triggered("ZT 建仓", item.code, price):
continue continue
order_id = run.orders.new_order_id("ZT","base") order_id = run.orders.new_order_id("zt", "base")
request = PlaceOrderRequest( request = PlaceOrderRequest(
OP_BUY, item.code, volume, order_id, "zt" OP_BUY, item.code, volume, order_id, "zt"
) )

View File

@@ -74,7 +74,7 @@ def manage_positions(
if not force_buy_back and not run.add_watch.triggered('ZT 买回', code, price): if not force_buy_back and not run.add_watch.triggered('ZT 买回', code, price):
continue continue
available -= amount available -= amount
request = PlaceOrderRequest(OP_BUY, code, volume, run.orders.new_order_id('t-buy'), 'zt') request = PlaceOrderRequest(OP_BUY, code, volume, run.orders.new_order_id('zt', 't-buy'), 'zt')
if run.orders.place(run.client, request): if run.orders.place(run.client, request):
run.add_watch.forget(code) run.add_watch.forget(code)
log.info('[ZT 买回] %s %d%s', code, volume, ',尾盘买回' if force_buy_back else '') log.info('[ZT 买回] %s %d%s', code, volume, ',尾盘买回' if force_buy_back else '')
@@ -93,7 +93,7 @@ def manage_positions(
volume = int(min(position.can_use_volume, recorded * run.account_cfg.zt_sell_ratio)) // 100 * 100 volume = int(min(position.can_use_volume, recorded * run.account_cfg.zt_sell_ratio)) // 100 * 100
if volume < (200 if code.startswith('688') else 100): if volume < (200 if code.startswith('688') else 100):
continue continue
request = PlaceOrderRequest(OP_SELL, code, volume, run.orders.new_order_id('t-sell'), 'zt') request = PlaceOrderRequest(OP_SELL, code, volume, run.orders.new_order_id('zt', 't-sell'), 'zt')
if run.orders.place(run.client, request): if run.orders.place(run.client, request):
log.info('[ZT 卖出] %s %d 股,按实际成交买回', code, volume) log.info('[ZT 卖出] %s %d 股,按实际成交买回', code, volume)
except Exception: except Exception:

View File

@@ -52,6 +52,20 @@ class ArchivingTests(unittest.TestCase):
self.assertAlmostEqual(row['added_price'], 1620 / 150) self.assertAlmostEqual(row['added_price'], 1620 / 150)
self.assertEqual(row['added_order_local_id'], 'late') self.assertEqual(row['added_order_local_id'], 'late')
def test_no_argument_archiving_recognizes_base_and_added_orders(self):
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=23)
self.insert_deal('zt-base-second', 100, 1200, '10:01:00', flag=48)
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=23)
self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['base_price']), (200, 11))
self.assertEqual((row['added_qty'], row['added_price']), (100, 9))
self.assertEqual(row['base_order_local_id'], 'zt-base-second')
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
saved = dict(row)
self.assertIsNone(self.book.archiving())
self.assertEqual(self.book.state['600000.SH'], saved)
def test_sell_added_then_clear_base(self): def test_sell_added_then_clear_base(self):
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)]) self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
self.insert_deal('buy1', 40, 400, '09:59:00') self.insert_deal('buy1', 40, 400, '09:59:00')
@@ -88,8 +102,9 @@ class ArchivingTests(unittest.TestCase):
def test_excess_sell_rolls_back(self): def test_excess_sell_rolls_back(self):
self.insert_deal('buy', 50, 500, '10:00:00') self.insert_deal('buy', 50, 500, '10:00:00')
self.insert_deal('sell', 100, 1200, '10:01:00', flag=49) self.insert_deal('sell', 100, 1200, '10:01:00', flag=49)
errors = self.book.archiving() with self.assertLogs(level='WARNING') as logs:
self.assertIn('600000.SH', errors) self.assertIsNone(self.book.archiving())
self.assertIn('600000.SH', '\n'.join(logs.output))
restarted = State(self.book.path) restarted = State(self.book.path)
self.assertEqual(restarted.state, {}) self.assertEqual(restarted.state, {})
self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.deals.values())) self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.deals.values()))
@@ -98,7 +113,7 @@ class ArchivingTests(unittest.TestCase):
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)]) self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
self.insert_deal('buy', 100, 1000, '10:00:00', flag=23) self.insert_deal('buy', 100, 1000, '10:00:00', flag=23)
self.insert_deal('sell', 50, 600, '10:01:00', flag=24) self.insert_deal('sell', 50, 600, '10:01:00', flag=24)
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH'] row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty']), (100, 50)) self.assertEqual((row['base_qty'], row['added_qty']), (100, 50))
self.assertEqual(self.book.deals['buy']['offset_flag'], 23) self.assertEqual(self.book.deals['buy']['offset_flag'], 23)
@@ -113,8 +128,9 @@ class ArchivingTests(unittest.TestCase):
db.execute("""CREATE TRIGGER fail_archive BEFORE UPDATE OF is_arch ON deals db.execute("""CREATE TRIGGER fail_archive BEFORE UPDATE OF is_arch ON deals
WHEN OLD.order_sys_id = 'second' WHEN OLD.order_sys_id = 'second'
BEGIN SELECT RAISE(ABORT, 'test failure'); END""") BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
errors = self.book.archiving() with self.assertLogs(level='WARNING') as logs:
self.assertIn('600001.SH', errors) self.assertIsNone(self.book.archiving())
self.assertIn('600001.SH', '\n'.join(logs.output))
restarted = State(self.book.path) restarted = State(self.book.path)
self.assertEqual(set(restarted.state), {'600000.SH'}) self.assertEqual(set(restarted.state), {'600000.SH'})
self.assertEqual(restarted.state, self.book.state) self.assertEqual(restarted.state, self.book.state)
@@ -134,17 +150,17 @@ class ArchivingTests(unittest.TestCase):
offset_flag=48, volume=100, price=8, trade_amount=800, offset_flag=48, volume=100, price=8, trade_amount=800,
trade_date='20260908', trade_time='100000', trade_date='20260908', trade_time='100000',
)]) )])
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH'] row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100)) self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
self.assertEqual(self.book.deals['first']['is_arch'], 1) self.assertEqual(self.book.deals['first']['is_arch'], 1)
restarted = State(self.book.path) restarted = State(self.book.path)
self.assertEqual(restarted.archiving(), {}) self.assertIsNone(restarted.archiving())
self.assertEqual(restarted.state, self.book.state) self.assertEqual(restarted.state, self.book.state)
self.insert_deal('new_buy', 100, 1000, '10:01:00') self.insert_deal('new_buy', 100, 1000, '10:01:00')
with closing(self.book._connect()) as db, db: with closing(self.book._connect()) as db, db:
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'") db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH'] row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM')) self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM'))
@@ -153,8 +169,7 @@ class ArchivingTests(unittest.TestCase):
with closing(self.book._connect()) as db, db: with closing(self.book._connect()) as db, db:
db.execute('UPDATE deals SET is_arch = 1') db.execute('UPDATE deals SET is_arch = 1')
self.insert_deal('new', 50, 500, '10:01:00') self.insert_deal('new', 50, 500, '10:01:00')
errors = self.book.archiving() self.assertIsNone(self.book.archiving())
self.assertEqual(errors, {})
self.assertEqual(self.book.state['600000.SH']['added_qty'], 50) self.assertEqual(self.book.state['600000.SH']['added_qty'], 50)
self.assertEqual(self.book.deals['old']['is_arch'], 1) self.assertEqual(self.book.deals['old']['is_arch'], 1)
self.assertEqual(self.book.deals['new']['is_arch'], 1) self.assertEqual(self.book.deals['new']['is_arch'], 1)
@@ -166,7 +181,7 @@ class ArchivingTests(unittest.TestCase):
self.assertEqual(self.book.state, {}) self.assertEqual(self.book.state, {})
self.book = State(self.book.path) self.book = State(self.book.path)
self.insert_deal('late', 100, 2000, '100100') self.insert_deal('late', 100, 2000, '100100')
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH'] row = self.book.state['600000.SH']
self.assertEqual(row['added_qty'], 100) self.assertEqual(row['added_qty'], 100)
self.assertEqual(row['added_price'], 20) self.assertEqual(row['added_price'], 20)
@@ -174,21 +189,23 @@ class ArchivingTests(unittest.TestCase):
def test_archive_sell_before_syncing_empty_positions(self): def test_archive_sell_before_syncing_empty_positions(self):
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)]) self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
self.insert_deal('sell', 100, 1000, '10:01:00', flag=49) self.insert_deal('sell', 100, 1000, '10:01:00', flag=49)
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
self.book.sync_state([]) self.book.sync_state([])
self.book = State(self.book.path) self.book = State(self.book.path)
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
self.assertEqual(self.book.state, {}) self.assertEqual(self.book.state, {})
self.assertEqual(self.book.deals['sell']['is_arch'], 1) self.assertEqual(self.book.deals['sell']['is_arch'], 1)
def test_bad_stock_does_not_block_good_stock_and_can_retry(self): def test_bad_stock_does_not_block_good_stock_and_can_retry(self):
self.insert_deal('bad_sell', 100, 1500, '10:02:00', flag=49) self.insert_deal('bad_sell', 100, 1500, '10:02:00', flag=49)
self.insert_deal('good_buy', 100, 1000, '10:00:00', code='600001.SH') self.insert_deal('good_buy', 100, 1000, '10:00:00', code='600001.SH')
self.assertIn('600000.SH', self.book.archiving()) with self.assertLogs(level='WARNING') as logs:
self.assertIsNone(self.book.archiving())
self.assertIn('600000.SH', '\n'.join(logs.output))
self.assertEqual(self.book.deals['good_buy']['is_arch'], 1) self.assertEqual(self.book.deals['good_buy']['is_arch'], 1)
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 0) self.assertEqual(self.book.deals['bad_sell']['is_arch'], 0)
self.insert_deal('late_buy', 100, 1000, '10:01:00') self.insert_deal('late_buy', 100, 1000, '10:01:00')
self.assertEqual(self.book.archiving(), {}) self.assertIsNone(self.book.archiving())
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 1) self.assertEqual(self.book.deals['bad_sell']['is_arch'], 1)
self.assertNotIn('600000.SH', self.book.state) self.assertNotIn('600000.SH', self.book.state)

View File

@@ -1,5 +1,6 @@
import tempfile import tempfile
import unittest import unittest
from unittest.mock import patch
from pathlib import Path from pathlib import Path
from libs.state import State from libs.state import State
@@ -49,6 +50,32 @@ class ZTStateTests(unittest.TestCase):
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100) self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
self.assertEqual(self.state.deals['sell']['is_arch'], 0) self.assertEqual(self.state.deals['sell']['is_arch'], 0)
def test_failed_initialization_leaves_original_database_empty(self):
invalid = self.position(100)
invalid.open_price = float('inf')
with self.assertRaises(ValueError):
sync_account_state(self.state, [invalid], [self.deal('old', 100)], initialize=True)
restarted = State(self.state.path)
self.assertEqual((restarted.state, restarted.deals), ({}, {}))
sync_account_state(restarted, [self.position(100)], [self.deal('old', 100)], initialize=True)
self.assertEqual(restarted.deals['old']['is_arch'], 1)
def test_initialization_cannot_overwrite_existing_database(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [], initialize=True)
self.assertEqual(State(self.state.path).state['600000.SH']['base_qty'], 100)
def test_no_new_deals_still_retries_failed_archiving(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
sell = self.deal('sell', 100, flag=24)
with patch.object(self.state, 'archiving'):
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [sell])
sync_account_state(self.state, [], [sell])
self.assertEqual(self.state.state, {})
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()

View File

@@ -7,6 +7,7 @@ from unittest.mock import Mock, patch
from config import AccountConfig from config import AccountConfig
from libs.grid_take_profit import GridState from libs.grid_take_profit import GridState
from libs.order import OrderBook
from libs.state import State from libs.state import State
from sdk import Assets, DealItem, PositionItem, Tick from sdk import Assets, DealItem, PositionItem, Tick
from strategy.zt import boot from strategy.zt import boot
@@ -24,7 +25,7 @@ class ZTTradingTests(unittest.TestCase):
self.run = SimpleNamespace(account_cfg=self.cfg, orders=Mock(), client=Mock(), self.run = SimpleNamespace(account_cfg=self.cfg, orders=Mock(), client=Mock(),
profit_tracker=Mock(), add_watch=Mock(), open_watch=Mock()) profit_tracker=Mock(), add_watch=Mock(), open_watch=Mock())
self.run.orders.busy.return_value = False self.run.orders.busy.return_value = False
self.run.orders.new_order_id.side_effect = lambda kind: f'zt-{kind}-order' self.run.orders.new_order_id.side_effect = lambda prefix, kind: f'{prefix}-{kind}-order'
self.run.profit_tracker.observe.return_value.state = GridState.RETREAT self.run.profit_tracker.observe.return_value.state = GridState.RETREAT
self.run.add_watch.triggered.return_value = True self.run.add_watch.triggered.return_value = True
self.run.open_watch.triggered.return_value = True self.run.open_watch.triggered.return_value = True
@@ -131,6 +132,21 @@ class ZTTradingTests(unittest.TestCase):
[SimpleNamespace(code='688001.SH')], 2000) [SimpleNamespace(code='688001.SH')], 2000)
self.run.orders.place.assert_not_called() self.run.orders.place.assert_not_called()
def test_real_order_id_is_recognized_by_state_sync(self):
orders = OrderBook('zt')
self.run.orders.new_order_id.side_effect = orders.new_order_id
with patch('strategy.zt.open.datetime') as clock:
clock.now.return_value = datetime(2026, 9, 9, 10)
open_signal(self.run, {self.code: Tick(last_price=10)},
[SimpleNamespace(code=self.code)], 2000)
request = self.run.orders.place.call_args.args[1]
self.assertTrue(request.order_id.startswith('zt-base-'))
deal = self.fill('base', 'b1', 100)
deal.remark = request.order_id + '|zt'
self.store.sync_state([])
boot.sync_account_state(self.store, [], [deal])
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()