from __future__ import annotations

import logging
from datetime import datetime
from typing import Any

logger = logging.getLogger(__name__)


class ErpAiModelsUpdateMixin:
    def update_ai_agent_model(
        self,
        agent_code: str,
        model_name: str,
        provider: str = "openrouter",
        temperature: float = 0.20,
        max_tokens: int = 2048,
        system_prompt_custom: str | None = None,
        is_stealth_promo: bool | None = None,
        stealth_expires_at: str | None = None,
        fallback_model_name: str | None = None,
        fallback_provider: str | None = None,
    ) -> dict[str, Any]:
        """Cập nhật cấu hình LLM Model, Provider, System Prompt và vòng đời Stealth cho AI Agent."""
        self._ensure_stealth_columns()
        is_stealth = (
            is_stealth_promo
            if is_stealth_promo is not None
            else model_name.startswith("stealth/")
        )

        # Nếu là stealth model và chưa có fallback, chỉ định fallback phù hợp
        if not fallback_model_name:
            if agent_code == "thao":
                fallback_model_name = "deepseek/deepseek-r1:free"
            elif agent_code == "quynh":
                fallback_model_name = "qwen/qwen-2.5-coder-32b-instruct:free"
            elif agent_code == "phuc":
                fallback_model_name = "qwen2.5-7b-instruct"
            else:
                fallback_model_name = "google/gemini-2.0-flash-exp:free"

        if not fallback_provider:
            fallback_provider = (
                "lmstudio" if agent_code in ["phuc", "hung"] else "openrouter"
            )

        expires_at_val: datetime | None = None
        if stealth_expires_at:
            if isinstance(stealth_expires_at, datetime):
                expires_at_val = stealth_expires_at
            else:
                try:
                    expires_at_val = datetime.fromisoformat(
                        stealth_expires_at.replace("Z", "+00:00")
                    )
                except Exception:
                    expires_at_val = None
        elif is_stealth:
            from datetime import timedelta, timezone

            expires_at_val = datetime.now(timezone.utc) + timedelta(days=30)

        sql = """
            UPDATE erp_ai_agent_budgets
            SET 
                model_name = %(model_name)s,
                provider = %(provider)s,
                temperature = %(temperature)s,
                max_tokens = %(max_tokens)s,
                system_prompt_custom = %(system_prompt_custom)s,
                is_stealth_promo = %(is_stealth_promo)s,
                stealth_expires_at = %(stealth_expires_at)s,
                fallback_model_name = %(fallback_model_name)s,
                fallback_provider = %(fallback_provider)s,
                updated_at = NOW()
            WHERE agent_code = %(agent_code)s
            RETURNING *;
        """
        params = {
            "agent_code": agent_code,
            "model_name": model_name.strip(),
            "provider": provider.strip().lower(),
            "temperature": temperature,
            "max_tokens": max_tokens,
            "system_prompt_custom": system_prompt_custom,
            "is_stealth_promo": is_stealth,
            "stealth_expires_at": expires_at_val,
            "fallback_model_name": fallback_model_name,
            "fallback_provider": fallback_provider,
        }

        with self.get_connection() as conn:
            with conn.cursor() as cur:
                cur.execute(sql, params)
                row = cur.fetchone()
                conn.commit()
                if not row:
                    raise ValueError(f"Không tìm thấy AI Agent với mã: {agent_code}")
                return {
                    "agent_code": row["agent_code"],
                    "display_name": row["display_name"],
                    "role_title": row["role_title"],
                    "avatar_icon": row["avatar_icon"],
                    "model_name": str(row["model_name"]),
                    "provider": str(row["provider"]),
                    "temperature": float(row["temperature"]),
                    "max_tokens": int(row["max_tokens"]),
                    "system_prompt_custom": row["system_prompt_custom"],
                    "is_active": bool(row["is_active"]),
                    "updated_at": row["updated_at"].isoformat()
                    if row.get("updated_at")
                    else None,
                    "is_stealth_promo": bool(row.get("is_stealth_promo")),
                    "stealth_expires_at": row["stealth_expires_at"].isoformat()
                    if row.get("stealth_expires_at")
                    else None,
                    "fallback_model_name": str(
                        row.get("fallback_model_name", fallback_model_name)
                    ),
                    "fallback_provider": str(
                        row.get("fallback_provider", fallback_provider)
                    ),
                }

    def check_and_fallback_expired_stealth_models(self) -> list[dict[str, Any]]:
        """Quét và tự động chuyển đổi toàn bộ AI Agents đang dùng Stealth đã hết hạn về Fallback model."""
        self._ensure_stealth_columns()
        sql_find = """
            SELECT agent_code, display_name, model_name, provider, fallback_model_name, fallback_provider, stealth_expires_at
            FROM erp_ai_agent_budgets
            WHERE is_stealth_promo = true 
              AND stealth_expires_at IS NOT NULL 
              AND stealth_expires_at <= NOW();
        """

        sql_update = """
            UPDATE erp_ai_agent_budgets
            SET 
                model_name = fallback_model_name,
                provider = fallback_provider,
                is_stealth_promo = false,
                updated_at = NOW()
            WHERE agent_code = %s;
        """
        reverted: list[dict[str, Any]] = []
        try:
            with self.get_connection() as conn:
                with conn.cursor() as cur:
                    cur.execute(sql_find)
                    expired_agents = cur.fetchall()
                    for ag in expired_agents:
                        cur.execute(sql_update, (ag["agent_code"],))
                        reverted.append(
                            {
                                "agent_code": ag["agent_code"],
                                "display_name": ag["display_name"],
                                "old_model": ag["model_name"],
                                "new_model": ag["fallback_model_name"],
                                "new_provider": ag["fallback_provider"],
                                "expired_at": ag["stealth_expires_at"].isoformat()
                                if ag.get("stealth_expires_at")
                                else None,
                            }
                        )
                    if reverted:
                        conn.commit()
                        logger.info(
                            "Đã tự động chuyển đổi %d agents từ Stealth hết hạn về Fallback model.",
                            len(reverted),
                        )
        except Exception as exc:
            logger.debug("Lỗi khi kiểm tra hết hạn stealth models: %s", exc)

        return reverted
