# tasks.py
from collections import Counter
import multiprocessing
import os
import json
import glob
import re
import shutil
import datetime
import tempfile
import subprocess
from typing import List
import math
import requests
import time

from celery import Celery
from langchain_ollama import OllamaLLM
from langchain_core.prompts import PromptTemplate
from langchain_core.output_parsers import StrOutputParser
import numpy as np
from sqlalchemy import create_engine, text
from sentence_transformers import SentenceTransformer
from langchain.text_splitter import RecursiveCharacterTextSplitter
import faiss
from celery.schedules import crontab

from config.settings import Config
from services.cache_service import CacheService
try:
    multiprocessing.set_start_method("spawn", force=True)
except RuntimeError:
    pass



broker_url = f"redis://{Config.REDIS_HOST}:{Config.REDIS_PORT}/0"
backend_url = f"redis://{Config.REDIS_HOST}:{Config.REDIS_PORT}/1"
celery_app = Celery("eltrion_tasks", broker=broker_url, backend=backend_url)   


cache_service = CacheService()
_embedder = None
def get_embedder():
    global _embedder
    if _embedder is None:
        _embedder = SentenceTransformer("/opt/models/bge-m3", device="cpu")
    return _embedder
embedder = get_embedder()
splitter = RecursiveCharacterTextSplitter(chunk_size=512, chunk_overlap=64)

vec_path = "vector_store/faiss.index"
weekly_vec_root = "vector_store/users_vectors"
sqlCacheKey = "sql_agent_cache"
table_vec_path = "vector_store/table_descriptions.index"
column_vec_path = "vector_store/column_descriptions.index"

SECRETS_FILE = "/home/jenkins/ai-assistant/jenkins_secrets.json"

def load_jenkins_credentials():
    with open(SECRETS_FILE, "r") as f:
        return json.load(f)

def trigger_jenkins():

    jenkins_config = load_jenkins_credentials()
    
    url = f"{jenkins_config['JENKINS_URL']}/job/{jenkins_config['JENKINS_JOB']}/build"
    response = requests.post(url, auth=(jenkins_config['JENKINS_USER'], jenkins_config['JENKINS_TOKEN']))

    if response.status_code == 201:
        print(" ===== Jenkins job triggered successfully ===== ")
    else:
        print(f" ===== Failed to trigger Jenkins: {response.status_code}, {response.text} ===== ")



def _load_index(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 _save_index(index: faiss.IndexFlatL2, path: str) -> None:
    os.makedirs(os.path.dirname(path), exist_ok=True)
    faiss.write_index(index, path)

#generate a unique path for users based on their ID
def _get_index_path(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")


@celery_app.task
def nightly_training():
    if os.path.exists(vec_path):  # Check if file exists delete it so that we can create a new one with the updated data
        os.remove(vec_path)

    start = datetime.datetime.now()
    index = _load_index(vec_path)

    # 1. ingest new docs
    new_files: List[str] = glob.glob("data/new_docs/**/*.*", recursive=True)
    for fpath in new_files:
        print(f"processing {fpath}")
        with open(fpath, "r", encoding="utf-8") as fp:
            raw = fp.read()
        
        if not raw.strip():
            print(f"Skipping empty file: {fpath}")
            continue
            
        chunk_no = 0
        chunks = splitter.split_text(raw) #split long text into number of chunks
        embs = embedder.encode(chunks, show_progress_bar=False)
        index.add(embs)
        for vec_id in range(index.ntotal - len(chunks), index.ntotal):
            cache_service.set(f"vec_meta:{vec_id}", {"chunk": chunks[chunk_no], "source": fpath}, expiry_seconds=3600* 24 ) 
            chunk_no += 1
        

    #currently we are not using the LLM to fine-tune the model, but it can be used later after testing and ensuring the system works well

    # #2. fine‑tune with QA logs
    # qa_entries = cache_service.xrange("qa_log")
    # if qa_entries:
    #     tmp_jsonl = tempfile.NamedTemporaryFile(delete=False, suffix=".jsonl").name
    #     with open(tmp_jsonl, "w", encoding="utf-8") as fp:
    #         for _, d in qa_entries:
    #             fp.write(json.dumps({"prompt": d["prompt"], "completion": d["answer"]}) + "\n")

    #     subprocess.run(
    #         [
    #             "ollama", "create", "eltrion-nightly",
    #             "--from", Config.OLLAMA_MODEL,
    #             "--adapter", "lora",
    #             "--train", tmp_jsonl,
    #         ],
    #         check=True,
    #     )
    #     cache_service.client.delete("qa_log")
    #     os.remove(tmp_jsonl)

    _save_index(index, vec_path)
    duration = datetime.datetime.now() - start
    print(f"nightly pipeline OK in {duration}")

    #delay the build 2 min after the task
    time.sleep(120)

    print("building now")
    trigger_jenkins()


@celery_app.task
def weekly_training():
    """task to clean up old queries and it's answers also add new ones to index files. the main idea is to act as cache for sql agent."""
    start = datetime.datetime.now()

    
    #remove old vector files and folders, then create a new folder for them
    if os.path.exists(weekly_vec_root):
        shutil.rmtree(weekly_vec_root)
    os.makedirs(weekly_vec_root, exist_ok=True)

    #get cache entries for SQL queries
    sql_entries = cache_service.xrange(sqlCacheKey)

    if sql_entries:
        #get user id for each user from the entries
        user_ids = {entry["company_id"] for _, entry in sql_entries if "company_id" in entry}

        for user_id in user_ids:
            user_vec_path = _get_index_path(user_id)
            user_index = _load_index(user_vec_path)

            #handle duplicates by using a set of questions already added
            user_questions = set()

            #filter SQL entries for the current user id
            user_sql_entries = [entry for _, entry in sql_entries if entry.get("company_id", "N/A") == user_id]

            for entry in user_sql_entries:
                question = entry.get("question", "")
                sql_query = entry.get("sql_query", "")
                result = entry.get("result", "")

                if not question or not result or not sql_query:
                    continue

                if question.strip().lower() in user_questions: #in case the same question was asked multiple times
                    continue
                
                #add only the question part to the index as it will be used to find similar questions when a user ask a question
                chunks = splitter.split_text(question)
                no_of_chunks = len(chunks)
                embs = embedder.encode(chunks, show_progress_bar=False)
                user_index.add(embs)
                
                #cache the question, sql_query and result based on the vector id
                chunk_no = 0
                for vec_id in range(user_index.ntotal - len(chunks), user_index.ntotal):
                    cache_service.set(f"{user_id}_vec_meta:{vec_id}", {"sql_query": sql_query, "question": chunks[chunk_no], "answer": result, "no_of_chunks":no_of_chunks}, expiry_seconds=3600 * 24)
                    chunk_no += 1

                user_questions.add(question.strip().lower())

            _save_index(user_index, user_vec_path)

    #delete the old stream
    cache_service.client.delete(sqlCacheKey)
    duration = datetime.datetime.now() - start
    print(f"weekly pipeline OK in {duration}")

    #delay the build 2 min after the task
    time.sleep(120)

    print("building now")
    trigger_jenkins()


def fix_table_description_response(tbl_name: str, errors: list, ai_response: str) -> dict:
    """Attempt to extract JSON by locating first { and last } -> remove extra comments in case of error extracting the json return json obj converted to dict"""

    start = ai_response.find("{") #if it didn't find it then it will return -1
    end = ai_response.rfind("}") #if it didn't find it then it will return -1
    if start != -1 and end != -1 and start < end:
        try:
            table_desc_data = json.loads(ai_response[start:end+1])
        except json.JSONDecodeError:
            table_desc_data = {"description": "", "columns": {}}
            errors.append(f"Failed to parse JSON for table: {tbl_name}, ai response could not be fixed")
    else:
        table_desc_data = {"description": "", "columns": {}}
        errors.append(f"Failed to parse JSON for table: {tbl_name}, no JSON found in response")
                
    # Always log the table and the full LLM response in errors
    errors.append(f"problem in parsing JSON for table: {tbl_name}, here is the ai response: {ai_response.replace('\n', ' ')}")

    return table_desc_data

def merge_table_descriptions(descriptions: list) -> str:
    """
    Merge multiple descriptions of the same table into one,
    removing duplicate sentences or repeated ideas.
    
    Args:
        descriptions: List of strings (AI outputs for each part of the table)
    
    Returns:
        A single merged description string
    """
    seen_sentences = set()
    merged_sentences = []

    for desc in descriptions:
        # Split into sentences (ends with ',' or '.')
        sentences = re.split(r'(?<=[.,])\s+', desc)
        for sentence in sentences:
            normalized = sentence.strip().lower()  # normalize for duplication
            if normalized and normalized not in seen_sentences:
                seen_sentences.add(normalized)
                merged_sentences.append(sentence.strip())
    
    return " ".join(merged_sentences)

def merge_table_classes(classes: List[str], errors: List[str], table_name: str) -> str:
    """returns tha most repeated class"""
    #if the llm failed to classifiy 
    if len(classes) == 0:
        errors.append(f"The table {table_name} does not have a class, it will be assigned with the class 'UNKNOWN'")
        return 'UNKNOWN'
    
    #decrease UNKNOWN weight by half
    counter = Counter(classes)
    if "UNKNOWN" in counter:
        counter["UNKNOWN"] = math.floor(counter["UNKNOWN"] *0.5)  #reduce influence

    #return the most common
    return counter.most_common(1)[0][0]


def get_database_scheme_info() -> tuple[any,any]:
        """
        connects to the db and return table/column names and more info about them

        Returns:
            (tuple)
            columns_result: info about the table/column names.
            fk_result: info about the fk of tables (relationships)
        """
        #connect to the database to get all table/column names and more info
        DATABASE_URI = Config.DATABASE_URI
        engine = create_engine(DATABASE_URI)

        with engine.connect() as conn:
            # Query to get table/column info and enum values if any
            columns_result = conn.execute(text("""
                SELECT
                    c.table_name,
                    c.column_name,
                    c.data_type,
                    CASE WHEN tc.constraint_type = 'PRIMARY KEY' THEN true ELSE false END AS is_primary_key,
                    c.is_nullable,
                    c.udt_name,
                    CASE
                        WHEN c.data_type = 'USER-DEFINED' THEN (
                            SELECT json_agg(e.enumlabel ORDER BY e.enumsortorder)
                            FROM pg_type t
                            JOIN pg_enum e ON t.oid = e.enumtypid
                            JOIN pg_namespace n ON n.oid = t.typnamespace
                            WHERE t.typname = c.udt_name
                            AND n.nspname = c.table_schema
                        )
                        ELSE NULL
                    END AS enum_values
                FROM information_schema.columns c
                LEFT JOIN information_schema.key_column_usage kcu
                    ON c.table_name = kcu.table_name
                    AND c.column_name = kcu.column_name
                    AND c.table_schema = kcu.table_schema
                LEFT JOIN information_schema.table_constraints tc
                    ON kcu.constraint_name = tc.constraint_name
                    AND kcu.table_schema = tc.table_schema
                    AND tc.constraint_type = 'PRIMARY KEY'
                WHERE c.table_schema = 'public'
                ORDER BY c.table_name, c.ordinal_position;
            """))

            # Foreign keys query
            fk_result = conn.execute(text("""
                SELECT
            kcu.table_name AS table_name,
            kcu.column_name AS column_name,
            ccu.table_name AS fk_table,
            ccu.column_name AS fk_column
        FROM information_schema.table_constraints tc
        JOIN information_schema.key_column_usage kcu
            ON tc.constraint_name = kcu.constraint_name
        AND tc.table_schema = kcu.table_schema
        JOIN information_schema.referential_constraints rc
            ON tc.constraint_name = rc.constraint_name
        JOIN information_schema.key_column_usage ccu
            ON rc.unique_constraint_name = ccu.constraint_name
        AND rc.unique_constraint_schema = ccu.table_schema
        WHERE tc.constraint_type = 'FOREIGN KEY'
        AND tc.table_schema = 'public';
            """))

        return columns_result, fk_result

def setup_ollama_model():
    """setup the prompt for generating descriptions and connect to the ollama model"""
    #connect to the main model to generate description for each table/column
    sql_llm = OllamaLLM(model= Config.OLLAMA_MODEL)

    prompt_for_table_descriptions = ("""                                   
You are an expert database engineer, AI assistant, and an accounting domain expert.

Your task is to generate a structured JSON description of a database table to optimize it for natural language query understanding and retrieval, and to classify the table.

Input:
Table Name: {table_name}
                                         
Columns:
{column_list}
                                         
Additional Info (must be used only inside table/column descriptions if relevant):
{more_info}

Classify each table into one of the following categories:
{classes}
                                         

    Requirements:
    1. Generate a **table-level description** (min 3 sentences, max 5 sentences) that explains:
    - What the table represents in business/domain terms.  
    - The type of data it stores, inferred from its columns.  
    - How it relates to other business concepts or processes.  
    - A summarized explanation of the columns (group them logically by purpose, such as identifiers, personal details, financial data, references to other entities, etc.), without listing every column name individually.
    - Use the table name in each sentence to explain.
                                          
    2. Generate a **column-level description** (min 1 sentence, max 4 sentences each) for every column. 
    - Focus on the column's meaning, purpose, and how it contributes to the table. 
    - Ignore technical details such as data type, nullability, or length constraints.
    - **If a column has ENUM values (shown inside square brackets [ENUM: ...]), include their meaning and order.**
      - Example: `[ENUM: 0: active, 1: inactive, 2: suspended]` → should be described as `(0: active — meaning, 1: suspended — meaning, ...)`.
      - Clearly explain what each ENUM value represents in business terms.
      - If ENUM values exist, describe them in one concise sentence immediately after the column meaning.
                                                                                                          
    3. **Do NOT rename, correct, or alter** the table name or column names. Use them exactly as provided.  This is very important.
                                         
    4. Descriptions must be **clear, human-readable, and unambiguous**. Optimize them for retrieval (include synonyms or business terms when helpful).
    
    5. Only use the classes provided, do not invent any new class.
                                           
    6. Output must be **strict JSON**, no extra text, comments, or explanations. Do NOT include any text outside the json object.  

    Output Format (STRICT JSON ONLY):
    {{
    "table_name": "{table_name}",
    "description": "Concise business/domain description of the table (3-6 sentences). Must explain what the table represents, the type of data it stores based on its columns, its relation to business concepts, and include a summarized explanation of the columns grouped logically by their role (identifiers, references, financial data, personal details, etc.).",
    "class": "choose the best suited class based on the table description.",
    "columns": [
        {{
        "name": "{table_name}.column_name_1",
        "description": "Concise explanation of this column's meaning, purpose, and relation to the table. If this column has ENUM values, describe each possible value in the format (0: value1, 1: value2, ...), and explain what each value represents."
        }},
        {{
        "name": "{table_name}.column_name_2",
        "description": "Concise explanation of this column's meaning, purpose, and relation to the table. If this column has ENUM values, describe each possible value in the format (0: value1, 1: value2, ...), and explain what each value represents."
        }}
    ]
    }}

                                            
    Strict Rules:
    - Incorporate any relevant Additional Info into the table-level or column-level descriptions.
    - Each sentance must contain the table name.
    - Do NOT output Additional Info outside of the JSON structure.
    - Include all columns, even if their meaning is unclear (make a best effort).
    - Descriptions must not be empty.
    - Do NOT add SQL code, queries, implementation details, comments, or any extra text outside of JSON.
    - Respond ONLY with valid JSON. No explanations, no extra text.
    - Columns MUST be in '{table_name}.column_name' form exactly. Do not rename, correct, or alter them.
    - If a column's meaning is unclear, infer its likely purpose from context or name. If still uncertain, describe it as 'represents unspecified or unclear data related to {table_name}' rather than leaving it blank.
    - Class must be directly relevant, in case not use the class "UNKNOWN"
    """)

    table_description_prompt = PromptTemplate.from_template(prompt_for_table_descriptions)

    table_description_query_chain = (
                            table_description_prompt
                            | sql_llm
                            | StrOutputParser()
                    )
    
    return table_description_query_chain

#used to help the llm classify the table
TABLE_CLASSES =  {
    "ASSETS": "anything the company owns or expects to receive (e.g., cash, receivables, inventory, property, equipment, assets).",
    "LIABILITIES": "anything the company owes to others (e.g., loans, payables, expenses, invoice).",
    "EQUITY": "owners’ residual interest (e.g., capital, retained earnings, shares).",
    "UNKNOWN": "anything that doesn’t directly represent past classes, such as customers, employees, logs, settings, etc."
}

@celery_app.task
def monthly_training():
    """create embeddings for table descriptions to be retrived easily and used in sql agent."""
    print("Starting monthly training...")
    skipped = 0
    errors = [] #to print all errors in the process
    max_num_columns = 15 # the max number of columns allowed to be processed at once -> this will make the ai work better in case of long tables 
    try:
        if os.path.exists(table_vec_path):  # Check if file exists delete it so that we can create a new one with the updated data  
            os.remove(table_vec_path)
        
        if os.path.exists(column_vec_path):
            os.remove(column_vec_path)

        start_time = datetime.datetime.now()
        index = _load_index(table_vec_path)
        columns_index = _load_index(column_vec_path)

        with open('synonyms.json', 'r') as f: #this file MUST have new tables with the same format as it is there, also it Must have fks for company named xxx_company_id in all tables that have it
            synonyms : dict = json.load(f)

        columns_result, fk_result = get_database_scheme_info()
        

        # Build a dict to quickly map foreign keys
        fk_map = {}
        for table, column, fk_table, fk_column in fk_result:
            #check if the table isn't in synonyms and print it
            if table not in synonyms.keys():
                errors.append(f"Warning: Table '{table}' not found in synonyms.json")
                skipped += 1

            mapped_table = synonyms.get(table, {}).get('synonym', table)
            mapped_column = mapped_table + "." + synonyms.get(table, {}).get('columns', {}).get(column, column) #after update: now each column name should be tbl_name.col_name
            mapped_fk_table = synonyms.get(fk_table, {}).get('synonym', fk_table)
            mapped_fk_column = mapped_table + "." + synonyms.get(fk_table, {}).get('columns', {}).get(fk_column, fk_column)
            fk_map[(mapped_table, mapped_column)] = {"table": mapped_fk_table, "column": mapped_fk_column}

        # Build final schema dictionary based on the mappings. if there isn't a synonym then use the original name, it will take it from synonyms.json 
        schema = {}
        for table, col, dtype, is_primary_key, is_nullable, udt_name, enum_values in columns_result:
            mapped_table = synonyms.get(table, {}).get('synonym', table)
            mapped_col = mapped_table + "." + synonyms.get(table, {}).get('columns', {}).get(col, col)
            if mapped_table not in schema:
                schema[mapped_table] = {"description": "", "columns": []}
            col_info = {
                "column_name": mapped_col,
                "data_type": dtype,
                "is_primary_key": is_primary_key,
                "is_nullable": is_nullable,
                "foreign_key": fk_map.get((mapped_table, mapped_col)),  # None if not FK
                "enum_values": synonyms.get(table, {}).get("enum_columns", {}).get(col, [])  #get the enum values if there was, from the synonyms json file. the enum list MUST be in the same name of the column 
            }

            schema[mapped_table]["columns"].append(col_info)

            # Include extra info from synonyms.json if available
            if synonyms.get(table, {}).get('more_info'):
                schema[mapped_table]['more_info'] = "\n".join(synonyms.get(table, {}).get('more_info', []))

        #create a reverse mapping to remap before executing a query
        reverse_mapping = {}
        for real_table, syn_info in synonyms.items():
            friendly_table = syn_info.get("synonym", real_table)
            reverse_mapping[friendly_table] = {
                "original": real_table,
                "columns": {  # friendly → real
                    friendly_table + "." + syn_col: real_table + "." + col for col, syn_col in syn_info.get("columns", {}).items()
                }
            }

        #setup AI call
        table_description_query_chain = setup_ollama_model()

        table_description_schema = {}
        column_dict = {} # to save columns and it's info in a seperate columns file

        for table, info in schema.items():
            print(f"Processing mapped table: {table}")

            #save columns in the dict
            column_dict[table] = {}
            column_info_lines = []
            for column in info["columns"]:
                # Initialize column entry in column_dict (description and column info will be filled later)
                column_dict[table][column["column_name"]] = {
                        "description": "",
                        "column_info": ""
                    }
                col_txt = f"({column['data_type']})"
                if str(column["is_nullable"]).lower() == "yes":
                    col_txt += ", nullable"
                if column["is_primary_key"]:
                    col_txt += " PRIMARY KEY"
                if column.get("foreign_key"):
                    fk = column["foreign_key"]
                    col_txt += f" -> FK {fk['table']}.{fk['column'].split('.')[1]}"
                
                column_dict[table][column["column_name"]]["column_info"] = col_txt # contain info like if it is the pk or fk... (for later uses)

                #if the column has enum values then add what is the values for the llm to descripe 
                enum_text = ""
                enums = []
                if column.get("enum_values"):
                    for value, idx in enumerate(column["enum_values"]): #make it in form of 1: available, 2: not available ..
                        enums.append(f"{idx}: {value}")
                    enum_values = ", ".join(enums)
                    enum_text = f" [ENUM: {enum_values}]"
                    print(f"\n\n\n\n {enum_text} \n\n\n\n")

                column_info_lines.append(f"- {column['column_name']} ({column['data_type']}){enum_text}")

            #after saving the column in a way than can be converted to a string we can start the ai description process        
            if len(column_info_lines) <= max_num_columns: #as schema has the columns in list form of dicts (we will at most process tables with less than the max number of columns allowed (in start of the function), otherwise we will split it)

                columns_str = "\n".join(column_info_lines) #make it in one string to pass it with the table name

                description_json = table_description_query_chain.invoke({
                        "table_name": table,
                        "column_list": columns_str,
                        "more_info": info.get("more_info", "no info provided"), #more info like what is the meaning of some strange or domain terms (if it was there)
                        "classes": '\n'.join([f"{table_class.upper()}: {description}" for table_class, description in TABLE_CLASSES.items()])
                        }).strip()
                    
                #Parse JSON from LLM
                try:
                    table_desc_data = json.loads(description_json)
                except json.JSONDecodeError:
                    table_desc_data = fix_table_description_response(tbl_name= table, errors= errors, ai_response= description_json)

                #save this table's class 
                if table_desc_data.get('class'):
                    current_class = table_desc_data.get('class').upper()
                else: 
                    current_class = 'UNKNOWN'
                    errors.append(f"The table: {table} does not have a class, it will be assigned with the class 'UNKNOWN'")

                #print only in case of enum values (test)
                for col in info["columns"]: 
                    if col.get("enum_values"):
                        print(table_desc_data)
                

                #Save table description
                table_description_schema[table] = {
                        "description": table_desc_data.get("description", ""),
                        "class": current_class
                    }

                for col_data in table_desc_data["columns"]:
                    col_name = col_data["name"]
                    if col_name in column_dict[table]:
                        column_dict[table][col_name]["description"] = col_data.get("description", "")



            else: #means that our table has lots of columns so we need to process it in parts

                # find the number of processes as we will make it in steps (based on the number of column, we can use this list as it has already text form of the column and the data type)
                num_of_processes = math.ceil(len(column_info_lines) / max_num_columns) 
                errors.append(f"Warning.. table {table} has large number of columns {len(column_info_lines)} so it will be processed in {num_of_processes} parts")

                #this will contain the table description generated for each part of the table (a description will be generated more than once)
                table_whole_description = [] 

                #saves the class of each part of the table
                table_part_class = []

                # for each iteration process part of the table 
                for process_num in range (num_of_processes): 
                    # quick note: (there won't be out of index error in clicing lists as python will just convert it to the last element index in processing)

                    #for clicing based on the process number 
                    start_column = max_num_columns*process_num
                    end_column = max_num_columns*(process_num+1)

                    columns_str = "\n".join(column_info_lines[start_column: end_column]) #each process take only parts of the column not all of it

                    description_json = table_description_query_chain.invoke({
                            "table_name": table,
                            "column_list": columns_str,
                            "more_info": info.get("more_info", "no info provided"), #more info like what is the meaning of some strange or domain terms (if it was there)
                            "classes": '\n'.join([f"{table_class}: {description}" for table_class, description in TABLE_CLASSES.items()])
                            }).strip()
                    
                    #Parse JSON from LLM
                    try:
                        table_desc_data = json.loads(description_json)
                    except json.JSONDecodeError:
                        table_desc_data = fix_table_description_response(tbl_name= table, errors= errors, ai_response= description_json)

                    #print only in case of enum values (test)
                    for col in info["columns"]: 
                        if col.get("enum_values"):
                            print(table_desc_data)
                        
                    #Save this part table description 
                    table_whole_description.append(table_desc_data.get("description", ""))

                    #save this part class 
                    if table_desc_data.get('class'):
                        current_class = table_desc_data.get('class').upper()
                    else: 
                        current_class = 'UNKNOWN'
                        errors.append(f"a part of {table} does not have a class, it will be assigned with the class 'UNKNOWN'")

                    table_part_class.append(current_class)

                    for col_data in table_desc_data["columns"]:
                        col_name = col_data["name"]
                        if col_name in column_dict[table]:
                            column_dict[table][col_name]["description"] = col_data.get("description", "")
                        else: errors.append(f"{col_name} isn't in the columns list of {table}")
                
                #save the whole table description once after getting rid of the duplication 
                table_description_schema[table] = {
                        "description": merge_table_descriptions(table_whole_description),
                        "class": merge_table_classes(table_part_class, errors, table)
                    }
        
            #print the table name and class after the process finished for this table 
            print(f"table {table} is classified as {table_description_schema[table].get('class', 'error - no class after the classification finished')}")
                

        #save the tables in case of future updates (as the retrieval will be using redis)
        with open("table_description.json", "w", encoding="utf-8") as f:    #it contain the mapped names not the original
            json.dump(table_description_schema, f, indent=2, ensure_ascii=False)
            
        #save the columns in case of future updates (as the retrieval will be using redis)
        with open("table_columns.json", "w", encoding="utf-8") as f:    #it contain the mapped names not the original
            json.dump(column_dict, f, indent=2, ensure_ascii=False)

        #save the reversed mapped names in case of future updates  (as the retrieval will be using redis)
        with open("reverse_mapping.json", "w", encoding="utf-8") as f:    #it contains the mapped names as key and the original in 'original'
            json.dump(reverse_mapping, f, indent=2, ensure_ascii=False)

        #chunking only table description as it will have rows description in it, if we added columns description it will create noise within retrieval as lots of columns are similar, the columns and it's info will be saved only in the cache so that when a table is found then only go through its columns
        table_info = table_description_schema
        for tbl_name, tbl_data in table_info.items():  #the name will be the mapped version not the real one (we could use it as key in reverse dict)
            if not tbl_name.strip():
                continue
            
            #in case of no description generated by ai
            if not tbl_data["description"]: 
                errors.append(f"warning {tbl_name} does not have any description!")
                tbl_data["description"] = "No description"

            #split the table info into chunks
            chunks = splitter.split_text(tbl_data["description"])
            embs = embedder.encode(chunks, show_progress_bar=False)
            embs = np.array(embs)
            if embs.ndim == 1:
                embs = embs.reshape(1, -1)
            index.add(embs)
            chunks_len = len(chunks)
                
            #now place everything on the cache based on the number of vector (each table has one) description will mostly be less than 512 chars so when retrival we wont need (chunk no), but just in case in the future we may need it
            #update: save the revesed mappings here to be updated each time the training happens (the original form) Note: the key in it is the friendly name not the original
            #save the reverse as dict where the key is the mapped vesrion
            #in case of the table isn't in synonyms then use:  original value -> original value
            chunk_no = 0
            for vec_id in range(index.ntotal - len(chunks), index.ntotal):
                cache_service.set(f"table_vec:{vec_id}", {"table_description": chunks[chunk_no], "total_chunks": chunks_len, 
                                                          "table_name": tbl_name, "table_class": tbl_data["class"], "chunk_no": chunk_no, "columns": column_dict[tbl_name], 
                                                          "reverse_mapping": reverse_mapping.get(tbl_name, { #in case of it isn't in the synonyms json file (means it isn't in reverse), so create a reverse dict with the original values
                                                                                                            "original": tbl_name,
                                                                                                            "columns": {col_name: col_name for col_name in column_dict[tbl_name].keys()}
                                                                                                        })}, expiry_seconds=3600* 24* 7 ) 
                chunk_no += 1
            
            #save columns and it's info in index file to be used in semantic search, make it in batches for faster process
            #no need for chunking as it will be less than a chunk
            column_names = []
            column_descriptions = []
            for column_name, col_info in column_dict[tbl_name].items():
                description = col_info.get('description', '')

                #in case of the ai not generating description of the column
                if not description.strip():
                    description = f'{column_name} is used in {tbl_name} table'
                    errors.append(f"Warning, {column_name} from the {tbl_name} table is empty")

                column_names.append(column_name)
                column_descriptions.append(description)

            embs = embedder.encode(column_descriptions, show_progress_bar = False, batch_size= 16)
            embs = np.array(embs) 

            columns_index.add(embs)  # add all vectors at once
            start_id = columns_index.ntotal - len(column_names)

            for i, (column_name, description) in enumerate(zip(column_names, column_descriptions)):
                col_vec_id = start_id + i
                cache_service.set(f"column_vec:{col_vec_id}", {"table_name": tbl_name, "column_name":column_name, "column_description": description}, expiry_seconds=3600* 24* 7)




        
        #save vectors (of the table descriptions)
        _save_index(index, table_vec_path)

        #save vectors (of the column descriptions)
        _save_index(columns_index, column_vec_path)
            
        duration = datetime.datetime.now() - start_time

        for e in errors:
            print(e)

        print(f"monthly pipeline OK in {duration} skipped {skipped} tables")

        #delay the build 2 min after the task
        time.sleep(120)

        print("building now")
        trigger_jenkins()

    except Exception as e:
        print(f"monthly pipeline failed: {e}")
        raise 


celery_app.conf.timezone = "Asia/Riyadh"
# Beat schedule 
celery_app.conf.beat_schedule = {
    "nightly-task": {
        "task": "tasks.nightly_training",
        "schedule": crontab(minute=00, hour=23),  # every day at x
    },
    "weekly-task": {
        "task": "tasks.weekly_training",
        "schedule": crontab(minute=0, hour=2),  # every day at y
    },
    "monthly-task": {
        "task": "tasks.monthly_training",
        "schedule": crontab(minute=5, hour=14),  # every day at z
    },
}

