from google.cloud import firestore
from app.core.config import settings
import logging
import os
from datetime import datetime

logger = logging.getLogger(__name__)

# Ensure credentials are set for the environment
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = settings.GOOGLE_APPLICATION_CREDENTIALS

db = firestore.Client(
    project=settings.FIRESTORE_PROJECT,
    database=settings.FIRESTORE_DATABASE
)

def get_chat_history(tenant_id: str, thread_id: str = "default", user_id: str = None):
    """
    Fetches chat history for a given tenant, thread, and user isolation.
    """
    try:
        doc_id = f"{tenant_id}_{thread_id}"
        doc_ref = db.collection("chat_history").document(doc_id)
        doc = doc_ref.get()
        if doc.exists:
            data = doc.to_dict()
            # Enforce user isolation if user_id is provided
            if user_id and data.get("user_id") and data.get("user_id") != user_id:
                logger.warning(f"Unauthorized access attempt for thread {thread_id} by user {user_id}")
                return []
            return data.get("messages", [])
        return []
    except Exception as e:
        logger.error(f"Error fetching chat history from Firestore: {e}", exc_info=True)
        return []

def save_chat_message(tenant_id: str, role: str, message: str, thread_id: str = "default", user_name: str = None, summary: str = None, user_id: str = None):
    """
    Saves a new chat message to Firestore with optional user isolation.
    """
    try:
        doc_id = f"{tenant_id}_{thread_id}"
        doc_ref = db.collection("chat_history").document(doc_id)
        
        new_msg = {
            "role": role, 
            "content": message,
            "timestamp": datetime.utcnow().isoformat()
        }
        
        now = datetime.utcnow().isoformat()
        update_data = {
            "messages": firestore.ArrayUnion([new_msg]),
            "updated_at": now,
            "tenant_id": tenant_id,
            "thread_id": thread_id
        }

        try:
            if not doc_ref.get().exists:
                update_data["created_at"] = now
        except Exception as read_err:
            logger.warning(f"Could not check existing chat document before save: {read_err}")
        
        if user_id:
            update_data["user_id"] = user_id
        if user_name:
            update_data["user_name"] = user_name
        if summary:
            update_data["summary"] = summary
            
        doc_ref.set(update_data, merge=True)
        
        logger.info(f"Saved {role} message to Firestore for {tenant_id} (User: {user_id}, Thread: {thread_id})")
        
        # Increment global counter if it's an AI response
        if role == "assistant":
            increment_ai_response_count(tenant_id)
    except Exception as e:
        logger.error(f"Error saving chat message to Firestore: {e}", exc_info=True)

def increment_ai_response_count(tenant_id: str):
    """
    Increments the global AI response counter AND a daily bucket for a tenant.
    """
    try:
        now = datetime.utcnow()
        date_str = now.strftime("%Y-%m-%d")
        
        # 1. Update Global Counter
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        stats_ref.set({"total_ai_responses": firestore.Increment(1)}, merge=True)
        
        # 2. Update Daily Bucket
        daily_ref = stats_ref.collection("daily_stats").document(date_str)
        daily_ref.set({
            "ai_responses": firestore.Increment(1),
            "date": date_str,
            "updated_at": now.isoformat()
        }, merge=True)
        
        logger.info(f"Incremented AI response count (Global & Daily: {date_str}) for {tenant_id}")
    except Exception as e:
        logger.error(f"Error incrementing AI response count for {tenant_id}: {e}")

def get_aggregated_ai_responses(tenant_id: str, start_date: str = None, end_date: str = None):
    """
    Sums AI responses for a tenant within a date range using daily buckets.
    """
    try:
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        query = stats_ref.collection("daily_stats")
        
        if start_date:
            query = query.where("date", ">=", start_date)
        if end_date:
            query = query.where("date", "<=", end_date)
            
        docs = query.stream()
        total = 0
        for doc in docs:
            total += doc.to_dict().get("ai_responses", 0)
        return total
    except Exception as e:
        logger.error(f"Error aggregating AI responses for {tenant_id}: {e}")
        raise

def get_ai_response_message_count(tenant_id: str, start_date=None, end_date=None):
    """
    Counts assistant messages directly from chat_history.
    This is the authoritative source for enforcing AI response limits.
    """
    try:
        start_iso = start_date.isoformat() if start_date else None
        end_iso = end_date.isoformat() if end_date else None
        threads_ref = db.collection("chat_history").where("tenant_id", "==", tenant_id)
        docs = threads_ref.stream()

        count = 0
        for doc in docs:
            data = doc.to_dict() or {}
            for msg in data.get("messages", []):
                if msg.get("role") != "assistant":
                    continue
                timestamp = msg.get("timestamp")
                if start_iso and (not timestamp or timestamp < start_iso):
                    continue
                if end_iso and timestamp and timestamp > end_iso:
                    continue
                count += 1
        return count
    except Exception as e:
        logger.error(f"Error counting AI response messages for {tenant_id}: {e}", exc_info=True)
        raise

def get_conversation_thread_count(tenant_id: str, start_date=None):
    """
    Counts tenant conversation threads from Firestore, the same store used to detect new chat threads.
    """
    try:
        start_iso = start_date.isoformat() if start_date else None
        threads_ref = db.collection("chat_history").where("tenant_id", "==", tenant_id)
        docs = threads_ref.stream()

        count = 0
        for doc in docs:
            data = doc.to_dict() or {}
            created_at = data.get("created_at") or data.get("updated_at")
            if start_iso and (not created_at or created_at < start_iso):
                continue
            count += 1
        return count
    except Exception as e:
        logger.error(f"Error counting conversation threads for {tenant_id}: {e}", exc_info=True)
        raise

def log_storage_snapshot(tenant_id: str, size_bytes: int):
    """
    Logs the current storage size into the daily bucket for historical tracking.
    """
    try:
        now = datetime.utcnow()
        date_str = now.strftime("%Y-%m-%d")
        
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        daily_ref = stats_ref.collection("daily_stats").document(date_str)
        
        daily_ref.set({
            "storage_bytes": size_bytes,
            "date": date_str,
            "updated_at": now.isoformat()
        }, merge=True)
        
        logger.info(f"Logged storage snapshot for {tenant_id} ({date_str}): {size_bytes} bytes")
    except Exception as e:
        logger.error(f"Error logging storage snapshot for {tenant_id}: {e}")

def get_latest_storage_snapshot(tenant_id: str, end_date: str = None):
    """
    Retrieves the latest available storage snapshot before or on the end_date.
    """
    try:
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        query = stats_ref.collection("daily_stats").order_by("date", direction=firestore.Query.DESCENDING)
        
        if end_date:
            query = query.where("date", "<=", end_date)
            
        docs = query.limit(1).stream()
        for doc in docs:
            return doc.to_dict().get("storage_bytes", 0)
        return 0
    except Exception as e:
        logger.error(f"Error getting latest storage snapshot for {tenant_id}: {e}")
        return 0

def get_total_ai_responses_count(tenant_id: str):
    """
    Retrieves the total AI response count from the global counter.
    Triggers a sync if the counter document doesn't exist.
    """
    try:
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        doc = stats_ref.get()
        if doc.exists:
            count = doc.to_dict().get("total_ai_responses", 0)
            return count
        else:
            # Trigger one-time sync if metadata document is missing
            return sync_ai_responses_from_history(tenant_id)
    except Exception as e:
        logger.error(f"Error getting total AI response count for {tenant_id}: {e}")
        return 0

def sync_ai_responses_from_history(tenant_id: str):
    """
    Scans all historical threads to calculate and initialize the global counter.
    This handles "past ones" by calculating the sum once and saving it.
    """
    try:
        logger.info(f"Initializing/Syncing AI response count from history for {tenant_id}...")
        threads_ref = db.collection("chat_history").where("tenant_id", "==", tenant_id)
        docs = threads_ref.stream()
        
        total_count = 0
        for doc in docs:
            messages = doc.to_dict().get("messages", [])
            total_count += sum(1 for m in messages if m.get("role") == "assistant")
            
        # Save the initial total
        stats_ref = db.collection("tenant_stats").document(tenant_id)
        stats_ref.set({"total_ai_responses": total_count}, merge=True)
        logger.info(f"Sync complete for {tenant_id}: {total_count} responses found.")
        return total_count
    except Exception as e:
        logger.error(f"Error syncing AI responses for {tenant_id}: {e}")
        return 0

def get_all_threads(tenant_id: str, user_id: str = None):
    """
    Retrieves all chat threads for a specific tenant, optionally filtered by user_id.
    """
    try:
        threads_ref = db.collection("chat_history").where("tenant_id", "==", tenant_id)
        if user_id:
            threads_ref = threads_ref.where("user_id", "==", user_id)
            
        docs = threads_ref.stream()
        
        threads = []
        for doc in docs:
            data = doc.to_dict()
            thread_id_val = data.get("thread_id") or data.get("session_id")
            if not thread_id_val:
                continue
                
            last_msg = data.get("messages", [])[-1] if data.get("messages") else None
            threads.append({
                "thread_id": thread_id_val,
                "user_id": data.get("user_id"),
                "user_name": data.get("user_name", "Anonymous"),
                "summary": data.get("summary", ""),
                "takeover_active": data.get("takeover_active", False),
                "intervention_requested": data.get("intervention_requested", False),
                "created_at": data.get("created_at"),
                "updated_at": data.get("updated_at"),
                "last_message_preview": last_msg["content"][:100] + "..." if last_msg else ""
            })
        
        # Sort by updated_at descending
        threads.sort(key=lambda x: x["updated_at"] or "", reverse=True)
        return threads
    except Exception as e:
        logger.error(f"Error listing threads from Firestore: {e}", exc_info=True)
        return []

def set_takeover_status(tenant_id: str, thread_id: str, status: bool, user_id: str = None):
    """
    Toggles the human takeover status for a specific thread.
    """
    try:
        doc_id = f"{tenant_id}_{thread_id}"
        doc_ref = db.collection("chat_history").document(doc_id)
        
        update_data = {
            "takeover_active": status,
            "updated_at": datetime.utcnow().isoformat()
        }
        # Clear intervention request if a human actually takes over
        if status is True:
            update_data["intervention_requested"] = False
            
        doc_ref.set(update_data, merge=True)
        logger.info(f"Updated takeover_active to {status} for {doc_id}")
        return True
    except Exception as e:
        logger.error(f"Error setting takeover status in Firestore: {e}", exc_info=True)
        return False

def get_takeover_status(tenant_id: str, thread_id: str, user_id: str = None):
    """
    Retrieves the current takeover status for a thread.
    """
    try:
        doc_id = f"{tenant_id}_{thread_id}"
        doc_ref = db.collection("chat_history").document(doc_id)
        doc = doc_ref.get()
        if doc.exists:
            return doc.to_dict().get("takeover_active", False)
        return False
    except Exception as e:
        logger.error(f"Error getting takeover status from Firestore: {e}", exc_info=True)
        return False

def set_intervention_status(tenant_id: str, thread_id: str, status: bool, user_id: str = None):
    """
    Toggles the intervention_requested status for a specific thread.
    """
    try:
        doc_id = f"{tenant_id}_{thread_id}"
        doc_ref = db.collection("chat_history").document(doc_id)
        
        update_data = {
            "tenant_id": tenant_id,
            "thread_id": thread_id,
            "intervention_requested": status,
            "updated_at": datetime.utcnow().isoformat()
        }
        if user_id:
            update_data["user_id"] = user_id
            
        doc_ref.set(update_data, merge=True)
        logger.info(f"Updated intervention_requested to {status} for {doc_id}")
        return True
    except Exception as e:
        logger.error(f"Error setting intervention status in Firestore: {e}", exc_info=True)
        return False
