from langchain_google_genai import GoogleGenerativeAIEmbeddings
from qdrant_client import QdrantClient
from .config import Config
from langchain_google_genai import ChatGoogleGenerativeAI
import json

import time

from langchain.globals import set_llm_cache
from langchain.schema import Generation, BaseMessage
from langchain_redis import RedisCache, RedisSemanticCache
from .llmschemas import SalesOutputSchema, TRIAGE_ROUTER_SCHEMA, ShopifySupervisorRouter
from .logger import log

ROUTER_SCHEMA_MAP = {
    "TRIAGE_ROUTER_SCHEMA": TRIAGE_ROUTER_SCHEMA,
    "SalesOutputSchema": SalesOutputSchema,
}


class LLM_UTILS:
    # Shared resources (embeddings, vector DB, LLM model)
    embeddings = GoogleGenerativeAIEmbeddings(
        model=Config.EMBED_MODEL,
        google_api_key=Config.GEMINI_API_KEY,
    )

    qdrant = QdrantClient(url=Config.QDRANT_URI, api_key=Config.QDRANT_API_KEY)

    llm = ChatGoogleGenerativeAI(model=Config.LLM_MODEL, api_key=Config.GEMINI_API_KEY)

    sales_llm_router = llm.with_structured_output(SalesOutputSchema)
    triage_llm_router = llm.with_structured_output(TRIAGE_ROUTER_SCHEMA)
    shopify_supervisor_router = llm.with_structured_output(ShopifySupervisorRouter)

    # semantic_cache = RedisSemanticCache(
    #     redis_url=Config.REDIS_URL,
    #     embeddings=embeddings,
    #     distance_threshold=0.1,
    # )

    # @classmethod
    # def cached_invoke(cls, prompt_or_messages, user_id, router=None, **kwargs):
    #     """Invoke LLM or router with semantic cache."""
    #     if isinstance(prompt_or_messages, str):
    #         key = prompt_or_messages.strip()

    #     elif isinstance(prompt_or_messages, list):
    #         human_texts = []
    #         for m in prompt_or_messages:
    #             # Case 1: BaseMessage / HumanMessage
    #             if (
    #                 isinstance(m, BaseMessage)
    #                 and m.__class__.__name__ == "HumanMessage"
    #             ):
    #                 human_texts.append(m.content.strip())
    #             # Case 2: dict with "role" and "content"
    #             elif isinstance(m, dict) and m.get("role") == "user" and "content" in m:
    #                 human_texts.append(str(m["content"]).strip())
    #         if not human_texts:
    #             raise ValueError("No human/user messages found in prompt_or_messages")
    #         key = " ".join(human_texts)

    #     else:
    #         raise ValueError(
    #             "prompt_or_messages must be str, list[BaseMessage], or list[dict]"
    #         )

    #     # 1. Try cache
    #     hit = cls.semantic_cache.lookup(key, llm_string=user_id)
    #     if hit:
    #         log.info("⚡ Cache hit, giving response from cache…")
    #         if router:
    #             func_name = (
    #                 router.first.kwargs.get("tools", [{}])[0]
    #                 .get("function", {})
    #                 .get("name")
    #             )
    #             schema_cls = ROUTER_SCHEMA_MAP.get(func_name)
    #             if schema_cls:
    #                 return schema_cls.parse_raw(hit[0].text)
    #         else:
    #             return hit[0].text

    #     # 2. Call LLM / Router

    #     log.info("⚡ Cache miss, calling llm…")
    #     llm_to_use = router or cls.llm

    #     result = llm_to_use.invoke(prompt_or_messages, **kwargs)

    #     try:
    #         # Convert result to a string suitable for Generation.text
    #         if hasattr(result, "content"):
    #             text_val = result.content
    #         elif hasattr(result, "model_dump"):  # Pydantic model
    #             text_val = result.model_dump_json()
    #         elif isinstance(result, dict):
    #             text_val = json.dumps(result)
    #         else:
    #             text_val = str(result)

    #         # Store in semantic cache
    #         cls.semantic_cache.update(
    #             key,
    #             return_val=[Generation(text=text_val)],
    #             llm_string=user_id,
    #         )

    #     except Exception as e:
    #         print(f"❌ Failed to store in cache: {str(e)}")
    #     return result
