"""Script utilities.""" import os import pandas as pd import pymysql from secrets_manager.flask_ext import FlaskSecretsManager import sqlalchemy from collaborator import config from collaborator.connectors.mysql import _db_engine from collaborator.connectors.snowflake import _db_engine as _db_engine_snowflake def _print_and_get_cred_wrapper(get_cred_func): def wrapper(cred_string): return get_cred_func(cred_string) return wrapper def _db_connection(user, pword, host, db, env): """Get a connection to a provided db based on the environment. Args: environment (str): Environment to base connection on. Returns: pymsql.Connection: Database connection. """ get_cred = _print_and_get_cred_wrapper(os.environ.get) if env in [config.PROD_ENVIRONMENT, config.QA_ENVIRONMENT]: secrets_manager_client = FlaskSecretsManager( application_context=False, environment=env, service_name=config.SERVICE_NAME ) get_cred = _print_and_get_cred_wrapper(secrets_manager_client.get_cred) return pymysql.connect( host=get_cred(host), user=get_cred(user), password=get_cred(pword), database=get_cred(db), cursorclass=pymysql.cursors.DictCursor, ) def art_relations_connection(env): """Get a connection to art_relations based on the environment. Returns: pymsql.Connection: Database connection. """ return _db_connection( user="ART_DB_USER", pword="ART_DB_PASSWORD", host="ART_DB_HOST", db="ART_DB_DATABASE", env=env, ) def ows_collaborator_connection() -> sqlalchemy.engine.Connection: """Get a connection to ows_collaborator. Returns: sqlalchemy.engine.Connection: Database connection. """ return _db_engine.connect() def snowflake_connection(): """Get a connection to snowflake. Returns: sqlalchemy.engine.Connection: Database connection. """ return _db_engine_snowflake.connect() def read_file_and_convert_to_json(file_path, num_rows=0, skip_num_rows=1): """Read an excel spreadsheet and convert it to json. Args: file_path (str): The location of the file num_rows (int): number of rows to process skip_num_rows (int): The number of rows to skip Returns: dict with the parsed file """ _, file_extention = os.path.splitext(file_path) # passing None to pandas makes it read all rows it detects nrows = None if num_rows < 1 else num_rows if file_extention == ".xlsx": df = pd.read_excel( file_path, engine="openpyxl", nrows=nrows, skiprows=list(range(1, skip_num_rows)), ) df = df.where(pd.notnull(df), None) else: df = pd.read_csv(file_path, nrows=nrows, skiprows=list(range(1, skip_num_rows))) return df.to_dict()