import argparse import codecs import csv import gzip from io import BytesIO from time import sleep from typing import Dict from apollo_main_db.apollo.models import UserMarket from sqlalchemy.dialects.mysql import insert from auth0 import Auth0Client from db import init_db, session_scope parser = argparse.ArgumentParser(description="Set all user's emails to the MySQL DB.") parser.add_argument("-d", "--domain", type=str, help="Auth0 domain", required=True) parser.add_argument("-i", "--client_id", type=str, help="Auth0 M2M client ID", required=True) parser.add_argument("-s", "--client_secret", type=str, help="Auth0 M2M client secret", required=True) parser.add_argument("-c", "--connection_id", nargs='+', help="Auth0 connection ID list", required=True) parser.add_argument("-t", "--host", type=str, help="MySQL host", required=True) parser.add_argument("-p", "--port", type=int, help="MySQL port", default=3306) parser.add_argument("-n", "--name", type=str, help="MySQL database name", required=True) parser.add_argument("-u", "--user", type=str, help="MySQL user", required=True) parser.add_argument("-w", "--password", type=str, help="MySQL password", required=True) parser.add_argument("-b", "--batch", type=int, help="Batch size", default=100) args = parser.parse_args() client = Auth0Client(domain=args.domain, client_id=args.client_id, client_secret=args.client_secret) def get_location(connection_id: str) -> str: job_id = client.init_export_job(connection_id) location = None while not location: location = client.get_job(job_id) if not location: sleep(1) return location def decompress(data: bytes) -> gzip.GzipFile: compressed_file = BytesIO() compressed_file.write(data) compressed_file.seek(0) return gzip.GzipFile(fileobj=compressed_file, mode="rb") def get_user_mapping(location: str) -> Dict[str, str]: data = client.get_file(location) data = decompress(data) return { row["user_id"].replace("'auth0|", ""): row["email"][1:] for row in csv.DictReader(codecs.iterdecode(data, "utf-8")) } def set_user_email(user_mapping: dict): user_list = list(user_mapping.items()) for i in range(0, len(user_list), args.batch): print(f"batch {i}-{i + args.batch}") user_chunk = [{"UserId": i[0], "Email": i[1]} for i in user_list[i: i + args.batch]] insert_stmt = insert(UserMarket).values(user_chunk) on_duplicate_key_stmt = insert_stmt.on_duplicate_key_update(Email=insert_stmt.inserted.Email) with session_scope() as session: session.execute(on_duplicate_key_stmt) def main(): print("Initializing db") init_db(args.host, args.port, args.user, args.password, args.name) user_id_to_email_mapping = {} for connection_id in args.connection_id: location = get_location(connection_id) print(f"{connection_id}:{location}") user_chunk = get_user_mapping(location) user_id_to_email_mapping.update(user_chunk) print(f"{connection_id}:{len(user_chunk)}") set_user_email(user_id_to_email_mapping) main()