from typing import Dict import hashlib import boto3 import json import psycopg2 from datetime import datetime from psycopg2._psycopg import ProgrammingError from psycopg2.extras import RealDictCursor from internal.commons import REGION, DB_SECRET import logging log = logging.getLogger() ## # Helpers ## SECRET = {} class DataProcessorError(RuntimeError): pass def get_secret(secret_name): global SECRET if secret_name not in SECRET: # Create a Secrets Manager client session = boto3.session.Session() client = session.client( service_name='secretsmanager', region_name=REGION ) SECRET[secret_name] = json.loads(client.get_secret_value(SecretId=secret_name)['SecretString']) return SECRET[secret_name] def get_rds_connection(cursor_factory=RealDictCursor): try: secret = get_secret(DB_SECRET) return psycopg2.connect(host=secret['host'], port=secret['port'], database=secret['dbname'], user=secret['username'], password=secret['password'], cursor_factory=cursor_factory) except Exception as e: log.exception(e) def rds_query(query, vars=None, fetch=True): """ :param fetch: :param query: either query string or list of query strings :param vars: :return: either result of query or list of results """ with get_rds_connection() as conn: with conn.cursor() as cur: if isinstance(query, list): res = [] for statement in query: cur.execute(statement, vars=vars) if fetch: try: res.append(cur.fetchall()) except ProgrammingError as e: res.append([]) return res else: cur.execute(query, vars=vars) if fetch: try: return cur.fetchall() except ProgrammingError as e: return [] def make_random_hash(event: Dict) -> str: random_hash = hashlib.sha224('SCHEMA094u'.encode()) random_hash.update(make_iso_date_now().encode()) if event and 'request' in event and 'userAttributes' in event['request'] and \ 'sub' in event['request']['userAttributes']: random_hash.update(event['request']['userAttributes']['sub'].encode()) if event and 'userSub' in event: random_hash.update(event['userSub'].encode()) return random_hash.hexdigest() def get_schema_id(event: Dict) -> str: try: user_sub = get_user_id(event) for user in rds_query("SELECT default_schema FROM commons.user_company WHERE user_id = %(user_id)s", {'user_id': user_sub}): return user['default_schema'] else: raise RuntimeError('Could not identify user') except Exception as e: log.error("Could not get schema", exc_info=e) raise RuntimeError('Could not identify user, please provide userSub as part of function payload.') def get_email(event: Dict) -> str: try: user_sub = get_user_id(event) for user in rds_query("SELECT email FROM commons.user_company WHERE user_id = %(user_id)s", {'user_id': user_sub}): return user['email'] else: raise RuntimeError('Could not identify user') except Exception as e: log.error("Could not get schema", exc_info=e) raise RuntimeError('Could not identify user, please provide userSub as part of function payload.') def generate_schema_id(event: Dict, marker='c') -> str: return marker + make_random_hash(event) def make_iso_date_now(): return f'{datetime.utcnow().isoformat(timespec="microseconds")}Z' def get_user_id(event: Dict) -> str: return event['userSub'] if 'userSub' in event else \ event['request']['userAttributes']['sub'] def get_id_and_management_schema_for_workspace(default_schema, workspace_id): try: for workspace in rds_query(f"SELECT id FROM {default_schema}.workspace " f"WHERE id = %(workspace_id)s " f"LIMIT 1", {'workspace_id': workspace_id}): return workspace['id'], default_schema else: log.error("FLAGGED: Illegal workspace access attempt!") raise RuntimeError('Access denied.') except Exception as e: log.error("Illegal workspace access attempt", exc_info=e) raise RuntimeError('Access denied.')