fix zt&state.py
This commit is contained in:
@@ -4,12 +4,16 @@ import math
|
||||
import logging as log
|
||||
import sqlite3
|
||||
from contextlib import closing
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import asdict, dataclass, fields
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
FLAG_BUY = 48
|
||||
FLAG_SELL = 49
|
||||
UNATTRIBUTED_PREFIX = '__unattributed__:'
|
||||
|
||||
SCHEMA = """
|
||||
-- 策略状态:base_ 表示底仓,added_ 表示补仓。
|
||||
CREATE TABLE IF NOT EXISTS state (
|
||||
@@ -69,14 +73,29 @@ class StateItem:
|
||||
added_created_at: str = '' # 补仓创建时间
|
||||
|
||||
|
||||
_STATE_COLUMNS = tuple(field.name for field in fields(StateItem))
|
||||
_UPSERT_STATE = (
|
||||
f"INSERT INTO state ({', '.join(_STATE_COLUMNS)}) "
|
||||
f"VALUES ({', '.join(':' + key for key in _STATE_COLUMNS)}) "
|
||||
"ON CONFLICT(stock_code) DO UPDATE SET "
|
||||
+ ', '.join(f'{key} = excluded.{key}' for key in _STATE_COLUMNS if key != 'stock_code')
|
||||
)
|
||||
|
||||
|
||||
class State:
|
||||
"""保存策略状态和只追加的成交记录,仅创建新表,不迁移旧数据。"""
|
||||
"""单写入者使用的 SQLite 存储;公开缓存仅在事务成功后替换。
|
||||
|
||||
sync_state 同步持仓基准,sync_deals 保存成交,archiving 记入增量。
|
||||
同证券未归档成交全部为买入且数量等于当前总持仓时,视为已计入
|
||||
快照,只标记归档;其他成交作为增量处理。本类不做表结构迁移。
|
||||
"""
|
||||
|
||||
def __init__(self, path: str | Path) -> None:
|
||||
self.path = Path(path)
|
||||
self.state: dict[str, dict] = {}
|
||||
self.deals: dict[str, dict] = {}
|
||||
self.deals_sys_ids: set[str] = set()
|
||||
self.blocked_codes: set[str] = set()
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with closing(self._connect()) as db:
|
||||
db.executescript(SCHEMA)
|
||||
@@ -87,176 +106,242 @@ class State:
|
||||
db.row_factory = sqlite3.Row
|
||||
return db
|
||||
|
||||
@staticmethod
|
||||
def _read_state(db: sqlite3.Connection) -> dict[str, dict]:
|
||||
return {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')}
|
||||
|
||||
@staticmethod
|
||||
def _read_deals(db: sqlite3.Connection) -> dict[str, dict]:
|
||||
return {
|
||||
row['order_sys_id']: dict(row)
|
||||
for row in db.execute('SELECT * FROM deals ORDER BY id')
|
||||
}
|
||||
|
||||
def get_by_code(self, code: str) -> dict:
|
||||
s = self.state.get(code,{})
|
||||
return s
|
||||
"""返回缓存中的状态;不存在时返回空字典。"""
|
||||
return self.state.get(code, {})
|
||||
|
||||
def load(self) -> None:
|
||||
"""缓存状态表和成交记录。"""
|
||||
self.load_state()
|
||||
self.load_deals()
|
||||
|
||||
def load_state(self) -> None:
|
||||
"""缓存状态表。"""
|
||||
"""在同一个读事务内加载两张表,全部成功后再发布缓存。"""
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN')
|
||||
state = {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')}
|
||||
state = self._read_state(db)
|
||||
deals = self._read_deals(db)
|
||||
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
|
||||
|
||||
def load_state(self) -> None:
|
||||
"""只刷新状态缓存。"""
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN')
|
||||
state = self._read_state(db)
|
||||
self.state = state
|
||||
|
||||
def load_deals(self) -> None:
|
||||
"""缓存成交记录。"""
|
||||
"""只刷新成交缓存及其去重编号集合。"""
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN')
|
||||
deals = {row['order_sys_id']: dict(row) for row in db.execute('SELECT * FROM deals ORDER BY id')}
|
||||
self.deals = deals
|
||||
self.deals_sys_ids = set(deals)
|
||||
deals = self._read_deals(db)
|
||||
self.deals, self.deals_sys_ids = deals, set(deals)
|
||||
|
||||
def sync_deals(self, deals: list[DealItem]) -> None:
|
||||
"""按成交编号去重;初始化底仓时,已包含在快照内的成交可直接标记归档。"""
|
||||
new_deals = {}
|
||||
@staticmethod
|
||||
def _insert_deals(db: sqlite3.Connection, deals: list[DealItem]) -> None:
|
||||
"""共享事务内保存成交;空本地编号使用明确的待核对标记。"""
|
||||
existing = {row[0] for row in db.execute('SELECT order_sys_id FROM deals')}
|
||||
new_deals: dict[str, DealItem] = {}
|
||||
for deal in deals:
|
||||
if deal.order_sys_id not in self.deals_sys_ids and deal.order_sys_id not in new_deals:
|
||||
new_deals[deal.order_sys_id] = deal
|
||||
if deal.order_sys_id not in existing:
|
||||
new_deals.setdefault(deal.order_sys_id, deal)
|
||||
if not new_deals:
|
||||
return
|
||||
with closing(self._connect()) as db, db:
|
||||
today = datetime.now().date().isoformat()
|
||||
values = []
|
||||
for deal in new_deals.values():
|
||||
order_id = deal.get_local_order_id
|
||||
if not order_id:
|
||||
raise ValueError('Local order ID is required')
|
||||
amount = deal.trade_amount if deal.trade_amount > 0 else deal.price * deal.volume
|
||||
if not math.isfinite(amount) or amount <= 0:
|
||||
raise ValueError('Trade amount must be positive and finite')
|
||||
date = deal.trade_date or datetime.now().date().isoformat()
|
||||
if not math.isfinite(deal.price) or not math.isfinite(amount):
|
||||
raise ValueError('Trade price and amount must be finite')
|
||||
date = deal.trade_date or today
|
||||
if len(date) == 8 and date.isdigit():
|
||||
date = f'{date[:4]}-{date[4:6]}-{date[6:]}'
|
||||
# 直接读取模型字段,金额和日期的补全不修改传入模型。
|
||||
db.execute(
|
||||
'INSERT INTO deals (stock_code, order_sys_id, order_local_id, ref, '
|
||||
'order_ref, direction, offset_flag, price, volume, trade_amount, '
|
||||
'trade_date, trade_time, remark, close_profit, is_arch) '
|
||||
'VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)',
|
||||
(deal.stock_code, deal.order_sys_id, order_id, deal.ref,
|
||||
deal.order_ref, deal.direction, deal.offset_flag, deal.price, deal.volume,
|
||||
amount, date, deal.trade_time, deal.remark, deal.close_profit, 0),
|
||||
)
|
||||
self.load_deals()
|
||||
values.append((
|
||||
deal.stock_code, deal.order_sys_id,
|
||||
deal.get_local_order_id.strip() or f'{UNATTRIBUTED_PREFIX}{deal.order_sys_id}',
|
||||
deal.ref, deal.order_ref, deal.direction, deal.offset_flag,
|
||||
deal.price, deal.volume, amount, date, deal.trade_time,
|
||||
deal.remark, deal.close_profit,
|
||||
))
|
||||
# 唯一键负责最终去重,避免旧缓存造成重复插入失败。
|
||||
db.executemany("""
|
||||
INSERT INTO deals (
|
||||
stock_code, order_sys_id, order_local_id, ref, order_ref,
|
||||
direction, offset_flag, price, volume, trade_amount,
|
||||
trade_date, trade_time, remark, close_profit, is_arch
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0)
|
||||
ON CONFLICT(order_sys_id) DO NOTHING
|
||||
""", values)
|
||||
|
||||
def sync_deals(self, deals: list[DealItem]) -> None:
|
||||
"""按成交编号只追加;同批保留首笔,任一无效成交回滚整批。"""
|
||||
if all(deal.order_sys_id in self.deals_sys_ids for deal in deals):
|
||||
return
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN IMMEDIATE')
|
||||
self._insert_deals(db, deals)
|
||||
cached_deals = self._read_deals(db)
|
||||
self.deals, self.deals_sys_ids = cached_deals, set(cached_deals)
|
||||
|
||||
def sync_state(self, positions: list[PositionItem]) -> None:
|
||||
"""同步完整持仓:无状态则插入底仓,已有则保留,清仓则删除。
|
||||
"""同步完整持仓:新增底仓、保留已有状态、删除已清仓证券。
|
||||
|
||||
此接口不会标记成交。已有库存对应的卖出须先归档,再传入
|
||||
清仓快照,以免删除归档所需的库存。
|
||||
"""
|
||||
# 传入完整账户持仓;同步时间作为新增底仓的创建时间。
|
||||
created_at = datetime.now().isoformat(timespec='seconds')
|
||||
holdings = {item.stock_code: item for item in positions if item.volume > 0}
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN')
|
||||
db.execute('BEGIN IMMEDIATE')
|
||||
existing = {row['stock_code'] for row in db.execute('SELECT stock_code FROM state')}
|
||||
db.executemany(
|
||||
'DELETE FROM state WHERE stock_code = ?',
|
||||
[(code,) for code in existing if code not in holdings],
|
||||
)
|
||||
|
||||
values = []
|
||||
for code, item in holdings.items():
|
||||
if code in existing:
|
||||
continue
|
||||
if not math.isfinite(item.open_price):
|
||||
raise ValueError('Base price must be finite')
|
||||
# 只插入底仓字段,补仓字段使用数据库默认值。
|
||||
db.execute(
|
||||
'INSERT INTO state '
|
||||
'(stock_code, status, base_order_local_id, base_qty, base_price, base_created_at) '
|
||||
"VALUES (?, '', '', ?, ?, ?)",
|
||||
(code, item.volume, item.open_price, created_at),
|
||||
values.append((code, item.volume, item.open_price, created_at))
|
||||
db.executemany(
|
||||
'DELETE FROM state WHERE stock_code = ?',
|
||||
[(code,) for code in existing if code not in holdings],
|
||||
)
|
||||
self.load_state()
|
||||
db.executemany(
|
||||
'INSERT INTO state (stock_code, base_qty, base_price, base_created_at) '
|
||||
'VALUES (?, ?, ?, ?)', values,
|
||||
)
|
||||
state = self._read_state(db)
|
||||
self.state = state
|
||||
|
||||
def archiving(self) -> None:
|
||||
"""将未归档成交累加到当前底仓和加仓;失败证券保留记录供重试。
|
||||
|
||||
归档是把成交反映到持仓状态中,再标记为已处理,不会删除成交记录。
|
||||
本地委托号以 zt-base- 开头的买入计入底仓,其他买入计入补仓。
|
||||
不返回结果;失败原因记录到日志,对应成交保留未归档标记供重试。
|
||||
"""
|
||||
with closing(self._connect()) as db, db:
|
||||
# 提前取得数据库写入锁,让持仓更新和成交标记在同一事务内完成。
|
||||
db.execute('BEGIN IMMEDIATE')
|
||||
# 只找尚未处理的成交:23、48 是买入,24、49 是卖出。
|
||||
codes = db.execute(
|
||||
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
|
||||
).fetchall()
|
||||
for entry in codes:
|
||||
code = entry['stock_code']
|
||||
# 每只股票设一个回滚点;这只处理失败时,不撤销其他股票的结果。
|
||||
db.execute('SAVEPOINT archive_stock')
|
||||
try:
|
||||
@staticmethod
|
||||
def _archive_stock(db: sqlite3.Connection, code: str, snapshot_dedup: bool = True) -> None:
|
||||
"""在调用方的保存点内完成一只证券的记账及成交标记。"""
|
||||
current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone()
|
||||
# 已有持仓就接着计算;没有记录则从底仓、补仓均为零开始。
|
||||
state = dict(current) if current else asdict(StateItem(stock_code=code))
|
||||
# 按成交日期、时间、记录编号依次处理,保证先买后卖等顺序正确。
|
||||
deals = db.execute(
|
||||
'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,)
|
||||
).fetchall()
|
||||
for deal in deals:
|
||||
qty = deal['volume']
|
||||
if deal['offset_flag'] in (23, 48):
|
||||
# 按已有的 ZT 底仓委托号约定识别,无需调用方传入规则。
|
||||
pending = db.execute("""
|
||||
SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0
|
||||
ORDER BY trade_date, REPLACE(trade_time, ':', ''), id
|
||||
""", (code,)).fetchall()
|
||||
if any(deal['order_local_id'].startswith(UNATTRIBUTED_PREFIX) for deal in pending):
|
||||
raise ValueError('Unattributed trade requires reconciliation')
|
||||
total_qty = state['base_qty'] + state['added_qty']
|
||||
if (
|
||||
snapshot_dedup and total_qty > 0
|
||||
and all(deal['offset_flag'] == FLAG_BUY for deal in pending)
|
||||
and sum(deal['volume'] for deal in pending) == total_qty
|
||||
):
|
||||
# 按快照去重规则,仅标记这批成交,原仓位数量和成本均保留。
|
||||
# 卖出和买卖混合批次不能用此规则,否则会漏掉真实减仓。
|
||||
db.executemany(
|
||||
'UPDATE deals SET is_arch = 1 WHERE id = ? AND is_arch = 0',
|
||||
[(deal['id'],) for deal in pending],
|
||||
)
|
||||
return
|
||||
for deal in pending:
|
||||
if deal['offset_flag'] == FLAG_BUY:
|
||||
bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added'
|
||||
total = state[f'{bucket}_qty'] + qty
|
||||
# 新均价 =(原数量 × 原均价 + 本次成交金额)÷ 买入后总数量。
|
||||
total = state[f'{bucket}_qty'] + deal['volume']
|
||||
state[f'{bucket}_price'] = (
|
||||
state[f'{bucket}_qty'] * state[f'{bucket}_price'] + deal['trade_amount']
|
||||
) / total
|
||||
state[f'{bucket}_qty'] = total
|
||||
# 保存这一类仓位最近一次买入的委托号和成交时间。
|
||||
state[f'{bucket}_order_local_id'] = deal['order_local_id']
|
||||
state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip()
|
||||
else:
|
||||
elif deal['offset_flag'] == FLAG_SELL:
|
||||
qty = deal['volume']
|
||||
total = state['base_qty'] + state['added_qty']
|
||||
# 卖出不能超过本地记录的总持仓;不一致时留待排查后重试。
|
||||
if qty > total:
|
||||
raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}')
|
||||
# 卖出优先扣补仓;补仓未卖完时只减数量,保留原均价。
|
||||
if qty < state['added_qty']:
|
||||
state['added_qty'] -= qty
|
||||
else:
|
||||
# 补仓全部卖完后,剩余卖出量扣底仓,并清空补仓信息。
|
||||
state['base_qty'] = total - qty
|
||||
state['added_qty'] = 0
|
||||
state['added_price'] = 0.0
|
||||
state['added_order_local_id'] = state['added_created_at'] = ''
|
||||
if state['base_qty'] == 0:
|
||||
state.update(asdict(StateItem(stock_code=code)))
|
||||
else:
|
||||
raise ValueError(f"Unsupported trade flag: {deal['offset_flag']}")
|
||||
if state['base_qty'] + state['added_qty'] == 0:
|
||||
# 中途清仓先重置,后续若又买入,就从零重新累计。
|
||||
state = asdict(StateItem(stock_code=code))
|
||||
if state['base_qty'] + state['added_qty'] == 0:
|
||||
# 全部成交处理完仍无持仓,则删除该股票的策略状态。
|
||||
db.execute('DELETE FROM state WHERE stock_code = ?', (code,))
|
||||
else:
|
||||
# 更新持仓,保留策略状态及已有记录主键。
|
||||
state['status'] = current['status'] if current else state['status']
|
||||
state.pop('id', None)
|
||||
columns = tuple(state)
|
||||
# 没有该股票就新增,已有则更新仓位字段,不替换原记录主键。
|
||||
db.execute(
|
||||
f"INSERT INTO state ({', '.join(columns)}) "
|
||||
f"VALUES ({', '.join(':' + key for key in columns)}) "
|
||||
'ON CONFLICT(stock_code) DO UPDATE SET '
|
||||
+ ', '.join(f'{key} = excluded.{key}' for key in columns if key != 'stock_code'),
|
||||
state,
|
||||
)
|
||||
# 持仓处理成功后才标记成交,防止下次重复加仓或重复扣减。
|
||||
# 保留原记录主键及策略状态,包括同批清仓后重新建仓。
|
||||
if current:
|
||||
state['status'] = current['status']
|
||||
db.execute(_UPSERT_STATE, state)
|
||||
db.execute(
|
||||
'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', (code,),
|
||||
)
|
||||
|
||||
def _archive_pending(self, db: sqlite3.Connection, snapshot_dedup: bool = True) -> None:
|
||||
"""保留现有逐证券保存点;ZT 增量路径禁用数量相等推断。"""
|
||||
codes = db.execute(
|
||||
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
|
||||
).fetchall()
|
||||
for row in codes:
|
||||
code = row['stock_code']
|
||||
db.execute('SAVEPOINT archive_stock')
|
||||
try:
|
||||
self._archive_stock(db, code, snapshot_dedup)
|
||||
except (ValueError, sqlite3.IntegrityError) as exc:
|
||||
# 撤销这只股票的全部归档修改,成交仍保持未归档,供下次重试。
|
||||
db.execute('ROLLBACK TO archive_stock')
|
||||
log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc)
|
||||
finally:
|
||||
# 释放当前股票的回滚点;整个事务在退出外层 with 时提交。
|
||||
db.execute('RELEASE archive_stock')
|
||||
# 数据库提交完成后刷新内存缓存,让策略读到最新持仓和归档标记。
|
||||
self.load()
|
||||
|
||||
def archiving(self) -> None:
|
||||
"""归档未处理成交;单证券失败回滚并保留重试,不影响其他证券。"""
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN IMMEDIATE')
|
||||
self._archive_pending(db)
|
||||
state = self._read_state(db)
|
||||
deals = self._read_deals(db)
|
||||
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
|
||||
|
||||
def sync_account(self, positions: list[PositionItem], deals: list[DealItem], *, initialize: bool = False) -> None:
|
||||
"""ZT 原子同步。初始化使用已核对快照;恢复只应用增量、不覆盖库存。
|
||||
|
||||
空库存且无成交历史时允许初始化。快照缺失或
|
||||
未归档成交只隔离对应证券,后续成交到齐并核对一致后自动恢复。
|
||||
"""
|
||||
holdings = {p.stock_code: p for p in positions if p.volume > 0}
|
||||
with closing(self._connect()) as db, db:
|
||||
db.execute('BEGIN IMMEDIATE')
|
||||
if initialize:
|
||||
if (db.execute('SELECT 1 FROM state LIMIT 1').fetchone()
|
||||
or db.execute('SELECT 1 FROM deals LIMIT 1').fetchone()):
|
||||
raise ValueError('ZT state already initialized; cannot overwrite existing data')
|
||||
created_at = datetime.now().isoformat(timespec='seconds')
|
||||
for code, position in holdings.items():
|
||||
if not code or not math.isfinite(position.open_price) or position.open_price <= 0:
|
||||
raise ValueError('Initial position code and cost must be valid')
|
||||
db.execute(_UPSERT_STATE, asdict(StateItem(
|
||||
stock_code=code, base_qty=position.volume,
|
||||
base_price=position.open_price, base_created_at=created_at,
|
||||
)))
|
||||
self._insert_deals(db, deals)
|
||||
# 这些成交已包含在已核对的基准中,不再增减快照库存。
|
||||
db.execute('UPDATE deals SET is_arch = 1')
|
||||
else:
|
||||
self._insert_deals(db, deals)
|
||||
self._archive_pending(db, snapshot_dedup=False)
|
||||
state = self._read_state(db)
|
||||
cached_deals = self._read_deals(db)
|
||||
blocked = {d['stock_code'] for d in cached_deals.values() if d['is_arch'] != 1}
|
||||
for code in set(state) | set(holdings):
|
||||
row = state.get(code, {})
|
||||
recorded = row.get('base_qty', 0) + row.get('added_qty', 0)
|
||||
actual = holdings[code].volume if code in holdings else 0
|
||||
if recorded != actual:
|
||||
blocked.add(code)
|
||||
self.state, self.deals, self.deals_sys_ids = state, cached_deals, set(cached_deals)
|
||||
self.blocked_codes = blocked
|
||||
if blocked:
|
||||
log.warning('[ZT 同步] 以下证券状态待核对,暂停交易:%s', ', '.join(sorted(blocked)))
|
||||
|
||||
@@ -2,14 +2,11 @@
|
||||
|
||||
import logging as log
|
||||
import time
|
||||
from contextlib import closing
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
import config
|
||||
from libs.calc import trading_time
|
||||
from libs.collector import collector_push
|
||||
from libs.grid_take_profit import GridTrailingTracker
|
||||
from libs.market import market_allow_open
|
||||
from libs.order import OrderBook
|
||||
@@ -18,9 +15,9 @@ from libs.runtime import Runtime
|
||||
from libs.signal import SignalItem, init_signals
|
||||
from libs.state import State
|
||||
from libs.watch import DipWatch
|
||||
from sdk import Client, DealItem, PositionItem
|
||||
from sdk import Client
|
||||
from .open import open_signal
|
||||
from .positions import manage_positions, t_rounds
|
||||
from .positions import manage_positions
|
||||
from libs.snapshot import cache_portfolio
|
||||
|
||||
|
||||
@@ -32,15 +29,13 @@ def StartZT() -> None:
|
||||
)
|
||||
executor = None
|
||||
try:
|
||||
portfolio = client.portfolio()
|
||||
assets = portfolio.assets
|
||||
positions = list(portfolio.positions.values())
|
||||
|
||||
state = State(Path(config.global_config.qmt_data_dir) / f'zt_{config.account_config.account_id}_state.db')
|
||||
executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="zt")
|
||||
run = Runtime(
|
||||
client=client, global_cfg=config.global_config, account_cfg=config.account_config,
|
||||
orders=OrderBook('zt'), open_watch=DipWatch(), add_watch=DipWatch(),
|
||||
profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct),
|
||||
executor=executor
|
||||
)
|
||||
|
||||
# 获取本策略的信号开仓数据
|
||||
@@ -48,22 +43,21 @@ def StartZT() -> None:
|
||||
config.global_config,
|
||||
config.account_config.signal_allow,
|
||||
)
|
||||
log.info(
|
||||
"[启动] Trend策略已启动,账户=%s,信号=%d,持仓=%d",
|
||||
config.account_config.account_id,
|
||||
len(signals),
|
||||
len(positions),
|
||||
)
|
||||
|
||||
initialize = not state.state and not state.deals
|
||||
deals = client.deals()
|
||||
portfolio = client.portfolio()
|
||||
if initialize and {d.order_sys_id: d for d in deals} != {
|
||||
d.order_sys_id: d for d in client.deals()
|
||||
}:
|
||||
raise RuntimeError('ZT 初始化期间成交发生变化,请重新启动')
|
||||
assets = portfolio.assets
|
||||
positions = list(portfolio.positions.values())
|
||||
log.info('[启动] ZT策略已启动,账户=%s,信号=%d,持仓=%d',
|
||||
config.account_config.account_id, len(signals), len(positions))
|
||||
cache_portfolio(config.account_config.account_id, assets, positions, deals)
|
||||
state.load()
|
||||
state.sync_deals(deals)
|
||||
state.sync_state(positions)
|
||||
state.archiving()
|
||||
state.sync_account(positions, deals, initialize=initialize)
|
||||
run.orders.refresh(client, portfolio.orders)
|
||||
Overview(assets, positions, config.account_config)
|
||||
executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="zt")
|
||||
|
||||
DEFAULT_TICK_INTERVAL = 30
|
||||
while True:
|
||||
@@ -118,8 +112,7 @@ def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
|
||||
positions = list(portfolio.positions.values())
|
||||
position_codes = list(portfolio.positions)
|
||||
|
||||
state.sync_deals(deals)
|
||||
state.archiving()
|
||||
state.sync_account(positions, deals)
|
||||
run.orders.refresh(run.client, portfolio.orders)
|
||||
except Exception:
|
||||
log.exception("[Portfolio] 刷新账户快照失败")
|
||||
@@ -143,7 +136,7 @@ def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
|
||||
allow_open: list[SignalItem] = []
|
||||
allow_codes: list[str] = []
|
||||
for signal in signals:
|
||||
if signal.code not in portfolio.positions:
|
||||
if signal.code not in portfolio.positions and signal.code not in state.blocked_codes:
|
||||
allow_open.append(signal)
|
||||
allow_codes.append(signal.code)
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""趋势策略持仓止盈与分级补仓。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
|
||||
from libs.calc import calc_buy_volume, calculate_min_profit_rate
|
||||
from libs.grid_take_profit import GridState
|
||||
@@ -45,9 +46,9 @@ def manage_positions(
|
||||
continue
|
||||
if (
|
||||
not code
|
||||
or position.open_price <= 0
|
||||
or position.volume <= 0
|
||||
or tick is None
|
||||
or not math.isfinite(tick.last_price)
|
||||
or tick.last_price <= 0
|
||||
):
|
||||
log.warning(
|
||||
@@ -57,16 +58,24 @@ def manage_positions(
|
||||
)
|
||||
continue
|
||||
|
||||
if code in state.blocked_codes:
|
||||
log.warning('[Position] %s 状态待核对,暂停该证券交易', code)
|
||||
continue
|
||||
|
||||
posState = state.get_by_code(position.stock_code)
|
||||
if not posState:
|
||||
continue
|
||||
|
||||
volume = position.can_use_volume
|
||||
target_qty = posState.get('base_qty', 0)
|
||||
cost_price = position.open_price
|
||||
if posState.get('added_qty',0) >=100:
|
||||
volume = posState.get('added_qty',0)
|
||||
if posState.get('added_qty', 0) > 0:
|
||||
target_qty = posState['added_qty']
|
||||
cost_price = posState.get('added_price',0)
|
||||
|
||||
# 在途股份不影响已有可卖库存;补仓、底仓均受柜台可卖上限约束。
|
||||
volume = max(0, min(target_qty, position.can_use_volume, position.volume))
|
||||
if not math.isfinite(cost_price) or cost_price <= 0:
|
||||
log.warning('[Position] %s 成本无效,暂停该证券交易', code)
|
||||
continue
|
||||
|
||||
pnl_rate = round(
|
||||
(tick.last_price - cost_price) / cost_price * 100,
|
||||
@@ -86,7 +95,8 @@ def manage_positions(
|
||||
if runtime.account_cfg.enable_loss_add_position and market_ok:
|
||||
loss_decision = handle_loss(
|
||||
runtime=runtime,
|
||||
position=position,
|
||||
stock_code=position.stock_code,
|
||||
volume=volume,
|
||||
tick=tick,
|
||||
pnl_rate=pnl_rate,
|
||||
available=available,
|
||||
@@ -99,7 +109,7 @@ def manage_positions(
|
||||
strTag = "-"
|
||||
if pnl_rate >= minimum_profit:
|
||||
strTag = "↑"
|
||||
elif pnl_rate< LOSS_TIERS[0]:
|
||||
elif pnl_rate< LOSS_TIERS:
|
||||
strTag = "↓"
|
||||
|
||||
if strTag != "-":
|
||||
@@ -148,9 +158,8 @@ def handle_profit(
|
||||
if runtime.orders.busy(stock_code, "SELL"):
|
||||
return TradeDecision(False, "卖出委托处理中")
|
||||
|
||||
volume = volume % 100
|
||||
if volume <= 0:
|
||||
return TradeDecision(False, "无可用整手持仓")
|
||||
return TradeDecision(False, "无可用持仓")
|
||||
order_id = runtime.orders.new_order_id("zt","SELL")
|
||||
request = PlaceOrderRequest(
|
||||
op=OP_SELL,
|
||||
@@ -199,7 +208,7 @@ def handle_loss(
|
||||
if not runtime.orders.place(runtime.client, request):
|
||||
return TradeDecision(False, "补仓订单委托失败")
|
||||
|
||||
runtime.add_watch.forget(position.stock_code)
|
||||
runtime.add_watch.forget(stock_code)
|
||||
return TradeDecision(True, f"[补仓买入] {volume} 股,订单={order_id}", amount)
|
||||
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import ast
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from dataclasses import asdict, fields
|
||||
@@ -88,7 +89,7 @@ class ApiModelTests(unittest.TestCase):
|
||||
self.assertEqual(deal.trade_amount, 0)
|
||||
deal.order_sys_id = 'sys3'
|
||||
deal.price = 0
|
||||
with self.assertRaises(ValueError):
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
book.sync_deals([deal])
|
||||
self.assertEqual(set(State(path).deals), {'sys1', 'sys2'})
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import unittest
|
||||
from contextlib import closing
|
||||
from pathlib import Path
|
||||
|
||||
from libs.state import State
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(tables, {'state', 'deals', 'sqlite_sequence'})
|
||||
|
||||
def test_accumulates_once_and_preserves_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=200, open_price=8)])
|
||||
original = dict(self.book.state['600000.SH'])
|
||||
self.insert_deal('first', 40, 400, '10:00:00')
|
||||
self.insert_deal('second', 60, 720, '10:01:00')
|
||||
@@ -53,9 +53,9 @@ class ArchivingTests(unittest.TestCase):
|
||||
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-first', 100, 1000, '10:00:00', flag=FLAG_BUY)
|
||||
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.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=FLAG_BUY)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['base_price']), (200, 11))
|
||||
@@ -111,13 +111,13 @@ class ArchivingTests(unittest.TestCase):
|
||||
|
||||
def test_stock_buy_and_sell_flags(self):
|
||||
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('sell', 50, 600, '10:01:00', flag=24)
|
||||
self.insert_deal('buy', 100, 1000, '10:00:00', flag=FLAG_BUY)
|
||||
self.insert_deal('sell', 50, 600, '10:01:00', flag=FLAG_SELL)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 50))
|
||||
self.assertEqual(self.book.deals['buy']['offset_flag'], 23)
|
||||
self.assertEqual(self.book.deals['sell']['offset_flag'], 24)
|
||||
self.assertEqual(self.book.deals['buy']['offset_flag'], FLAG_BUY)
|
||||
self.assertEqual(self.book.deals['sell']['offset_flag'], FLAG_SELL)
|
||||
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
||||
|
||||
def test_new_state_and_failed_mark_roll_back_together(self):
|
||||
@@ -143,7 +143,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(self.book.state['600000.SH']['base_qty'], 0)
|
||||
self.assertEqual(self.book.state['600000.SH']['added_qty'], 100)
|
||||
|
||||
def test_equal_quantity_buy_is_added_and_preserves_status(self):
|
||||
def test_equal_quantity_buy_only_marks_and_preserves_status(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.book.sync_deals([DealItem(
|
||||
stock_code='600000.SH', order_sys_id='first', remark='base1|test',
|
||||
@@ -152,17 +152,17 @@ class ArchivingTests(unittest.TestCase):
|
||||
)])
|
||||
self.assertIsNone(self.book.archiving())
|
||||
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, 0))
|
||||
self.assertEqual(self.book.deals['first']['is_arch'], 1)
|
||||
restarted = State(self.book.path)
|
||||
self.assertIsNone(restarted.archiving())
|
||||
self.assertEqual(restarted.state, self.book.state)
|
||||
self.insert_deal('new_buy', 100, 1000, '10:01:00')
|
||||
self.insert_deal('new_buy', 50, 500, '10:01:00')
|
||||
with closing(self.book._connect()) as db, db:
|
||||
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
|
||||
self.assertIsNone(self.book.archiving())
|
||||
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, 50, 'CUSTOM'))
|
||||
|
||||
def test_archived_history_is_not_reapplied(self):
|
||||
self.insert_deal('old', 100, 1000, '10:00:00')
|
||||
|
||||
92
py-client/tests/test_state_snapshot.py
Normal file
92
py-client/tests/test_state_snapshot.py
Normal file
@@ -0,0 +1,92 @@
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from contextlib import closing
|
||||
from pathlib import Path
|
||||
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
|
||||
class SnapshotArchiveTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.store = State(Path(tmp.name) / 'state.db')
|
||||
self.code = '600000.SH'
|
||||
|
||||
def deal(self, identity, qty, flag=FLAG_BUY, code=None):
|
||||
return DealItem(
|
||||
stock_code=code or self.code, order_sys_id=identity,
|
||||
remark=f'zt-base-{identity}|zt', offset_flag=flag,
|
||||
volume=qty, price=12, trade_amount=qty * 12,
|
||||
trade_date='2026-09-12', trade_time='100000',
|
||||
)
|
||||
|
||||
def snapshot(self, qty):
|
||||
self.store.sync_state([PositionItem(stock_code=self.code, volume=qty, open_price=10)])
|
||||
|
||||
def test_matching_partial_fills_only_mark_and_survive_restart(self):
|
||||
self.snapshot(100)
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
restarted = State(self.store.path)
|
||||
restarted.archiving()
|
||||
self.assertEqual(restarted.state[self.code], saved)
|
||||
|
||||
def test_total_includes_added_holdings(self):
|
||||
self.snapshot(100)
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute("UPDATE state SET added_qty=50, added_price=9, status='CUSTOM'")
|
||||
self.store.load()
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 150)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertEqual(self.store.deals['one']['is_arch'], 1)
|
||||
|
||||
def test_nonmatching_and_other_stock_are_incremental(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('other', 60, code='600001.SH')])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 140)
|
||||
self.assertEqual(self.store.state['600001.SH']['base_qty'], 60)
|
||||
|
||||
def test_matching_sell_still_liquidates(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('sell', 100, FLAG_SELL)])
|
||||
self.store.archiving()
|
||||
self.assertNotIn(self.code, self.store.state)
|
||||
self.assertEqual(self.store.deals['sell']['is_arch'], 1)
|
||||
|
||||
def test_mixed_batch_with_matching_gross_volume_is_not_skipped(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('buy', 40), self.deal('sell', 60, FLAG_SELL)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 80)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
|
||||
def test_failed_mark_rolls_back_entire_stock_and_retries(self):
|
||||
self.snapshot(100)
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute("""CREATE TRIGGER fail_mark BEFORE UPDATE OF is_arch ON deals
|
||||
WHEN OLD.order_sys_id='two'
|
||||
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
||||
with self.assertLogs(level='WARNING'):
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 0 for d in self.store.deals.values()))
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute('DROP TRIGGER fail_mark')
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
86
py-client/tests/test_state_storage.py
Normal file
86
py-client/tests/test_state_storage.py
Normal file
@@ -0,0 +1,86 @@
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from libs.state import FLAG_BUY, State
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
|
||||
class StateStorageTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.store = State(Path(tmp.name) / 'state.db')
|
||||
|
||||
def deal(self, identity='first'):
|
||||
return DealItem(
|
||||
stock_code='600000.SH', order_sys_id=identity,
|
||||
remark=f'zt-base-{identity}|zt', offset_flag=FLAG_BUY,
|
||||
volume=100, price=10, trade_amount=1000,
|
||||
trade_date='20260912', trade_time='100000',
|
||||
)
|
||||
|
||||
def test_stale_cache_duplicate_preserves_original_and_imports_new_trade(self):
|
||||
writer = State(self.store.path)
|
||||
first = self.deal()
|
||||
writer.sync_deals([first])
|
||||
first.price = 20
|
||||
first.trade_amount = 2000
|
||||
self.store.sync_deals([first, self.deal('second')])
|
||||
self.assertEqual(self.store.deals['first']['price'], 10)
|
||||
self.assertEqual(self.store.deals_sys_ids, {'first', 'second'})
|
||||
self.assertEqual(State(self.store.path).deals, self.store.deals)
|
||||
|
||||
def test_load_failure_does_not_publish_partial_cache(self):
|
||||
writer = State(self.store.path)
|
||||
writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
|
||||
with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
|
||||
with self.assertRaises(sqlite3.OperationalError):
|
||||
self.store.load()
|
||||
self.assertEqual((self.store.state, self.store.deals, self.store.deals_sys_ids), ({}, {}, set()))
|
||||
self.store.load()
|
||||
self.assertEqual(self.store.state['600000.SH']['base_qty'], 100)
|
||||
|
||||
def test_cache_read_failure_rolls_back_archive_and_can_retry(self):
|
||||
self.store.sync_deals([self.deal()])
|
||||
with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
|
||||
with self.assertRaises(sqlite3.OperationalError):
|
||||
self.store.archiving()
|
||||
restarted = State(self.store.path)
|
||||
self.assertEqual(restarted.state, {})
|
||||
self.assertEqual(restarted.deals['first']['is_arch'], 0)
|
||||
self.assertEqual(self.store.deals, restarted.deals)
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(self.store.deals['first']['is_arch'], 1)
|
||||
|
||||
def test_invalid_snapshot_preserves_existing_holdings(self):
|
||||
self.store.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
|
||||
saved = self.store.state
|
||||
with self.assertRaises(ValueError):
|
||||
self.store.sync_state([PositionItem(stock_code='600001.SH', volume=100, open_price=float('inf'))])
|
||||
self.assertEqual(self.store.state, saved)
|
||||
self.assertEqual(State(self.store.path).state, saved)
|
||||
|
||||
def test_normalization_preserves_input_and_rejects_nonfinite_price(self):
|
||||
deal = self.deal()
|
||||
deal.trade_amount = 0
|
||||
original = asdict(deal)
|
||||
self.store.sync_deals([deal])
|
||||
self.assertEqual(asdict(deal), original)
|
||||
self.assertEqual(self.store.deals['first']['trade_amount'], 1000)
|
||||
self.assertEqual(self.store.deals['first']['trade_date'], '2026-09-12')
|
||||
for price in (float('inf'), float('-inf'), float('nan')):
|
||||
with self.subTest(price=price):
|
||||
invalid = self.deal('invalid')
|
||||
invalid.price = price
|
||||
with self.assertRaises(ValueError):
|
||||
self.store.sync_deals([self.deal('second'), invalid])
|
||||
self.assertEqual(State(self.store.path).deals_sys_ids, {'first'})
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
@@ -1,81 +1,131 @@
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from libs.state import State
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State, UNATTRIBUTED_PREFIX
|
||||
from sdk import DealItem, PositionItem
|
||||
from strategy.zt.boot import sync_account_state
|
||||
|
||||
|
||||
class ZTStateTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.state = State(Path(tmp.name) / 'zt_test_state.db')
|
||||
self.state = State(Path(tmp.name) / 'state.db')
|
||||
self.code = '600000.SH'
|
||||
|
||||
def position(self, qty):
|
||||
return PositionItem(stock_code='600000.SH', volume=qty, open_price=10)
|
||||
def position(self, qty, code=None):
|
||||
return PositionItem(stock_code=code or self.code, volume=qty, open_price=10)
|
||||
|
||||
def deal(self, order, qty, flag=23, strategy='zt'):
|
||||
return DealItem(
|
||||
stock_code='600000.SH', order_sys_id=order,
|
||||
remark=f'{strategy}-buy-{order}|{strategy}', offset_flag=flag,
|
||||
volume=qty, price=10, trade_amount=qty * 10,
|
||||
trade_date='20260909', trade_time='100000',
|
||||
)
|
||||
def deal(self, identity, qty=100, flag=FLAG_BUY, code=None, remark=None):
|
||||
return DealItem(stock_code=code or self.code, order_sys_id=identity,
|
||||
remark=f'zt-base-{identity}|zt' if remark is None else remark,
|
||||
offset_flag=flag, volume=qty, price=10, trade_amount=qty*10,
|
||||
trade_date='20260912', trade_time='100000')
|
||||
|
||||
def test_initial_snapshot_and_incremental_deals_after_restart(self):
|
||||
historical = self.deal('old', 100)
|
||||
unrelated = self.deal('trend', 100, strategy='trend')
|
||||
sync_account_state(self.state, [self.position(100)], [historical, unrelated], initialize=True)
|
||||
self.assertEqual(set(self.state.deals), {'old'})
|
||||
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(self.state.state['600000.SH']['added_qty'], 0)
|
||||
def test_initial_snapshot_and_equal_size_increment_are_distinct(self):
|
||||
old = self.deal('old')
|
||||
self.state.sync_account([self.position(100)], [old], initialize=True)
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
self.assertEqual(self.state.deals['old']['is_arch'], 1)
|
||||
self.state = State(self.state.path)
|
||||
bought = self.deal('new', 100, flag=48)
|
||||
new = self.deal('new', remark='zt-added-new|zt')
|
||||
for _ in range(2):
|
||||
sync_account_state(self.state, [self.position(200)], [historical, bought, unrelated])
|
||||
row = self.state.state['600000.SH']
|
||||
self.state.sync_account([self.position(200)], [old, new, new])
|
||||
row = self.state.state[self.code]
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
|
||||
sold = self.deal('sell', 200, flag=24)
|
||||
sync_account_state(self.state, [], [historical, bought, sold])
|
||||
self.assertEqual(self.state.blocked_codes, set())
|
||||
|
||||
def test_initial_mixed_trades_are_already_in_snapshot(self):
|
||||
self.state.sync_account([self.position(150)],
|
||||
[self.deal('buy', 200), self.deal('sell', 50, FLAG_SELL)],
|
||||
initialize=True)
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 150)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.state.deals.values()))
|
||||
|
||||
def test_restart_full_sell_is_archived_before_reconciliation(self):
|
||||
self.state.sync_account([self.position(100)], [], initialize=True)
|
||||
self.state = State(self.state.path)
|
||||
self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)])
|
||||
self.assertEqual(self.state.state, {})
|
||||
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
|
||||
self.assertEqual(self.state.blocked_codes, set())
|
||||
|
||||
def test_archive_failure_preserves_holdings_for_retry(self):
|
||||
sync_account_state(self.state, [self.position(100)], [], initialize=True)
|
||||
with self.assertRaisesRegex(ValueError, 'ZT'):
|
||||
sync_account_state(self.state, [], [self.deal('sell', 200, flag=49)])
|
||||
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(self.state.deals['sell']['is_arch'], 0)
|
||||
def test_empty_initialized_account_survives_restart(self):
|
||||
self.state.sync_account([], [], initialize=True)
|
||||
self.state = State(self.state.path)
|
||||
# 空账户重复初始化仍为空,无需额外标记表。
|
||||
self.state.sync_account([], [], initialize=True)
|
||||
self.state.sync_account([self.position(100)], [self.deal('new')])
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
|
||||
def test_failed_initialization_leaves_original_database_empty(self):
|
||||
def test_initialization_cannot_overwrite_existing_holdings(self):
|
||||
self.state.sync_account([self.position(100)], [], initialize=True)
|
||||
with self.assertRaises(ValueError):
|
||||
self.state.sync_account([], [], initialize=True)
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
|
||||
def test_initialization_failure_is_atomic(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)
|
||||
self.state.sync_account([invalid], [self.deal('one')], 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)
|
||||
with patch.object(self.state, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
|
||||
with self.assertRaises(sqlite3.OperationalError):
|
||||
self.state.sync_account([self.position(100)], [self.deal('one')], initialize=True)
|
||||
self.assertEqual(State(self.state.path).deals, {})
|
||||
self.assertEqual(self.state.state, {})
|
||||
|
||||
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_lagging_snapshot_never_deletes_or_recreates_inventory(self):
|
||||
self.state.sync_account([self.position(100)], [], initialize=True)
|
||||
self.state.sync_account([], [])
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
self.assertEqual(self.state.blocked_codes, {self.code})
|
||||
self.state.sync_account([self.position(100)], [])
|
||||
self.assertEqual(self.state.blocked_codes, set())
|
||||
sell = self.deal('sell', flag=FLAG_SELL)
|
||||
self.state.sync_account([self.position(100)], [sell])
|
||||
self.assertNotIn(self.code, self.state.state)
|
||||
self.assertEqual(self.state.blocked_codes, {self.code})
|
||||
self.state.sync_account([], [sell])
|
||||
self.assertEqual(self.state.blocked_codes, set())
|
||||
|
||||
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])
|
||||
def test_blank_remark_persists_and_isolates_only_affected_stock(self):
|
||||
self.state.sync_account([self.position(100)], [], initialize=True)
|
||||
manual = self.deal('manual', remark=' |')
|
||||
good = self.deal('good', code='600001.SH')
|
||||
positions = [self.position(200), self.position(100, '600001.SH')]
|
||||
for _ in range(2):
|
||||
self.state.sync_account(positions, [manual, good])
|
||||
self.assertEqual(self.state.blocked_codes, {self.code})
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
self.assertEqual(self.state.state['600001.SH']['base_qty'], 100)
|
||||
self.assertEqual(self.state.deals['manual']['is_arch'], 0)
|
||||
self.assertTrue(self.state.deals['manual']['order_local_id'].startswith(UNATTRIBUTED_PREFIX))
|
||||
self.assertEqual(self.state.deals['manual']['remark'], ' |')
|
||||
self.state = State(self.state.path)
|
||||
|
||||
def test_legacy_database_is_restored_without_reinitialization(self):
|
||||
self.state.sync_state([self.position(100)])
|
||||
self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)])
|
||||
self.assertEqual(self.state.state, {})
|
||||
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
|
||||
|
||||
def test_failed_archive_blocks_only_stock_and_retries(self):
|
||||
self.state.sync_account([self.position(100)], [], initialize=True)
|
||||
sell = self.deal('sell', qty=200, flag=FLAG_SELL)
|
||||
self.state.sync_account([], [sell])
|
||||
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
|
||||
self.assertEqual(self.state.blocked_codes, {self.code})
|
||||
buy = self.deal('buy', remark='zt-added-buy|zt')
|
||||
buy.trade_time = '095900'
|
||||
self.state.sync_account([], [sell, buy])
|
||||
self.assertEqual(self.state.state, {})
|
||||
self.assertEqual(self.state.blocked_codes, set())
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
@@ -1,151 +1,149 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import closing
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from types import SimpleNamespace as NS
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from config import AccountConfig
|
||||
from libs.grid_take_profit import GridState
|
||||
from libs.order import OrderBook
|
||||
from libs.state import State
|
||||
from libs.state import FLAG_BUY, State
|
||||
from sdk import Assets, DealItem, PositionItem, Tick
|
||||
from strategy.zt import boot
|
||||
from strategy.zt.open import open_signal
|
||||
from strategy.zt.positions import manage_positions, t_rounds
|
||||
from strategy.zt.positions import manage_positions
|
||||
|
||||
|
||||
class ZTTradingTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.store = State(Path(tmp.name) / 'state.db')
|
||||
self.code = '600000.SH'
|
||||
self.cfg = AccountConfig(account_id='test', buy_value=2000, zt_sell_ratio=0.5)
|
||||
self.run = SimpleNamespace(account_cfg=self.cfg, orders=Mock(), client=Mock(),
|
||||
profit_tracker=Mock(), add_watch=Mock(), open_watch=Mock())
|
||||
self.run = NS(account_cfg=NS(account_id='test', strategy='zt', buy_value=1000,
|
||||
excluded_codes=[], enable_loss_add_position=False,
|
||||
min_cash_ratio=0.1),
|
||||
orders=Mock(), client=Mock(), profit_tracker=Mock(), add_watch=Mock())
|
||||
self.run.orders.busy.return_value = False
|
||||
self.run.orders.new_order_id.side_effect = lambda prefix, kind: f'{prefix}-{kind}-order'
|
||||
self.run.orders.place.return_value = True
|
||||
self.run.profit_tracker.observe.return_value.state = GridState.RETREAT
|
||||
self.run.add_watch.triggered.return_value = True
|
||||
self.run.open_watch.triggered.return_value = True
|
||||
self.position = PositionItem(stock_code=self.code, volume=200, can_use_volume=200, open_price=10)
|
||||
boot.sync_account_state(self.store, [self.position], [], initialize=True)
|
||||
|
||||
def fill(self, kind, order, qty, price=10, date='2026-09-09'):
|
||||
return DealItem(stock_code=self.code, order_sys_id=order, remark=f'zt-{kind}-{order}|zt',
|
||||
offset_flag=24 if kind == 't-sell' else 23,
|
||||
volume=qty, price=price, trade_amount=qty * price,
|
||||
trade_date=date, trade_time='100000')
|
||||
def manage(self, added=0, usable=500, road=0, cost=10, added_cost=10, price=11):
|
||||
position = PositionItem(stock_code=self.code, volume=1000, can_use_volume=usable,
|
||||
on_road_volume=road, open_price=cost)
|
||||
state = NS(blocked_codes=set(), get_by_code=lambda code: dict(
|
||||
base_qty=500, added_qty=added, added_price=added_cost))
|
||||
manage_positions(self.run, {self.code: Tick(last_price=price)}, [position], True, 1500, state)
|
||||
|
||||
def manage(self, price=11, available=10000, positions=None, force=False, today='2026-09-09'):
|
||||
return manage_positions(self.run, self.store, {self.code: Tick(last_price=price)},
|
||||
[self.position] if positions is None else positions,
|
||||
t_rounds(self.store), available, today, force)
|
||||
|
||||
def test_sell_only_available_shares_and_no_loss_sell(self):
|
||||
self.position.can_use_volume = 0
|
||||
self.manage()
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.position.can_use_volume = 100
|
||||
self.manage(price=9)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage()
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertEqual((request.op, request.volume), (24, 100))
|
||||
|
||||
def test_full_sale_restart_and_force_buyback_without_price_or_market_gate(self):
|
||||
sell = self.fill('t-sell', 's1', 200, price=11)
|
||||
boot.sync_account_state(self.store, [], [sell])
|
||||
self.store = State(self.store.path)
|
||||
self.cfg.zt_max_price = 10
|
||||
self.run.add_watch.triggered.return_value = False
|
||||
remaining = self.manage(price=12, positions=[], force=True)
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertEqual((request.op, request.volume), (23, 200))
|
||||
self.assertAlmostEqual(remaining, 10000 - 12 * 200 * 1.01)
|
||||
|
||||
def test_partial_fills_once_and_completed_round_blocks_same_day_sale(self):
|
||||
deals = [self.fill('t-sell', 's1', 40, 11), self.fill('t-sell', 's2', 60, 12)]
|
||||
self.position.volume = 100
|
||||
boot.sync_account_state(self.store, [self.position], deals + deals)
|
||||
item = t_rounds(self.store)[self.code]
|
||||
self.assertEqual(item['sold'], 100)
|
||||
self.assertEqual(item['amount'], 1160)
|
||||
self.manage(price=10)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
|
||||
deals.append(self.fill('t-buy', 'b1', 100))
|
||||
self.position.volume = 200
|
||||
boot.sync_account_state(self.store, [self.position], deals)
|
||||
def test_added_position_is_capped_by_sellable_inventory(self):
|
||||
for added, usable, expected in [(500, 100, 100), (100, 500, 100), (0, 500, 500)]:
|
||||
with self.subTest(added=added, usable=usable):
|
||||
self.run.orders.place.reset_mock()
|
||||
self.manage(price=11)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage(price=11, today='2026-09-10')
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].op, 24)
|
||||
self.manage(added=added, usable=usable)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, expected)
|
||||
|
||||
def test_cross_day_debt_and_insufficient_cash(self):
|
||||
boot.sync_account_state(self.store, [], [self.fill('t-sell', 's1', 200, date='2026-09-08')])
|
||||
self.manage(positions=[], available=100, force=True)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage(positions=[], force=True)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 200)
|
||||
|
||||
def test_delayed_snapshot_does_not_delete_or_recreate_holdings(self):
|
||||
boot.sync_account_state(self.store, [], [])
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 200)
|
||||
sell = self.fill('t-sell', 's1', 200)
|
||||
boot.sync_account_state(self.store, [self.position], [sell])
|
||||
self.assertNotIn(self.code, self.store.state)
|
||||
self.manage()
|
||||
def test_zero_sellable_does_not_divide_by_default_added_cost(self):
|
||||
with patch('strategy.zt.positions.log.exception') as error:
|
||||
self.manage(usable=0, added_cost=0)
|
||||
error.assert_not_called()
|
||||
self.run.orders.place.assert_not_called()
|
||||
|
||||
def test_base_fills_stay_in_base_bucket(self):
|
||||
self.store.sync_state([])
|
||||
deals = [self.fill('base', 'b1', 100), self.fill('base', 'b2', 100, 12)]
|
||||
boot.sync_account_state(self.store, [self.position], deals)
|
||||
row = self.store.state[self.code]
|
||||
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty']), (200, 11, 0))
|
||||
|
||||
def test_run_once_queries_sold_out_code_and_never_opens_with_debt(self):
|
||||
sell = self.fill('t-sell', 's1', 200, 11)
|
||||
self.run.client.deals.return_value = [sell]
|
||||
self.run.client.portfolio.return_value = SimpleNamespace(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
self.run.client.full_tick.return_value = {self.code: Tick(last_price=12)}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'collector_push'), \
|
||||
patch.object(boot, 'open_signal') as opened, patch.object(boot, 'market_allow_open') as market:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 14, 50)
|
||||
boot.RunOnce(self.run, self.store, [])
|
||||
self.run.client.full_tick.assert_called_once_with([self.code])
|
||||
opened.assert_not_called()
|
||||
market.assert_not_called()
|
||||
def test_unavailable_shares_do_not_disable_loss_management(self):
|
||||
self.run.account_cfg.enable_loss_add_position = True
|
||||
self.manage(usable=0, cost=20, price=10, added_cost=0)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].op, 23)
|
||||
|
||||
def test_open_budget_includes_buffer_and_star_minimum(self):
|
||||
with patch('strategy.zt.open.datetime') as clock:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 10)
|
||||
remaining = open_signal(self.run, {self.code: Tick(last_price=10)},
|
||||
[SimpleNamespace(code=self.code)], 2000)
|
||||
def test_on_road_shares_do_not_disable_available_base(self):
|
||||
self.manage(road=100)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 500)
|
||||
|
||||
def test_added_cost_is_used_even_if_base_cost_is_higher(self):
|
||||
self.manage(added=100, cost=20, added_cost=10, price=11)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
|
||||
self.assertEqual(remaining, 990)
|
||||
self.run.orders.place.reset_mock()
|
||||
open_signal(self.run, {'688001.SH': Tick(last_price=10)},
|
||||
[SimpleNamespace(code='688001.SH')], 2000)
|
||||
|
||||
def test_invalid_selected_cost_never_trades(self):
|
||||
for cost in [0, -1, float('nan'), float('inf')]:
|
||||
with self.subTest(cost=cost), patch('strategy.zt.positions.log.exception') as error:
|
||||
self.manage(added=100, added_cost=cost)
|
||||
error.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)
|
||||
def test_run_once_quarantines_manual_trade_but_manages_good_stock(self):
|
||||
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
|
||||
store = State(Path(tmp) / 'state.db')
|
||||
good = '600001.SH'
|
||||
positions = [PositionItem(stock_code=c, volume=100, can_use_volume=100, open_price=10)
|
||||
for c in [self.code, good]]
|
||||
store.sync_account(positions, [], initialize=True)
|
||||
manual = DealItem(stock_code=self.code, order_sys_id='manual', remark='',
|
||||
offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000)
|
||||
self.run.executor = executor
|
||||
self.run.client.deals.return_value = [manual]
|
||||
self.run.client.portfolio.return_value = NS(
|
||||
assets=Assets(10000, 10000), positions={p.stock_code: p for p in positions}, orders=[])
|
||||
self.run.client.full_tick.return_value = {p.stock_code: Tick(last_price=11) for p in positions}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True):
|
||||
clock.now.return_value = datetime(2026, 9, 11, 10)
|
||||
boot.RunOnce(self.run, store, [])
|
||||
self.assertEqual(store.blocked_codes, {self.code})
|
||||
self.run.orders.refresh.assert_called_once()
|
||||
self.assertEqual(self.run.orders.place.call_count, 1)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].code, good)
|
||||
|
||||
def test_run_once_does_not_reopen_quarantined_sold_out_code(self):
|
||||
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
|
||||
store = State(Path(tmp) / 'state.db')
|
||||
store.sync_account([], [], initialize=True)
|
||||
self.run.executor = executor
|
||||
self.run.client.deals.return_value = [DealItem(
|
||||
stock_code=self.code, order_sys_id='manual', remark='', offset_flag=FLAG_BUY,
|
||||
volume=100, price=10, trade_amount=1000)]
|
||||
self.run.client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
self.run.client.full_tick.return_value = {}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True), \
|
||||
patch.object(boot, 'open_signal') as opened:
|
||||
clock.now.return_value = datetime(2026, 9, 11, 10)
|
||||
boot.RunOnce(self.run, store, [NS(code=self.code)])
|
||||
opened.assert_not_called()
|
||||
|
||||
def start(self, client, directory):
|
||||
self.run.account_cfg.grid_step_pct = 1
|
||||
global_cfg = NS(qmt_base_url='unused', qmt_token='', qmt_data_dir=directory)
|
||||
self.run.account_cfg.signal_allow = []
|
||||
with patch.object(boot, 'Client', return_value=client), \
|
||||
patch.object(boot.config, 'global_config', global_cfg), \
|
||||
patch.object(boot.config, 'account_config', self.run.account_cfg), \
|
||||
patch.object(boot, 'init_signals', return_value=[]), \
|
||||
patch.object(boot, 'cache_portfolio'), patch.object(boot, 'Overview'), \
|
||||
patch.object(boot.time, 'localtime', return_value=NS(tm_hour=15, tm_min=0, tm_sec=0)):
|
||||
boot.StartZT()
|
||||
|
||||
def test_start_initializes_once_without_snapshot_retry_loop(self):
|
||||
client = Mock()
|
||||
client.deals.return_value = []
|
||||
client.portfolio.return_value = NS(assets=Assets(10000, 10000),
|
||||
positions={self.code: PositionItem(stock_code=self.code, volume=100, open_price=10)}, orders=[])
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
self.start(client, tmp)
|
||||
store = State(Path(tmp) / 'zt_test_state.db')
|
||||
self.assertEqual(store.state[self.code]['base_qty'], 100)
|
||||
with closing(store._connect()) as db:
|
||||
self.assertIsNone(db.execute("SELECT 1 FROM sqlite_master WHERE name='state_meta'").fetchone())
|
||||
self.assertEqual(client.portfolio.call_count, 1)
|
||||
self.assertEqual(client.deals.call_count, 2)
|
||||
client.reset_mock()
|
||||
self.start(client, tmp)
|
||||
self.assertEqual(client.portfolio.call_count, 1)
|
||||
self.assertEqual(client.deals.call_count, 1)
|
||||
|
||||
def test_start_rejects_changed_deals_without_writing_baseline(self):
|
||||
client = Mock()
|
||||
client.deals.side_effect = [[], [DealItem(order_sys_id='new')]]
|
||||
client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.start(client, tmp)
|
||||
store = State(Path(tmp) / 'zt_test_state.db')
|
||||
self.assertEqual((store.state, store.deals), ({}, {}))
|
||||
client.close.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
Reference in New Issue
Block a user