import json
import os
from typing import Any, Dict, Optional

import boto3
import pandas as pd
import snowflake.connector
from snowflake.connector.pandas_tools import write_pandas
from .snowflake.constants import Warehouse, Role
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.backends import default_backend
from dotenv import load_dotenv


def get_snowflake_private_key():
    raw_private_key = os.environ["SNOWFLAKE_PRIVATE_KEY"]

    # Handle newlines: replace literal '\n' with actual newlines
    if "\\n" in raw_private_key:
        raw_private_key = raw_private_key.replace("\\n", "\n")

    try:
        private_key = serialization.load_pem_private_key(
            raw_private_key.encode("utf-8"), password=None, backend=default_backend()
        )
        return private_key
    except ValueError as e:
        print(f"Failed to load private key: {str(e)}")
        raise


def get_snowflake_connection(
    warehouse: Optional[Warehouse] = Warehouse.SMALL, role: Optional[Role] = None
):
    """Create and return a Snowflake connection.

    Args:
        warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL)
        role: Snowflake role to use (default: None, uses user's default role)
    """
    connection_params = {
        "user": os.environ["SNOWFLAKE_ACCOUNT_USER"],
        "private_key": get_snowflake_private_key(),
        "account": os.environ["SNOWFLAKE_ACCOUNT"],
        "warehouse": warehouse.value if warehouse else Warehouse.SMALL.value,
        "database": "SUNO_PROD",
        "schema": "PROD",
    }

    if role:
        connection_params["role"] = role.value

    return snowflake.connector.connect(**connection_params)


def execute_snowflake_ddl(
    ddl_statement: str,
    warehouse: str = Warehouse.LARGE,  # default to a larger warehouse for DDL operations
    role: Optional[Role] = Role.ACCOUNTADMIN,
) -> None:
    """Execute a DDL statement against Snowflake.

    Args:
        ddl_statement: DDL statement to execute (CREATE, ALTER, DROP, etc.)
        warehouse: Snowflake warehouse to use (default: SUNO_PROD_LARGE)
        role: Snowflake role to use (default: ACCOUNTADMIN)
    """
    with get_snowflake_connection(warehouse=warehouse, role=role) as conn:
        try:
            cursor = conn.cursor()
            cursor.execute(ddl_statement)
        except Exception as e:
            raise Exception(f"Error executing Snowflake DDL: {e!s}")


def query_snowflake(
    query: str,
    params: Optional[Dict[str, Any]] = None,
    warehouse: Optional[Warehouse] = Warehouse.SMALL,
    role: Optional[Role] = None,
) -> pd.DataFrame:
    """Execute a query against Snowflake and return results as a DataFrame.

    Args:
        query: SQL query string to execute
        params: Optional dictionary of query parameters
        warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL)
        role: Snowflake role to use (default: None, uses user's default role)

    Returns:
        DataFrame containing query results
    """
    with get_snowflake_connection(warehouse=warehouse, role=role) as conn:
        try:
            cursor = conn.cursor()
            if params:
                cursor.execute(query, params)
            else:
                cursor.execute(query)
            results = cursor.fetchall()
            columns = [desc[0] for desc in cursor.description]
            return pd.DataFrame(results, columns=columns)
        except Exception as e:
            raise Exception(f"Error executing Snowflake query: {e!s}")


def write_to_snowflake(
    df: pd.DataFrame,
    table_name: str,
    schema: Optional[str] = None,
    warehouse: Optional[Warehouse] = Warehouse.SMALL,
    role: Optional[Role] = None,
) -> None:
    """Write a DataFrame to a Snowflake table.

    Args:
        df: DataFrame to write
        table_name: Target table name
        schema: Optional schema name (defaults to connection schema)
        warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL)
        role: Snowflake role to use (default: None, uses user's default role)
    """
    with get_snowflake_connection(warehouse=warehouse, role=role) as conn:
        try:
            return write_pandas(conn, df, table_name, schema=schema, on_error="continue")
        except Exception as e:
            raise Exception(f"Error writing to Snowflake: {e!s}")


def get_lambda_client():
    """Create and return an AWS Lambda client."""
    return boto3.client(
        "lambda",
        aws_access_key_id=os.environ["AWS_ACCESS_KEY_ID"],
        aws_secret_access_key=os.environ["AWS_SECRET_ACCESS_KEY"],
        region_name="us-east-2",
    )


def get_s3_client():
    """Create and return an AWS S3 client."""
    return boto3.client(
        "s3",
        aws_access_key_id=os.environ["AWS_ACCESS_KEY_ID"],
        aws_secret_access_key=os.environ["AWS_SECRET_ACCESS_KEY"],
        region_name="us-east-2",
    )


def invoke_lambda(function_name: str, payload: Dict[str, Any]) -> Dict[str, Any]:
    """Invoke an AWS Lambda function and return the response."""
    lambda_client = get_lambda_client()
    response = lambda_client.invoke(
        FunctionName=function_name,
        InvocationType="RequestResponse",
        Payload=json.dumps(payload),
    )
    return json.loads(response["Payload"].read().decode("utf-8"))


def invoke_postgres_lambda(query, values=None, is_reader=True):
    """Invoke the Postgres Lambda, which  proxies queries from Dagster."""
    payload = {"query": query, "values": values, "is_reader": is_reader}
    return invoke_lambda(
        function_name="database-accessor",
        payload=payload,
    )


def invoke_redis_lambda(
    operation: str, keys: list[str], values: Optional[list[dict]] = None, ttl: Optional[int] = None
) -> Any:
    """Invoke the Redis Lambda function to perform Redis operations.

    Args:
        operation: The Redis operation to perform ('get', 'set_json', or 'delete')
        keys: A list of keys to operate on
        values: The values to set. This should be a list of values matching the length of keys.
        ttl: Optional time-to-live in seconds (only used for 'set' operation)
    """
    valid_operations = ["get", "set_json", "delete"]
    operation = operation.lower()
    if operation not in valid_operations:
        raise ValueError(f"Invalid operation: {operation}. Must be one of {valid_operations}")

    # Validate values for set operation
    if operation == "set_json":
        if values is None:
            raise ValueError("Values must be provided for 'set_json' operation")
        if len(values) != len(keys):
            raise ValueError(
                f"Number of values ({len(values)}) must match number of keys ({len(keys)})"
            )

    payload = {
        "operation": operation,
        "keys": keys,
    }

    if operation == "set_json":
        payload["values"] = values
        if ttl is not None:
            payload["ttl"] = ttl

    return invoke_lambda(
        function_name="redis-generic-accessor",
        payload=payload,
    )
