# services/sql_agent.py

from collections import Counter
import json
import os
import random
import re
import logging
from typing import List, Optional, Dict, Set, Any
import ast

import faiss
from fastapi import HTTPException
from sentence_transformers import SentenceTransformer
from sqlalchemy import inspect, text
from langchain_ollama import OllamaLLM
from langchain_community.utilities import SQLDatabase
from langchain_core.prompts import PromptTemplate
from langchain_core.runnables import RunnablePassthrough
from langchain_core.output_parsers import StrOutputParser
import transformers
import sqlparse
from sqlparse.sql import Identifier, IdentifierList
import time


from utils.erp_context import get_table_info
from services.cache_service import CacheService

logger = logging.getLogger(__name__)

sqlCacheKey = "sql_agent_cache"
_embedder = None
def get_embedder():
    global _embedder
    if _embedder is None:
        _embedder = SentenceTransformer("/opt/models/bge-m3", device="cpu")
    return _embedder
embedder = get_embedder()
weekly_vec_root = "vector_store/users_vectors"
cache_service = CacheService()
table_vec_path = "vector_store/table_descriptions.index"
column_vec_path = "vector_store/column_descriptions.index"


class SQLAgent:
    """Handles database queries for a single authenticated company."""

    def __init__(self, db: SQLDatabase, sql_llm: OllamaLLM, ollama_llm: OllamaLLM, translator: transformers.pipelines.Pipeline, cache_service: CacheService, *, max_retries: int = 2):
        self.db = db
        self.sql_llm = sql_llm
        self.ollama_llm = ollama_llm
        self.translator = translator
        self.cache_service = cache_service
        self.max_retries = max_retries
        self.search_config = {
            'schema_top_k': 10,  
            'cache_top_k': 6,   # Increased from 2
            'cache_similarity_threshold': 0.75,  # Minimum similarity score for cache
            'similarity_threshold': 0.2,  # Minimum similarity score for schema
            'query_expansion_enabled': True, #used to search for tables using different words
            'cache_enabled': False, #enable cache response
            'column_search_top_k': 15 #number of columns passed to the column selection llm for each table selected
        }
        self._setup_chains()
        self.table_index = self._load_index(table_vec_path)
        self.column_index = self._load_index(column_vec_path)
        self.table_selection_chain_set_up()
        self.column_selection_chain_set_up()
        self.cache_selection_chain_set_up()



    def _setup_chains(self) -> None:
        """Initialise the LangChain chain for complex SQL generation."""
        sql_writer_template = (
                        """
You are an expert SQL assistant. You will generate a SQL query to answer the user's question.
Use only the tables and columns provided in the table information below.

Relevant table information (table description + selected columns):
{table_info}

RULES:
1. **Only generate SELECT queries.**  
   - Never generate SHOW, INSERT, UPDATE, DELETE, DROP, CREATE, ALTER, or TRUNCATE.  
   - Aggregations like COUNT, SUM, MIN, MAX, and AVG must always be written as part of a SELECT statement, never with SHOW.  
   - Example (valid):  
     ```sql  
     SELECT COUNT(employees.id)  
     FROM employees  
     WHERE employees.company_id = {company_id};  
     ```  
   - Example (invalid):  
     ```sql  
     SHOW COUNT(employees.id) ...  
     ```  

2. Table and column names must appear EXACTLY as shown in the schema, even if they contain typos, misspellings, or unusual formatting.
   - Do NOT correct, shorten, pluralize, or modify them.
   - Example: if a table is named `companies`, you must use `companies` (not `users_table`).
   - Example: if a column is `employee.id`, you must use `employee.id` (not just `id`).


3. Do NOT assume or invent values. Only use the provided company_id value: {company_id}.

4. Use proper JOINs only when necessary, following the foreign key relationships.

5. Retrieve ONLY the columns explicitly provided in the table information. Do NOT add extra columns.

6. Output format:
   - Only SQL inside a markdown block:
     ```sql
     SELECT ...
     FROM ...
     WHERE ...;
     ```
   - No explanations, no comments, no extra words.

7. **Always include the primary key column in the SELECT clause if provided.**
   - If a table has a primary key (PK) listed in the table information, it must appear in the SELECT statement even if the user did not explicitly ask for it.

**User Question:** {question}
**SQL Query:**
"""
        )
        sql_prompt = PromptTemplate.from_template(sql_writer_template)

        self.sql_query_chain = (
                sql_prompt
                | self.sql_llm
                | StrOutputParser()
        )

    def _get_index_path(self, user_id): 
        prefix = str(user_id)[:2]
        folder = os.path.join(weekly_vec_root, prefix)
        os.makedirs(folder, exist_ok=True)
        return os.path.join(folder, f"{user_id}.index")
    
    def _load_index(self, path: str) -> faiss.IndexFlatL2:
        if os.path.exists(path):
            return faiss.read_index(path)
        dim = embedder.get_sentence_embedding_dimension()
        return faiss.IndexFlatL2(dim)

     
    def cache_selection_chain_set_up(self) -> None:
        """Set up the LLM chain for selecting relevant cached question based on the user's question."""
        selection_prompt_template = ("""
You are an expert AI assistant. 

Task:
You are given a user's new question and a list of cached questions with their corresponding SQL queries. Your goal is to identify which SQL query would best answer the user's question.

Input format:
- question: {{cached question a}}
  query: {{cached sql a}}

- question: {{cached question b}}
  query: {{cached sql b}}

output example:
{{cached question a}}
...

questions list:
{questions_list}

User's Question:
{user_question}

Output:
- Return **only the cached question** that corresponds to the SQL query most likely to answer the user's question.
- If none of the cached queries is suitable, return exactly: "no match".
- Do not modify the question or the query. No extra text or explanation.


""")
        
        cache_selection_prompt = PromptTemplate.from_template(selection_prompt_template)

        self.cache_selection_chain = ( 
                cache_selection_prompt
                | self.ollama_llm
                | StrOutputParser()
        )
    
    def table_selection_chain_set_up(self) -> None:
        selection_prompt_template = ("""
            You are an expert database assistant.  
            Your task is to select which tables are needed to answer the user’s question based on the table's description provided.

                                                
            The user asks: "{user_question}"

            Candidate Tables and Descriptions:
            {table_description}


            Your task:
            1. Select the tables likely to contain information relevant to the user’s question  
                - If you are unsure, include the table.  
                - If multiple tables likely to be relevant, include all of them.  
            2. Use table names, column names (if included), and descriptions to decide relevance.  
            3. Match semantic meaning, synonyms, or related terms. Examples:  
                - "phone number" → "contact info" or "phone" in description  
                - "employee" → "user", "staff", "personnel"  
                - "address" → "location", "city", "residence"  
            4. Do not invent or hallucinate tables or columns. Only choose from the candidate tables.  
            5. Table names must appear EXACTLY as in the schema, even if they contain typos.  
            6. Output only a valid **JSON array of strings** (table names), with no extra text, explanation, or formatting.
            
            - If multiple tables have the same column name, choose the table that logically stores the data for the entity mentioned in the question.
            - if a user asks about a company, you must use the "companies" table, not "users", even if both have a "name" column.

            Example:
            ["companies", "customers", "orders"]
                                                
            Now return ONLY the JSON array:
            """)
        
        selection_prompt = PromptTemplate.from_template(selection_prompt_template)

        self.selection_chain = (
                selection_prompt
                | self.ollama_llm
                | StrOutputParser()
        )

    def column_selection_chain_set_up(self) -> None:
        """
    Set up the LLM chain for selecting relevant columns from the selected tables
    based on the user's question.
    """

        selection_prompt_template = ("""
You are an expert database assistant.
Your task is to select the minimum set of columns from the provided tables
that are required to answer the user's question.

User question: "{user_question}"

Candidate Tables, Descriptions, and Columns:
{table_column_description}

Selection Rules:
1. Choose **only columns that are clearly required** to answer the question.  
   - If not directly relevant, exclude it.  
   - If uncertain, exclude it.  
   - Do not use columns outside the provided tables.
2. Match semantic meaning, synonyms, and related terms.  
   Examples:  
      - "phone number" → column "phone" or "contact_number"  
      - "employee" → columns like "user_id", "employee_name", "staff_id"
3. **Do not invent or guess** columns, values, or tables.
4. **Mandatory columns:**  
   - Always include the **primary key** (`id`) of each table.  
   - Always include **foreign key columns** (columns ending with `_id`) if they link to other selected tables.  
   - Always include a column ending in **company_id** if it exists in that table.
5. Output format (strict):  
   - A single JSON object mapping table names to arrays of fully qualified column names.  
   - Example:  
     {{
       "customers": ["customers.id", "customers.name", "customers.phone"],
       "orders": ["orders.id", "orders.customer_id", "orders.total"]
     }}
   - No explanations, comments, Markdown, or extra text.
   - Do not output multiple JSON objects or nested structures.

Output:  
Return **only** one JSON object in the exact format specified above.

        """)

        selection_prompt = PromptTemplate.from_template(selection_prompt_template)

        self.column_selection_chain = (
            selection_prompt
            | self.ollama_llm
            | StrOutputParser()
        )


    def table_selection_using_llm(self, question: str, response_dict: dict[str,dict[str,dict[str,any]]]) -> dict[str, dict[str, dict[str, any]]]:
        """
        Use an LLM to understand and select tables that contain data we need from the retrieval function.

        Args:
            question (str): The user's question.
            response_dict (dict): A dictionary containing tables similar to the question.
                It follows the same format as the output of `retrieve_relevant_schema_chunks`.

        Returns:
            dict: A dictionary in the same format as `response_dict`, 
                but containing only the selected tables.
        """

        table_description = ""
        for table_name, table_content in response_dict.items():
            table_description+= f'\n\nTable name: {table_name}\nTable Description: {table_content['table_description']}'

        #llm call will return a list of tables in a string type 
        table_names = self.selection_chain.invoke({
                "user_question": question,
                "table_description": table_description #string form of table name and description 
            })
        #make it ready to filter the rsponse
        table_names = table_names.removeprefix("[").removesuffix("]").replace('"', '').replace("'", '') #get it ready to convert into acutal list

        table_names = table_names.split(", ") #convert to a list

        table_names = [tbl.strip() for tbl in table_names if tbl.strip()] #remove extra white spaces

        selected_with_columns = {}
        for tbl_name in table_names:
            selected_with_columns[tbl_name] = response_dict.get(tbl_name, {"table_description": "N/A", "table_columns": {}, "reverse_mapping": {}}) #pass the table info for selected ones 
        
        #for debugging
        #logger.info(f"selected tables from tables selection {", ".join(list(selected_with_columns.keys()))}")


        return selected_with_columns 

    def column_selection_using_index(self, question: str, selected_tables: dict[str, dict[str, dict[str, any]]]) -> dict[str, dict[str, dict[str, any]]]:
        """
        Uses the semantic and keyword search to select relevant columns for the tables already chosen.

        Args:
            question (str): The user’s question.
            selected_tables (dict): Output from table_selection_using_llm. 

        Returns:
            response_dict (dict): A dictionary containing all info of the selected table.
            
            Example structure:
                tbl_name: { 
                            'table_description': '-',
                            'table_columns': {
                                table_column_name:{
                                    'description': '-',
                                    'column_info': '-'
                                }
                            },
                            'reverse_mapping': {
                                    'original': '-',
                                    'columns': {'-': '-'}
                                }
                            }
                        }
                    }
        """
        #get all table names that are selected to add related FKs in ids step
        tables_used = [tbl for tbl in selected_tables.keys()]

        #for debugging 
        #logger.info(f'this is the columns before index retrieval {'\n'.join([f"table name: {table} columns: {', '.join(list(info.get('table_columns', {}).keys()))}" for table, info in selected_tables.items()])}')

        columns_top_k = self.search_config.get('column_search_top_k', 20) #number of columns that will be passed to the llm

        q_embedding = embedder.encode(question, show_progress_bar=False).reshape(1, -1)

        if self.column_index.ntotal < 1:
            logger.warning("There isn't columns in the fiass file" )
            return selected_tables #return it as is
        
        #semantic search
        #get lots of columns as it will get all columns from all tables that are related to our question and we want to filter them later to get columns_top_k * n for each table
        D, I = self.column_index.search(q_embedding, k= min(columns_top_k * 20, self.column_index.ntotal)) #take the min value in case we got more than the ntotal of the index
        
        if I[0][0] == -1:
            logger.warning("No column found in column selection using index file")

        #it will be a list of metadatas which will be filtered based on the table selected
        all_results = []
        for dist, vec_id in zip(D[0], I[0]):
            metadata = cache_service.get(f"column_vec:{vec_id}")
            all_results.append((dist, metadata))

        #for each table take the columns related to it and return it in the same format of the input dict
        for tbl_name, tbl_info in selected_tables.items():
            new_columns = {} #this will replace the original columns info with only selected ones

            #get all column results related to the table
            tbl_results_columns = [(dist, column_result) for dist, column_result in all_results if column_result and column_result.get('table_name') == tbl_name]

            boosted_results = []
            #add a simple keyword boost to the distances
            for dist, metadata in tbl_results_columns:
                boost = 0

                #get the table name versions to reduce it's boost as all columns is from the same table
                table_name: str = metadata.get('table_name', '')
                table_name_versions = [table_name, table_name.removesuffix('s')]

                column_name: str = metadata.get('column_name')
                column_name = column_name.replace('_', ' ')
                column_name = column_name.split('.')[1] #the column name in -> tbl_name.column_name
                column_name_version = column_name.split() #get the column name in parts to improve the boost

                #the column description which will help us get the number of terms
                column_description: str = metadata.get('column_description', '')

                #count how many a term is repeated in a column
                all_terms = Counter(re.findall( r'\w+', column_description))

                #normalize question and remove any non alphabet char
                question_norm = re.sub(r'[^a-zA-Z\s]', '', question.lower())

                question_set = set(question_norm.split()) #as we don't want duplicates

                for term in question_set:
                    if not term.strip(): continue
                    #as we do not want to calculate as all comes from the same table
                    if term.strip() in table_name_versions: continue
                    #keyword boost for term in the question
                    boost += min(all_terms.get(term, 0) * 0.05, 0.2)

                    #only apply the boost ones for each part of the column name
                    if term.strip() in column_name_version: 
                        boost += 0.1
                
                #finally add the boost (we need to reduce the distance)
                boosted_results.append((dist - boost, metadata))
                

            #sort it based on distance for the current table in this iteration
            sorted_results = sorted(boosted_results, key= lambda x: x[0])

            #take only partion of the results of each table
            sorted_results = sorted_results[:columns_top_k]

            #complete the columns randomly if it was less than columns_top_k
            if len(sorted_results) < columns_top_k:
                column_keys = list(tbl_info.get('table_columns', {}).keys())
                random.shuffle(column_keys) #shuffle in place 

                for column_name in column_keys:
                    #we completed the list
                    if len(sorted_results) == columns_top_k:
                        break

                    #we will add the ids later so we wont do it now
                    if column_name.endswith('id'):
                        continue

                    #add the column
                    sorted_results.append(( 0, column_name)) #the distance will be 0 as we don't need it anymore

            #to remove duplicates, places as is if it was a string with the name, if it was a dict then take the column name only
            unique_columns = {col.get('column_name') if (isinstance(col, dict)) else col for _, col in sorted_results}

            #ensure that the pk and any fk is there
            for col_name in tbl_info.get('table_columns', {}).keys():
                if not col_name: continue
                if col_name in unique_columns: continue #no duplicates

                if col_name.endswith('_id') and any(
                    re.search(fr'\b{tbl}\b', selected_tables.get(tbl_name, {}).get('table_columns', {}).get(col_name, {}).get('column_info', ''), re.IGNORECASE) 
                    for tbl in tables_used
                    ): 
                    #all FKs for only retrieved tables using re search function, search in this particular column info (fk info if there was)
                    sorted_results.append((0, col_name)) #the distance will be 0 as we don't need it anymore
                    unique_columns.add(col_name) #to make sure there isn't any duplicates
                    continue

                if col_name.endswith('.id'): #as column names is tbl_name.col_name - > tbl_name.id
                        sorted_results.insert(0, (0, col_name)) #the distance will be 0 as we don't need it anymore
                        unique_columns.add(col_name) #to make sure there isn't any duplicates
                        continue


                if col_name.endswith('company_id'):
                    sorted_results.append((0, col_name)) #the distance will be 0 as we don't need it anymore
                    unique_columns.add(col_name) #to make sure there isn't any duplicates
                

            #construct new columns dict with it's info to replace original one it will be in the same format of the original one
            for _, col in sorted_results:
                #make sure it is the name not a result from metadata
                col = col.get('column_name') if (isinstance(col, dict)) else col
                if col and col.strip():
                    new_columns[col] = tbl_info.get('table_columns', {}).get(col, {})
        
            #replace the original list of columns 
            selected_tables[tbl_name]['table_columns'] = new_columns
        
        #for debugging
        #logger.info(f'this is the columns index retrieval {'\n'.join([f"table name: {table} columns: {', '.join(list(info.get('table_columns', {}).keys()))}" for table, info in selected_tables.items()])}')

        return selected_tables


    def column_selection_using_llm(self, question: str, selected_tables: dict[str, dict[str, dict[str, any]]]) -> dict[str, dict[str, dict[str, any]]]:
        """
        Uses the LLM to select relevant columns for the tables already chosen.

        Args:
            question (str): The user’s question.
            selected_tables (dict): Output from table_selection_using_llm. 

        Returns:
            dict: in same format of the selected_tables but with only selected tables (only reverse_mapping will be the same)
        """

        # Build the table + column description string for the prompt
        table_column_description = ""
        for table_name, table_info in selected_tables.items():
            cols_desc = "\n\n".join(f"{col_name}\n  info: {col_desc['column_info']}\n  Description: {col_desc['description']}" for col_name, col_desc in table_info["table_columns"].items())
            table_column_description += f"Table: {table_name}\nColumns:\n{cols_desc}\n\n"

        # Prepare the LLM prompt
        prompt_input = {
            "user_question": question,
            "table_column_description": table_column_description
        }

        # Call the LLM; it should return a JSON like:
        #        "customers": ["customers.id", "customers.name", "customers.phone"],
        #         "orders": ["orders.id", "orders.customer_id", "orders.total"]

        selected_columns_json = self.column_selection_chain.invoke(prompt_input)
        # Parse the JSON safely
        try:
            selected_columns = json.loads(selected_columns_json)
        except Exception as e:
            logger.error(f"Error parsing LLM output of the column selection: {e}")
            selected_columns = {}

        # Ensure selected_columns is a dict of table -> list[str]
        if isinstance(selected_columns, list):
            normalized = {}
            for entry in selected_columns:
                if isinstance(entry, dict):
                    for table, cols in entry.items():
                        if table and isinstance(cols, list):
                            normalized[table] = cols
            selected_columns = normalized

        final_selection = {}
        for table, cols in selected_columns.items():
            # normalize to list of strings if LLM returned dicts (some type of hallucination i encountered) [{'name': col}, ...]
            if isinstance(cols, list) and all(isinstance(c, dict) and 'name' in c for c in cols):
                cols = [c['name'] for c in cols]

            if not isinstance(cols, list): #in case of wrong outputs from the ai
                logger.warning(f"Skipping table {table}: invalid columns format")
                continue

            col_info = selected_tables.get(table, {}).get("table_columns", {})
            reverse_mapping = selected_tables.get(table, {}).get("reverse_mapping", {}) #get the dict that contain the real names to be replaced in the last step

            if table in selected_tables:
                #get only valid column names
                valid_cols = [col for col in cols if col in selected_tables[table]["table_columns"].keys()]

                #inject anything with company_id if it wasn't there or 'id' in case of the company id "uses mapped names"
                company_cols = [c for c in selected_tables[table]["table_columns"].keys() if (c.endswith("company_id")) or (table == 'companies' and c.endswith('.id'))]
                for c in company_cols:
                    if c not in valid_cols:
                        valid_cols.append(c)

                #ensure the pk is in the columns
                pk = f'{table}.id'
                if pk not in valid_cols:
                    valid_cols.append(pk)
                
                #build a dict of columns with the column description to be used in the sql generation prompt
                pk_cols = {c: cinfo.get("column_info", "").lower() for c, cinfo in col_info.items() if cinfo.get("column_info", "") and c in valid_cols}
                
                final_selection[table] = {
                    "table_description": selected_tables[table]['table_description'],
                    "table_columns": pk_cols,
                    "reverse_mapping": reverse_mapping
                }
            else:
                logger.warning(f"Table {table} from LLM output not in selected tables.") #do not place it then

        
        #for debugging
        #logger.info(f"last selection for column llm")
        return final_selection 


    STOP_WORDS = {'a', 'an', 'the', 'is', 'in', 'it', 'for', 'of', 'show', 'me', 'list', 'what', 'who', 'and', 'in', 'on', 'at', 'by', 'for', 'of', 'with'}

    def retrieve_relevant_schema_chunks(self, question: str, top_k: int = None) -> dict[str,dict[str,any]]:
        """
        Hybrid retrieval: semantic search on table embeddings + keyword-based scoring on cached columns.
        Works even without column embeddings/index—only tables. Then uses semantic and keyword search on descriptions.
        Important note: it will only use one chunk of the table description (mostly enough, but keep in mind).

        Args:
            question (str): The user's question used to search for a relevant table.
            top_k (int): How many tables to return for AI selection.

        Returns:
            response_dict (dict): A dictionary containing all info of the selected table.
            
            Example structure:
                {
                    tbl_name: { 
                        'table_description': '-',
                        'table_columns': {
                                table_column_name:{
                                    'description': '-',
                                    'column_info': '-'
                                }
                            },
                        'reverse_mapping': {
                                'original': '-',
                                'columns': {'-': '-'}
                            }
                        }
                    }
                }
        """
        if top_k is None:
            top_k = self.search_config['schema_top_k']

        all_results = []
        all_distances = []
        table_columns = {}
        tables_found = {}

        # Query expansion if enabled
        queries = [question]
        if self.search_config.get('query_expansion_enabled', False):
            queries = self._expand_query(question)

        #use semantic search and calculate how similar it is to the text and exclude less similar ones
        for query in queries:
            q_embedding = embedder.encode(query, show_progress_bar=False).reshape(1, -1)
            D, I = self.table_index.search(q_embedding, k=top_k * 3)  # more candidates for boosting and getting more results to apply keyword search

            if I[0][0] == -1: #no similar tables found
                continue

            #for each table found from the search method, take it and add it to the result to be filtered more later
            for i, vec_id in enumerate(I[0]):
                distance = D[0][i]
                similarity = 1.0 / (1.0 + distance)
                if similarity < self.search_config.get('similarity_threshold', 0.2):
                    continue

                metadata = cache_service.get(f"table_vec:{vec_id}")
                if not metadata:
                    continue

                table_name = metadata.get("table_name", "")
                if table_name in tables_found:
                    continue

                all_results.append(metadata) #table info (saved in redis -> task.py)
                all_distances.append(distance) #result of how close it is to the table?
                tables_found[table_name] = metadata.get("table_description", "-")
                table_columns[table_name] = metadata.get("columns", {})

        if not tables_found:
            logger.error("\n---No relevant schema found---\n")
            return {}

        #added Column Keyword Match to the semantic results
        # convert the question to a version that does not have stop words (used allot in any sentance, as we are going to search within the description of the column not only the name)
        question_search_version = ' '.join([word for word in question.strip().lower().split() if word not in self.STOP_WORDS])
        question_terms = re.findall(r'\w+', question_search_version)
        boosted_results = []

        for result, distance in zip(all_results, all_distances):
            table_name: str = result["table_name"]
            table_name_versions = [table_name, table_name.removesuffix('s')] #to get the table name and a version without 's' such as companys -> company
            columns = table_columns[table_name]
            boost = 0.0

            # Keyword matching on column names and descriptions (BM25-style simple boost)
            col_texts = " ".join([f"{col} {desc.get('description', '')}" for col, desc in columns.items() if desc])
            col_texts = col_texts.lower()
            term_counts = Counter(re.findall(r'\w+', col_texts))
            for term in question_terms:
                boost += min(0.2, term_counts.get(term, 0) * 0.05)  # max 0.2 boost per term

            #huge boost in case of table name found in the question
            for version in table_name_versions:
                if re.search(fr'\b{version}\b', question_search_version):
                    boost += (0.3)
                    break

            boosted_results.append((result, distance - boost))  # lower distance is better

        # Sort by adjusted distance
        boosted_results.sort(key=lambda x: x[1])

        # Build final response, you could find out the keys in task.py cache service set
        response_dict = {}
        for result, _ in boosted_results[:top_k]: #take best k values and make it into dict to be passed into the table selection
            table_name = result["table_name"]
            response_dict[table_name] = {
                "table_description": result.get("table_description", "-"),
                "table_columns": table_columns[table_name],
                "reverse_mapping": result.get("reverse_mapping", {}) #to remap to original values
            }
        
        #for debugging        
        #logger.info(f"selected tables from vectors {", ".join(list(response_dict.keys()))}")

        return response_dict 


    def _expand_query(self, question: str) -> List[str]:
        """ Expand query with synonyms and related terms for better search coverage."""
        expanded_queries = [question]
        
        # Common business term expansions
        expansions = {
            'sales': ['revenue', 'income', 'earnings', 'transactions'],
            'customers': ['clients', 'buyers', 'accounts'],
            'products': ['items', 'goods', 'merchandise', 'inventory'],
            'orders': ['purchases', 'transactions', 'sales'],
            'employees': ['staff', 'workers', 'personnel', 'team'],
            'profit': ['earnings', 'income', 'revenue', 'gains'],
            'cost': ['expense', 'spending', 'expenditure', 'price'],
            'resident' : ['muqeem', 'muqeem', 'iqama'],
            'company' : ['organization']
        }
        
        question_lower = question.lower()
        for term, synonyms in expansions.items():
            if term in question_lower:
                for synonym in synonyms:
                    expanded_queries.append(question.replace(term, synonym))
                    
        return expanded_queries[:3]  # Limit to avoid too many queries

    def get_table_aliases(self, parsed_statement):
        """this function will be used to get table names/alias which contain a recursive function in case of subqueries"""
        alias_map = {}

        def _extract(tokens):
            for token in tokens:
                if isinstance(token, Identifier): #in cse it was a simple query (or a sub query after a recursive call) it will map table names/alias
                    if token.get_real_name().lower() in ['count', 'sum', 'min', 'max', 'avg']:
                        continue
                    table_name = token.get_real_name()
                    alias = token.get_alias()  # get alias if exists, else None. For the table remapping to work
                    if table_name:
                        alias_map[table_name] = alias
                #other cases (functions, sub quereies, etc..)
                elif isinstance(token, IdentifierList): 
                    _extract(token.get_identifiers())
                elif token.is_group:
                    _extract(token.tokens)

        _extract(parsed_statement.tokens)
        return alias_map

    def remap_sql_using_selected_tables(self, sql_query: str, selected_tables: dict) -> str:
        """
        Remap friendly table/column names to original database names using selected tables.
        Handles aliases, subqueries and more.

        Args:
            sql_query (str): the one that is generated by the SQL LLM
            selected_tables (dict): contains info of the table selected like the columns, reversed mappings, 
                it will be in the same format of 'retrieve_relevant_schema_chunks' function output example
            
        Returns:
            sql_query(str): corrected sql query with the correct table/column names
        """


        #normalize SQL and aliases (remove quotes for consistency)
        sql_query_clean = sql_query.replace('"', '').replace('`', '')

        parsed = sqlparse.parse(sql_query_clean)
        statement = parsed[0]
        
        #get alias map for remapping in case of alias (also it will get the table names in case of alias to map the columns related)
        alias_map = self.get_table_aliases(statement) 

        #remapping to original values
        for table, tbl_info in selected_tables.items():
            real_table = tbl_info["reverse_mapping"]["original"]
            col_map = tbl_info["reverse_mapping"]["columns"]

            #replace column names using re 
            #columns after the last change will always be in form of tbl_name.col which will simplify remapping and sql agent results
            for friendly_col, real_col in col_map.items():
                alias = alias_map.get(table) #if not there then return None
                if alias:
                    friendly_col = friendly_col.replace(f"{table}.",f"{alias}." )
                    real_col = real_col.replace(f"{real_table}.",f"{alias}." )
                    sql_query = re.sub(rf'\b{re.escape(friendly_col)}\b', real_col, sql_query, re.IGNORECASE)
                else:
                    sql_query = re.sub(rf'\b{re.escape(friendly_col)}\b', real_col, sql_query, re.IGNORECASE)

            #replace table names using word boundaries to avoid partial matches
            sql_query = re.sub(rf'\b{re.escape(table)}\b', real_table, sql_query, re.IGNORECASE)
        #return the remapped sql query    
        return sql_query



    def get_company_name(self, company_id: str) -> str:
        """
        Executes a hard-coded, guaranteed-safe query to get the company name.
        This method does NOT use the AI.
        """
        if not company_id:
            raise HTTPException(status_code=401, detail="Company ID is required for this operation.")
        try:
            query = text("SELECT companyname FROM company WHERE id = :company_id")
            logger.info(f"Executing hard-coded query for company ID: {company_id}")

            with self.db._engine.connect() as connection:
                result = connection.execute(query, {"company_id": int(company_id)})
                row = result.fetchone()

            if row and row[0]:
                return row[0]
            return "Company name not found."
        except Exception as e:
            logger.error(f"Error in get_company_name: {e}", exc_info=True)
            raise HTTPException(status_code=500, detail="Could not retrieve company name.")
        
    def get_answer_from_cache(self, question: str, company_id: str) -> Optional[dict[str, str]]: #retrun the metadata that is retrived from redis
        """Check if the question has been answered before and return cached answer."""

        #encode the question to be searched in the index for the answer
        q_embedding = embedder.encode(question, show_progress_bar=False)
        q_embedding = q_embedding.reshape(1, -1)

        path = self._get_index_path(company_id)

        index = self._load_index(path)

        D, I = index.search(q_embedding, k=self.search_config.get('cache_top_k', 4))  #get the closest matches that will be reduced based on the simqilarity threshold and the llm

        if I[0][0] == -1:  # no match found
            return None
    
        all_results = []
        
        for i, vec_id in enumerate(I[0]):
                distance = D[0][i]
                similarity = 1.0 / (1.0 + distance)
                if similarity < self.search_config.get('cache_similarity_threshold', 0.75):
                    continue

                metadata = self.cache_service.get(f"{company_id}_vec_meta:{vec_id}") # Get metadata for the a match

                if not metadata:
                    continue

                all_results.append((metadata, similarity))
        
        if len(all_results) == 0:  #no matches above threshold
            return None
        
        # arrange by similarity
        all_results.sort(key=lambda x: x[1], reverse=True)

        
        # Otherwise, use the LLM to decide if any of the matches are relevant
        questions_list = "\n\n".join([f"- question: {res[0].get('question', 'N/A')}\n query: {res[0].get('sql_query', "")}" for res in all_results])

        selected_question = self.cache_selection_chain.invoke({
            "questions_list": questions_list,
            "user_question": question
        })

        # If LLM says "no match", return None
        if selected_question.strip().lower() == "no match":
            logger.info("LLM determined no relevant cached answer.")
            return None
        
        # Find the corresponding answer of the selected question by the LLM
        for res, _ in all_results: 
            if res.get("question", "").strip().lower() == selected_question.strip().lower():
                return res
                   
        return None

    def ensure_one_query(self, query: str) -> str | None:
        """check only the first statement to ensure there isn't more than one"""
        #if it was empty
        if not query:
            logger.error('query is an empty string')
            return None

        #parse it to see the type of the query and to only use one
        statements = sqlparse.parse(query)

        #get the first sql query if there was more than one
        statement = statements[0]

        #if it wan't sql query at all then it will return UNKNOWN
        if statement.get_type().strip().lower() != 'select':
            logger.error("The sql llm failed to generate a select function")
            return None
        
        return statement

    def run_query_in_db(self, db: SQLDatabase, query: str, row_limit = 200) -> list:
        """
        this function simply return the results of the query but with column names as this will make the llm results better, 
        also ensures that we only use the first sql query, also ensures that it uses select in it
        """
        #remove the leading <s> that is generated by llm sometimes also normalize the string
        query = re.sub(r'<s>', '', query.strip(), re.IGNORECASE)
        query = re.sub(r"\s+", " ", query.strip()) #make it one space rather than 2+

        #takes the first query if it has more than one, also ensures it is select
        first_correct_query = self.ensure_one_query(query)

        if first_correct_query is None:
            return []

        try:
            with db._engine.connect() as conn:
                result = conn.execution_options(timeout = 10).execute(text(str(first_correct_query)))
                rows = result.fetchmany(row_limit)
                cols = result.keys()
            
            return [dict(zip(cols, row)) for row in rows]
        
        except Exception as e:
            logger.error(f"SQL execution error: {e}")
            return []



    def get_sql_data(self, question: str, company_id: Optional[str] = None) -> Dict[str, Any]:
        """Generate, validate, and run a SQL query for complex questions."""
        #inject 'where comapny id = {company_id}' in the question to make sure the llm will always consider it
        if company_id:
            if "company id" not in question.lower():
                question += f"\n(where the company id is {company_id})"

        #before using the query to generate an answer, check if the question is already answered in the cache, if yes return it directly without any llm and sql generation calls
        if company_id and self.search_config.get('cache_enabled', False):
            cached_response = self.get_answer_from_cache(question, company_id)
            if cached_response:
                logger.info("Returning cached response for question: %s", question)
                return {
                    "sql_query": cached_response["sql_query"],
                    "sql_result": cached_response["answer"],
                    "from_cache": True
                    }

        clean_sql = ""
        raw_sql = ""
        attempts = 0
        last_error = ""

        while attempts <= self.max_retries:
            prompt_question = question
            if last_error:
                prompt_question = (
                    f"{question}\n\n"
                    f"Your previous attempt had an issue. Hint: {self._generate_friendly_hint(last_error)}\n"
                    f"Here is your previous query:\n{clean_sql}\n"
                    "Please correct it using only the tables and columns provided, always using table aliases."
                )

            relievant_dict = self.retrieve_relevant_schema_chunks(question) #get relievant tables 
            selected_tables = self.table_selection_using_llm(question, relievant_dict) #select tables (same output form of relievant_dict)
            reduced_selected_tables_and_columns = self.column_selection_using_index(question, selected_tables)
            selected_tables_and_columns = self.column_selection_using_llm(question, reduced_selected_tables_and_columns) #select columns within tables (same output form of selected_tables)

            # Convert selected tables and columns to string for the SQL prompt
            tables_info_for_prompt = ""
            for table_name, table_info in selected_tables_and_columns.items():
                cols = table_info.get("table_columns", [])
                if not cols:
                    continue
                cols_str = "\n".join([f"- {col}: {info}" for col, info in table_info["table_columns"].items()])
                tables_info_for_prompt += f"Table: {table_name}\ncolumns:\n{cols_str}\n\n"
            raw_sql = self.sql_query_chain.invoke({
                "table_info": tables_info_for_prompt.strip(),
                "question": prompt_question,
                "company_id": company_id
            })
            #for debugging
            #logger.info(f"this is the final llm table appearnce: {tables_info_for_prompt.strip()}")
            try:
                clean_sql = self._extract_select(raw_sql) #before remapping to be used in error feedback for the next itteration
                #for debugging
                #logger.info(f'before mapping {clean_sql}') 
                correct_clean_sql = self.remap_sql_using_selected_tables(clean_sql, selected_tables_and_columns) #remapping to use in db
                #for debugging
                #logger.info(f"aftter mapping {correct_clean_sql}")
                #Note: the first one we are using the query after remapping to correct values but the second function we are using the query before remapping to correct values
                validation_error = self._validate_identifiers(correct_clean_sql) or self._ensure_company_filter(clean_sql,company_id) #check if the remapped validated or the mapped has company_id that is correct

                if validation_error:
                    friendly_hint = self._generate_friendly_hint(validation_error)
                    last_error = friendly_hint
                    raise ValueError(validation_error)
                
                logger.info("Executing SQL: %s", correct_clean_sql)


                #update: changed it to execute
                #run -> faster but does not return any info about the columns that it uses
                #execute -> slower but returns info we need to show the results to the llm and make it understand
                sql_result = self.run_query_in_db(self.db, correct_clean_sql) # -> list of  dicts with column_name: value values

                #convert the list to a string 
                sql_result = json.dumps(sql_result)


                if not sql_result:
                    logger.warning("SQL query returned no results.")
                else: #save to cache before executing 
                    if self.cache_service:
                        self.cache_service.xadd (sqlCacheKey, {
                            "company_id": company_id,
                            "question": question,
                            "sql_query": clean_sql, #save the friendly sql for better understanding for the llm in cache selection
                            "result": sql_result
                        })


                logger.info("SQL query executed successfully.")
                #for debugging
                #logger.info(sql_result)
                return {"sql_query": correct_clean_sql, "sql_result": sql_result, "from_cache": False}

            except (ValueError, Exception) as e:
                attempts += 1
                last_error = str(e)
                logger.warning(f"Attempt {attempts}/{self.max_retries + 1} failed. Retrying due to: {last_error}")
                if attempts > self.max_retries:
                    logger.error("SQLAgent failed after max retries.", exc_info=True)
                    if raw_sql: logger.error("Last problematic SQL from LLM:\n%s", raw_sql)
                    raise HTTPException(status_code=500,
                                        detail="I'm sorry, I'm experiencing technical difficulties. Please try again later.") from e

        raise HTTPException(status_code=500, detail="An unexpected error occurred in the AI agent.")

    _SELECT_RE = re.compile(r"SELECT.*?;", re.I | re.S)
    _COLUMN_REF_RE = re.compile(r"([\w\"]+)\.([\w\"]+)", re.I)

    def _extract_select(self, text: str) -> str:
        """Extracts a SQL query from a markdown block or raw text."""
        # --- Start of Final Fix for SQL Extraction ---
        # Look for a SQL markdown block first
        match = re.search(r"```sql\n(.*?)\n```", text, re.DOTALL)
        if match:
            return match.group(1).strip()

        # Fallback for raw SQL without markdown
        match = self._SELECT_RE.search(text)
        if not match:
            raise ValueError("AI failed to generate a valid SELECT query.")
        return match.group(0).strip()
        # --- End of Final Fix ---

    def _generate_friendly_hint(self, validation_msg: str) -> str:
        """Convert technical validation errors into user-friendly hints."""
        if "Unknown table alias" in validation_msg:
            return "The query uses an unknown table alias. Check the table aliases."
        if "doesn't exist on table" in validation_msg:
            return "A column in your query does not exist in the referenced table. Verify column names."
        if "Could not find any tables" in validation_msg:
            return "No tables could be found in the query. Ensure you are using the provided tables."
        if "but it does not match the current company" in validation_msg:
            return validation_msg
        if "must include a company filter" in validation_msg:
            return validation_msg
        return "There is an error in the SQL query. Please check table and column names and ensure that you are using aliases correctly."




    def _parse_aliases(self, sql: str) -> Dict[str, str]:
        alias_map: Dict[str, str] = {}
        pattern = re.compile(r'\b(?:FROM|JOIN)\s+([^\s,]+)(?:\s+(?:AS\s+)?([^\s,;]+))?', re.IGNORECASE)
        keywords = {'on', 'where', 'group', 'order', 'left', 'right', 'inner', 'join', 'limit', 'using'}

        for table, alias in pattern.findall(sql):
            table = table.strip('"')
            if alias and alias.lower() not in keywords:
                alias = alias.strip('"')
                alias_map[alias.lower()] = table.lower()
            alias_map[table.lower()] = table.lower()
        return alias_map

    def _validate_identifiers(self, sql: str) -> Optional[str]:
        alias_map = self._parse_aliases(sql)
        if not alias_map: return "Could not find any tables in the query."
        inspector = inspect(self.db._engine)
        cache: Dict[str, Set[str]] = {}
        for alias, column in self._COLUMN_REF_RE.findall(sql):
            alias_clean = alias.strip('"').lower()
            column_clean = column.strip('"').lower()
            tbl = alias_map.get(alias_clean)
            if tbl is None: return f"Unknown table alias '{alias}'"
            if tbl not in cache:
                try:
                    cache[tbl] = {c["name"].lower() for c in inspector.get_columns(tbl)}
                except Exception:
                    return f"Could not inspect columns for table '{tbl}'."
            if column_clean not in cache[tbl]: return f"Column '{column}' doesn't exist on table '{tbl}'."
        return None

    def _ensure_company_filter(self, sql: str, company_id: Optional[str]) -> Optional[str]:
        """
        Ensure that the SQL query filters by the authenticated company_id. ether in company table where it's column "id" or any other table where it ends with "company_id"
        
        Args:
            sql (str): The SQL query to validate.
            company_id (Optional[str]): The authenticated company's ID.
        
        Returns:
            Optional[str]: An error message if validation fails, or None if it passes.
        """
        if not company_id:
            return None

        sql_lower = sql.lower()

        # Find all occurrences of columns ending with 'company_id' and their assigned values
        # This regex captures:
        #   - col: the column name ending in 'company_id' (e.g., "user_company_id")
        #   - value: the value assigned in the SQL query (e.g., "1234567")
        other_table_matches = re.findall(r"\b(?:\w+\.)?(\w*company_id)\s*=\s*['\"]?(\d+)['\"]?", sql_lower)
        # get table names and related alias to use in case of using the company table
        parsed = sqlparse.parse(sql)
        statement = parsed[0]
        new_alias = self.get_table_aliases(statement)

        #important note: we are using the names that are **mapped** (you can know the names in reverse_mapping.json values **not** keys) so 'companies' is the used not 'company' which is the original name
        using_company = "companies" in new_alias.keys() #are we using the company table (the name is hard coded based on the database mapping in synonyms.json)
        company_table_matches = [] #initial val
        if using_company: #if we used company table as it has "id" not "company_id"
            company_table_or_alias = new_alias.get("companies") or "companies" #if companies table does not have an alias then use the table name itself (companies), as we are comparing None with a value (the value will be used rathe than None)
            company_table_or_alias = company_table_or_alias.lower() #ensure it's lowercase
            company_table_matches = re.findall(fr"({company_table_or_alias}\.)(id)\s*=\s*['\"]?(\d+)['\"]?", sql_lower, flags=re.IGNORECASE)

        if (not other_table_matches) and (not using_company or not company_table_matches):
            # No 'company_id' column found in the query, or "id" in case of "company" table
            return f"Query must include a company filter with a column ending in 'company_id'. or 'id' in case of using company table"

        # Check matches for other tables
        if any(value == str(company_id) for _, value in other_table_matches):
            return None

        # Check matches for company table
        for alias, col, value in company_table_matches:
            if value == str(company_id):
                return None

        # If we found 'company_id' columns but none has the correct value
        return f"The query uses a company_id column, but it does not match the current company with id: ({company_id})."
