from __future__ import annotations

import logging
from typing import Any

logger = logging.getLogger(__name__)


class ErpAiBudgetMixin:
    def log_ai_agent_usage(
        self,
        agent_code: str,
        prompt_tokens: int,
        completion_tokens: int,
        cost_vnd: Any | None = None,
        task_type: str = "chat",
        user_email: str | None = None,
        query_preview: str | None = None,
    ) -> dict[str, Any]:
        """Ghi nhận vết sử dụng token và chi phí ngân sách cho nhân viên AI."""
        total_tokens = prompt_tokens + completion_tokens

        # Nếu chưa truyền chi phí, tính theo đơn giá định mức của agent
        if cost_vnd is None:
            cost_per_1k = 50.0  # 50 VNĐ / 1.000 tokens
            cost_vnd = round((total_tokens / 1000.0) * cost_per_1k, 4)

        sql = """
            INSERT INTO erp_ai_agent_usage_logs (
                agent_code, task_type, prompt_tokens, completion_tokens,
                total_tokens, cost_vnd, user_email, query_preview
            ) VALUES (
                %(agent_code)s, %(task_type)s, %(prompt_tokens)s, %(completion_tokens)s,
                %(total_tokens)s, %(cost_vnd)s, %(user_email)s, %(query_preview)s
            ) RETURNING *;
        """
        params = {
            "agent_code": agent_code,
            "task_type": task_type,
            "prompt_tokens": prompt_tokens,
            "completion_tokens": completion_tokens,
            "total_tokens": total_tokens,
            "cost_vnd": cost_vnd,
            "user_email": user_email,
            "query_preview": (query_preview[:250] + "...")
            if query_preview and len(query_preview) > 250
            else query_preview,
        }

        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, params)
            conn.commit()
            return cur.fetchone()

    def get_ai_agent_budget_summary(
        self, agent_code: str | None = None
    ) -> dict[str, Any]:
        """Tổng hợp ngân sách đã sử dụng theo ngày, tuần, tháng, năm cho nhân viên AI."""
        where_clause = "WHERE b.is_active = TRUE"
        params: list[Any] = []
        if agent_code:
            where_clause += " AND b.agent_code = %s"
            params.append(agent_code)

        sql = f"""
            SELECT 
                b.agent_code,
                b.display_name,
                b.role_title,
                b.avatar_icon,
                b.daily_budget_vnd,
                b.weekly_budget_vnd,
                b.monthly_budget_vnd,
                b.yearly_budget_vnd,
                b.cost_per_1k_tokens_vnd,
                b.model_name,
                b.provider,
                COALESCE(SUM(l.cost_vnd) FILTER (WHERE l.created_at >= date_trunc('day', NOW())), 0) AS day_spent_vnd,
                COALESCE(SUM(l.total_tokens) FILTER (WHERE l.created_at >= date_trunc('day', NOW())), 0) AS day_tokens,
                COALESCE(COUNT(l.id) FILTER (WHERE l.created_at >= date_trunc('day', NOW())), 0) AS day_requests,
                COALESCE(SUM(l.cost_vnd) FILTER (WHERE l.created_at >= date_trunc('week', NOW())), 0) AS week_spent_vnd,
                COALESCE(SUM(l.total_tokens) FILTER (WHERE l.created_at >= date_trunc('week', NOW())), 0) AS week_tokens,
                COALESCE(COUNT(l.id) FILTER (WHERE l.created_at >= date_trunc('week', NOW())), 0) AS week_requests,
                COALESCE(SUM(l.cost_vnd) FILTER (WHERE l.created_at >= date_trunc('month', NOW())), 0) AS month_spent_vnd,
                COALESCE(SUM(l.total_tokens) FILTER (WHERE l.created_at >= date_trunc('month', NOW())), 0) AS month_tokens,
                COALESCE(COUNT(l.id) FILTER (WHERE l.created_at >= date_trunc('month', NOW())), 0) AS month_requests,
                COALESCE(SUM(l.cost_vnd) FILTER (WHERE l.created_at >= date_trunc('year', NOW())), 0) AS year_spent_vnd,
                COALESCE(SUM(l.total_tokens) FILTER (WHERE l.created_at >= date_trunc('year', NOW())), 0) AS year_tokens,
                COALESCE(COUNT(l.id) FILTER (WHERE l.created_at >= date_trunc('year', NOW())), 0) AS year_requests
            FROM erp_ai_agent_budgets b
            LEFT JOIN erp_ai_agent_usage_logs l ON b.agent_code = l.agent_code
            {where_clause}
            GROUP BY b.agent_code, b.display_name, b.role_title, b.avatar_icon, b.daily_budget_vnd, b.weekly_budget_vnd, b.monthly_budget_vnd, b.yearly_budget_vnd, b.cost_per_1k_tokens_vnd, b.model_name, b.provider
            ORDER BY b.agent_code ASC;
        """

        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, params)
            rows = cur.fetchall()

        agents_data = []
        total_today = 0.0
        total_week = 0.0
        total_month = 0.0
        total_year = 0.0

        for r in rows:
            day_spent = float(r["day_spent_vnd"])
            day_budget = float(r["daily_budget_vnd"])
            day_pct = (
                round((day_spent / day_budget * 100.0), 1) if day_budget > 0 else 0.0
            )

            week_spent = float(r["week_spent_vnd"])
            week_budget = float(r["weekly_budget_vnd"])
            week_pct = (
                round((week_spent / week_budget * 100.0), 1) if week_budget > 0 else 0.0
            )

            month_spent = float(r["month_spent_vnd"])
            month_budget = float(r["monthly_budget_vnd"])
            month_pct = (
                round((month_spent / month_budget * 100.0), 1)
                if month_budget > 0
                else 0.0
            )

            year_spent = float(r["year_spent_vnd"])
            year_budget = float(r["yearly_budget_vnd"])
            year_pct = (
                round((year_spent / year_budget * 100.0), 1) if year_budget > 0 else 0.0
            )

            total_today += day_spent
            total_week += week_spent
            total_month += month_spent
            total_year += year_spent

            agents_data.append(
                {
                    "agent_code": r["agent_code"],
                    "display_name": r["display_name"],
                    "role_title": r["role_title"],
                    "avatar_icon": r["avatar_icon"],
                    "cost_per_1k_tokens_vnd": float(r["cost_per_1k_tokens_vnd"]),
                    "model_name": r["model_name"],
                    "provider": r["provider"],
                    "day": {
                        "spent_vnd": day_spent,
                        "budget_vnd": day_budget,
                        "percentage": day_pct,
                        "tokens": int(r["day_tokens"]),
                        "requests_count": int(r["day_requests"]),
                    },
                    "week": {
                        "spent_vnd": week_spent,
                        "budget_vnd": week_budget,
                        "percentage": week_pct,
                        "tokens": int(r["week_tokens"]),
                        "requests_count": int(r["week_requests"]),
                    },
                    "month": {
                        "spent_vnd": month_spent,
                        "budget_vnd": month_budget,
                        "percentage": month_pct,
                        "tokens": int(r["month_tokens"]),
                        "requests_count": int(r["month_requests"]),
                    },
                    "year": {
                        "spent_vnd": year_spent,
                        "budget_vnd": year_budget,
                        "percentage": year_pct,
                        "tokens": int(r["year_tokens"]),
                        "requests_count": int(r["year_requests"]),
                    },
                }
            )

        return {
            "agents": agents_data,
            "total_spent_today_vnd": round(total_today, 2),
            "total_spent_week_vnd": round(total_week, 2),
            "total_spent_month_vnd": round(total_month, 2),
            "total_spent_year_vnd": round(total_year, 2),
        }
