import pymysql
import datetime
import pytz
import os

class MySQLHandler:
    def __init__(self, config):
        # 修正配置路径：mysql 应该在 pm 下面
        self.db_conf = config.get("pm", {}).get("mysql", {})
        self.enabled = bool(self.db_conf.get("host"))
        if self.enabled:
            self.init_db()

    def get_conn(self):
        return pymysql.connect(
            host=self.db_conf.get("host"),
            port=int(self.db_conf.get("port", 3306)),
            user=self.db_conf.get("user"),
            password=self.db_conf.get("password"),
            database=self.db_conf.get("database"),
            charset='utf8mb4',
            cursorclass=pymysql.cursors.DictCursor
        )

    def init_db(self):
        try:
            conn = pymysql.connect(
                host=self.db_conf.get("host"),
                port=int(self.db_conf.get("port", 3306)),
                user=self.db_conf.get("user"),
                password=self.db_conf.get("password"),
                charset='utf8mb4'
            )
            with conn.cursor() as cursor:
                cursor.execute(f"CREATE DATABASE IF NOT EXISTS `{self.db_conf.get('database')}` CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci")
                conn.select_db(self.db_conf.get("database"))
                
                create_table_sql = """
                CREATE TABLE IF NOT EXISTS `tg_manual_reply` (
                    `id` INT AUTO_INCREMENT PRIMARY KEY,
                    `server_code` VARCHAR(50),
                    `self_name` VARCHAR(100),
                    `tg_account` VARCHAR(50),
                    `peer_id` VARCHAR(50),
                    `username` VARCHAR(100),
                    `first_name` VARCHAR(100),
                    `self_avatar` VARCHAR(255),
                    `avatar_file` VARCHAR(255),
                    `last_messages` TEXT,
                    `status` TINYINT DEFAULT 0,
                    `reply_content` TEXT,
                    `create_time` DATETIME,
                    INDEX (`server_code`),
                    INDEX (`tg_account`),
                    INDEX (`status`)
                ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
                """
                cursor.execute(create_table_sql)
                
                # 检查并增加新字段 (防止旧表存在时报错)
                cursor.execute("SHOW COLUMNS FROM `tg_manual_reply` LIKE 'server_code'")
                if not cursor.fetchone():
                    cursor.execute("ALTER TABLE `tg_manual_reply` ADD COLUMN `server_code` VARCHAR(50) AFTER `id`")
                    cursor.execute("ALTER TABLE `tg_manual_reply` ADD COLUMN `self_name` VARCHAR(100) AFTER `server_code`")
                    cursor.execute("ALTER TABLE `tg_manual_reply` ADD INDEX (`server_code`)")
                
                cursor.execute("SHOW COLUMNS FROM `tg_manual_reply` LIKE 'self_avatar'")
                if not cursor.fetchone():
                    cursor.execute("ALTER TABLE `tg_manual_reply` ADD COLUMN `self_avatar` VARCHAR(255) AFTER `first_name`")
                
                print(f"[MySQL] ✅ 数据库表 `tg_manual_reply` 初始化成功")
            conn.commit()
            conn.close()
        except Exception as e:
            print(f"[MySQL] 初始化失败: {e}")

    def get_china_time(self):
        tz = pytz.timezone('Asia/Shanghai')
        return datetime.datetime.now(tz).strftime('%Y-%m-%d %H:%M:%S')

    def report_unread(self, server_code, self_name, tg_account, peer_id, username, first_name, self_avatar, avatar_file, messages):
        if not self.enabled: return
        try:
            conn = self.get_conn()
            with conn.cursor() as cursor:
                # 检查是否已存在该账号该用户的待处理记录
                check_sql = "SELECT id FROM tg_manual_reply WHERE tg_account=%s AND peer_id=%s AND status IN (0, 1)"
                cursor.execute(check_sql, (tg_account, peer_id))
                result = cursor.fetchone()
                
                # 截取消息长度，防止超过数据库限制
                processed_msgs = []
                for msg in messages:
                    processed_msgs.append(msg[:200])
                
                full_text = "\n".join(processed_msgs)
                if len(full_text) > 2000:
                    full_text = full_text[:2000] + "...(内容过长已截断)"
                
                if result:
                    update_sql = "UPDATE tg_manual_reply SET server_code=%s, self_name=%s, self_avatar=%s, avatar_file=%s, last_messages=%s, create_time=%s WHERE id=%s"
                    cursor.execute(update_sql, (server_code, self_name, self_avatar, avatar_file, full_text, self.get_china_time(), result['id']))
                else:
                    insert_sql = """
                    INSERT INTO tg_manual_reply (server_code, self_name, tg_account, peer_id, username, first_name, self_avatar, avatar_file, last_messages, create_time, status)
                    VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, 0)
                    """
                    cursor.execute(insert_sql, (server_code, self_name, tg_account, peer_id, username, first_name, self_avatar, avatar_file, full_text, self.get_china_time()))
            conn.commit()
            conn.close()
        except Exception as e:
            print(f"[MySQL] 上报未读失败: {e}")

    def check_replies(self, tg_account):
        if not self.enabled: return []
        try:
            conn = self.get_conn()
            replies = []
            with conn.cursor() as cursor:
                sql = "SELECT id, peer_id, reply_content FROM tg_manual_reply WHERE tg_account=%s AND status=2"
                cursor.execute(sql, (tg_account,))
                replies = cursor.fetchall()
            conn.close()
            return replies
        except Exception as e:
            print(f"[MySQL] 检查回复失败: {e}")
            return []

    def delete_record(self, record_id):
        if not self.enabled: return
        try:
            conn = self.get_conn()
            with conn.cursor() as cursor:
                sql = "DELETE FROM tg_manual_reply WHERE id=%s"
                cursor.execute(sql, (record_id,))
            conn.commit()
            conn.close()
        except Exception as e:
            print(f"[MySQL] 删除记录失败: {e}")
