from __future__ import annotations
from typing import Dict, Any, Optional
from pymongo import MongoClient
from bson import ObjectId
from ..logger import log


class MongoRepository:
    def __init__(self, uri: str, db_name: str) -> None:
        """Initialize MongoDB client."""
        self._client = MongoClient(uri)
        self._db = self._client[db_name]

    def collection(self, collection_name: str):
        """Get a specific collection."""
        return self._db[collection_name]

    def get_document(
        self, collection_name: str, document_id: str
    ) -> Optional[Dict[str, Any]]:
        """Retrieve a document by its ObjectId."""
        col = self.collection(collection_name)
        try:
            _id = ObjectId(document_id)
        except Exception:
            log.error("Invalid ObjectId: %s", document_id)
            return None

        doc = col.find_one({"_id": _id})
        if not doc:
            return None

        # Normalize fields
        return {
            "_id": str(doc["_id"]),
            "knowledge_base_id": str(doc.get("knowledge_base_id", "")),
            "name": doc.get("name"),
            "type": doc.get("type"),
            "content": doc.get("content", ""),
        }

    def get_agent_id_by_call_id(
        self, collection_name: str, call_id: str
    ) -> Optional[str]:
        """Retrieve agent_id using a call_id."""
        col = self.collection(collection_name)
        try:
            call = col.find_one({"call_id": call_id}, {"agent_id": 1})
            if call and "agent_id" in call:
                return str(call["agent_id"])
            return None
        except Exception as e:
            log.error("Error fetching agent_id for call_id %s: %s", call_id, e)
            return None

    def get_agent_knowledge_base_ids(self, user_id: str) -> list[str]:
        try:
            doc = self.collection("knowledgebases").find_one(
                {"user_id": user_id},
                {"_id": 1}  # only fetch main id
            )

            if not doc:
                return []

            return [str(doc["_id"])]

        except Exception as e:
            log.error(
                "Error fetching knowledgebase _id for user_id %s: %s",
                user_id,
                str(e)
            )
            return []

    def get_agent_system_prompt(self, agent_id: str) -> str:
        try:
            agent = self.collection("agents").find_one(
                {"_id": ObjectId(agent_id)},
                {"chat_prompt": 1},
            )
            if not agent:
                return ""  # No agent found, return empty string

            kb = agent.get("chat_prompt")
            # Ensure we always return a single string
            if isinstance(kb, list) and kb:
                return str(kb[0])  # Take first element if it's a list
            elif isinstance(kb, str):
                return kb
            else:
                return ""  # kb is None or empty

        except Exception as e:
            log.error("Error fetching chat_prompt for agent %s: %s", agent_id, e)
            return ""

    def get_product_by_id(self, product_id: str):
        # replace with your DB name
        collection = self.collection("shopifyproducts")

        # Fetch product by product_id
        product = collection.find_one({"product_id": product_id}, {"_id": 0})

        if not product:
            return None

        # If price is array, take only first item
        # if isinstance(product.get("prices"), list) and product["prices"]:
        #     product["price"] = product["prices"][0]

        # Drop unwanted fields if they exist
        for field in [
            "__v",
            "createdAt",
            "updatedAt",
            "store_id",
            # "prices",
        ]:
            product.pop(field, None)

        return product

    def get_response_product_by_id(self, product_id: str):
        # replace with your DB name
        collection = self.collection("shopifyproducts")

        # Fetch product by product_id
        product = collection.find_one({"product_id": product_id}, {"_id": 0})

        if not product:
            return None

        # Add hardcoded currency field
        product["currency"] = "INR"

        # If price is array, take only first item
        # if isinstance(product.get("prices"), list) and product["prices"]:
        #     product["price"] = product["prices"][0]

        # Drop unwanted fields if they exist
        for field in [
            "__v",
            "createdAt",
            "updatedAt",
            "store_id",
            "options",
            "tags",
            # "prices",
        ]:
            product.pop(field, None)

        return product

    def get_all_products(self, store_id: str):
        # replace with your DB name
        collection = self.collection("shopifyproducts")

        # Fetch all products, exclude MongoDB internal _id field
        products = list(collection.find({"store_id": store_id}, {"_id": 0}))

        print(f"Total products fetched: {len(products)}")
        return products
