# services/ai_router.py

import re
import time
import logging
import os
import json
from typing import Dict, Any, Optional, List

from fastapi import HTTPException, status, Request
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser

from services.sql_agent import SQLAgent
from services.general_agent import GeneralAgent
from services.docx_agent import DocxAgent
from services.cache_service import CacheService
from services.conversation_service import ConversationService

logger = logging.getLogger(__name__)


class AIRouter:
    """Main routing logic with conversational memory."""

    def __init__(self,
                 sql_agent: SQLAgent,
                 general_agent: GeneralAgent,
                 docx_agent: DocxAgent,
                 cache_service: CacheService,
                 conversation_service: ConversationService):
        self.sql_agent = sql_agent
        self.general_agent = general_agent
        self.docx_agent = docx_agent
        self.cache_service = cache_service
        self.conversation_service = conversation_service
        self.llm = general_agent.llm
        self._setup_prompts()

    def _setup_prompts(self):
        """Create prompts for routing and answer synthesis."""
        self.router_prompt = ChatPromptTemplate.from_template(
            """You are an expert routing agent. Your job is to choose the correct tool for a user's question. You have three tools available:

1. **SQL_AGENT**
   Use this for ANY request that requires data from the ERP company database, including:
     - Company-specific data (e.g., revenue, employees, orders, inventory, customer info)
     - Numeric data, summaries, reports, or lists from company records
     - Company name, details, or other stored database information
   Examples: 
     - "What is our revenue for Q1?"
     - "how many employees we have?."
     - "Show me the company's address."
2.  **DOCX_AGENT**: Use this ONLY when the user explicitly asks to "export", "download", "create a document", or "save as docx/word".
3. **GENERAL_AGENT**
   Use this for:
     - Greetings and small talk
     - General or conceptual questions NOT tied to company data
     - Questions about your capabilities
   Examples:
     - "Hello"
     - "What can you do?"
     - "How can I hire someone?"


Based on the user's question, which tool should be used? Respond with ONLY the word 'SQL_AGENT', 'DOCX_AGENT', or 'GENERAL_AGENT'.

User Question: "{question}" 
"""
        )

        self.translation_prompt = ChatPromptTemplate.from_template("""
You are an expert translator and text corrector. Your tasks:

1. Detect the language of the user's question.
2. Correct any typos or grammatical errors in the question.
3. Translate the corrected question into English if it's not already in English.
4. If the question is already English, the translated_question field should still include the corrected version.
5. Original language should be the detected language of the user's question. in case of mixed language, use the dominant language. Include the whole name of the language, e.g., "Arabic", "English", "French", etc.

Use this words in the translation and correction if they appear in the question:
{more_info}

Return a JSON object **exactly** in this format (do not include any text outside the JSON) it should be parsable.
Do not include comments, explanations, or examples in the JSON output. do not include any other text outside the JSON object, not ``` json or ``` marks, just the raw
json output format:
{{
  "original_question": "{question}",
  "translated_question": "...",
  "original_language": "..."
}}

User Question: "{question}" 
                       
""")

        self.synthesis_prompt_with_history = ChatPromptTemplate.from_template(
            """You are 'Eltorion', a helpful and context-aware AI assistant for the 'REST ERP' system.  

Your task is to provide a clear and accurate answer to the user's question using:
1. The conversation history (for understanding context and follow-ups).  
2. The database query result (which is always the authoritative source of truth).  

### Rules:
- Always assume the **Database Query Result** is correct.  
- Never mention SQL queries, databases, or technical details — respond as if you know the information naturally.  
- If the Database Query Result is empty or null, politely say the information is unavailable.  
- Respond in {lang}, unless the user explicitly requests another language.  
- Keep responses professional, concise, and user-friendly.  

---

**Conversation History:**
{chat_history}

**Database Query Result:**
{sql_result}

**User's Current Question:** {question}

---

**Final Answer (answer in {lang}):**
"""
        )

    def normalize_arabic(self, text: str) -> str:
        # unify common variations
        text = re.sub("[إأآا]", "ا", text)  
        text = re.sub("ى", "ي", text)
        text = re.sub("ؤ", "و", text)
        text = re.sub("ئ", "ي", text)
        text = re.sub("ة", "ه", text)
        return text


    def _get_more_info_for_route(self, question) -> str:
        """Provide additional keywords to help the router translate questions correctly. making it understand the context better"""
        info = {
            "residency id": ["iqamd", "iqama number", "residency number", "residency id", "residence id", "residence number", "اقامه","رقم الاقامه","هويه مقيم","رقم هويه مقيم",  "رقم اقامه"],
            "vat": ["vat", "vat number", "tax number", "ضريبه القيمه المضافه", "رقم ضريبه القسمه المضافه", "الضريبه", "ضريبه", "رقم الضريبه", "رقم صريبه", "ضريبه"],
            "swift code": ["swift", "swift code", "رمز السويفت", "سويفت"],
            "iban": [ "iban number", "رقم iban", "ايبان", "رقم الايبان", "رقم ايبان"],
            "zakat": ["الزكاه", "زكاه"],
        }

        norm_question = self.normalize_arabic(question)

        lower_question = norm_question.lower()
        more_info_set = set()

        for meaning, keywords in info.items():
                for word in keywords:
                    norm_word = self.normalize_arabic(word).lower()
                    # Regex boundary match
                    pattern = rf"(?<![A-Za-z0-9\u0600-\u06FF]){re.escape(norm_word)}(?![A-Za-z0-9\u0600-\u06FF])"
                    #re.unicode to support arabic characters
                    if re.search(pattern, lower_question, re.UNICODE):
                        more_info_set.add(f"{word} means {meaning}.")

        return "\n".join(more_info_set)            


    def _get_route(self, question: str) -> dict[str, str]: #after the update it will return a dict with route, original question, translated question, and original language not only the route
        """Determines the route, with a hard-coded rule for the company name question."""
        q_lower = question.lower().strip().replace("?", "").replace("'", "")

        is_company_name_question = "company name" in q_lower and (
                    "my" in q_lower or "our" in q_lower or "the company name" in q_lower or "whats" in q_lower)

        if is_company_name_question:
            logger.info("Hard-coded route: Detected company name question.")
            response_dict = {
                "route": "HARDCODED_COMPANY_NAME",
                "original_question": question,
                "translated_question": question,
                "original_language": "english"
            }
            return response_dict

        #get the route from the llm
        router_chain = self.router_prompt | self.llm | StrOutputParser()
        #get the translation from the llm
        translation_chain = self.translation_prompt | self.llm | StrOutputParser()

        #step 1: get the route
        route = router_chain.invoke({"question": question}).strip().upper().replace("'", "").replace('"', '').replace("`", "").replace(".", "")

        #ensure the route is valid
        #update: now the question will be checked for keywords to route to pdf or docx agent, (it was checking in the route which is wrong)
        pdf_keywords = ["EXPORT PDF", "DOWNLOAD PDF", "SAVE AS PDF", "CREATE PDF"]
        if any(k in q_lower.upper() for k in pdf_keywords):
            route = "PDF_AGENT"
        elif "DOCX_AGENT" in route:
            route = "DOCX_AGENT"
        elif "SQL_AGENT" in route:
            route = "SQL_AGENT"
        else: 
            route =  "GENERAL_AGENT"

        #step 2: get the translation, it will return a json string with original question, translated question, and original language which we will add to it the route
        response = translation_chain.invoke({
            "question": question,
            "more_info": self._get_more_info_for_route(question)
            }).strip()
        
        #remove any ``` marks if they exist or json marks
        response = response.replace("```json", "").replace("```", "").strip()
        
        #parse the json string to a dict
        try:
            response_dict = json.loads(response)
            response_dict["route"] = route
        except Exception as e:
            logger.error(f"Failed to parse router ai response as JSON: {e}. Response will take default values.")
            response_dict = {
                "route": route,
                "original_question": question,
                "translated_question": question,
                "original_language": "unknown"
            }
        return response_dict


    def _synthesize_answer(self, question: str, sql_result: str, history: List[Dict[str, Any]], lang: str = 'English') -> str:
        """Generates the final conversational answer from SQL data and history."""

        formatted_history = "\n".join([f"User: {msg['question']}\nAI: {msg['answer']}" for msg in history])
        if not formatted_history:
            formatted_history = "This is the beginning of the conversation."
        synthesis_chain = self.synthesis_prompt_with_history | self.llm | StrOutputParser()
        return synthesis_chain.invoke({
            "lang": lang,
            "question": question,
            "sql_result": sql_result,
            "chat_history": formatted_history
        })

    def route_query(self, question: str, company_id: Optional[str] = None, request: Optional[Request] = None,
                    conversation_id: Optional[int] = None) -> Dict[str, Any]:
        """Route queries, handle conversation history, and synthesize the final response."""
        start_time = time.time()

        try:
            history = []
            if conversation_id:
                history = self.conversation_service.get_conversation_history(conversation_id)
            lang_code = 'ar' if re.search(r'[\u0600-\u06FF]', question) else 'en'
            #Update: handle the new response from the router which is a dict with route, original question, translated question, and original language
            #Update: replace the question with the translated question from the router response, also replace the route with the route from the response
            #Update: pass the original language to the synthesis function to generate the answer in the same language

            router_response_dict = self._get_route(question)
            route = router_response_dict.get("route", "GENERAL_AGENT").strip().upper()
            question = router_response_dict.get("translated_question", question)
            original_language = router_response_dict.get("original_language", "english")
            original_question = router_response_dict.get("original_question", question)

            response_data = {}

            if route == "HARDCODED_COMPANY_NAME":
                logger.info("Executing hard-coded path for company name.")
                company_name = self.sql_agent.get_company_name(company_id)
                answer = f"Your company name is {company_name}."
                response_data = {"answer": answer, "metadata": {"query_type": "hard-coded",
                                                                "sql_query": f"SELECT companyname FROM company WHERE id = {company_id};"}}

            elif route == "SQL_AGENT":
                logger.info(f"Routing to SQL_AGENT for company_id: {company_id}")
                data = self.sql_agent.get_sql_data(question, company_id)
                answer = self._synthesize_answer(question, data["sql_result"], history, original_language)
                response_data = {"answer": answer,
                                 "metadata": {"query_type": "sql", "sql_query": data.get("sql_query"), "from_cache": data.get("from_cache", False)}}

            elif route == "DOCX_AGENT":
                logger.info(f"Routing to DOCX_AGENT for company_id: {company_id}")
                data = self.sql_agent.get_sql_data(question, company_id)
                doc_info = self.docx_agent.create_document_from_sql_result(data["sql_result"], title=question)

                download_token = os.path.basename(doc_info["filepath"]).replace('.docx', '')
                self.cache_service.set(f"download:{download_token}", doc_info["filepath"], expiry_seconds=86400)

                download_url = f"{str(request.base_url)}download/{download_token}" if request else None

                answer = doc_info["message"]
                response_data = {
                    "answer": answer,
                    "metadata": {
                        "query_type": "docx",
                        "download_token": download_token,
                        "download_filename": doc_info["filename"],
                        "download_url": download_url
                    }
                }
            elif route == "PDF_AGENT":
                logger.info("Routing to PDF_AGENT")
                data = self.sql_agent.get_sql_data(question, company_id)
                pdf_info = self.docx_agent.create_document_from_sql_result(  # reuse existing docx_agent logic for table
                    data["sql_result"], title=question
                )
                # just convert docx -> text string for simplicity
                from services.pdf_agent import PdfAgent
                pdf_agent = PdfAgent()
                pdf_info = pdf_agent.generate_pdf(question, pdf_info["filepath"])

                token = os.path.basename(pdf_info["filepath"]).replace(".pdf", "")
                self.cache_service.set(f"download:{token}", pdf_info["filepath"], expiry_seconds=86400)
                download_url = f"{str(request.base_url)}download/{token}" if request else None

                answer = pdf_info["message"]
                response_data = {
                    "answer": answer,
                    "metadata": {
                        "query_type": "pdf",
                        "download_token": token,
                        "download_filename": pdf_info["filename"],
                        "download_url": download_url
                    }
                }

            else:  # GENERAL_AGENT
                logger.info(f"Routing to GENERAL_AGENT")
                answer = self._synthesize_answer(question, self.general_agent.respond(question=question), history, original_language)
                response_data = {"answer": answer, "metadata": {"query_type": "general"}}

            if conversation_id:
                self.conversation_service.add_message_to_conversation(
                    conversation_id=conversation_id,
                    question=original_question,
                    answer=response_data["answer"],
                    metadata=response_data["metadata"]
                )

            final_response = {
                "answer": response_data["answer"],
                "query_type": response_data["metadata"].get("query_type"),
                "company_id": str(company_id) if company_id is not None else None,
                "execution_time_ms": round((time.time() - start_time) * 1000, 2),
                "from_cache": False,
                **response_data["metadata"]
            }
            try:
                self.cache_service.xadd(
                    "qa_log",
                    {"prompt": original_question, "answer": response_data["answer"]}
                )
            except Exception as exc:
                logger.warning("could not push qa_log: %s", exc)
            return final_response

        except HTTPException:
            raise
        except Exception as e:
            logger.error(f"Router error: {e}", exc_info=True)
            raise HTTPException(
                status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
                detail="AI processing failed. Please try again."
            )