from snowflake.snowpark import Session
from dotenv import load_dotenv
import os
import argparse

# Load environment variables from .env file
load_dotenv()

# Connection parameters from environment variables
connection_parameters = {
    "account": "fu90569.us-east-2.aws",
    "user": os.getenv("SNOWFLAKE_USERNAME"),
    "password": os.getenv("SNOWFLAKE_PASSWORD"),
    "role": "ACCOUNTADMIN",
    "warehouse": "SUNO_PROD_X_SMAL",
    "database": "SUNO_PROD",
    "schema": "PROD",  # Default to PUBLIC if not specified
}


def create_session():
    try:
        # Create Snowflake session
        session = Session.builder.configs(connection_parameters).create()
        print("Successfully connected to Snowflake!")
        session.sql("USE WAREHOUSE SUNO_PROD_X_SMAL").collect()
        return session
    except Exception as e:
        print(f"Error connecting to Snowflake: {str(e)}")
        raise


def run_sample_query(session):
    try:
        # Example query - modify as needed
        query = "SELECT CURRENT_WAREHOUSE(), CURRENT_DATABASE(), CURRENT_SCHEMA()"
        df = session.sql(query).collect()
        print("\nQuery Results:")
        for row in df:
            print(row)
    except Exception as e:
        print(f"Error executing query: {str(e)}")
        raise


def load_file(file_path):
    with open(file_path, "r") as file:
        return [s.strip() for s in file.readlines()]


def get_stripe_customer_ids_by_user_ids(session, user_ids):
    query = f"""
        SELECT stripe_customer_id
        FROM suno_prod.prod.rds_discord_info
        WHERE user_id IN ({', '.join(user_ids)})
        AND stripe_customer_id IS NOT NULL
    """
    results = session.sql(query).collect()
    return [row["STRIPE_CUSTOMER_ID"] for row in results]


def get_user_ids_by_stripe_customer_id(session, stripe_customer_ids):
    customer_id_strs = [f"'{customer_id}'" for customer_id in stripe_customer_ids]
    query = f"""
        SELECT user_id
        FROM suno_prod.prod.rds_discord_info
        WHERE stripe_customer_id IN ({', '.join(customer_id_strs)})
        AND stripe_customer_id IS NOT NULL
    """
    results = session.sql(query).collect()
    # print(results)
    return [row["USER_ID"] for row in results]


def parse_arguments():
    parser = argparse.ArgumentParser(description='Process Stripe customer and user IDs using Snowflake.')
    parser.add_argument('--customer-id-file', help='Input file containing customer IDs')
    parser.add_argument('--user-id-file', help='Input file containing user IDs')
    parser.add_argument('--output-file', required=True, help='Output file for results')
    return parser.parse_args()


def main():
    args = parse_arguments()

    if (args.customer_id_file and args.user_id_file) or (not args.customer_id_file and not args.user_id_file):
        print("Either --customer-id-file or --user-id-file must be provided")
        return

    session = None
    try:
        session = create_session()
        run_sample_query(session)

        if args.user_id_file:
            user_ids = load_file(args.user_id_file)
            print(f"Loaded {len(user_ids)} user IDs to query for.")
            # print(user_ids)

            if not user_ids:
                print("No user IDs provided, check your input file!")
                return

            print("Querying for stripe customer IDs by user IDs")
            output_ids = get_stripe_customer_ids_by_user_ids(session, user_ids)
            print(f"Found {len(output_ids)} customer IDs")

        elif args.customer_id_file:
            customer_ids = load_file(args.customer_id_file)
            print(f"Loaded {len(customer_ids)} customer IDs to query for.")

            if not customer_ids:
                print("No customer IDs provided, check your input file!")
                return

            print("Querying for user IDs by stripe customer IDs")
            output_ids = get_user_ids_by_stripe_customer_id(session, customer_ids)
            print(f"Found {len(output_ids)} user IDs")
            # print(output_ids)

        else:
            print("Either --customer-id-file or --user-id-file must be provided")
            return

        # Write results to output file
        print(f"Writing {len(output_ids)} results to {args.output_file}")
        with open(args.output_file, "w") as f:
            for output_id in output_ids:
                f.write(f"{output_id}\n")
                # print(output_id)
        print(f"Wrote {len(output_ids)} results to {args.output_file}")

    finally:
        if session:
            session.close()
            print("\nSnowflake session closed.")


if __name__ == "__main__":
    main()
