from typing import Any, Dict, Optional, Union, List
from bson import ObjectId
from pymongo.errors import PyMongoError
from langchain.schema import AIMessage
from typing import Optional
import time
from qdrant_client.http.models import NamedVector, SearchRequest
from qdrant_client import models
from qdrant_client.models import Filter, FieldCondition, MatchValue
import time
import ast
from langchain_core.runnables import RunnableConfig
from langchain_core.tools import tool
import random
import requests
import json

# Internal imports
from .shopify_tools import ShopifyClient
from .init_repo import agents_col, agent_domains_col, transcript_col, products_col
from ..llm_utils import LLM_UTILS
from ..repositories import repo
from ..config import Config
from ..logger import log


def get_store_id_from_agent(user_id: str) -> Optional[str]:
    """Return store_id for a given user_id, or None if not found/invalid."""
    if not user_id:
        log.warning("Empty user_id passed to get_store_id_from_agent")
        return None

    try:
        agent_doc = agents_col.find_one(
            {"user_id": ObjectId(user_id)}, {"store_id": 1, "_id": 0}
        )
    except PyMongoError as e:
        log.error("DB error fetching store_id for user_id=%s: %s", user_id, e)
        return None

    if not agent_doc:
        log.info("No store found for user_id=%s", user_id)
        return None

    store_id = agent_doc.get("store_id")
    if store_id is None:
        log.warning("Agent doc missing store_id for user_id=%s", user_id)
        return None

    return str(store_id)


def get_store_credentials(store_id: str) -> Optional[dict]:
    """Return Shopify credentials for a given store_id, or None if missing."""
    if not store_id:
        log.warning("Empty store_id passed to get_store_credentials")
        return None

    try:
        store_doc = agents_col.find_one(
            {"store_id": ObjectId(store_id)},
            {"metadata.access_token": 1, "metadata.domain": 1},
        )
    except PyMongoError as e:
        log.error("DB error fetching credentials for store_id=%s: %s", store_id, e)
        return None

    if not store_doc:
        log.info("No store found for store_id=%s", store_id)
        return None

    metadata = store_doc.get("metadata", {})
    token = metadata.get("access_token")
    domain = metadata.get("domain")

    if not token or not domain:
        log.warning("Incomplete credentials for store_id=%s", store_id)
        return None

    return {"access_token": token, "shop_domain": domain}


def save_message(role: str, text: Any, chat_id: str) -> None:
    """Persist chat messages into transcripts collection."""
    # if isinstance(text, AIMessage):
    #     text_to_store = text
    # elif isinstance(text, dict):
    #     text_to_store = text
    # else:
    #     text_to_store = text

    transcript_col.insert_one(
        {
            "role": role,
            "text": text,
            "call_id": ObjectId(chat_id),
            "created_at": time.time(),  # system time, avoids loop issues
        }
    )


def _hybrid_search(
    query_text: str, store_id: Optional[str] = None, top_k: int = 5
) -> List[Dict]:
    """Hybrid dense-vector search from Qdrant."""
    log.info(f"Hybrid search: '{query_text}', store_id:{store_id}")
    dense_vector = LLM_UTILS.embeddings.embed_query(query_text,output_dimensionality=768)
    log.info(f"Generated embedding vector | length={len(dense_vector)}")

    query_filter = (
        Filter(must=[FieldCondition(key="store_id", match=MatchValue(value=store_id))])
        if store_id
        else None
    )
    request = SearchRequest(
        vector=NamedVector(name="dense", vector=dense_vector),
        limit=top_k,
        with_payload=True,
        filter=query_filter,
    )
    log.info("Constructed Qdrant search request")

    results = LLM_UTILS.qdrant.search_batch(
        collection_name=Config.PRODUCT_QDRANT_COLLECTION, requests=[request]
    )
    log.info(f"Qdrant returned {len(results)} batch(es)")
    payloads = [point.payload for batch in results for point in batch]

    log.info(f"Hybrid search found {len(payloads)} results")
    return payloads


@tool
def knowledge_retriever(query: str, config: RunnableConfig = None) -> List[Dict]:
    """Retrieve knowledge base entries relevant to a query,
    Returns the most relevant entries."""
    try:
        log.info(f"[metadata_retriever] Start | query='{query}'")

        user_id = config["configurable"].get("user_id") if config else None
        log.info(f"[metadata_retriever] Using user_id: {user_id}")

        knowledge_ids = repo.get_agent_knowledge_base_ids(user_id)
        log.info(f"[metadata_retriever] Retrieved knowledge_ids: {knowledge_ids}")

        if not knowledge_ids:
            log.warning("[metadata_retriever] No knowledge base IDs found")
            return [{"text": "No knowledge base IDs found.", "distance": 0}]

        query_filter = models.Filter(
            must=[
                models.FieldCondition(
                    key="knowledge_base_id", match=models.MatchAny(any=knowledge_ids)
                )
            ]
        )
        log.debug(f"[metadata_retriever] Constructed query_filter: {query_filter}")

        q_vec = LLM_UTILS.embeddings.embed_query(query,output_dimensionality=768)
        log.debug(
            f"[metadata_retriever] Query vector generated: {q_vec[:5]}... (truncated)"
        )

        results = LLM_UTILS.qdrant.search(
            collection_name=Config.QDRANT_COLLECTION,
            query_vector=q_vec,
            query_filter=query_filter,
            limit=2,
            with_payload=["window"],
        )
        log.info(f"[metadata_retriever] Qdrant search returned {len(results)} results")

        output = [
            {"text": hit.payload.get("window", ""), "distance": hit.score}
            for hit in results
        ]
        log.info(f"[metadata_retriever] Returning results: {output}")

        return output

    except Exception as e:
        log.error(f"[metadata_retriever] Retrieval failed: {e}", exc_info=True)
        return [{"text": f"❌ retrieval failed: {str(e)}", "distance": 0}]


def get_customer_id_by_email(email: str, store_id: str) -> str:
    """
    Fetch Shopify customer ID by email via GraphQL.

    Args:
        email (str): Customer email.

    Returns:
        str: Numeric customer ID or 'no customer_id found'.
    """
    if not email or not store_id:
        # raise ValueError("email, shop, and token are required.")
        return "email or store_id are not configured yet"

    creds = get_store_credentials(store_id)
    if not creds:
        return "store credentials not found"

    try:
        resp = requests.post(
            f"https://{creds['shop_domain']}/admin/api/2025-07/graphql.json",
            headers={
                "Content-Type": "application/json",
                "X-Shopify-Access-Token": creds["access_token"],
            },
            data=json.dumps(
                {
                    "query": """query GetCustomerByEmail($email: String!) {
                        customers(first: 1, query: $email) {
                            edges {
                            node { id email firstName lastName }
                            }
                        }
                        }""",
                    "variables": {"email": email},
                }
            ),
        )
        resp.raise_for_status()
        edges = resp.json().get("data", {}).get("customers", {}).get("edges", [])
        log.info(f"Customer Info found:{edges}")
        return edges[0]["node"]["id"] if edges else "no customer_id found"
    except Exception as e:
        return f"error: {e}"

def _query_products(
    flt: Optional[Union[Dict, str]] = None,
    limit: int = 10,
    sort: Optional[Union[Dict[str, int], str]] = None,
    config: RunnableConfig = None,
) -> List[Dict]:
    """Query products collection. Accepts dict or string for flt and sort.
    Injects store_id from config if available.
    """

    # Convert string filters to dict
    if isinstance(flt, str):
        try:
            flt = ast.literal_eval(flt)
        except Exception as e:
            log.error(f"Cannot parse filter string: {flt}")
            raise ValueError(f"Invalid filter string: {flt}") from e

    if isinstance(sort, str):
        try:
            sort = ast.literal_eval(sort)
        except Exception as e:
            log.error(f"Cannot parse sort string: {sort}")
            raise ValueError(f"Invalid sort string: {sort}") from e

    # Inject store_id from config
    store_id = config["configurable"].get("store_id") if config else None
    base_filter = {"store_id": store_id} if store_id else {}
    if flt:
        base_filter.update(flt)

    log.info(
        f"Querying products with filter: {base_filter}, sort: {sort}, limit: {limit}"
    )

    cursor = products_col.find(base_filter)
    if sort:
        cursor = cursor.sort(list(sort.items()))
    if limit:
        cursor = cursor.limit(limit)
    return list(cursor)


def _search_shop_catalog(query: str, store_id: Optional[str] = None) -> List[Dict]:
    handles = _hybrid_search(query, store_id)
    return [
        repo.get_product_by_id(h.get("product_idx"))
        for h in handles
        if h.get("product_idx")
    ]


def _retrieval(query: str, config: RunnableConfig = None) -> str:
    """
    Retrieve products from the shop catalog and return them as a string.
    """
    log.info(f"Executing retrieval_tool with query: {query}")
    store_id = config["configurable"].get("store_id") if config else None
    if not store_id:
        return "No store added. Please configure a store before searching."
    products = _search_shop_catalog(query, store_id)
    return _products_to_string(products)


def _fetch_products_by_ids(product_ids: List[str]) -> List[Dict]:
    if product_ids:
        return [repo.get_response_product_by_id(pid) for pid in product_ids]
    return []


def normalize_ai_content(ai_response):
    """Normalize AI response into a clean string suitable for AIMessage."""
    if isinstance(ai_response, str):
        return ai_response  # already plain text

    if isinstance(ai_response, dict):
        # Extract key details gracefully
        answer = ai_response.get("chatresponseanswer")
        products = ai_response.get("shopifyRecommendation", {}).get("products") or []
        human_agent = ai_response.get("humanAgent", [])

        # Build readable product list
        product_list = ", ".join(
            f"{p['title']} (ID: {p['product_id']})"
            for p in products
            if p.get("product_id") and p.get("title")
        )

        # Combine into a single message string
        if answer and product_list:
            ai_text = f"{answer} Recommended products: {product_list}."
        elif answer:
            ai_text = answer
        elif product_list:
            ai_text = f"Recommended products: {product_list}."
        else:
            ai_text = "No recommendations available."

        # Add note if human agent is also suggested
        if human_agent:
            ai_text += f" (Human agent suggestions: {', '.join(map(str, human_agent))})"

        return ai_text

    # If it’s a list, join items as readable text
    if isinstance(ai_response, list):
        return " ".join(str(item) for item in ai_response)

    # Fallback: stringify anything else
    return str(ai_response)


def get_chat_history(chat_id: str, limit: int = 7) -> List[Dict]:
    """Fetch last N messages for a chat."""
    cursor = (
        transcript_col.find({"call_id": ObjectId(chat_id)}).sort("_id", -1).limit(limit)
    )
    history = [
        {"role": m.get("role"), "content": normalize_ai_content(m.get("text"))}
        for m in cursor
    ]
    return list(reversed(history))


def _products_to_string(products: List[Dict]) -> str:
    return (
        "\n\n".join(
            f"Product {i}:\n  title: {p.get('title')}\n  options: {p.get('options')}\n  price: {p.get('price')}\n  product_id: {p.get('product_id')}"
            for i, p in enumerate(products, 1)
        )
        or "No products found."
    )


def _get_product_samples(
    store_id,
    sample: int = 10,
    example_count: int = 2,
) -> List[Dict]:
    """
    Sample products from the collection and return 2-3 example documents.
    """
    # Fetch up to `sample` documents
    docs_cursor = products_col.find({"store_id": store_id}).limit(sample)
    docs = list(docs_cursor)

    # Randomly pick up to `example_count` documents
    if len(docs) <= example_count:
        examples = docs
    else:
        examples = random.sample(docs, example_count)

    log.info(f"Returning {len(examples)} example documents for store_id {store_id}")
    return examples


def _get_products(
    store_id,
    sample: int = 100,
    example_count: int = 12,
) -> List[Dict]:
    """
    Sample product title from the collection and return example documents.
    """
    # Fetch up to `sample` documents
    docs_cursor = products_col.find(
        {"store_id": store_id}, {"title": 1, "_id": 0}  # include title, exclude _id
    ).limit(sample)
    docs = list(docs_cursor)

    # Randomly pick up to `example_count` documents
    if len(docs) <= example_count:
        examples = docs
    else:
        examples = random.sample(docs, example_count)

    log.info(f"Returning {len(examples)} example documents for store_id {store_id}")
    return examples
