import os import json import boto3 import psycopg2 from psycopg2._psycopg import ProgrammingError from psycopg2.extras import RealDictCursor import logging log = logging.getLogger() log.setLevel(logging.INFO) REGION = os.environ["AWS_REGION"] if "AWS_REGION" in os.environ else "eu-west-1" S3_BUCKET_DEVEL = "frontend-api-devel-filestore" def get_secret(secret_name): # Create a Secrets Manager client session = boto3.session.Session() client = session.client( service_name='secretsmanager', region_name=REGION ) get_secret_value_response = json.loads(client.get_secret_value(SecretId=secret_name)['SecretString']) return get_secret_value_response DB_CONN = {'fansifter-rds': None, 'fansifter-rds-test': None, 'fansifter-rds-live': None, 'sandbox': None} DB_PORT = {'fansifter-rds': 5432, 'fansifter-rds-test': 5433, 'fansifter-rds-live': 5434, 'sandbox': None} def get_rds_connection(conn_name): global DB_CONN if DB_CONN[conn_name] is None: try: secret = get_secret(conn_name) # Toomas: this port mapping only exists for ssh tunneling setup in my virtual machine. # Anywhere else you can just use secret['port'] instead of DB_PORT[conn_name] in get_rds_connection() DB_CONN[conn_name] = psycopg2.connect(host=secret['host'], #'localhost', port=DB_PORT[conn_name], database=secret['dbname'], user=secret['username'], password=secret['password'], cursor_factory=RealDictCursor) except Exception as e: log.exception(e) return DB_CONN[conn_name] def rds_query(connection_name, query, vars=None, fetch=False): with get_rds_connection(connection_name) as conn: with conn.cursor() as cur: if isinstance(query, list): for statement in query: cur.execute(statement, vars=vars) else: cur.execute(query, vars=vars) try: if fetch: return cur.fetchall() except ProgrammingError as e: return [] # select here which database connections should be included. databases = [ 'fansifter-rds', 'fansifter-rds-test', 'fansifter-rds-live', ] only_schemas = ['c43e12e843c97d73ce49d2d0632c31c032874407a4375255b26fe7638'] new_version = 'mvp 2.2' other_buckets = set() for connection in databases: # all_schemas = rds_query(connection, # "select table_schema, table_name from information_schema.tables where table_name IN ('meta_data' )", fetch=True) all_schemas = rds_query(connection, "select default_schema, user_id, email from commons.user_company", fetch=True) print(f"Found {len(all_schemas)} schemas in database {connection}") for schema, user_id, email in [sch.values() for sch in all_schemas]: print(f'Working on: {schema} {user_id} {email}') # if schema in only_schemas: # continue try: ws = rds_query(connection, f"SELECT trial FROM {schema}.company WHERE company_id = '{schema}'", fetch=True) except Exception as e: ver = None if True: statements=[ #f"alter table {schema}.fb_audience alter column parent_id type int using parent_id::int;" #f"alter table {schema}.fb_audience alter column parent_id type varchar(255) using parent_id::varchar(255);" #f"alter table {schema}.company add trial bool default FALSE;" #f"alter table {schema}.company add billing_cycle varchar(100) default 'Monthly';" f"""CREATE TABLE IF NOT EXISTS {schema}.company ( company_id varchar(60), name varchar(255), current_package_id int, subscription_valid_until_date date, trial boolean default false, billing_cycle varchar(100) default 'Monthly' );""", f"""INSERT INTO {schema}.company (company_id, name, current_package_id) VALUES ('{schema}', (SELECT company_name FROM commons.user_company WHERE default_schema = '{schema}'), (SELECT id FROM commons.packages WHERE package_name = 'lite'));""" ] try: rds_query(connection, statements) # , vars={'new_version': new_version}) except Exception as e: log.exception("oops", exc_info=e) # else: # if ver == 'mvp 2.2': # print('Already version mvp 2.2') # if ver is None: # print('Too old to upgrade')