from datetime import datetime
from threading import Timer
from typing import Any, Literal

from pythongo.base import BaseParams, BaseState, Field
from pythongo.classdef import KLineData, OrderData, TickData, TradeData
from pythongo.core import KLineStyleType
from pythongo.ui import BaseStrategy
from pythongo.utils import KLineGenerator


class Params(BaseParams):
    """参数映射模型"""
    exchange: str = Field(default="SHFE", title="交易所代码")
    instrument_id: str = Field(default="rb2610", title="合约代码")
    kline_style: KLineStyleType = Field(default="M5", title="小周期K线")
    big_kline_style: KLineStyleType = Field(default="M30", title="大周期K线")
    expma_fast: int = Field(default=12, title="EXPMA快线")
    expma_slow: int = Field(default=26, title="EXPMA慢线")
    ma40_period: int = Field(default=40, title="大周期MA周期")
    lots: int = Field(default=1, title="开仓手数")
    over_price: int = Field(default=1, title="委托超价(跳数)")
    enable_short: int = Field(default=1, title="允许做空(1是0否)")


class State(BaseState):
    """状态映射模型"""
    status: str = Field(default="等待数据", title="运行状态")
    expma_fast: float = Field(default=0.0, title="EXPMA快线")
    expma_slow: float = Field(default=0.0, title="EXPMA慢线")
    ma40_30m: float = Field(default=0.0, title="30分钟MA40")
    signal_desc: str = Field(default="", title="信号说明")


PendingAction = Literal[
    "",
    "open_long",
    "open_short",
    "close_long",
    "close_short",
    "reverse_to_short",
    "reverse_to_long",
]


class EXPMA多周期MA40(BaseStrategy):
    """
    小周期 EXPMA 快慢线金叉/死叉产生交易信号，30 分钟周期 MA40 作为价格位置过滤。

    系统要素：
    1. 小周期（主图）EXPMA 快线、慢线
    2. 30 分钟周期 MA40 作为大周期价格过滤线
    3. 金叉/死叉及价格过滤均用前一根 K 线确认值，避免信号闪烁

    入场：
    - 开多：小周期 EXPMA 金叉，且收盘价在 30 分钟 MA40 上方
    - 开空：小周期 EXPMA 死叉，且收盘价在 30 分钟 MA40 下方（需开启做空）

    出场：
    - 持多遇死叉：MA40 上方只平多；MA40 下方平多并反手开空
    - 持空遇金叉：MA40 下方只平空；MA40 上方平空并反手开多

    K 线收线判定信号，下一跳 tick 对手价委托。
    """

    def __init__(self) -> None:
        super().__init__()
        self.params_map = Params()
        self.state_map = State()

        self.kline_generator: KLineGenerator | None = None
        self.big_kline_generator: KLineGenerator | None = None

        self.last_price = 0.0
        self.bid_price1 = 0.0
        self.ask_price1 = 0.0
        self.price_tick = 1.0

        self.long_pos = 0
        self.short_pos = 0

        self.expma_fast_val = 0.0
        self.expma_slow_val = 0.0
        self.big_ma40_val = 0.0

        self.pending_action: PendingAction = ""

        self.buy_orderid: int | None = None
        self.sell_orderid: int | None = None
        self.short_orderid: int | None = None
        self.cover_orderid: int | None = None

        self.signal_price = 0.0

    @property
    def main_indicator_data(self) -> dict[str, float]:
        return {
            "EXPMA快": self.expma_fast_val,
            "EXPMA慢": self.expma_slow_val,
            "MA40_30M": self.big_ma40_val,
        }

    def on_init(self) -> None:
        super().on_init()
        self.output("EXPMA多周期MA40策略初始化")

    def on_start(self) -> None:
        p = self.params_map

        self.kline_generator = KLineGenerator(
            callback=self.on_bar,
            real_time_callback=self.on_bar_realtime,
            exchange=p.exchange,
            instrument_id=p.instrument_id,
            style=p.kline_style,
        )
        self.kline_generator.push_history_data()

        self.big_kline_generator = KLineGenerator(
            callback=self.on_big_bar,
            real_time_callback=self.on_big_bar_realtime,
            exchange=p.exchange,
            instrument_id=p.instrument_id,
            style=p.big_kline_style,
        )
        self.big_kline_generator.push_history_data()

        try:
            inst = self.get_instrument_data(p.exchange, p.instrument_id)
            if inst:
                self.price_tick = inst.price_tick
        except Exception as e:
            self.output(f"[警告] 获取合约信息失败: {e}")

        super().on_start()

        if p.expma_fast >= p.expma_slow:
            self.output("错误: EXPMA快线周期须小于慢线周期")
        self.output(
            f"策略启动 {p.exchange}.{p.instrument_id} "
            f"小周期{p.kline_style} 大周期{p.big_kline_style} "
            f"EXPMA{p.expma_fast}/{p.expma_slow} MA{p.ma40_period}"
        )

    def on_stop(self) -> None:
        super().on_stop()
        self.output("策略停止")

    def on_tick(self, tick: TickData) -> None:
        super().on_tick(tick)
        if tick.instrument_id != self.params_map.instrument_id:
            return

        self.last_price = tick.last_price
        self.bid_price1 = tick.bid_price1
        self.ask_price1 = tick.ask_price1

        if self.is_trading_time():
            self._update_position()
            if not self._has_pending_order():
                self._execute_pending_action()

        if self.big_kline_generator:
            self.big_kline_generator.tick_to_kline(tick)
        if self.kline_generator:
            self.kline_generator.tick_to_kline(tick)

        if self.trading:
            self.update_status_bar()

    def on_bar(self, kline: KLineData) -> None:
        """小周期K线收线：用前一根确认值判定信号"""
        self._calc_small_indicators()
        self._calc_big_ma40()

        if self.trading and self.is_trading_time():
            self._evaluate_signals(kline)

        self._update_chart(kline)

    def on_bar_realtime(self, kline: KLineData) -> None:
        self._calc_small_indicators()
        self._calc_big_ma40()
        self._update_chart(kline)

    def on_big_bar(self, kline: KLineData) -> None:
        self._calc_big_ma40()

    def on_big_bar_realtime(self, kline: KLineData) -> None:
        self._calc_big_ma40()

    def _calc_small_indicators(self) -> None:
        if self.kline_generator is None:
            return

        p = self.params_map
        bars = self.kline_generator.producer
        need = max(p.expma_fast, p.expma_slow) + 3
        if len(bars.close) < need:
            return

        fast_arr = bars.ema(timeperiod=p.expma_fast, array=True)
        slow_arr = bars.ema(timeperiod=p.expma_slow, array=True)

        self.expma_fast_val = float(fast_arr[-1])
        self.expma_slow_val = float(slow_arr[-1])
        self.state_map.expma_fast = round(self.expma_fast_val, 4)
        self.state_map.expma_slow = round(self.expma_slow_val, 4)

    def _calc_big_ma40(self) -> None:
        if self.big_kline_generator is None:
            return

        p = self.params_map
        bars = self.big_kline_generator.producer
        need = p.ma40_period + 2
        if len(bars.close) < need:
            return

        ma_arr = bars.sma(timeperiod=p.ma40_period, array=True)
        self.big_ma40_val = float(ma_arr[-1])
        self.state_map.ma40_30m = round(self.big_ma40_val, 4)

    def _get_confirmed_signal(self) -> dict[str, Any] | None:
        """取前一根K线确认值，避免当前未完成K线闪烁"""
        if self.kline_generator is None or self.big_kline_generator is None:
            return None

        p = self.params_map
        small = self.kline_generator.producer
        big = self.big_kline_generator.producer

        small_need = max(p.expma_fast, p.expma_slow) + 3
        big_need = p.ma40_period + 2
        if len(small.close) < small_need or len(big.close) < big_need:
            return None

        fast_arr = small.ema(timeperiod=p.expma_fast, array=True)
        slow_arr = small.ema(timeperiod=p.expma_slow, array=True)
        ma40_arr = big.sma(timeperiod=p.ma40_period, array=True)

        confirm_close = float(small.close[-2])
        confirm_fast = float(fast_arr[-2])
        confirm_slow = float(slow_arr[-2])
        prev_fast = float(fast_arr[-3])
        prev_slow = float(slow_arr[-3])
        confirm_ma40 = float(ma40_arr[-2])

        if min(confirm_close, confirm_fast, confirm_slow, prev_fast, prev_slow, confirm_ma40) <= 0:
            return None

        golden = prev_fast <= prev_slow and confirm_fast > confirm_slow
        death = prev_fast >= prev_slow and confirm_fast < confirm_slow
        above_ma40 = confirm_close > confirm_ma40
        below_ma40 = confirm_close < confirm_ma40

        return {
            "golden": golden,
            "death": death,
            "above_ma40": above_ma40,
            "below_ma40": below_ma40,
            "confirm_close": confirm_close,
            "confirm_ma40": confirm_ma40,
            "confirm_fast": confirm_fast,
            "confirm_slow": confirm_slow,
        }

    def _evaluate_signals(self, kline: KLineData) -> None:
        sig = self._get_confirmed_signal()
        if sig is None:
            return

        self.pending_action = ""
        dt_str = kline.datetime.strftime("%Y-%m-%d %H:%M:%S")
        allow_short = self.params_map.enable_short == 1

        if self.long_pos > 0 and sig["death"]:
            if sig["above_ma40"]:
                self.pending_action = "close_long"
                self.state_map.signal_desc = "持多死叉(MA40上)只平多"
                self.output(
                    f"{dt_str} 持多遇死叉，收盘{sig['confirm_close']:.2f}>"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳平多"
                )
            elif sig["below_ma40"] and allow_short:
                self.pending_action = "reverse_to_short"
                self.state_map.signal_desc = "持多死叉(MA40下)平多反手空"
                self.output(
                    f"{dt_str} 持多遇死叉，收盘{sig['confirm_close']:.2f}<"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳平多反手开空"
                )
            elif sig["below_ma40"]:
                self.pending_action = "close_long"
                self.state_map.signal_desc = "持多死叉(MA40下)只平多"
                self.output(
                    f"{dt_str} 持多遇死叉且收盘低于MA40，未开启做空，等待下一跳平多"
                )
            return

        if self.short_pos > 0 and sig["golden"]:
            if sig["below_ma40"]:
                self.pending_action = "close_short"
                self.state_map.signal_desc = "持空金叉(MA40下)只平空"
                self.output(
                    f"{dt_str} 持空遇金叉，收盘{sig['confirm_close']:.2f}<"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳平空"
                )
            elif sig["above_ma40"]:
                self.pending_action = "reverse_to_long"
                self.state_map.signal_desc = "持空金叉(MA40上)平空反手多"
                self.output(
                    f"{dt_str} 持空遇金叉，收盘{sig['confirm_close']:.2f}>"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳平空反手开多"
                )
            return

        if self.long_pos == 0 and self.short_pos == 0:
            if sig["golden"] and sig["above_ma40"]:
                self.pending_action = "open_long"
                self.state_map.signal_desc = "EXPMA金叉且MA40上方开多"
                self.output(
                    f"{dt_str} EXPMA金叉，收盘{sig['confirm_close']:.2f}>"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳开多 "
                    f"快={sig['confirm_fast']:.2f} 慢={sig['confirm_slow']:.2f}"
                )
            elif sig["death"] and sig["below_ma40"] and allow_short:
                self.pending_action = "open_short"
                self.state_map.signal_desc = "EXPMA死叉且MA40下方开空"
                self.output(
                    f"{dt_str} EXPMA死叉，收盘{sig['confirm_close']:.2f}<"
                    f"MA40({sig['confirm_ma40']:.2f})，等待下一跳开空 "
                    f"快={sig['confirm_fast']:.2f} 慢={sig['confirm_slow']:.2f}"
                )

    def _execute_pending_action(self) -> None:
        action = self.pending_action
        if not action:
            return

        if action == "open_long" and self.long_pos == 0 and self.short_pos == 0:
            self._open_long("EXPMA金叉开多")
            self.pending_action = ""
        elif action == "open_short" and self.long_pos == 0 and self.short_pos == 0:
            self._open_short("EXPMA死叉开空")
            self.pending_action = ""
        elif action == "close_long" and self.long_pos > 0:
            self._close_long("EXPMA死叉平多")
            self.pending_action = ""
        elif action == "close_short" and self.short_pos > 0:
            self._close_short("EXPMA金叉平空")
            self.pending_action = ""
        elif action == "reverse_to_short" and self.long_pos > 0:
            self._close_long("死叉平多准备反手空")
        elif action == "reverse_to_long" and self.short_pos > 0:
            self._close_short("金叉平空准备反手多")
        else:
            self.pending_action = ""

    def _schedule_reverse_open(self, direction: Literal["long", "short"]) -> None:
        def _open_after_delay() -> None:
            if not self.trading:
                return
            self._update_position()
            if self._has_pending_order():
                return
            if direction == "long" and self.long_pos == 0 and self.short_pos == 0:
                self._open_long("金叉反手开多")
                self.pending_action = ""
            elif direction == "short" and self.long_pos == 0 and self.short_pos == 0:
                if self.params_map.enable_short == 1:
                    self._open_short("死叉反手开空")
                self.pending_action = ""

        Timer(0.1, _open_after_delay).start()

    def _open_long(self, reason: str) -> None:
        if self.buy_orderid is not None:
            return
        price = self._opp_price_buy()
        if price is None:
            self.output(f"开多跳过：{reason}，卖一价无效")
            return

        self.buy_orderid = self.send_order(
            exchange=self.params_map.exchange,
            instrument_id=self.params_map.instrument_id,
            volume=self.params_map.lots,
            price=price,
            order_direction="buy",
        )
        if self.buy_orderid is not None:
            self.signal_price = price
            self.state_map.status = "开多委托中"
            self.output(f"{reason} 对手价{price:.2f} {self.params_map.lots}手")

    def _open_short(self, reason: str) -> None:
        if self.short_orderid is not None:
            return
        price = self._opp_price_sell()
        if price is None:
            self.output(f"开空跳过：{reason}，买一价无效")
            return

        self.short_orderid = self.send_order(
            exchange=self.params_map.exchange,
            instrument_id=self.params_map.instrument_id,
            volume=self.params_map.lots,
            price=price,
            order_direction="sell",
        )
        if self.short_orderid is not None:
            self.signal_price = -price
            self.state_map.status = "开空委托中"
            self.output(f"{reason} 对手价{price:.2f} {self.params_map.lots}手")

    def _close_long(self, reason: str) -> None:
        if self.sell_orderid is not None or self.long_pos <= 0:
            return
        price = self._opp_price_sell()
        if price is None:
            self.output(f"平多跳过：{reason}，买一价无效")
            return

        vol = self.long_pos
        self.sell_orderid = self.auto_close_position(
            exchange=self.params_map.exchange,
            instrument_id=self.params_map.instrument_id,
            volume=vol,
            price=price,
            order_direction="sell",
        )
        if self.sell_orderid is not None:
            self.signal_price = -price
            self.state_map.status = "平多委托中"
            self.output(f"{reason} 对手价{price:.2f} {vol}手")

    def _close_short(self, reason: str) -> None:
        if self.cover_orderid is not None or self.short_pos <= 0:
            return
        price = self._opp_price_buy()
        if price is None:
            self.output(f"平空跳过：{reason}，卖一价无效")
            return

        vol = self.short_pos
        self.cover_orderid = self.auto_close_position(
            exchange=self.params_map.exchange,
            instrument_id=self.params_map.instrument_id,
            volume=vol,
            price=price,
            order_direction="buy",
        )
        if self.cover_orderid is not None:
            self.signal_price = price
            self.state_map.status = "平空委托中"
            self.output(f"{reason} 对手价{price:.2f} {vol}手")

    def _opp_price_buy(self) -> float | None:
        if self.ask_price1 <= 0:
            return None
        return self._round_price(self.ask_price1 + self.params_map.over_price * self.price_tick)

    def _opp_price_sell(self) -> float | None:
        if self.bid_price1 <= 0:
            return None
        return self._round_price(self.bid_price1 - self.params_map.over_price * self.price_tick)

    def _round_price(self, price: float) -> float:
        tick = self.price_tick if self.price_tick > 0 else 1.0
        return round(round(price / tick) * tick, 10)

    def _has_pending_order(self) -> bool:
        return any(
            oid is not None
            for oid in (self.buy_orderid, self.sell_orderid, self.short_orderid, self.cover_orderid)
        )

    def _update_position(self) -> None:
        position = self.get_position(instrument_id=self.params_map.instrument_id)
        self.long_pos = position.long.close_available
        self.short_pos = position.short.close_available

        if self.long_pos > 0:
            self.state_map.status = f"持多{self.long_pos}手"
        elif self.short_pos > 0:
            self.state_map.status = f"持空{self.short_pos}手"
        elif not self._has_pending_order():
            self.state_map.status = "空仓"

    def _update_chart(self, kline: KLineData) -> None:
        if not self.widget:
            return
        self.widget.recv_kline({
            "kline": kline,
            "signal_price": self.signal_price,
            **self.main_indicator_data,
        })
        self.signal_price = 0.0

    def on_order_cancel(self, order: OrderData) -> None:
        super().on_order_cancel(order)
        if order.order_id == self.buy_orderid:
            self.buy_orderid = None
        elif order.order_id == self.sell_orderid:
            self.sell_orderid = None
            if self.pending_action in ("close_long", "reverse_to_short"):
                self.pending_action = ""
        elif order.order_id == self.short_orderid:
            self.short_orderid = None
        elif order.order_id == self.cover_orderid:
            self.cover_orderid = None
            if self.pending_action in ("close_short", "reverse_to_long"):
                self.pending_action = ""

    def on_order(self, order: OrderData) -> None:
        super().on_order(order)
        if order.status in ("全部成交", "已撤销", "部分撤销"):
            if order.order_id == self.buy_orderid:
                self.buy_orderid = None
            elif order.order_id == self.sell_orderid:
                self.sell_orderid = None
            elif order.order_id == self.short_orderid:
                self.short_orderid = None
            elif order.order_id == self.cover_orderid:
                self.cover_orderid = None

    def on_trade(self, trade: TradeData) -> None:
        super().on_trade(trade)

        direction_text = "买入" if trade.direction == "0" else "卖出"
        offset_text = "开仓" if trade.offset == "0" else ("平今" if trade.offset == "3" else "平仓")
        self.output(f"成交：{direction_text}{offset_text} {trade.volume}手 价格{trade.price}")

        if trade.offset in ("1", "3"):
            if self.pending_action == "reverse_to_short" and trade.direction == "1":
                self._schedule_reverse_open("short")
            elif self.pending_action == "reverse_to_long" and trade.direction == "0":
                self._schedule_reverse_open("long")
            elif self.pending_action == "close_long":
                self.pending_action = ""
            elif self.pending_action == "close_short":
                self.pending_action = ""

    def is_trading_time(self) -> bool:
        from datetime import time

        now = datetime.now().time()
        periods = [
            (time(21, 0), time(2, 30)),
            (time(9, 0), time(10, 15)),
            (time(10, 30), time(11, 30)),
            (time(13, 30), time(15, 0)),
        ]
        if periods[0][0] <= now or now < periods[0][1]:
            return True
        for start, end in periods[1:]:
            if start <= now < end:
                return True
        return False
