Files
eth_hedge_sim/backend/app/strategy/engine.py
T

365 lines
14 KiB
Python

"""策略状态机:选向开仓 / 盯盘平仓 / 休息(无开仓窗、无轮次上限)。"""
from __future__ import annotations
import asyncio
import logging
import time
from typing import Any
from ..config import get_settings
from .session import get_session
from ..models.db import get_db
from ..sim.ledger import Ledger
from ..sim.matcher import Matcher
from .clock import can_open_new, window_key
from .exits import check_expiry_close, check_exits, resolve_exit_target
from .group import next_group_id
logger = logging.getLogger(__name__)
class StrategyEngine:
def __init__(self) -> None:
self.db = get_db()
self.matcher = Matcher(self.db)
self.ledger = Ledger(self.db)
self._task: asyncio.Task[None] | None = None
self._lock = asyncio.Lock()
def state(self) -> dict[str, Any]:
row = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert row is not None
upl = self.matcher.unrealized()
s = get_settings()
exit_mode = self.ledger.get_setting_str("exit_mode", s.exit_mode)
net_target = self.ledger.get_setting_float(
"net_profit_target", s.net_profit_target
)
prem_mult = self.ledger.get_setting_float(
"premium_exit_multiple", s.premium_exit_multiple
)
exit_amt, _ = resolve_exit_target(
exit_mode=exit_mode,
net_profit_target=net_target,
premium_exit_multiple=prem_mult,
initial_premium=float(upl.get("initial_premium") or 0),
)
rest_sec = self.ledger.get_setting_int("rest_seconds", s.rest_seconds)
skip_weekends = self.ledger.get_setting_bool("skip_weekends", s.skip_weekends)
leverage = self.ledger.get_setting_float("leverage", s.leverage)
min_hours = self.ledger.get_setting_float("min_option_hours", s.min_option_hours)
min_opt_lev = self.ledger.get_setting_float(
"min_option_leverage", s.min_option_leverage
)
atm_off_on = self.ledger.get_setting_bool(
"atm_open_offset_enabled", s.atm_open_offset_enabled
)
max_atm_off = self.ledger.get_setting_float(
"max_atm_open_offset", s.max_atm_open_offset
)
rest_until = row["rest_until_ms"]
rest_left = 0
if rest_until:
rest_left = max(0, int((int(rest_until) - time.time() * 1000) / 1000))
last_error = row["last_error"]
if last_error and "PriceResult" in str(last_error) and "__dict__" in str(last_error):
self._set_state(last_error=None)
last_error = None
allow_open = can_open_new(skip_weekends=skip_weekends)
return {
"running": bool(row["running"]),
"phase": row["phase"],
"rounds_done": int(row["rounds_done"] or 0),
"window_key": row["window_key"],
"rest_until_ms": rest_until,
"rest_left_sec": rest_left,
"rest_seconds": rest_sec,
"skip_weekends": skip_weekends,
"exit_mode": exit_mode,
"net_profit_target": net_target,
"premium_exit_multiple": prem_mult,
"exit_target_usdt": exit_amt,
"leverage": leverage,
"min_option_hours": min_hours,
"min_option_leverage": min_opt_lev,
"atm_open_offset_enabled": atm_off_on,
"max_atm_open_offset": max_atm_off,
"can_open": allow_open,
"last_error": last_error,
"position": upl,
"ledger": self.ledger.snapshot(),
}
def _set_state(self, **kwargs: Any) -> None:
cols = []
vals: list[Any] = []
for k, v in kwargs.items():
cols.append(f"{k}=?")
vals.append(v)
cols.append("updated_at_ms=?")
vals.append(int(time.time() * 1000))
sql = f"UPDATE strategy_state SET {', '.join(cols)} WHERE id=1"
self.db.execute(sql, tuple(vals))
async def pause(self) -> dict[str, Any]:
self._set_state(running=0, phase="paused", last_error=None)
return self.state()
async def start(self) -> dict[str, Any]:
self._set_state(running=1, last_error=None, phase="idle")
self.ensure_loop()
return self.state()
def ensure_loop(self) -> None:
"""保证后台循环在跑(即使策略暂停,也要盯到期全平)。"""
if self._task is None or self._task.done():
self._task = asyncio.create_task(self._loop(), name="strategy-engine")
async def emergency_close(self) -> dict[str, Any]:
async with self._lock:
# 紧急全平:绕过期权流动性/偏差校验
r = self.matcher.close_group(reason="emergency", bypass_liquidity=True)
if r.ok:
self._after_close()
return {
"close": {
"ok": r.ok,
"detail": r.detail,
"liquidity_wait": r.liquidity_wait,
"data": r.data,
},
"state": self.state(),
}
def _after_close(self) -> None:
s = get_settings()
row = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert row is not None
rounds = int(row["rounds_done"] or 0) + 1
rest_sec = self.ledger.get_setting_int("rest_seconds", s.rest_seconds)
rest_until = int(time.time() * 1000) + rest_sec * 1000
self._set_state(
rounds_done=rounds,
phase="resting",
rest_until_ms=rest_until,
)
def _count_groups_for_day(self, wkey: str) -> int:
rows = self.db.fetchall(
"SELECT group_id FROM groups WHERE group_id LIKE ?",
(f"G-{wkey}-%",),
)
return len(rows)
def _position_expiry_ms(self, upl: dict[str, Any]) -> int | None:
raw = upl.get("expiry_ms")
if raw is not None:
try:
return int(raw)
except (TypeError, ValueError):
pass
ymd = upl.get("expiry_ymd")
if ymd:
try:
from ..exchange.expiry import expiry_ms_from_ymd
return int(expiry_ms_from_ymd(str(ymd)))
except Exception:
return None
return None
async def _close_open_position(
self,
*,
reason: str,
bypass_liquidity: bool,
pending_close: bool,
) -> None:
if not pending_close:
self._set_state(phase="closing", last_error=None)
r = await asyncio.to_thread(
self.matcher.close_group,
reason=reason,
bypass_liquidity=bypass_liquidity,
)
if r.ok:
self._after_close()
elif r.liquidity_wait and not bypass_liquidity:
self._set_state(phase="liquidity_wait", last_error=r.detail)
else:
self._set_state(phase="closing", last_error=r.detail)
async def _maybe_expiry_close(self) -> bool:
"""若持仓已到期则强制全平。返回是否触发到期平仓。"""
pos = self.matcher.current_position()
if pos.get("status") != "open":
return False
upl = self.matcher.unrealized()
expired = check_expiry_close(expiry_ms=self._position_expiry_ms(upl))
if not expired.should_close:
return False
st = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert st is not None
pending = st["phase"] in ("liquidity_wait", "closing")
await self._close_open_position(
reason="expiry",
bypass_liquidity=True,
pending_close=pending,
)
return True
async def _loop(self) -> None:
logger.info("strategy engine loop started")
while True:
try:
row = self.db.fetchone("SELECT running FROM strategy_state WHERE id=1")
running = bool(row and int(row["running"]))
if not running:
# 暂停时仍执行到期全平,避免拖过期
async with self._lock:
await self._maybe_expiry_close()
await asyncio.sleep(1)
continue
async with self._lock:
try:
await get_session().ensure_atm_async(force=False)
except Exception as e:
logger.warning("ATM ensure before tick failed: %s", e)
await self._tick_async()
except asyncio.CancelledError:
raise
except Exception as e:
err = str(e)
# 限流时勿刷屏;拉长休眠给 eapi 冷却
if "418" in err or "429" in err or "cooldown" in err.lower():
logger.warning("strategy tick rate-limited: %s", err[:200])
self._set_state(last_error="币安期权接口限流,稍后自动重试")
await asyncio.sleep(15)
continue
logger.exception("strategy tick failed")
self._set_state(last_error=err)
await asyncio.sleep(1)
async def _tick_async(self) -> None:
s = get_settings()
st = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert st is not None
wkey = window_key()
if st["window_key"] != wkey:
self._set_state(window_key=wkey, phase="idle")
st = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert st is not None
exit_mode = self.ledger.get_setting_str("exit_mode", s.exit_mode)
net_target = self.ledger.get_setting_float(
"net_profit_target", s.net_profit_target
)
prem_mult = self.ledger.get_setting_float(
"premium_exit_multiple", s.premium_exit_multiple
)
pos = self.matcher.current_position()
# 有未平仓:只盯平仓,绝不开下一组
if pos.get("status") == "open":
upl = self.matcher.unrealized()
expired = check_expiry_close(expiry_ms=self._position_expiry_ms(upl))
decision = check_exits(
net_pnl=float(upl.get("net_pnl") or 0),
exit_mode=exit_mode,
net_profit_target=net_target,
premium_exit_multiple=prem_mult,
initial_premium=float(upl.get("initial_premium") or 0),
)
pending_close = st["phase"] in ("liquidity_wait", "closing")
if expired.should_close or decision.should_close or pending_close:
if expired.should_close:
reason = "expiry"
bypass = True
else:
reason = decision.reason or "liquidity_retry"
bypass = False
await self._close_open_position(
reason=reason,
bypass_liquidity=bypass,
pending_close=pending_close,
)
else:
self._set_state(phase="open", last_error=None)
return
if st["phase"] == "resting" and st["rest_until_ms"]:
if int(time.time() * 1000) < int(st["rest_until_ms"]):
return
self._set_state(phase="idle", rest_until_ms=None)
st = self.db.fetchone("SELECT * FROM strategy_state WHERE id=1")
assert st is not None
if st["phase"] in ("paused",):
return
# 旧「轮次停开」状态:自动恢复为空闲以便继续
if st["phase"] in ("stopped", "outside_window"):
self._set_state(phase="idle")
skip_weekends = self.ledger.get_setting_bool("skip_weekends", s.skip_weekends)
if not can_open_new(skip_weekends=skip_weekends):
self._set_state(
phase="weekend_skip",
last_error="周六/周日跳过开仓(上海时区);持仓仍可平仓",
)
return
if st["phase"] == "weekend_skip":
self._set_state(phase="idle", last_error=None)
# 双保险:账本仍显示有仓则不开
if self.matcher.has_open_position():
self._set_state(phase="open", last_error="有未平仓,禁止开下一组")
return
self._set_state(phase="wait_signal")
pick = await get_session().pick_for_open_async()
if pick is None:
self._set_state(
last_error="无合格期权:需剩余时长、杠杆(及已开启的ATM偏差)同时满足"
)
return
self._set_state(phase="opening", last_error=None)
count = self._count_groups_for_day(wkey)
gid = next_group_id(count)
option_inst = (
pick.pair.call_inst_id if pick.option_side == "call" else pick.pair.put_inst_id
)
entry_idx = pick.underlying_px
r = await asyncio.to_thread(
self.matcher.open_group,
group_id=gid,
bias=pick.bias,
option_side=pick.option_side,
perp_side=pick.perp_side,
option_inst_id=option_inst,
entry_index_px=float(entry_idx),
strike=pick.pair.strike,
expiry_ymd=pick.pair.expiry_ymd,
)
if r.ok:
self._set_state(phase="open", last_error=None)
else:
self._set_state(phase="idle", last_error=r.detail)
_engine: StrategyEngine | None = None
def get_engine() -> StrategyEngine:
global _engine
if _engine is None:
_engine = StrategyEngine()
return _engine
def set_engine(engine: StrategyEngine | None) -> None:
global _engine
_engine = engine