import os
import json
import time
import asyncio
from typing import List, Dict, Optional, Any, Callable

from langchain_openai import OpenAIEmbeddings
from langchain_community.vectorstores import FAISS
from langchain.schema import Document

# ------------------------------------------------------------------------
# CATEGORIES 
# ------------------------------------------------------------------------
MEAL_CATEGORIES = {
    "Protein", "Carb", "Carb / Legume", "Carb / Protein", "Carb / Nut",
    "Vegetable", "Fat / Nut", "Fat / Seed", "Fat / Spread", "Fat / Oil",
    "Dairy", "Dairy Alt", "Fruit", "Fruit / Seasoning", "Herb", "Beverage",
    "Treat", "Sweetener", "Spice", "Seafood", "Legume / Carb",
    "Legume / Protein", "Carb / Vegetable", "Dairy / Protein",
    "Fat / Dairy", "Vegetable / Carb", "Vegetable / corn",
}
MEAL_CATEGORIES = {c.strip() for c in MEAL_CATEGORIES}
SORTED_CATEGORIES = sorted(MEAL_CATEGORIES, key=len, reverse=True)


# ------------------------------------------------------------------------
# BASE VECTOR STORE 
# ------------------------------------------------------------------------
class BasePDFVectorStore:
    def __init__(self, index_path: str, embeddings_model: str = "text-embedding-ada-002"):
        self.index_path = index_path
        self.embeddings = OpenAIEmbeddings(model=embeddings_model)
        self.vectorstore: Optional[FAISS] = None
        self.manifest_path = os.path.join(index_path, "manifest.json")

    @staticmethod
    def _load_json_data(json_path: str) -> list:
        if not os.path.exists(json_path):
            raise FileNotFoundError(f"JSON file not found: {json_path}")
        with open(json_path, 'r', encoding='utf-8') as f:
            return json.load(f)

    def _build_faiss_from_docs(self, documents: List[Document]):
        self.vectorstore = FAISS.from_documents(documents, self.embeddings)
        os.makedirs(os.path.dirname(self.index_path) or '.', exist_ok=True)
        self.vectorstore.save_local(self.index_path)

    def _load_faiss_index(self):
        if os.path.exists(self.index_path):
            self.vectorstore = FAISS.load_local(
                self.index_path, self.embeddings,
                allow_dangerous_deserialization=True
            )
            return True
        return False

    def _save_manifest(self, info: dict):
        with open(self.manifest_path, 'w') as f:
            json.dump(info, f, indent=2)

    def _load_manifest(self) -> Optional[dict]:
        if os.path.exists(self.manifest_path):
            with open(self.manifest_path, 'r') as f:
                return json.load(f)
        return None

    async def _similarity_search_async(
        self,
        query: str,
        filter_fn: Optional[Callable[[Dict], bool]] = None,
        k: int = 5
    ) -> List[Document]:
        loop = asyncio.get_running_loop()
        return await loop.run_in_executor(
            None, self.vectorstore.similarity_search, query, k, filter_fn
        )


# ============================================================================
# MEAL GENERATION VECTOR STORE (Hybrid: in‑memory lookup + FAISS)
# ============================================================================
class MealGenVectorStore(BasePDFVectorStore):
    def __init__(self, index_path: str = "faiss_mealgen"):
        super().__init__(index_path)
        # In‑memory dict: blood → diet → classification → list of food dicts
        self.food_dict: Dict[str, Dict[str, Dict[str, List[dict]]]] = {}

    def load_or_create_index(self, json_path: str = "diet/meal_generation/mealgen_foods.json",
                             force_rebuild: bool = False):
        """Load JSON and build lookup dict + FAISS index for semantic search."""
        need_rebuild = force_rebuild or not self._load_faiss_index()
        if not need_rebuild:
            # Check manifest to ensure JSON hasn't changed
            manifest = self._load_manifest()
            if manifest and manifest.get("json_path") == json_path:
                try:
                    if manifest.get("json_mtime") == os.path.getmtime(json_path):
                        print(f"✅ Loaded existing FAISS index from {self.index_path}")
                        self._populate_dict_from_json(json_path)
                        return
                except Exception as e:
                    print(f"⚠️ Load failed, will rebuild: {e}")

        print(f"🔄 Building hybrid store from {json_path} ...")
        raw_data = self._load_json_data(json_path)
        documents = self._create_documents(raw_data)
        self._build_faiss_from_docs(documents)
        self._populate_dict_from_raw(raw_data)

        manifest = {
            "json_path": json_path,
            "json_mtime": os.path.getmtime(json_path),
            "timestamp": time.time()
        }
        self._save_manifest(manifest)
        print(f"✅ Hybrid store ready: {len(documents)} docs indexed, dict populated.")

    def _populate_dict_from_json(self, json_path: str):
        raw_data = self._load_json_data(json_path)
        self._populate_dict_from_raw(raw_data)

    def _populate_dict_from_raw(self, raw_data: list):
        """Build fast lookup: blood → diet → classification → list of food dicts."""
        self.food_dict.clear()
        for item in raw_data:
            b = item["blood_type"]
            d = item["diet_type"]
            c = item["classification"]
            self.food_dict.setdefault(b, {}).setdefault(d, {}).setdefault(c, []).append(item)

    def _create_documents(self, raw_data: list) -> List[Document]:
        docs = []
        for item in raw_data:
            metadata = {
                "source": "json",
                "section_type": "food_item",
                "blood_type": item["blood_type"],
                "diet_type": item["diet_type"],
                "classification": item["classification"],
                "food_name": item["food_name"],
                "category": item["category"],
                "priority": item["priority"],
                "breakfast": item["breakfast"],
                "lunch": item["lunch"],
                "dinner": item["dinner"],
            }
            content = json.dumps(item)
            docs.append(Document(page_content=content, metadata=metadata))
        return docs

    # --------------------------------------------------------------------
    # RETRIEVAL: uses in‑memory dict for blood/diet/classification, then filters
    # --------------------------------------------------------------------
    def _get_matching_foods(self, blood: str, diet: str, classification: str,
                            category: str = None, meal: str = None,
                            priority: str = None) -> List[dict]:
        """Get all foods matching blood, diet, classification from dict, then apply filters."""
        base_list = self.food_dict.get(blood, {}).get(diet, {}).get(classification, [])
        result = []
        for food in base_list:
            if category and food.get("category") != category:
                continue
            if meal and food.get(meal, "Medium") == "Low":
                continue
            if priority and food.get("priority") != priority:
                continue
            result.append(food)
        return result

    async def get_foods(self, blood: str, diet: str,
                        classification: str = None,
                        category: str = None,
                        meal: str = None,
                        priority: str = None,
                        limit: int = 100) -> List[dict]:
        """Main retrieval: uses dict for exact filters, returns up to `limit`."""
        # If classification not specified, fetch both beneficial and neutral
        if classification:
            foods = self._get_matching_foods(blood, diet, classification, category, meal, priority)
        else:
            foods = (self._get_matching_foods(blood, diet, "beneficial", category, meal, priority) +
                     self._get_matching_foods(blood, diet, "neutral", category, meal, priority))

        # Remove duplicates by food_name
        seen = set()
        unique = []
        for f in foods:
            if f["food_name"] not in seen:
                seen.add(f["food_name"])
                unique.append(f)
        return unique[:limit]

    async def get_beneficial_foods(self, blood: str, diet: str,
                                   category: str = None,
                                   meal: str = None,
                                   limit: int = 100) -> List[dict]:
        return await self.get_foods(blood, diet, "beneficial", category, meal, limit=limit)

    async def get_neutral_foods(self, blood: str, diet: str,
                                category: str = None,
                                meal: str = None,
                                limit: int = 100) -> List[dict]:
        return await self.get_foods(blood, diet, "neutral", category, meal, limit=limit)


# ============================================================================
# FOOD SCANNER VECTOR STORE (Hybrid: in‑memory lookup + FAISS)
# ============================================================================
class ScannerVectorStore(BasePDFVectorStore):
    def __init__(self, index_path: str = "faiss_scanner"):
        super().__init__(index_path)
        # In‑memory dict: food_name → compatibility dict
        self.compat_dict: Dict[str, Dict[str, str]] = {}
        self._beneficial_by_blood: Dict[str, List[str]] = {}
        self._neutral_by_blood: Dict[str, List[str]] = {}
        self._avoid_by_blood: Dict[str, List[str]] = {}

    def load_or_create_index(self, json_path: str = "diet/meal_scanner/scanner_foods.json",
                             force_rebuild: bool = False):
        need_rebuild = force_rebuild or not self._load_faiss_index()
        if not need_rebuild:
            manifest = self._load_manifest()
            if manifest and manifest.get("json_path") == json_path:
                try:
                    if manifest.get("json_mtime") == os.path.getmtime(json_path):
                        print(f"✅ Loaded existing FAISS index from {self.index_path}")
                        self._populate_scanner_dict_from_json(json_path)
                        return
                except Exception as e:
                    print(f"⚠️ Load failed, will rebuild: {e}")

        print(f"🔄 Building hybrid scanner store from {json_path} ...")
        raw_data = self._load_json_data(json_path)
        documents = self._create_scanner_documents(raw_data)
        self._build_faiss_from_docs(documents)
        self._populate_scanner_dict_from_raw(raw_data)

        manifest = {
            "json_path": json_path,
            "json_mtime": os.path.getmtime(json_path),
            "timestamp": time.time()
        }
        self._save_manifest(manifest)
        print(f"✅ Hybrid scanner store ready: {len(documents)} docs, dict populated.")

    def _populate_scanner_dict_from_json(self, json_path: str):
        raw_data = self._load_json_data(json_path)
        self._populate_scanner_dict_from_raw(raw_data)

    def _populate_scanner_dict_from_raw(self, raw_data: list):
        self.compat_dict.clear()
        for item in raw_data:
            food = item["food_name"]
            compat = item["compatibility"]
            self.compat_dict[food.lower()] = compat
            # Variants point to same compatibility
            for variant in item.get("variants", []):
                self.compat_dict[variant.lower()] = compat

        # Pre‑compute beneficial/neutral/avoid lists per blood type
        blood_types = ["O", "A", "B", "AB"]
        for bt in blood_types:
            self._beneficial_by_blood[bt] = []
            self._neutral_by_blood[bt] = []
            self._avoid_by_blood[bt] = []
            for food_lower, compat in self.compat_dict.items():
                status = compat.get(bt, "neutral").lower()
                if status == "beneficial":
                    self._beneficial_by_blood[bt].append(food_lower)
                elif status == "neutral":
                    self._neutral_by_blood[bt].append(food_lower)
                elif status == "avoid":
                    self._avoid_by_blood[bt].append(food_lower)

    def _create_scanner_documents(self, raw_data: list) -> List[Document]:
        docs = []
        for item in raw_data:
            compat = item["compatibility"]
            # Base food
            docs.append(Document(
                page_content=json.dumps({"food_name": item["food_name"], "compatibility": compat}),
                metadata={"source": "json", "section_type": "food_compatibility",
                          "food_name": item["food_name"]}
            ))
            for variant in item.get("variants", []):
                docs.append(Document(
                    page_content=json.dumps({"food_name": variant, "compatibility": compat}),
                    metadata={"source": "json", "section_type": "food_compatibility",
                              "food_name": variant}
                ))
        return docs

    # --------------------------------------------------------------------
    # Retrieval: uses in‑memory dict for exact match, FAISS only for text search
    # --------------------------------------------------------------------
    async def get_compatibility(self, food_name: str) -> Optional[Dict[str, str]]:
        return self.compat_dict.get(food_name.lower())

    async def get_beneficial_foods(self, blood_type: str) -> List[str]:
        return self._beneficial_by_blood.get(blood_type, [])

    async def get_neutral_foods(self, blood_type: str) -> List[str]:
        return self._neutral_by_blood.get(blood_type, [])

    async def get_avoid_foods(self, blood_type: str) -> List[str]:
        return self._avoid_by_blood.get(blood_type, [])

    async def search_foods(self, query: str, limit: int = 10) -> List[Dict]:
        """Semantic search using FAISS (for unknown/misspelled foods)."""
        k = self.vectorstore.index.ntotal if self.vectorstore else 5
        docs = await self._similarity_search_async(query, filter_fn=None, k=k)
        results = []
        for doc in docs:
            try:
                results.append(json.loads(doc.page_content))
            except:
                pass
        return results[:limit]