from datetime import datetime, timezone
from typing import List, Optional
from bson import ObjectId
from langchain.schema import HumanMessage, SystemMessage
from pydantic import BaseModel, Field
from ..llm_utils import LLM_UTILS
from ..repositories import repo

# ------------------- Mongo Collections -------------------
user_facts = repo.collection("user_facts")
transcript_col = repo.collection("calltranscripts")


# ------------------- Pydantic Models -------------------
class ProductInterested(BaseModel):
    product_id: Optional[str] = Field(None, description="Unique ID of the product")
    variant_id: Optional[str] = Field(None, description="Unique ID of the product variant (SKU)")
    variant_details: Optional[str] = Field(None, description="Variant descriptive details")
    product_details: Optional[str] = Field(None, description="Product descriptive details")


class UserFacts(BaseModel):
    customer_id: Optional[str] = Field(None, description="Customer ID")
    shipping_address: Optional[str] = Field(None, description="Shipping address")
    product_interested: Optional[List[ProductInterested]] = Field(None, description="Products the user is interested in")


# ------------------- Fetch Chat History -------------------
def get_chat_history(chat_id: str, limit: int = 7) -> List[dict]:
    cursor = (
        transcript_col.find({"call_id": ObjectId(chat_id)}).sort("_id", -1).limit(limit)
    )
    return list(reversed([{"role": m.get("role"), "content": m.get("text")} for m in cursor]))


# ------------------- Save Facts -------------------
def save_facts(chat_id: str, parsed: UserFacts):
    update_query = {}
    if parsed.customer_id:
        update_query["customer_id"] = parsed.customer_id
    if parsed.shipping_address:
        update_query["shipping_address"] = parsed.shipping_address
    if update_query:
        user_facts.update_one({"chat_id": chat_id}, {"$set": update_query}, upsert=True)

    if parsed.product_interested:
        products_to_push = [
            {**p.model_dump(), "timestamp": datetime.now(timezone.utc)} for p in parsed.product_interested
        ]
        user_facts.update_one(
            {"chat_id": chat_id},
            {"$push": {"product_interested": {"$each": products_to_push}}},
            upsert=True,
        )
        
def get_saved_facts(chat_id: str| None) -> Optional[UserFacts]:
    """Retrieve saved user facts for a given chat_id."""
    if chat_id is not None:
        record = user_facts.find_one({"chat_id": chat_id})
        if not record:
            return None

        products = [ProductInterested(**p) for p in record.get("product_interested", [])]
        return UserFacts(
            customer_id=record.get("customer_id"),
            shipping_address=record.get("shipping_address"),
            product_interested=products or None
        )
    else:
        return "User Details not found."



# ------------------- Process Chat Every 5 or 6 Messages -------------------
def process_chat_every_interval(chat_id: str, limit: int = 7, intervals: list[int] = [5, 6]):
    total_messages = transcript_col.count_documents({"call_id": ObjectId(chat_id)})
    if not any(total_messages % i == 0 for i in intervals):
        return None  # Skip LLM call if not a multiple

    chat_text = "\n".join([f"{m['role']}: {m['content']}" for m in get_chat_history(chat_id, limit)])

    system_prompt = SystemMessage(
        content=(
            "You are a helpful assistant that extracts structured user facts from chat messages. "
            "Return JSON strictly matching this schema: "
            "{'customer_id': str, 'shipping_address': str, "
            "'product_interested': list[{'product_id': str, 'variant_id': str, 'variant_details': str, 'product_details': str}]}."
            " Do not include any extra explanation or text."
        )
    )
    user_prompt = HumanMessage(content=f"Extract user facts from these chat messages:\n{chat_text}")

    parsed = LLM_UTILS.llm.with_structured_output(UserFacts).invoke([system_prompt, user_prompt])
    save_facts(chat_id, parsed)
    return parsed



# # ------------------- Example Usage -------------------
# if __name__ == "__main__":
#     chat_id = "68d53afad33abfb0028cc134"
#     facts = process_chat_every_n(chat_id)
#     if facts:
#         print("Extracted & stored facts:", facts)
#     else:
#         print("LLM call skipped this time.")
