This commit is contained in:
2026-09-11 14:38:24 +08:00
parent 2888126a3a
commit af790bca28
4 changed files with 186 additions and 87 deletions

View File

@@ -137,7 +137,7 @@ def RunOnce(run: Runtime, signals: list[SignalItem]) -> None:
# 3. 获取大盘状态,只有大盘信号允许时才执行开仓。 # 3. 获取大盘状态,只有大盘信号允许时才执行开仓。
market_ok = market_allow_open() market_ok = market_allow_open()
# 4. 验证有效开仓信号:排除已有持仓和未决订单 # 4. 验证有效开仓信号:排除已有持仓。
allow_open: list[SignalItem] = [] allow_open: list[SignalItem] = []
allow_codes: list[str] = [] allow_codes: list[str] = []
for signal in signals: for signal in signals:

View File

@@ -6,7 +6,7 @@ from contextlib import closing
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from tempfile import TemporaryDirectory from tempfile import TemporaryDirectory
from concurrent.futures import Future, ThreadPoolExecutor
import config import config
from libs.calc import trading_time from libs.calc import trading_time
from libs.collector import collector_push from libs.collector import collector_push
@@ -21,98 +21,179 @@ from libs.watch import DipWatch
from sdk import Client, DealItem, PositionItem from sdk import Client, DealItem, PositionItem
from .open import open_signal from .open import open_signal
from .positions import manage_positions, t_rounds from .positions import manage_positions, t_rounds
from libs.snapshot import cache_portfolio
def StartZT() -> None: def StartZT() -> None:
with Client(config.global_config.qmt_base_url, config.global_config.qmt_token, config.HTTP_TIMEOUT) as client: client = Client(
config.global_config.qmt_base_url,
config.global_config.qmt_token,
config.HTTP_TIMEOUT,
)
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') state = State(Path(config.global_config.qmt_data_dir) / f'zt_{config.account_config.account_id}_state.db')
run = Runtime( run = Runtime(
client=client, global_cfg=config.global_config, account_cfg=config.account_config, client=client, global_cfg=config.global_config, account_cfg=config.account_config,
orders=OrderBook('zt'), open_watch=DipWatch(), add_watch=DipWatch(), orders=OrderBook('zt'), open_watch=DipWatch(), add_watch=DipWatch(),
profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct), profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct),
) )
# 先读取成交,再读取持仓,减少成交已入账而快照仍未更新的情况。
# 获取本策略的信号开仓数据
signals = init_signals(
config.global_config,
config.account_config.signal_allow,
)
log.info(
"[启动] Trend策略已启动账户=%s,信号=%d,持仓=%d",
config.account_config.account_id,
len(signals),
len(positions),
)
deals = client.deals() deals = client.deals()
portfolio = client.portfolio() cache_portfolio(config.account_config.account_id, assets, positions, deals)
positions = list(portfolio.positions.values()) state.load()
sync_account_state(state, positions, deals, initialize=not state.state and not state.deals) state.sync_deals(deals)
state.sync_state(positions)
state.archiving()
run.orders.refresh(client, portfolio.orders) run.orders.refresh(client, portfolio.orders)
signals = init_signals(config.global_config, ['dcm']) Overview(assets, positions, config.account_config)
Overview(portfolio.assets, positions, config.account_config) executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="zt")
log.info('[ZT] 启动,账户=%s,信号=%d', config.account_config.account_id, len(signals))
while datetime.now().hour < 15: DEFAULT_TICK_INTERVAL = 30
while True:
lt = time.localtime()
if (lt.tm_hour, lt.tm_min, lt.tm_sec) >= (15, 0, 0):
log.info("[ZT] 已到 15:00结束趋势策略")
return
current_sec = lt.tm_sec
# 计算距离下一个目标时间点0秒或30秒的等待时间
if current_sec < DEFAULT_TICK_INTERVAL:
wait_seconds = DEFAULT_TICK_INTERVAL - current_sec
elif current_sec < 60:
wait_seconds = 60 - current_sec
else:
wait_seconds = DEFAULT_TICK_INTERVAL
# 等待到目标时间点
time.sleep(wait_seconds)
# 单轮失败不能杀死唯一的交易定时线程。
try: try:
RunOnce(run, state, signals) RunOnce(run, state, signals)
except Exception: except Exception as e:
log.exception('[ZT] 本轮失败,下一轮重试') log.error(
time.sleep(30 - time.time() % 30) f"[ZT] 本 tick 执行失败,下一 tick 继续: {e}", exc_info=True
# 收盘后补记最后一轮成交,不再下单。 )
deals = client.deals() finally:
sync_account_state(state, list(client.portfolio().positions.values()), deals) try:
if executor is not None:
executor.shutdown(wait=True)
finally:
client.close()
def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None: def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
now = datetime.now() now = datetime.now()
if not trading_time(now): if not trading_time(now):
return return
deals = run.client.deals()
portfolio = run.client.portfolio() print(
assets = portfolio.assets "=" * 40 + f" Ticker {datetime.now().strftime('%Y-%m-%d %H:%M:%S')} " + "=" * 40
positions = list(portfolio.positions.values()) )
run.orders.refresh(run.client, portfolio.orders)
sync_account_state(state, positions, deals) started_at = time.monotonic()
# 收盘集合竞价前停止提交新委托,继续保存成交。 futures: list[tuple[str, Future]] = []
if (now.hour, now.minute) >= (14, 57):
return
rounds = t_rounds(state)
pending = {code for code, item in rounds.items() if item['sold'] > item['bought']}
candidates = {s.code: s for s in signals if s.code not in portfolio.positions
and s.code not in state.state and s.code not in pending}
codes = list(dict.fromkeys(list(state.state) + sorted(pending) + list(candidates)))
ticks = run.client.full_tick(codes) if codes else {}
force = (now.hour, now.minute) >= (14, 50)
available = manage_positions(run, state, ticks, positions, rounds, assets.available, now.date().isoformat(), force)
# 尚未买回时不分走资金;买回与新建仓使用同一份剩余资金。
if not force and not pending and available >= assets.total * run.account_cfg.min_cash_ratio:
if candidates and market_allow_open():
budget = max(0.0, available - assets.total * run.account_cfg.min_cash_ratio)
open_signal(run, ticks, list(candidates.values()), budget)
try: try:
collector_push(run.account_cfg.account_id, assets, positions) deals = run.client.deals()
portfolio = run.client.portfolio()
assets = portfolio.assets
positions = list(portfolio.positions.values())
position_codes = list(portfolio.positions)
state.sync_deals(deals)
state.archiving()
run.orders.refresh(run.client, portfolio.orders)
except Exception: except Exception:
log.exception('[ZT] 持仓上报失败') log.exception("[Portfolio] 刷新账户快照失败")
def sync_account_state(
state: State, positions: list[PositionItem], deals: list[DealItem], *, initialize: bool = False,
) -> None:
"""初次持仓作为底仓;后续只按成交减仓,避免延迟快照删除持仓。"""
zt_deals = [d for d in deals if d.get_local_order_id.startswith('zt-')]
if initialize:
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
state.sync_deals(zt_deals)
# 在外层统一归档:新增成交写入后处理,也重试此前失败的未归档成交 # 2. 验证可用资金;低于资金安全线时禁止开新仓
state.archiving() allow_open_by_cash = (
# 从刷新后的成交缓存检查失败记录,避免归档未完成时继续交易。 assets.available >= assets.total * run.account_cfg.min_cash_ratio
pending = sorted({d['stock_code'] for d in state.deals.values() if d['is_arch'] == 0}) )
if pending: if not allow_open_by_cash:
raise ValueError(f'ZT 成交归档未完成:{pending},原因见归档日志') log.info(
traded = {d['stock_code'] for d in state.deals.values()} "[Status] 禁止开仓:可用资金不足,可用=%.2f,总资产=%.2f",
# sync_state 要求完整持仓。保留成交账本中的现有仓位,避免延迟快照删仓; assets.available,
# 只从账户快照补入没有策略成交历史的股票,避免刚卖完又被旧快照重建。 assets.total,
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) # 3. 获取大盘状态,只有大盘信号允许时才执行开仓。
state.sync_state(holdings) market_ok = market_allow_open()
# 4. 验证有效开仓信号:排除已有持仓。
allow_open: list[SignalItem] = []
allow_codes: list[str] = []
for signal in signals:
if signal.code not in portfolio.positions:
allow_open.append(signal)
allow_codes.append(signal.code)
if allow_open and not market_ok:
log.info("[开仓] 禁止开仓:大盘信号不允许,候选=%d", len(allow_open))
# 5. 获取持仓和待开仓证券的实时行情 tick。
all_codes = list(dict.fromkeys(position_codes + allow_codes))
try:
ticks = run.client.full_tick(all_codes)
except Exception:
log.exception("[行情] 获取行情失败,代码数量=%d", len(all_codes))
return
log.info(
"[RunOnce] 本轮就绪,持仓=%d,候选=%d,大盘允许=%s,资金允许=%s",
len(positions),
len(allow_open),
market_ok,
allow_open_by_cash,
)
# 启动线程,开始计算
# 7. 持仓计算。
futures.append(
(
"持仓计算",
run.executor.submit(
manage_positions, run, ticks, positions, market_ok, assets.available
),
)
)
# 8. 开仓计算:必须同时存在有效信号且大盘允许开仓。
if allow_open and market_ok and allow_open_by_cash:
futures.append(
("开仓计算", run.executor.submit(open_signal, run, ticks, allow_open))
)
# 9. 开始执行
for name, future in futures:
_wait_worker(name, future)
log.info(
"[RunOnce] 本轮完成,耗时=%d毫秒", int((time.monotonic() - started_at) * 1000)
)
def _wait_worker(name: str, future: Future) -> None:
"""保留单轮继续运行的语义,分别记录工作线程异常。"""
try:
future.result()
except Exception:
log.exception("[运行] %s线程失败", name)

View File

@@ -81,7 +81,7 @@ class PerformanceRegressionTests(unittest.TestCase):
signals = [SignalItem(code=c) for c in ('new-b', 'held', 'new-a', 'new-b')] signals = [SignalItem(code=c) for c in ('new-b', 'held', 'new-a', 'new-b')]
with patch.object(boot, 'trading_time', return_value=True), \ with patch.object(boot, 'trading_time', return_value=True), \
patch.object(boot, 'market_allow_open', return_value=True), \ patch.object(boot, 'market_allow_open', return_value=True), \
patch.object(boot, '_cache_portfolio'), patch('builtins.print'): patch.object(boot, 'cache_portfolio'), patch('builtins.print'):
boot.RunOnce(run, signals) boot.RunOnce(run, signals)
run.client.full_tick.assert_called_once_with(['held', 'new-b', 'new-a']) run.client.full_tick.assert_called_once_with(['held', 'new-b', 'new-a'])
self.assertEqual(run.executor.submit.call_args_list[1].args, self.assertEqual(run.executor.submit.call_args_list[1].args,

View File

@@ -7,8 +7,8 @@ from contextlib import ExitStack, redirect_stdout
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import Mock, patch from unittest.mock import Mock, patch
from libs import collector from libs import collector, snapshot
from sdk import Assets, PositionItem from sdk import Assets, DealItem, PositionItem
from strategy.trend import boot from strategy.trend import boot
@@ -19,25 +19,28 @@ class TrendCollectorTests(unittest.TestCase):
cls.app = importlib.import_module('main') cls.app = importlib.import_module('main')
def setUp(self): def setUp(self):
old_snapshot = boot._collector_snapshot old_snapshot = snapshot._collector_snapshot
self.addCleanup(setattr, boot, '_collector_snapshot', old_snapshot) self.addCleanup(setattr, snapshot, '_collector_snapshot', old_snapshot)
boot._collector_snapshot = None snapshot._collector_snapshot = None
def test_submission_reads_latest_cache_and_skips_empty(self): def test_submission_reads_latest_cache_and_skips_empty(self):
with patch.object(collector, 'collector_push') as push: with patch.object(collector, 'collector_push') as push:
collector.submit_trend_data() collector.submit_trend_data()
push.assert_not_called() push.assert_not_called()
boot._cache_portfolio('account', Assets(available=100), []) snapshot.cache_portfolio('account', Assets(available=100), [], [])
assets = Assets(available=200) assets = Assets(available=200)
positions = [PositionItem(stock_code='600000.SH', volume=100)] positions = [PositionItem(stock_code='600000.SH', volume=100)]
boot._cache_portfolio('account', assets, positions) deals = [DealItem(stock_code='600000.SH', volume=100)]
snapshot.cache_portfolio('account', assets, positions, deals)
collector.submit_trend_data() collector.submit_trend_data()
push.assert_called_once_with('account', assets, positions) push.assert_called_once_with('account', assets, positions, deals)
uploaded = push.call_args.args uploaded = push.call_args.args
uploaded[1].available = 0 uploaded[1].available = 0
uploaded[2].clear() uploaded[2].clear()
self.assertEqual(boot.get_collector_snapshot()[1].available, 200) uploaded[3][0].volume = 0
self.assertEqual(len(boot.get_collector_snapshot()[2]), 1) self.assertEqual(snapshot.get_collector_snapshot()[3][0].volume, 100)
self.assertEqual(snapshot.get_collector_snapshot()[1].available, 200)
self.assertEqual(len(snapshot.get_collector_snapshot()[2]), 1)
def test_run_once_caches_portfolio_without_submitting_data(self): def test_run_once_caches_portfolio_without_submitting_data(self):
completed = Future() completed = Future()
@@ -48,16 +51,31 @@ class TrendCollectorTests(unittest.TestCase):
) )
assets = Assets(available=100, total=1000) assets = Assets(available=100, total=1000)
run.client.portfolio.return_value = SimpleNamespace(assets=assets, positions={}, orders=[]) run.client.portfolio.return_value = SimpleNamespace(assets=assets, positions={}, orders=[])
deals = [DealItem(stock_code='600000.SH', volume=100)]
run.client.deals.return_value = deals
run.client.full_tick.return_value = {} run.client.full_tick.return_value = {}
run.executor.submit.return_value = completed run.executor.submit.return_value = completed
with patch.object(boot, 'trading_time', return_value=True), \ with patch.object(boot, 'trading_time', return_value=True), \
patch.object(boot, 'market_allow_open', return_value=True), \ patch.object(boot, 'market_allow_open', return_value=True), \
patch.object(collector, 'collector_push') as push, redirect_stdout(io.StringIO()): patch.object(collector, 'collector_push') as push, redirect_stdout(io.StringIO()):
boot.RunOnce(run, []) boot.RunOnce(run, [])
self.assertEqual(boot.get_collector_snapshot(), ('account', assets, [])) self.assertEqual(snapshot.get_collector_snapshot(), ('account', assets, [], deals))
run.client.deals.assert_called_once_with()
push.assert_not_called() push.assert_not_called()
run.executor.submit.assert_called_once_with(boot.manage_positions, run, {}, [], True, 100) run.executor.submit.assert_called_once_with(boot.manage_positions, run, {}, [], True, 100)
def test_submission_serializes_deals(self):
snapshot.cache_portfolio(
'account', Assets(available=100), [],
[DealItem(stock_code='600000.SH', volume=100)],
)
with patch.object(collector.httpx, 'post') as post:
collector.submit_trend_data()
payload = post.call_args.kwargs['json']
self.assertEqual(payload['account_id'], 'account')
self.assertEqual(payload['deals'][0]['stock_code'], '600000.SH')
self.assertEqual(payload['deals'][0]['volume'], 100)
def test_main_registers_five_minute_collector_job(self): def test_main_registers_five_minute_collector_job(self):
for strategy in ('trend', 'zt'): for strategy in ('trend', 'zt'):
with self.subTest(strategy=strategy), ExitStack() as stack: with self.subTest(strategy=strategy), ExitStack() as stack: