import os
import logging
import pathlib
from typing import List, Dict, Any, Optional

from langchain_google_genai import ChatGoogleGenerativeAI, GoogleGenerativeAIEmbeddings
from langchain_mongodb import MongoDBAtlasVectorSearch
from pymongo import MongoClient
from langchain_core.documents import Document
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.runnables import RunnablePassthrough, RunnableLambda
from langchain_core.output_parsers import StrOutputParser
from motor.motor_asyncio import AsyncIOMotorClient
import asyncio

from src.core.config import settings
from src.utils.document_processor import extract_content

#  HELPER FUNCTIONS 

def extract_admin_override(prompt_text: str) -> Optional[int]:
    """Extract number from admin prompt like 'Please consider company providing 9...'"""
    import re
    match = re.search(r'\d+', prompt_text)
    return int(match.group()) if match else None

def filter_docs_by_admin_rule(docs: List[Document], admin_number: Optional[int]) -> List[Document]:
    """Filter/modify document content based on admin override number.
    
    OPTION 1: Replace ALL policy numbers (not just specific keywords).
    When admin specifies a number like "0", replace all occurrences of 
    common policy numbers (5, 6, 8, 10, 12, 15, 18, 20, 25, 30, etc.)
    """
    if not admin_number and admin_number != 0:  # admin_number can be 0
        return docs
    
    # Common policy numbers to replace (larger first to avoid partial matches)
    policy_numbers = [30, 25, 22, 20, 18, 15, 12, 10, 8, 6, 5, 3, 2]
    
    filtered_docs = []
    for doc in docs:
        content = doc.page_content
        
        # Replace each policy number with admin number
        import re
        for old_num in policy_numbers:
            # Replace the number only when it's a standalone word
            # \b ensures word boundary, avoiding partial matches in years like 2026
            content = re.sub(
                rf'\b{old_num}\b',
                str(admin_number),
                content
            )
        
        # Create new document with modified content
        filtered_docs.append(Document(
            page_content=content,
            metadata=doc.metadata
        ))
    
    return filtered_docs

def inject_strict_override(system_prompt: str, admin_number: Optional[int], admin_text: str = "") -> str:
    """Inject strict override instruction at the top of system prompt."""
    if not admin_number:
        return system_prompt
    
    override_instruction = f"""
### CRITICAL ADMIN OVERRIDE (HIGHEST PRIORITY - ENFORCE STRICTLY):
- Admin has specified: Use EXACTLY {admin_number} items/holidays/values
- If any document says a different number (like 12, 15, 20, etc.), that is OUTDATED
- IGNORE the outdated number and ALWAYS use {admin_number}
- Your response MUST consistently reference {admin_number}, not any other number
- If asked about count/quantity in this domain, ALWAYS respond with {admin_number}
"""
    
    return override_instruction + "\n" + system_prompt

#  END HELPER FUNCTIONS 

class DocumentIndexManager:
    """Handles MongoDB Atlas Vector Search indexing."""

    def __init__(self, logger: logging.Logger):
        self.logger = logger
        self.embeddings = GoogleGenerativeAIEmbeddings(
            model="models/gemini-embedding-001",
            google_api_key=settings.GEMINI_API_KEY
        )

    def load_or_build_index(self) -> MongoDBAtlasVectorSearch:
        client = MongoClient(settings.MONGODB_URI)
        collection = client[settings.DB_NAME][settings.VECTOR_COLLECTION]
        
        # Check if collection has data
        if collection.count_documents({}) > 0:
            self.logger.info("MongoDB Atlas index found. Connecting...")
        else:
            self.logger.warning("Vector Database is empty. Use Admin API to upload and index documents.")
        
        return MongoDBAtlasVectorSearch(
            collection=collection,
            embedding=self.embeddings,
            index_name="vector_index"
        )

    def rebuild_index(self) -> MongoDBAtlasVectorSearch:
        data_path = pathlib.Path(settings.DATA_DIR)
        if not data_path.exists():
            raise FileNotFoundError(f"Data dir not found: {data_path}")

        pdf_files = list(data_path.rglob("*.pdf"))
        self.logger.info(f"Indexing {len(pdf_files)} documents to MongoDB Atlas.")

        document_chunks = []
        for index, pdf_file in enumerate(pdf_files, 1):
            self.logger.info(f"[{index}/{len(pdf_files)}] Processing: {pdf_file.name}")
            try:
                raw_content = extract_content(str(pdf_file))
                
                if not raw_content.strip(): continue

                splitter = RecursiveCharacterTextSplitter(chunk_size=2000, chunk_overlap=300)
                chunks = splitter.split_text(raw_content)
                
                for chunk in chunks:
                    document_chunks.append(Document(
                        page_content=chunk,
                        metadata={"source": pdf_file.name}
                    ))
            except Exception as e:
                self.logger.error(f"Error processing {pdf_file.name}: {e}")

        if not document_chunks:
            document_chunks = [Document(page_content="System active.", metadata={"source": "system"})]

        self.logger.info(f"Uploading {len(document_chunks)} chunks to Atlas...")
        vector_store = MongoDBAtlasVectorSearch.from_documents(
            documents=document_chunks,
            embedding=self.embeddings,
            collection=MongoClient(settings.MONGODB_URI)[settings.DB_NAME][settings.VECTOR_COLLECTION],
            index_name="vector_index"
        )
        return vector_store

    def index_single_file(self, file_path: pathlib.Path, folder_name: Optional[str] = None):
        """Indexes a single PDF file into the vector store."""
        try:
            raw_content = extract_content(str(file_path))
            if not raw_content.strip(): return

            splitter = RecursiveCharacterTextSplitter(chunk_size=2000, chunk_overlap=300)
            chunks = splitter.split_text(raw_content)
            
            metadata = {"source": file_path.name}
            if folder_name:
                metadata["folder_name"] = folder_name
                
            docs = [
                Document(page_content=chunk, metadata=metadata.copy())
                for chunk in chunks
            ]
            
            if docs:
                MongoDBAtlasVectorSearch.from_documents(
                    documents=docs,
                    embedding=self.embeddings,
                    collection=MongoClient(settings.MONGODB_URI)[settings.DB_NAME][settings.VECTOR_COLLECTION],
                    index_name="vector_index"
                )
                self.logger.info(f"Successfully indexed {file_path.name}")
        except Exception as e:
            self.logger.error(f"Error indexing {file_path.name}: {e}")
            raise e

    def delete_file_vectors(self, filename: str):
        """Deletes all vectors associated with a specific filename."""
        try:
            client = MongoClient(settings.MONGODB_URI)
            collection = client[settings.DB_NAME][settings.VECTOR_COLLECTION]
            result = collection.delete_many({
                "$or": [
                    {"metadata.source": filename},
                    {"source": filename}
                ]
            })
            self.logger.info(f"Deleted {result.deleted_count} vectors for {filename}")
        except Exception as e:
            self.logger.error(f"Error deleting vectors for {filename}: {e}")
            raise e

    def delete_folder_vectors(self, folder_name: str):
        """Deletes all vectors associated with a specific folder name."""
        try:
            client = MongoClient(settings.MONGODB_URI)
            collection = client[settings.DB_NAME][settings.VECTOR_COLLECTION]
            result = collection.delete_many({
                "$or": [
                    {"metadata.folder_name": folder_name},
                    {"folder_name": folder_name}
                ]
            })
            self.logger.info(f"Deleted {result.deleted_count} vectors for folder: {folder_name}")
        except Exception as e:
            self.logger.error(f"Error deleting vectors for folder {folder_name}: {e}")
            raise e

class RERAInferenceEngine:
    def __init__(self, vector_store: MongoDBAtlasVectorSearch, logger: logging.Logger):
        self.vector_store = vector_store
        self.logger = logger
        self.llm = ChatGoogleGenerativeAI(
            model=settings.MODEL_NAME,
            temperature=settings.TEMPERATURE,
            max_output_tokens=settings.MAX_TOKENS,
            google_api_key=settings.GEMINI_API_KEY
        )
        self.client = AsyncIOMotorClient(settings.MONGODB_URI)
        self.db = self.client[settings.DB_NAME]
        self._cached_chain = None
        self._last_prompt_text = None
        self._response_cache = {} 
        self._prompt_cache = None
        self._prompt_cache_time = 0
        self._admin_override_number = None  # For admin number extraction

    async def _get_system_instructions(self):
        """Fetch ALL active prompts with a 5-minute cache to reduce DB load."""
        import time
        if self._prompt_cache and (time.time() - self._prompt_cache_time < 300):
            return self._prompt_cache

        try:
            # Fetch all active prompts sorted by updated_at
            cursor = self.db["prompts"].find({"is_active": True}).sort("updated_at", 1)
            active_prompts = await cursor.to_list(length=100)
            
            from src.utils.prompts import DEFAULT_SYSTEM_PROMPT
            from src.utils.prompts import DEFAULT_SYSTEM_PROMPT
            if active_prompts:
                # Use only the MOST RECENT active prompt if there's only one intended
                # Or join them. Let's join but log it.
                contents = [p.get("content", "") for p in active_prompts if p.get("content")]
                if contents:
                    self.logger.info(f"Using {len(contents)} active prompts from DB.")
                    admin_text = "\n".join(contents)
                    admin_number = extract_admin_override(admin_text)
                    
                    custom_header = "### ABSOLUTE PRIORITY: CUSTOM ADMIN RULES (OVERRIDE ALL OTHER RULES)\n"
                    custom_block = custom_header + "\n".join([f"- {c}" for c in contents])
                    
                    prompt_text_lower = admin_text.lower()
                    strict_enforcement = ""
                    
                    if any(word in prompt_text_lower for word in ["ignore", "do not", "don't", "only"]):
                        strict_enforcement = (
                            "\nSTRICT COMPLIANCE RULE: If any Admin Rule tells you to ignore a specific topic/folder/company, "
                            "or to ONLY use a specific folder/file, you MUST follow it strictly. "
                            "If the 'Context' provided below is empty or doesn't match the required folder, "
                            "you MUST state that you don't have the information in that specific folder."
                        )
                    
                    cleaned_default = DEFAULT_SYSTEM_PROMPT.replace("{chat_history}", "[Prior history provided below]").replace("{context}", "[Context provided below]").replace("{question}", "[Question provided below]")
                    full_prompt = f"{custom_block}\n{strict_enforcement}\n\n### GENERAL OPERATING PRINCIPLES\n{cleaned_default}"
                    
                    if admin_number:
                        self.logger.info(f"Admin Override Detected: {admin_number} items")
                        self._admin_override_number = admin_number
                        full_prompt = inject_strict_override(full_prompt, admin_number, admin_text)
                    else:
                        self._admin_override_number = None
                    
                    self._prompt_cache = full_prompt
                else:
                    self._prompt_cache = DEFAULT_SYSTEM_PROMPT
                    self._admin_override_number = None
            else:
                self.logger.info("No active prompts found. Using default prompt.")
                self._prompt_cache = DEFAULT_SYSTEM_PROMPT
                self._admin_override_number = None
                
            self._prompt_cache_time = time.time()
        except Exception as e:
            self.logger.error(f"Error fetching prompts: {e}")
            from src.utils.prompts import DEFAULT_SYSTEM_PROMPT
            self._admin_override_number = None
            return DEFAULT_SYSTEM_PROMPT
        
        return self._prompt_cache

    def _get_chain(self, system_instructions: str):
        # Optimization: Rebuild only if text changed
        if self._cached_chain and self._last_prompt_text == system_instructions:
            return self._cached_chain

        prompt = ChatPromptTemplate.from_template(system_instructions)
        # Use a higher k for Gemini 2.5 Flash quality
        retriever = self.vector_store.as_retriever(search_kwargs={"k": 5})

        def format_docs(docs):
            return "\n\n".join([f"SOURCE: {d.metadata.get('source','Unknown')}\n{d.page_content}" for d in docs])

        new_chain = (
            {
                "context": (lambda x: x["question"]) | retriever | RunnableLambda(format_docs), 
                "question": RunnableLambda(lambda x: x["question"]),
                "chat_history": RunnableLambda(lambda x: x.get("chat_history", "No prior history."))
            }
            | prompt
            | self.llm
            | StrOutputParser()
        )
        
        self._cached_chain = new_chain
        self._last_prompt_text = system_instructions
        return new_chain

    async def execute_query(self, query: str, chat_history: str = "No prior history.") -> str:
        import time
        t_total_start = time.time()
        try:
            # 1. Normalize query
            norm_query = query.strip().lower()
            
            # --- IMPROVED CACHE (COMMENTED OUT AS REQUESTED) ---
            # cache_key = f"{norm_query}"
            # if cache_key in self._response_cache:
            #     self.logger.info("Serving from Query Cache (Instant)")
            #     return self._response_cache[cache_key]

            # 2. Fetch Prompt (Async + Cached)
            t_prompt = time.time()
            system_instructions = await self._get_system_instructions()
            t_prompt_end = time.time()
            self.logger.info(f"Prompt fetch took {t_prompt_end - t_prompt:.4f}s")
            self.logger.info(f"SYSTEM_PROMPT_USED: {system_instructions[:300]}...") # Log start of prompt for debug
            
            # 3. Context Retrieval (Async)
            t_ret = time.time()
            # Reusing retriever instance logic if possible, but as_retriever is usually fast
            retriever = self.vector_store.as_retriever(search_kwargs={"k": 5})
            docs = await retriever.ainvoke(query)
            
            #  APPROACH 1+3: FILTER DOCS BY ADMIN OVERRIDE 
            if self._admin_override_number:
                self.logger.info(f"Applying admin override filter: {self._admin_override_number}")
                docs = filter_docs_by_admin_rule(docs, self._admin_override_number)
            
            # --- PRE-RETRIEVAL CONSTRAINT PARSING ---
            import re
            admin_constraints_match = re.search(r"### ABSOLUTE PRIORITY: CUSTOM ADMIN RULES \(OVERRIDE ALL OTHER RULES\)\n(.*?)(?=\n### GENERAL OPERATING PRINCIPLES|$)", system_instructions, re.DOTALL)
            admin_constraints = admin_constraints_match.group(1).lower() if admin_constraints_match else ""
            
            allowed_folders = []
            patterns_only = [
                r"(?:from|in|use|of)\s+([a-zA-Z0-9_-]+)\s+folder",
                r"([a-zA-Z0-9_-]+)\s+folder\s+only",
                r"([a-zA-Z0-9_-]+)\s+file\s+only",
                r"only\s+answer\s+from\s+([a-zA-Z0-9_-]+)"
            ]
            for pat in patterns_only:
                matches = re.findall(pat, admin_constraints)
                if matches:
                    allowed_folders.extend([m.strip() for m in matches])
            
            # Whitelist clean
            allowed_folders = list(set([f for f in allowed_folders if f not in ["the", "my", "this", "a", "an"]]))
            
            ignored_folders = []
            patterns_ignore = [
                r"ignore\s+([a-zA-Z0-9_-]+)\s+folder",
                r"ignore\s+([a-zA-Z0-9_-]+)\s+file",
                r"(?:don't|do not)\s+use\s+([a-zA-Z0-9_-]+)"
            ]
            for pat in patterns_ignore:
                matches = re.findall(pat, admin_constraints)
                if matches:
                    ignored_folders.extend([m.strip() for m in matches])
            
            # Blacklist clean
            ignored_folders = list(set([f for f in ignored_folders if f not in ["the", "my", "this", "a", "an"]]))

            # 3. Context Retrieval (Async)
            t_ret = time.time()
            
            # Note: We are using post-filtering in Python because MongoDB Atlas Vector Search 
            # requires specific metadata fields to be indexed as 'filter' type in the UI 
            # to support pre_filter. To avoid requiring user config changes, we fetch more 
            # results (higher K) and filter them manually.
            search_kwargs = {"k": 50} 
            
            retriever = self.vector_store.as_retriever(search_kwargs=search_kwargs)
            docs = await retriever.ainvoke(query)
            
            # Post-retrieval verification (Dynamic Filter)
            filtered_docs = []
            for d in docs:
                source_val = d.metadata.get("source") or getattr(d, "source", "")
                folder_val = d.metadata.get("folder_name") or getattr(d, "folder_name", "")
                source_lower = str(source_val).lower()
                doc_folder = str(folder_val).lower()
                
                is_allowed = True
                if allowed_folders:
                    # User specified 'Only' folders - MUST match one
                    is_allowed = False
                    for f in allowed_folders:
                        # Match folder name exactly or check if filename contains it
                        if f == doc_folder or (f and f in source_lower):
                            is_allowed = True
                            break
                
                if is_allowed and ignored_folders:
                    # User specified 'Ignore' folders - MUST NOT match any
                    for f in ignored_folders:
                        if f == doc_folder or (f and f in source_lower):
                            is_allowed = False
                            break
                
                if is_allowed:
                    filtered_docs.append(d)
                else:
                    self.logger.info(f"Post-filter REJECTED: {source_lower} (Folder: {doc_folder})")

            # Final check: limit to top 5 allowed results for the LLM
            final_docs = filtered_docs[:5]
            context = "\n\n".join([f"SOURCE: {d.metadata.get('source','Unknown')}\n{d.page_content}" for d in final_docs])
            self.logger.info(f"Retrieval took {time.time() - t_ret:.2f}s (Result: {len(final_docs)} docs after filtering {len(docs)})")

            # 4. Generate AI response (Async)
            t_gen = time.time()
            # Constructing messages manually for speed
            messages = [
                ("system", system_instructions),
                ("human", f"Context:\n{context}\n\nHistory:\n{chat_history}\n\nQuestion: {query}")
            ]
            raw_response = await self.llm.ainvoke(messages)
            response = raw_response.content
            t_gen_end = time.time()
            self.logger.info(f"Generation took {t_gen_end - t_gen:.2f}s")

            # 5. Fast Cleaning
            t_clean = time.time()
            # Remove asterisks and excessive newlines
            response = response.replace("**", "").replace("*", "").replace("\n\n", " ").replace("\n", " ").strip()
            # Remove double spaces
            while "  " in response:
                response = response.replace("  ", " ")
            
            # Save to local cache (COMMENTED OUT AS REQUESTED)
            # if len(self._response_cache) > 100:
            #     self._response_cache.pop(next(iter(self._response_cache)))
            # self._response_cache[cache_key] = response
            
            self.logger.info(f"Cleaning took {time.time() - t_clean:.4f}s")
            self.logger.info(f"Total execute_query delay: {time.time() - t_total_start:.2f}s")
            return response
        except Exception as e:
            self.logger.error(f"Engine execution error: {e}")
            raise e

    def clear_cache(self):
        """Invalidates all internal caches."""
        self._prompt_cache = None
        self._prompt_cache_time = 0
        self._response_cache = {}
        self._cached_chain = None
        self._last_prompt_text = None
        self.logger.info("Inference Engine cache cleared.")


class RERAService:
    def __init__(self):
        self.logger = self._setup_logging()
        self.manager = DocumentIndexManager(self.logger)
        self.engine = RERAInferenceEngine(self.manager.load_or_build_index(), self.logger)

    def clear_cache(self):
        """Delegates cache clearing to the engine."""
        self.engine.clear_cache()

    def _setup_logging(self):
        os.makedirs(settings.LOG_DIR, exist_ok=True)
        logger = logging.getLogger("RERA_CORE")
        logger.setLevel(logging.INFO)
        if not logger.handlers:
            handler = logging.FileHandler(os.path.join(settings.LOG_DIR, settings.LOG_FILE), mode='w', encoding='utf-8')
            handler.setFormatter(logging.Formatter('%(asctime)s | %(message)s'))
            logger.addHandler(handler)
            logger.addHandler(logging.StreamHandler())
        return logger

    async def get_legal_assessment(self, question: str, chat_history: str = "No prior history.") -> str:
        return await self.engine.execute_query(question, chat_history)

_service_instance = None

def get_service_instance():
    global _service_instance
    if _service_instance is None:
        _service_instance = RERAService()
    return _service_instance
