# source: https://raw.githubusercontent.com/xiedidan/freqtrade/1fe7d439dc8d3b6801a4c19ac4984f32ed418f58/user_data/strategies/hourly_breakout_signal.py
import logging
import pandas as pd
import numpy as np
from datetime import datetime, timezone
import traceback
from typing import Dict, List, Optional, Tuple
from sqlalchemy import Column, String, Float, Integer, Boolean, DateTime, create_engine, func
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker, scoped_session
import talib.abstract as ta

from freqtrade.strategy import IStrategy
from freqtrade.strategy.interface import IStrategy
from freqtrade.persistence import Trade
from freqtrade.exchange import timeframe_to_minutes
from freqtrade.enums import SignalType

logger = logging.getLogger(__name__)

# 定义数据库模型
Base = declarative_base()

class SignalHistory(Base):
    """信号历史记录"""
    __tablename__ = 'hourly_breakout_signal_history'
    
    id = Column(Integer, primary_key=True)
    pair = Column(String, nullable=False)
    signal_type = Column(String, nullable=False)  # 'long' 或 'short'
    timeframe = Column(String, nullable=False)
    timestamp = Column(DateTime, nullable=False)
    price = Column(Float, nullable=False)
    high_level = Column(Float, nullable=True)  # 前1H K线高点
    low_level = Column(Float, nullable=True)   # 前1H K线低点
    
    @staticmethod
    def add_signal(session, pair, signal_type, timeframe, price, high_level=None, low_level=None):
        """添加信号记录"""
        signal = SignalHistory(
            pair=pair,
            signal_type=signal_type,
            timeframe=timeframe,
            timestamp=datetime.now(timezone.utc),
            price=price,
            high_level=high_level,
            low_level=low_level
        )
        session.add(signal)
        session.commit()
        return signal
    
    @staticmethod
    def get_signals(session, pair=None, limit=100):
        """获取信号记录"""
        query = session.query(SignalHistory)
        if pair:
            query = query.filter(SignalHistory.pair == pair)
        return query.order_by(SignalHistory.timestamp.desc()).limit(limit).all()


class Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123(IStrategy):
    """
    1H周期突破信号策略
    
    入场信号：
    1. 如果1H K线突破前1H K线高点，并且5Min K线收盘在前1H K线高点之上，则发出做多信号；
    2. 如果1H K线突破前1H K线低点，并且5Min K线收盘在前1H K线高点之下，则发出做空信号。
    
    监控币种范围：
    交易量前30的币种，只考虑与USDT的交易对
    
    配置选项:
    - timeframe: 策略使用的时间周期 (默认: 5m)
    - higher_timeframe: 用于突破判断的高级时间周期 (默认: 1h)
    - enable_debug_logs: 是否启用详细的调试日志 (默认: True)
    - top_volume_pairs: 监控交易量前多少的币种 (默认: 30)
    - auto_update_pairs: 是否自动更新监控的交易对列表 (默认: True)
    - pairs_update_interval: 更新交易对列表的间隔时间（秒）(默认: 3600，即1小时)
    """
    # 策略配置
    timeframe = '5m'  # 使用5分钟时间周期
    higher_timeframe = '1h'  # 用于突破判断的高级时间周期
    process_only_new_candles = True  # 只处理新的K线
    stoploss = -0.10  # 止损设置（必需）
    
    # 调试配置
    enable_debug_logs = True  # 是否启用详细的调试日志
    
    # 交易对配置
    top_volume_pairs = 30  # 监控交易量前多少的币种
    auto_update_pairs = True  # 是否自动更新监控的交易对列表
    pairs_update_interval = 3600  # 更新交易对列表的间隔时间（秒）
    
    # 数据库配置
    db_initialized = False
    db_url = None
    
    # 存储监控的交易对列表
    monitored_pairs = set()
    # 上次更新监控交易对列表的时间
    last_pairs_update_time = None
    
    # 存储每个交易对的1H K线数据
    hourly_candles = {}
    
    def __init__(self, config: dict) -> None:
        """
        初始化策略
        """
        super().__init__(config)
        
        # 初始化数据库连接
        if 'db_url' in config:
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_url = config['db_url']
        
        # 初始化数据库会话
        if not Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_initialized:
            self.init_db_session()
        
        # 更新监控的交易对列表
        self.update_monitored_pairs()
    
    @staticmethod
    def init_db_session():
        """初始化数据库会话"""
        try:
            if Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_url is None:
                Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_url = 'sqlite:///user_data/hourly_breakout_signal.sqlite'
            
            engine = create_engine(Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_url)
            Base.metadata.create_all(engine)
            session_factory = sessionmaker(bind=engine)
            Session = scoped_session(session_factory)
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.session = Session
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_initialized = True
            logger.info(f"数据库会话初始化成功，使用URL: {Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_url}")
        except Exception as e:
            logger.error(f"初始化数据库会话时出错: {e}")
            logger.debug(traceback.format_exc())
    
    def update_monitored_pairs(self) -> bool:
        """
        更新监控的交易对列表，获取交易量前30的USDT交易对
        
        :return: 是否成功更新
        """
        try:
            # 获取当前时间
            current_time = datetime.now()
            
            # 检查是否需要更新
            if not self.auto_update_pairs:
                return False
                
            if Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.last_pairs_update_time is not None:
                elapsed_seconds = (current_time - Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.last_pairs_update_time).total_seconds()
                if elapsed_seconds < self.pairs_update_interval:
                    return False
            
            # 确保数据提供者可用
            if not self.dp:
                logger.warning("数据提供者不可用，无法更新交易对列表")
                return False
            
            # 获取所有USDT交易对
            all_pairs = []
            if hasattr(self.dp, 'available_pairs') and callable(getattr(self.dp, 'available_pairs')):
                all_pairs = [p for p in self.dp.available_pairs if p.endswith('/USDT')]
            else:
                logger.warning("无法获取可用交易对列表")
                return False
            
            # 获取交易量数据
            volume_data = []
            for pair in all_pairs:
                try:
                    # 获取1天的K线数据来计算交易量
                    ohlcv = self.dp.ohlcv(pair, '1d', limit=1)
                    if ohlcv is not None and len(ohlcv) > 0:
                        # 计算交易量（以USDT计）
                        volume_usdt = ohlcv[0]['volume'] * ohlcv[0]['close']
                        volume_data.append((pair, volume_usdt))
                except Exception as e:
                    logger.debug(f"获取{pair}交易量数据时出错: {e}")
            
            # 按交易量排序并取前30名
            volume_data.sort(key=lambda x: x[1], reverse=True)
            top_pairs = [p[0] for p in volume_data[:self.top_volume_pairs]]
            
            # 更新监控的交易对列表
            old_pairs = Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs.copy()
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs = set(top_pairs)
            
            # 更新时间戳
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.last_pairs_update_time = current_time
            
            # 检查是否有变化
            added_pairs = Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs - old_pairs
            removed_pairs = old_pairs - Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs
            
            if added_pairs or removed_pairs or not old_pairs:
                logger.info(f"===== 监控的交易对列表已更新 =====")
                logger.info(f"当前监控的交易对: {len(Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs)} 个")
                for pair in sorted(Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs):
                    logger.info(f"  - {pair}")
                
                if added_pairs:
                    logger.info(f"新增的交易对: {len(added_pairs)} 个")
                    for pair in sorted(added_pairs):
                        logger.info(f"  + {pair}")
                
                if removed_pairs:
                    logger.info(f"移除的交易对: {len(removed_pairs)} 个")
                    for pair in sorted(removed_pairs):
                        logger.info(f"  - {pair}")
            
            return True
        except Exception as e:
            logger.error(f"更新监控的交易对列表时出错: {e}")
            logger.debug(traceback.format_exc())
            return False
    
    def fetch_hourly_candles(self, pair: str) -> Dict:
        """
        获取指定交易对的1小时K线数据
        
        :param pair: 交易对
        :return: 包含当前和前一个1小时K线数据的字典
        """
        try:
            # 获取1小时K线数据
            hourly_df = self.dp.get_pair_dataframe(pair=pair, timeframe=self.higher_timeframe)
            
            if hourly_df is None or len(hourly_df) < 2:
                logger.warning(f"无法获取{pair}的{self.higher_timeframe}K线数据，或数据不足")
                return {}
            
            # 获取当前和前一个1小时K线
            current_candle = hourly_df.iloc[-1].to_dict()
            previous_candle = hourly_df.iloc[-2].to_dict()
            
            return {
                'current': current_candle,
                'previous': previous_candle,
                'timestamp': datetime.now(timezone.utc).timestamp()
            }
        except Exception as e:
            logger.error(f"获取{pair}的{self.higher_timeframe}K线数据时出错: {e}")
            logger.debug(traceback.format_exc())
            return {}
    
    def populate_indicators(self, dataframe: pd.DataFrame, metadata: dict) -> pd.DataFrame:
        """
        计算技术指标
        
        :param dataframe: 输入数据
        :param metadata: 元数据
        :return: 带有技术指标的数据
        """
        # 更新监控的交易对列表
        self.update_monitored_pairs()
        
        # 获取当前交易对
        current_pair = metadata['pair']
        
        # 如果当前交易对不在监控列表中，则跳过
        if current_pair not in Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs:
            return dataframe
        
        try:
            # 获取1小时K线数据
            hourly_data = self.fetch_hourly_candles(current_pair)
            if not hourly_data:
                return dataframe
            
            # 存储1小时K线数据
            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.hourly_candles[current_pair] = hourly_data
            
            # 获取前一个1小时K线的高点和低点
            previous_high = hourly_data['previous']['high']
            previous_low = hourly_data['previous']['low']
            
            # 添加到数据中
            dataframe['prev_1h_high'] = previous_high
            dataframe['prev_1h_low'] = previous_low
            
            # 判断当前1小时K线是否突破前一个1小时K线的高点或低点
            current_high = hourly_data['current']['high']
            current_low = hourly_data['current']['low']
            
            dataframe['break_high'] = current_high > previous_high
            dataframe['break_low'] = current_low < previous_low
            
            # 判断5分钟K线收盘价是否在前一个1小时K线的高点之上或低点之下
            dataframe['close_above_prev_high'] = dataframe['close'] > previous_high
            dataframe['close_below_prev_low'] = dataframe['close'] < previous_low
            
            if self.enable_debug_logs:
                # 获取最新的5分钟K线
                if len(dataframe) > 0:
                    last_5m = dataframe.iloc[-1]
                    logger.info(f"===== {current_pair} 突破检测 =====")
                    logger.info(f"前1H高点: {previous_high}, 前1H低点: {previous_low}")
                    logger.info(f"当前1H高点: {current_high}, 当前1H低点: {current_low}")
                    logger.info(f"当前5M收盘价: {last_5m['close']}")
                    logger.info(f"突破高点: {current_high > previous_high}, 突破低点: {current_low < previous_low}")
                    logger.info(f"5M收盘价>前1H高点: {last_5m['close'] > previous_high}, 5M收盘价<前1H低点: {last_5m['close'] < previous_low}")
        except Exception as e:
            logger.error(f"计算{current_pair}的指标时出错: {e}")
            logger.debug(traceback.format_exc())
        
        return dataframe
    
    def populate_buy_trend(self, dataframe: pd.DataFrame, metadata: dict) -> pd.DataFrame:
        """
        生成做多信号
        
        :param dataframe: 输入数据
        :param metadata: 元数据
        :return: 带有做多信号的数据
        """
        current_pair = metadata['pair']
        
        # 如果当前交易对不在监控列表中，则跳过
        if current_pair not in Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs:
            return dataframe
        
        # 初始化信号列
        dataframe['buy'] = 0
        
        try:
            # 做多条件: 1H K线突破前1H K线高点，并且5Min K线收盘在前1H K线高点之上
            long_condition = (
                dataframe['break_high'] &
                dataframe['close_above_prev_high']
            )
            
            # 设置做多信号
            if long_condition.any():
                # 获取满足条件的最后一个索引
                last_index = long_condition[long_condition].index[-1]
                
                # 设置信号
                dataframe.loc[last_index, 'buy'] = 1
                
                # 记录信号
                if Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_initialized:
                    try:
                        last_row = dataframe.loc[last_index]
                        SignalHistory.add_signal(
                            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.session(),
                            current_pair,
                            'long',
                            self.timeframe,
                            last_row['close'],
                            last_row['prev_1h_high'],
                            last_row['prev_1h_low']
                        )
                        logger.info(f"LONG信号已记录: {current_pair} 价格: {last_row['close']}")
                    except Exception as e:
                        logger.error(f"记录做多信号时出错: {e}")
                
                logger.info(f"LONG信号已生成: {current_pair} 在前1H高点 {dataframe.loc[last_index, 'prev_1h_high']} 上方")
        except Exception as e:
            logger.error(f"生成{current_pair}的做多信号时出错: {e}")
            logger.debug(traceback.format_exc())
        
        return dataframe
    
    def populate_sell_trend(self, dataframe: pd.DataFrame, metadata: dict) -> pd.DataFrame:
        """
        生成做空信号
        
        :param dataframe: 输入数据
        :param metadata: 元数据
        :return: 带有做空信号的数据
        """
        current_pair = metadata['pair']
        
        # 如果当前交易对不在监控列表中，则跳过
        if current_pair not in Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.monitored_pairs:
            return dataframe
        
        # 初始化信号列
        dataframe['sell'] = 0
        
        try:
            # 做空条件: 1H K线突破前1H K线低点，并且5Min K线收盘在前1H K线低点之下
            short_condition = (
                dataframe['break_low'] &
                dataframe['close_below_prev_low']
            )
            
            # 设置做空信号
            if short_condition.any():
                # 获取满足条件的最后一个索引
                last_index = short_condition[short_condition].index[-1]
                
                # 设置信号
                dataframe.loc[last_index, 'sell'] = 1
                
                # 记录信号
                if Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.db_initialized:
                    try:
                        last_row = dataframe.loc[last_index]
                        SignalHistory.add_signal(
                            Github_xiedidan_freqtrade__hourly_breakout_signal__20250721_100123.session(),
                            current_pair,
                            'short',
                            self.timeframe,
                            last_row['close'],
                            last_row['prev_1h_high'],
                            last_row['prev_1h_low']
                        )
                        logger.info(f"SHORT信号已记录: {current_pair} 价格: {last_row['close']}")
                    except Exception as e:
                        logger.error(f"记录做空信号时出错: {e}")
                
                logger.info(f"SHORT信号已生成: {current_pair} 在前1H低点 {dataframe.loc[last_index, 'prev_1h_low']} 下方")
        except Exception as e:
            logger.error(f"生成{current_pair}的做空信号时出错: {e}")
            logger.debug(traceback.format_exc())
        
        return dataframe 