from typing import Dict, List, Tuple

import pandas as pd
import sqlparse
from snowflake.snowpark.session import Session


def get_snowflake_session(
    snowflake_configs: dict, warehouse: str, database: str, schema: str
) -> Session:
    return Session.builder.configs(
        {"warehouse": warehouse, "database": database, "schema": schema, **snowflake_configs}
    ).create()


def get_sql_result(session: Session, schema: str):
    df = session.sql(schema)
    local_df = df.collect()
    if len(local_df) == 0:
        raise Exception("No data found")
    return pd.DataFrame(local_df)


def extract_columns_and_data_types_from_sql(
    sql_text: str, required_dtypes: List[str] = ["number", "date", "varchar(255)"]
) -> List[Tuple[str, str]]:
    """
    Parse DDL to get column names and their exact SQL data types.
    Returns a list of (column_name, sql_data_type) tuples in order of definition.

    The parsing is done by selecting the SQL token that is parenthesized.
    To avoid incorrectly selecting tokens such as "CLUSTER BY (p_date)",
    we only select parenthesized tokens that contain all of the specified data types.
    These can be customized for tables with different data types.
    """
    # Parse SQL into tokens
    parsed = sqlparse.parse(sql_text)[0]

    columns = []
    for token in parsed.tokens:
        # Find parenthesized tokens that contain all of the specified data types
        if isinstance(token, sqlparse.sql.Parenthesis) and all(
            [dtype in token.value.lower() for dtype in required_dtypes]
        ):
            column_defs = token.value.strip("()").split(",")

            for col_def in column_defs:
                col_def = col_def.strip()
                # Remove empty lines and SQL comments
                if not col_def or col_def.startswith("--"):
                    continue

                # Split by whitespace for initial parsing
                parts = col_def.split()
                if len(parts) < 2:
                    continue

                # Extract column name
                column_name = parts[0]

                # Handle inline comments with COMMENT keyword and inline SQL comments
                remaining = col_def[len(column_name) :].strip()
                data_type = remaining.split("COMMENT")[0].split("--")[0].strip()

                columns.append((column_name, data_type))

    if columns:
        return columns
    else:
        raise Exception("No columns found in the parsed SQL: \n" + sql_text)


def create_schedule_task(session: Session, task_name: str, warehouse: str, schedule: str, func: str):
    schema = """
        CREATE OR REPLACE TASK {task_name}
            WAREHOUSE = {warehouse}
            SCHEDULE = 'USING CRON {schedule} UTC'
        AS
            DECLARE
                p_date string;
            begin
                p_date := DATE(SYSDATE() - INTERVAL '1 DAY');
                CALL {func}(:p_date);
            end;
    """.format(task_name=task_name, warehouse=warehouse, schedule=schedule, func=func)
    get_sql_result(session, schema)


def get_table_names(session: Session, database: str, schema: str, table_prefix: str = "") -> List[str]:
    """Get all table names from a given database and schema"""
    query = f"""
        SELECT table_name
        FROM {database}.INFORMATION_SCHEMA.TABLES
        WHERE table_schema = '{schema}'
        AND table_name LIKE '{table_prefix}%'
    """

    result = session.sql(query).collect()
    return sorted([row["TABLE_NAME"] for row in result])


def get_table_columns(
    session: Session, database: str, schema: str, table_names: List[str]
) -> Dict[str, List[str]]:
    """Get all table columns from a given database and schema"""

    table_names_quoted = [f"'{table_name}'" for table_name in table_names]
    query = f"""
        SELECT table_name, column_name
        FROM {database}.INFORMATION_SCHEMA.COLUMNS
        WHERE 1=1
        AND table_schema = '{schema}'
        AND table_name IN ({", ".join(table_names_quoted)})
        ORDER BY table_name, ordinal_position
    """
    result = session.sql(query).collect()

    table_columns = {}
    for row in result:
        table_name = row["TABLE_NAME"]
        column_name = row["COLUMN_NAME"]
        if table_name not in table_columns:
            table_columns[table_name] = []
        table_columns[table_name].append(column_name)

    return table_columns
