# utils/erp_context.py

import json
import logging
from functools import lru_cache
from typing import Any, Dict, List, Optional, Union

logger = logging.getLogger(__name__)


@lru_cache(maxsize=2)
def load_json_file(filename: str) -> Union[Dict[str, Any], List[Any]]:
    """Load & cache JSON; return empty container on failure."""
    try:
        with open(filename, "r", encoding="utf-8") as f:
            return json.load(f)
    except Exception as exc:
        load_json_file.cache_clear()
        logger.error("Failed to load %s – %s", filename, exc)
        return {} if filename.endswith(".json") else []


TABLE_DESCRIPTIONS = load_json_file("table_description.json")


def get_table_info(company_id: Optional[str] = None) -> str:
    """
    Assemble a schema description string for the LLM prompt, including dynamically
    inferred foreign key relationships for clarity.
    """
    if not TABLE_DESCRIPTIONS:
        return "No table information available."

    all_table_names = set(TABLE_DESCRIPTIONS.keys())
    parts: List[str] = []

    for tbl, meta in TABLE_DESCRIPTIONS.items():
        # --- Start of Final Fix for Context Generation ---
        column_parts = []
        foreign_key_parts = []

        for col in meta.get("columns", []):
            col_name = col.get("column_name")
            col_desc = col.get("description")

            # Add column with its hint
            if col_desc:
                column_parts.append(f"{col_name} (Hint: {col_desc})")
            else:
                column_parts.append(col_name)

            # Dynamically infer foreign keys based on naming convention
            if col_name.endswith('_id'):
                # Guess the target table name (e.g., 'customer_id' -> 'customer')
                potential_table_name = col_name[:-3]
                # Handle pluralization special cases (e.g., 'companies_id' -> 'company')
                if potential_table_name.endswith('s'):
                    potential_table_name = potential_table_name[:-1]

                # Check if the guessed table exists in the schema
                if potential_table_name in all_table_names:
                    # Assume the PK of the target table is 'id' (a common convention)
                    foreign_key_parts.append(
                        f"The `{col_name}` column is a foreign key to `{potential_table_name}.id`.")
                # Special case for users_table which has user_id as PK
                elif potential_table_name == "user" and "users_table" in all_table_names:
                    foreign_key_parts.append(f"The `{col_name}` column is a foreign key to `users_table.user_id`.")

        cols_str = ", ".join(column_parts)
        fk_str = " ".join(foreign_key_parts)

        # Build the full description for the table
        tbl_desc = meta.get("description")
        segment = f'Table "{tbl}"'
        if tbl_desc:
            segment += f' (Hint: {tbl_desc})'
        segment += f' has columns: {cols_str}.'
        if fk_str:
            segment += f' {fk_str}'
        # --- End of Final Fix ---

        if company_id and tbl != "company":
            segment += (
                f" IMPORTANT: When querying this table, you MUST filter by companies_id = '{company_id}'."
            )
        parts.append(segment)

    return "\n\n".join(parts)


def is_sql_query_needed(question: str) -> bool:
    """
    (This function is currently not used by the LLM-based router but is kept for potential future use).
    """
    if not question:
        return False
    lowered = question.lower()
    sql_keywords = load_json_file("sql_keywords.json")
    if isinstance(sql_keywords, list):
        return any(kw in lowered for kw in sql_keywords)
    return False