import boto3 # type: ignore
import json
from cryptography.hazmat.primitives import serialization # type: ignore
from snowflake.snowpark.session import Session # type: ignore
import snowflake.connector # type: ignore
from utils.util import get_environment

environment = get_environment()
database_name = f"SUNO_{environment}"

# Snowflake parameters
SNOWFLAKE_CONFIGS = {}
def get_secret(secret_name, region_name):
    session = boto3.session.Session() # type: ignore
    client = session.client(
        service_name='secretsmanager',
        region_name=region_name
    )
    try:
        get_secret_value_response = client.get_secret_value(
            SecretId=secret_name
        )
        return get_secret_value_response["SecretString"]
    except Exception as e:
        print(f"Error retrieving secret: {e}")
        raise e

def get_private_key():
    raw_private_key = get_secret("prod-snowflake-account-private-key", 'us-east-2')
    if raw_private_key:
        # Ensure the private key is properly formatted, Convert string to bytes
        raw_private_key = raw_private_key.strip().encode()

        # Load the private key into the correct format
        private_key = serialization.load_pem_private_key(
            raw_private_key,
            password=None,  # If encrypted, replace with the passphrase
        )

        # Convert the private key into DER format for Snowflake
        private_key_der = private_key.private_bytes(
            encoding=serialization.Encoding.DER,
            format=serialization.PrivateFormat.PKCS8,
            encryption_algorithm=serialization.NoEncryption(),
        )

        return private_key_der

def init_snowflake_configs():
    secrets = get_secret('prod-snowflake-account', 'us-east-2')
    if secrets:
        secrets = json.loads(secrets)
        SNOWFLAKE_CONFIGS["account"] = secrets.get("sfAccount")
        SNOWFLAKE_CONFIGS["user"] = secrets.get("sfUser")
        SNOWFLAKE_CONFIGS["private_key"] = get_private_key()
        SNOWFLAKE_CONFIGS["role"] = secrets.get("sfRole")
    else:
        raise Exception("Failed to get Snowflake secrets from AWS Secret Manager")

init_snowflake_configs()

def get_snowflake_session(warehouse: str, database: str = database_name, schema: str = environment):
    return Session.builder.configs(
        {"warehouse": warehouse, "database": database, "schema": schema, **SNOWFLAKE_CONFIGS}
    ).create() 

def get_snowflake_connection(warehouse: str, database: str, schema: str):
    return snowflake.connector.connect(
        user=SNOWFLAKE_CONFIGS["user"],
        account=SNOWFLAKE_CONFIGS["account"],
        private_key=SNOWFLAKE_CONFIGS["private_key"],
        warehouse=warehouse,
        database=database,
        schema=schema,
        role=SNOWFLAKE_CONFIGS["role"]
    )
