import logging
import json
import subprocess
import os
import shutil
import asyncio
import sys
from app.services.infra.database import get_db_connection, upsert_thread_analytics, insert_llm_usage, bootstrap_tenant
from app.services.integrations.firestore import (
    get_chat_history, 
    get_all_threads, 
    get_total_ai_responses_count,
    get_aggregated_ai_responses,
    get_latest_storage_snapshot
)
from app.services.infra.quotas import check_quota
from app.services.infra.embedding import client as openai_client
from app.core.config import settings
from app.core.prompts import ANALYTICS_ANALYSIS_PROMPT
from psycopg2 import sql

logger = logging.getLogger(__name__)

from datetime import datetime, timedelta

def format_storage_size(size_bytes: int):
    """
    Dynamically adapts units from GB down to Bytes for readability.
    """
    if size_bytes >= 1024**3:
        return f"{round(size_bytes / (1024**3), 1)} GB"
    elif size_bytes >= 1024**2:
        return f"{round(size_bytes / (1024**2), 1)} MB"
    elif size_bytes >= 1024:
        return f"{round(size_bytes / 1024, 1)} KB"
    else:
        return f"{size_bytes} Bytes"

def get_time_bounds(time_range: str, start_date_str: str = None, end_date_str: str = None):
    """
    Returns (current_start_sql, current_end_sql, prev_start_sql, prev_end_sql, bucket_size, label_format, iso_bounds)
    iso_bounds is a dict: {cs_iso, ce_iso, ps_iso, pe_iso}
    """
    now = datetime.utcnow()
    
    if start_date_str:
        try:
            cs_dt = datetime.strptime(start_date_str, "%Y-%m-%d")
            if end_date_str:
                ce_dt = datetime.strptime(end_date_str, "%Y-%m-%d") + timedelta(days=1, seconds=-1)
            else:
                ce_dt = now
            # If start_date is provided, we force time_range to custom so it doesn't get overwritten
            time_range = "custom"
        except Exception as e:
            logger.error(f"Error parsing custom dates: {e}. Falling back to 30d.")
            time_range = "30d"
    
    if time_range != "custom":
        if time_range == "1h":
            cs_dt, ce_dt = now - timedelta(hours=1), now
        elif time_range == "24h":
            cs_dt, ce_dt = now - timedelta(hours=24), now
        elif time_range == "7d":
            cs_dt, ce_dt = now - timedelta(days=7), now
        elif time_range == "30d":
            cs_dt, ce_dt = now - timedelta(days=30), now
        elif time_range == "90d":
            cs_dt, ce_dt = now - timedelta(days=90), now
        else: # all
            cs_dt, ce_dt = datetime(1970, 1, 1), now

    duration = ce_dt - cs_dt
    ps_dt = cs_dt - duration
    pe_dt = cs_dt
    
    # SQL format
    cs = f"TIMESTAMP '{cs_dt.strftime('%Y-%m-%d %H:%M:%S')}'"
    ce = f"TIMESTAMP '{ce_dt.strftime('%Y-%m-%d %H:%M:%S')}'"
    ps = f"TIMESTAMP '{ps_dt.strftime('%Y-%m-%d %H:%M:%S')}'"
    pe = f"TIMESTAMP '{pe_dt.strftime('%Y-%m-%d %H:%M:%S')}'"
    
    # ISO dates for Firestore/Buckets
    iso_bounds = {
        "cs": cs_dt.strftime("%Y-%m-%d"),
        "ce": ce_dt.strftime("%Y-%m-%d"),
        "ps": ps_dt.strftime("%Y-%m-%d"),
        "pe": pe_dt.strftime("%Y-%m-%d")
    }
    
    # Bucketing logic
    days = duration.days
    if days < 1:
        bucket, fmt = "minute" if time_range == "1h" else "hour", "HH24:MI" if time_range == "1h" else "HH24:00"
    elif days <= 2:
        bucket, fmt = "hour", "HH24:00"
    elif days <= 60:
        bucket, fmt = "day", "MM-DD"
    else:
        bucket, fmt = "month", "Mon YYYY"
        
    return cs, ce, ps, pe, bucket, fmt, iso_bounds

async def analyze_thread(tenant_id: str, thread_id: str):
    """
    Fetches history for a thread and performs a single batch AI analysis.
    """
    logger.info(f"Analyzing thread: {thread_id} for tenant: {tenant_id}")
    
    # Quota Enforcement: AI Tokens
    check_quota(tenant_id, "ai_tokens")
    
    history = get_chat_history(tenant_id, thread_id)
    
    if not history:
        logger.warning(f"No history found for thread {thread_id}")
        return False

    conversation_text = ""
    for msg in history:
        role = msg.get("role", "unknown")
        content = msg.get("content", "")
        conversation_text += f"{role.upper()}: {content}\n"

    analysis_prompt = ANALYTICS_ANALYSIS_PROMPT.format(conversation_text=conversation_text)

    try:
        response = openai_client.chat.completions.create(
            model=settings.LLM_MODEL,
            messages=[{"role": "user", "content": analysis_prompt}],
            max_completion_tokens=600,
            response_format={"type": "json_object"}
        )
        analysis_data = json.loads(response.choices[0].message.content)
        analysis_data['thread_id'] = thread_id
        
        # Log Token Usage
        usage = response.usage
        insert_llm_usage(
            tenant_id, 
            "Thread Analysis", 
            settings.LLM_MODEL, 
            usage.prompt_tokens, 
            usage.completion_tokens, 
            usage.total_tokens,
            {"thread_id": thread_id}
        )
        
        # Save to intermediary analysis table
        upsert_thread_analytics(tenant_id, analysis_data)
        logger.info(f"Successfully analyzed and saved thread {thread_id}")
        return True
    except Exception as e:
        logger.error(f"Failed to analyze thread {thread_id} for tenant {tenant_id}: {e}", exc_info=True)
        return False

async def run_dbt_for_tenant(tenant_id: str):
    """
    Executes dbt run for a specific tenant using non-blocking asyncio subprocess.
    """
    logger.info(f"Executing dbt run for tenant: {tenant_id}")
    project_root = os.getcwd()
    dbt_project_dir = os.path.join(project_root, "analytics")
    
    # Prepare environment variables
    env = os.environ.copy()
    env["DB_HOST"] = settings.DB_HOST
    env["DB_PORT"] = str(settings.DB_PORT)
    env["DB_USER"] = settings.DB_USER
    env["DB_PASSWORD"] = settings.DB_PASSWORD
    env["DB_NAME"] = settings.DB_NAME
    env["TENANT_ID"] = tenant_id

    try:
        configured_dbt = getattr(settings, "DBT_EXECUTABLE", "")
        dbt_candidates = [
            configured_dbt,
            os.path.join(project_root, "venv_dbt/bin/dbt"),
            os.path.join(project_root, "venv/bin/dbt"),
            shutil.which("dbt"),
        ]
        dbt_script = next((path for path in dbt_candidates if path and os.path.exists(path)), None)
        if not dbt_script:
            logger.error(
                "dbt executable not found. Set DBT_EXECUTABLE or install dbt in venv/bin/dbt."
            )
            return False
        
        process = await asyncio.create_subprocess_exec(
            dbt_script, "run", "--profiles-dir", ".",
            cwd=dbt_project_dir,
            env=env,
            stdout=asyncio.subprocess.PIPE,
            stderr=asyncio.subprocess.PIPE
        )
        
        stdout, stderr = await process.communicate()
        
        if process.returncode != 0:
            logger.error(f"dbt run failed. Return code: {process.returncode}")
            if stdout:
                for line in stdout.decode().splitlines():
                    logger.error(f"dbt STDOUT: {line}")
            if stderr:
                for line in stderr.decode().splitlines():
                    logger.error(f"dbt STDERR: {line}")
            return False
        
        if stdout:
            for line in stdout.decode().splitlines():
                logger.info(f"dbt STDOUT: {line}")
        logger.info("dbt run completed successfully")
        return True
    except Exception as e:
        logger.error(f"Failed to execute dbt: {e}", exc_info=True)
        return False

async def get_kpis(tenant_id: str):
    """
    Orchestrates the fresh analysis and returns KPIs from dbt Mart tables.
    """
    logger.info(f"Refreshing analytics discovery for tenant: {tenant_id}")
    
    # 0. Self-healing Schema Update (Ensure new columns exist)
    try:
        bootstrap_tenant(tenant_id)
    except Exception as e:
        logger.warning(f"Self-healing bootstrap failed for {tenant_id} (ignoring if schema is already correct): {e}")

    # 1. Thread Discovery from Firestore
    all_threads = get_all_threads(tenant_id)
    thread_ids = [t['thread_id'] for t in all_threads]

    # 2. Check which threads need analysis
    conn = get_db_connection()
    unanalyzed = []
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL("""
                SELECT thread_id FROM {}.strategist_thread_analytics
            """).format(sql.Identifier(tenant_id)))
            analyzed_ids = {row[0] for row in cur.fetchall()}
            unanalyzed = [tid for tid in thread_ids if tid not in analyzed_ids]
    except Exception as e:
        logger.warning(f"Error checking analyzed threads (likely table doesn't exist yet): {e}")
        unanalyzed = thread_ids # Process all if we can't check

    # 3. Perform Batch AI Analysis
    if unanalyzed:
        logger.info(f"Found {len(unanalyzed)} unanalyzed threads. Processing...")
        for tid in unanalyzed:
            await analyze_thread(tenant_id, tid)

    # 4. Trigger dbt Transformation
    dbt_success = await run_dbt_for_tenant(tenant_id)
    if not dbt_success:
        logger.error(f"dbt transformation failed for tenant {tenant_id}. KPIs might be stale.")

    # 5. Fetch KPIs from dbt Marts
    try:
        with conn.cursor() as cur:
            # Query all Marts and join/combine
            final_metrics = {}
            
            # General Metrics
            cur.execute(sql.SQL("SELECT * FROM {}.mart_general_chat_metrics").format(sql.Identifier(tenant_id)))
            row = cur.fetchone()
            if row:
                desc = [d[0] for d in cur.description]
                final_metrics["general"] = dict(zip(desc, row))

            # Intent & Demand (Aggregated)
            cur.execute(sql.SQL("SELECT * FROM {}.mart_customer_intent_demand").format(sql.Identifier(tenant_id)))
            intent_rows = cur.fetchall()
            if intent_rows:
                desc = [d[0] for d in cur.description]
                final_metrics["intent_and_demand"] = [dict(zip(desc, r)) for r in intent_rows]

            # Conversion
            cur.execute(sql.SQL("SELECT * FROM {}.mart_conversion_enablement").format(sql.Identifier(tenant_id)))
            row = cur.fetchone()
            if row:
                desc = [d[0] for d in cur.description]
                final_metrics["conversion"] = dict(zip(desc, row))

            # CX
            cur.execute(sql.SQL("SELECT * FROM {}.mart_customer_experience").format(sql.Identifier(tenant_id)))
            row = cur.fetchone()
            if row:
                desc = [d[0] for d in cur.description]
                final_metrics["customer_experience"] = dict(zip(desc, row))

            # VOC
            cur.execute(sql.SQL("SELECT * FROM {}.mart_voc_insights").format(sql.Identifier(tenant_id)))
            row = cur.fetchone()
            if row:
                desc = [d[0] for d in cur.description]
                final_metrics["voice_of_customer"] = dict(zip(desc, row))

            # Helpdesk
            cur.execute(sql.SQL("SELECT * FROM {}.mart_helpdesk_ticketing").format(sql.Identifier(tenant_id)))
            row = cur.fetchone()
            if row:
                desc = [d[0] for d in cur.description]
                final_metrics["helpdesk"] = dict(zip(desc, row))

            # Resource Usage (Aggregated Tokens & Storage)
            try:
                cur.execute(sql.SQL("SELECT * FROM {}.mart_resource_usage").format(sql.Identifier(tenant_id)))
                row = cur.fetchone()
                if row:
                    desc = [d[0] for d in cur.description]
                    final_metrics["resource_usage"] = dict(zip(desc, row))
            except Exception as re:
                logger.warning(f"Could not fetch resource usage for {tenant_id}: {re}")

            # 6. Total AI Responses (Real-time from Firestore)
            total_ai_responses = get_total_ai_responses_count(tenant_id)
            final_metrics["total_ai_responses"] = total_ai_responses

            return {
                "status": "success",
                "tenant_id": tenant_id,
                "metrics": final_metrics
            }
    except Exception as e:
        logger.error(f"Failed to fetch KPIs from Marts: {e}", exc_info=True)
        # Fallback to empty metrics if dbt hasn't run yet
        return {"status": "partial_success", "tenant_id": tenant_id, "metrics": {}}
    finally:
        conn.close()
async def get_general_kpis(tenant_id: str, time_range: str = "30d", start_date: str = None, end_date: str = None):
    """
    Returns a simplified set of KPIs specifically for UI dashboards, filtered by time range (relative or absolute).
    Ensures self-healing migration for old tenants.
    """
    logger.info(f"Fetching general UI KPIs for tenant: {tenant_id}, range: {time_range}, custom: {start_date} to {end_date}")
    
    # Self-healing migration for old tenants
    bootstrap_tenant(tenant_id)
    
    conn = get_db_connection()
    try:
        cur_start, cur_end, prev_start, prev_end, bucket, label_fmt, iso = get_time_bounds(time_range, start_date, end_date)
        
        with conn.cursor() as cur:
            # 1. Fetch Top KPIs using dynamic SQL on int_chat_sessions
            kpi_query = sql.SQL("""
                WITH metrics AS (
                    SELECT
                        count(distinct case when session_start >= {cs} and session_start <= {ce} then thread_id end) as current_conversations,
                        count(distinct case when session_start >= {ps} and session_start <= {pe} then thread_id end) as prev_conversations,
                        
                        -- Using explicit fallback for avg_response_time_seconds if missing
                        CASE WHEN EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = {cs_schema} AND table_name = 'int_chat_sessions' AND column_name = 'avg_response_time_seconds') 
                             THEN avg(case when session_start >= {cs} and session_start <= {ce} then avg_response_time_seconds end) 
                             ELSE 0 END as current_avg_response_time,
                        CASE WHEN EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = {cs_schema} AND table_name = 'int_chat_sessions' AND column_name = 'avg_response_time_seconds') 
                             THEN avg(case when session_start >= {ps} and session_start <= {pe} then avg_response_time_seconds end) 
                             ELSE 0 END as prev_avg_response_time,

                        (sum(case when session_start >= {cs} and session_start <= {ce} then is_resolved::int else 0 end)::float / 
                            nullif(count(case when session_start >= {cs} and session_start <= {ce} then thread_id end), 0) * 100) as current_resolution_rate,
                        (sum(case when session_start >= {ps} and session_start <= {pe} then is_resolved::int else 0 end)::float / 
                            nullif(count(case when session_start >= {ps} and session_start <= {pe} then thread_id end), 0) * 100) as prev_resolution_rate,

                        avg(case when session_start >= {cs} and session_start <= {ce} then sentiment_score end) as current_satisfaction,
                        avg(case when session_start >= {ps} and session_start <= {pe} then sentiment_score end) as prev_satisfaction
                    FROM {schema}.int_chat_sessions
                )
                SELECT * FROM metrics
            """).format(
                schema=sql.Identifier(tenant_id),
                cs=sql.SQL(cur_start),
                ce=sql.SQL(cur_end),
                ps=sql.SQL(prev_start),
                pe=sql.SQL(prev_end),
                cs_schema=sql.Literal(tenant_id)
            )
            
            try:
                cur.execute(kpi_query)
                row = cur.fetchone()
                if not row:
                    return None
                desc = [d[0] for d in cur.description]
                metrics = dict(zip(desc, row))
            except Exception as e:
                logger.warning(f"Failed to query int_chat_sessions directly for {tenant_id}: {e}. Trying raw table fallback...")
                conn.rollback()
                
                # --- RAW TABLE FALLBACK (Stage 2) ---
                # This calculates metrics directly from strategist_feedback and strategist_thread_analytics
                raw_query = sql.SQL("""
                    WITH session_raw AS (
                        SELECT 
                            thread_id,
                            MIN(created_at) as session_start
                        FROM {schema}.strategist_feedback
                        GROUP BY thread_id
                    ),
                    metrics_raw AS (
                        SELECT
                            s.thread_id,
                            s.session_start,
                            t.is_resolved,
                            t.sentiment_score
                        FROM session_raw s
                        LEFT JOIN {schema}.strategist_thread_analytics t ON s.thread_id = t.thread_id
                    ),
                    aggregated AS (
                        SELECT
                            count(distinct case when session_start >= {cs} then thread_id end) as current_conversations,
                            count(distinct case when session_start >= {ps} and session_start < {cs} then thread_id end) as prev_conversations,
                            
                            (sum(case when session_start >= {cs} then is_resolved::int else 0 end)::float / 
                                nullif(count(case when session_start >= {cs} then thread_id end), 0) * 100) as current_resolution_rate,
                            (sum(case when session_start >= {ps} and session_start < {cs} then is_resolved::int else 0 end)::float / 
                                nullif(count(case when session_start >= {ps} and session_start < {cs} then thread_id end), 0) * 100) as prev_resolution_rate,

                            avg(case when session_start >= {cs} then sentiment_score end) as current_satisfaction,
                            avg(case when session_start >= {ps} and session_start < {cs} then sentiment_score end) as prev_satisfaction
                        FROM aggregated_metrics
                    )
                    SELECT * FROM aggregated
                """)
                # Note: For simplicity in the fallback, we use 0 or default values for complex metrics 
                # like avg_response_time if not easily derivable from raw feedback pairs.
                
                # Re-executing a simplified but robust version of the above
                fallback_query = sql.SQL("""
                    SELECT 
                        (SELECT count(distinct thread_id) FROM {schema}.strategist_feedback WHERE created_at >= {cs} AND created_at <= {ce}) as current_conversations,
                        (SELECT count(distinct thread_id) FROM {schema}.strategist_feedback WHERE created_at >= {ps} AND created_at <= {pe}) as prev_conversations,
                        0 as current_avg_response_time, 
                        0 as prev_avg_response_time,
                        (SELECT (sum(is_resolved::int)::float / nullif(count(*), 0) * 100) FROM {schema}.strategist_thread_analytics WHERE created_at >= {cs} AND created_at <= {ce}) as current_resolution_rate,
                        (SELECT (sum(is_resolved::int)::float / nullif(count(*), 0) * 100) FROM {schema}.strategist_thread_analytics WHERE created_at >= {ps} AND created_at <= {pe}) as prev_resolution_rate,
                        (SELECT avg(sentiment_score) FROM {schema}.strategist_thread_analytics WHERE created_at >= {cs} AND created_at <= {ce}) as current_satisfaction,
                        (SELECT avg(sentiment_score) FROM {schema}.strategist_thread_analytics WHERE created_at >= {ps} AND created_at <= {pe}) as prev_satisfaction
                """).format(
                    schema=sql.Identifier(tenant_id),
                    cs=sql.SQL(cur_start),
                    ce=sql.SQL(cur_end),
                    ps=sql.SQL(prev_start),
                    pe=sql.SQL(prev_end)
                )
                
                try:
                    cur.execute(fallback_query)
                    row = cur.fetchone()
                    if not row: return None
                    desc = [d[0] for d in cur.description]
                    metrics = dict(zip(desc, row))
                except Exception as e2:
                    logger.error(f"Critical failure: Both view and raw fallback failed for {tenant_id}: {e2}")
                    return None
            
            # Extract and format Period-over-Period (PoP) metrics
            def calc_trend(current, prev):
                current = float(current or 0)
                prev = float(prev or 0)
                if prev == 0:
                    return current, "N/A", "flat"
                diff = current - prev
                pct = (diff / prev) * 100
                direction = "up" if diff > 0 else ("down" if diff < 0 else "flat")
                sign = "+" if pct > 0 else ""
                return current, f"{sign}{pct:.2f}%", direction

            cur_conv, tr_conv, dir_conv = calc_trend(metrics.get("current_conversations"), metrics.get("prev_conversations"))
            cur_resp, tr_resp, dir_resp = calc_trend(metrics.get("current_avg_response_time"), metrics.get("prev_avg_response_time"))
            cur_res, tr_res, dir_res = calc_trend(metrics.get("current_resolution_rate"), metrics.get("prev_resolution_rate"))
            
            # --- 2. Advanced Metrics (AI Tokens & Responses) ---
            # AI Tokens: Summing the DELTAS logged for this period to ensure absolute accuracy
            # Safety check: Table existence
            cur.execute("""
                SELECT EXISTS (
                    SELECT FROM information_schema.tables 
                    WHERE table_schema = %s AND table_name = 'strategist_llm_usage'
                )
            """, (tenant_id,))
            
            cur_tokens, prev_tokens = 0, 0
            if cur.fetchone()[0]:
                token_query = sql.SQL("""
                    SELECT 
                        COALESCE(SUM(CASE WHEN created_at >= {cs} AND created_at <= {ce} THEN total_tokens ELSE 0 END), 0) as cur_tokens,
                        COALESCE(SUM(CASE WHEN created_at >= {ps} AND created_at <= {pe} THEN total_tokens ELSE 0 END), 0) as prev_tokens
                    FROM {schema}.strategist_llm_usage
                """).format(
                    schema=sql.Identifier(tenant_id),
                    cs=sql.SQL(cur_start),
                    ce=sql.SQL(cur_end),
                    ps=sql.SQL(prev_start),
                    pe=sql.SQL(prev_end)
                )
                cur.execute(token_query)
                trow = cur.fetchone()
                if trow:
                    cur_tokens, prev_tokens = trow
            
            _, tr_tokens, dir_tokens = calc_trend(cur_tokens, prev_tokens)

            # AI Responses: Using Firestore Buckets for accuracy as requested
            cur_ai_resp_count = get_aggregated_ai_responses(tenant_id, iso["cs"], iso["ce"])
            prev_ai_resp_count = get_aggregated_ai_responses(tenant_id, iso["ps"], iso["pe"])
            cur_ai_resp, tr_ai_resp, dir_ai_resp = calc_trend(cur_ai_resp_count, prev_ai_resp_count)

            # Storage Capacity: Using latest snapshots from Firestore daily metrics
            cur_storage_bytes = get_latest_storage_snapshot(tenant_id, iso["ce"])
            prev_storage_bytes = get_latest_storage_snapshot(tenant_id, iso["pe"])
            
            # Fallback for storage: if no historical snapshot exists, use current relation size for Today
            if cur_storage_bytes == 0:
                from app.services.infra.database import get_vector_db_size
                cur_storage_bytes = get_vector_db_size(tenant_id)
            
            # Trends are calculated on raw bytes to ensure precision
            _, tr_storage, dir_storage = calc_trend(cur_storage_bytes, prev_storage_bytes)
            
            cur_storage_display = format_storage_size(cur_storage_bytes)
            
            # Sentiment mapped to 5-star
            cur_sat = float(metrics.get("current_satisfaction") or 0)
            prev_sat = float(metrics.get("prev_satisfaction") or 0)
            star_cur = round(((cur_sat + 1) / 2) * 4 + 1, 1) if metrics.get("current_satisfaction") is not None else 0.0
            star_prev = round(((prev_sat + 1) / 2) * 4 + 1, 1) if metrics.get("prev_satisfaction") is not None else 0.0
            _, tr_sat, dir_sat = calc_trend(star_cur, star_prev)

            top_kpis = {
                "total_conversations": {
                    "value": int(cur_conv),
                    "trend": tr_conv,
                    "trend_direction": dir_conv
                },
                "avg_response_time": {
                    "value": f"{round(cur_resp, 1)}s",
                    "trend": tr_resp,
                    "trend_direction": dir_resp
                },
                "resolution_rate": {
                    "value": f"{round(cur_res, 1)}%",
                    "trend": tr_res,
                    "trend_direction": dir_res
                },
                "user_satisfaction": {
                    "value": f"{star_cur}/5",
                    "trend": tr_sat,
                    "trend_direction": dir_sat
                },
                "total_ai_responses": {
                    "value": int(cur_ai_resp),
                    "trend": tr_ai_resp,
                    "trend_direction": dir_ai_resp
                },
                "ai_tokens": {
                    "value": int(cur_tokens),
                    "trend": tr_tokens,
                    "trend_direction": dir_tokens
                },
                "storage_capacity": {
                    "value": cur_storage_display,
                    "trend": tr_storage or "Latest",
                    "trend_direction": dir_storage
                }
            }
            
            # --- 2. Charting with Adaptive Bucketing ---
            # Total Users Chart
            labels, users_data, resp_time_data = [], [], []
            chart_from_raw_feedback = False
            try:
                chart_query = sql.SQL("""
                    SELECT 
                        to_char(date_trunc({bucket}, session_start), {fmt}) as label,
                        count(distinct thread_id) as users,
                        avg(avg_response_time_seconds) as resp_time
                    FROM {schema}.int_chat_sessions
                    WHERE session_start >= {cs} AND session_start <= {ce}
                    GROUP BY 1, date_trunc({bucket}, session_start)
                    ORDER BY date_trunc({bucket}, session_start) ASC
                """).format(
                    schema=sql.Identifier(tenant_id),
                    bucket=sql.Literal(bucket),
                    fmt=sql.Literal(label_fmt),
                    cs=sql.SQL(cur_start),
                    ce=sql.SQL(cur_end)
                )
                
                cur.execute(chart_query)
                chart_rows = cur.fetchall()
            except Exception as e:
                logger.warning(f"Chart query failed for {tenant_id}: {e}. Trying raw feedback fallback.")
                conn.rollback()
                chart_rows = []

            if not chart_rows:
                try:
                    raw_chart_query = sql.SQL("""
                        SELECT
                            to_char(date_trunc({bucket}, created_at), {fmt}) as label,
                            count(distinct thread_id) as users,
                            0 as resp_time,
                            date_trunc({bucket}, created_at) as bucket_start
                        FROM {schema}.strategist_feedback
                        WHERE created_at >= {cs} AND created_at <= {ce}
                        GROUP BY 1, 4
                        ORDER BY bucket_start ASC
                    """).format(
                        schema=sql.Identifier(tenant_id),
                        bucket=sql.Literal(bucket),
                        fmt=sql.Literal(label_fmt),
                        cs=sql.SQL(cur_start),
                        ce=sql.SQL(cur_end)
                    )
                    cur.execute(raw_chart_query)
                    chart_rows = cur.fetchall()
                    chart_from_raw_feedback = True
                except Exception as e:
                    logger.warning(f"Raw feedback chart fallback failed for {tenant_id}: {e}. Returning empty chart data.")
                    conn.rollback()

            if chart_rows:
                labels = [r[0] for r in chart_rows]
                users_data = [r[1] for r in chart_rows]
                resp_time_data = [round(r[2] or 0, 1) for r in chart_rows]

            chart_total_users = {
                "labels": labels,
                "datasets": [{"label": "Users", "data": users_data}]
            }
            
            chart_response_time = {
                "labels": labels,
                "datasets": [{"label": "Avg Response Time (s)", "data": resp_time_data}]
            }

            # --- 3. Conversation Channels (Real Proportions) ---
            # We filter by current range if applicable
            clabels, ccounts = [], []
            channels_data = []
            try:
                channel_query = sql.SQL("""
                    SELECT 'Website Chat' as channel, 
                        count(distinct thread_id) as conversations
                    FROM {schema}.int_chat_sessions
                    WHERE session_start >= {cs} AND session_start <= {ce}
                    GROUP BY 1
                """).format(
                    schema=sql.Identifier(tenant_id),
                    cs=sql.SQL(cur_start),
                    ce=sql.SQL(cur_end)
                )
                cur.execute(channel_query)
                channels_data = cur.fetchall()
            except Exception as e:
                logger.warning(f"Channel query failed for {tenant_id}: {e}. Trying raw feedback fallback.")
                conn.rollback()

            if chart_from_raw_feedback or not channels_data or not any(c[1] for c in channels_data):
                try:
                    raw_channel_query = sql.SQL("""
                        SELECT 'Website Chat' as channel,
                            count(distinct thread_id) as conversations
                        FROM {schema}.strategist_feedback
                        WHERE created_at >= {cs} AND created_at <= {ce}
                        GROUP BY 1
                    """).format(
                        schema=sql.Identifier(tenant_id),
                        cs=sql.SQL(cur_start),
                        ce=sql.SQL(cur_end)
                    )
                    cur.execute(raw_channel_query)
                    channels_data = cur.fetchall()
                except Exception as e:
                    logger.warning(f"Raw feedback channel fallback failed for {tenant_id}: {e}. Returning empty data.")
                    conn.rollback()

            if channels_data:
                clabels = [c[0] for c in channels_data]
                ccounts = [c[1] for c in channels_data]
            
            total_c = sum(ccounts) or 1
            cpercent = [round((cnt/total_c)*100, 1) for cnt in ccounts]

            chart_channels = {
                "labels": clabels,
                "datasets": [{
                    "data": cpercent,
                    "raw_counts": ccounts
                }]
            }

            return {
                "top_kpis": top_kpis,
                "charts": {
                    "total_users": chart_total_users,
                    "conversation_channels": chart_channels,
                    "response_time": chart_response_time
                }
            }
    except Exception as e:
        logger.error(f"Failed to fetch general KPIs for tenant {tenant_id}: {e}", exc_info=True)
        return None
    finally:
        conn.close()
