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()