# Retrieve job configuration from Parameter Store
def getConfigFromSSM(SSM_PARAMETER_NAME):
    ssm_param_values = json.loads(ssmClient.get_parameter(Name = SSM_PARAMETER_NAME)['Parameter']['Value'])

    res = [ssm_param_values['rawBucket'], ssm_param_values['rawBucketPrefix'], ssm_param_values['stageS3BucketName'], ssm_param_values['warehouses3Path']]
    return res

# Retrieve table configuration from JSON content retrieved from Parameter Store
def getTableInfoFromSSM(ssm_param_table_values_key):
    primary_condition = ' '
    primaryKey = ssm_param_table_values_key['primaryKey']
    dbName = ssm_param_table_values_key['domain']
    keylist = primaryKey.split(',')
    for key in  keylist:
        primary_condition += f'target.{key}=source.{key} and '
    primary_key_condition = primary_condition[ : -5]
    partitionCols = ssm_param_table_values_key.get('partitionCols', ' ')
    partitionStr = ""
    partitionStrSQL = ""
    if partitionCols != '':
        partitionStr = f'PARTITIONED BY({partitionCols})'
        partitionStrSQL = f'ORDER BY {partitionCols}'

    res = [primaryKey, partitionCols, dbName, keylist, primary_key_condition, partitionStr, partitionStrSQL]
    return res

# Read incoming data from Amazon S3
def readS3DF(rawS3BucketName, rawBucketPrefix, schemaName, tableName):
    inputDf = glueContext.create_dynamic_frame_from_options(
        connection_type = 's3', 
        connection_options = {
            'paths': [f's3://{rawS3BucketName}/{rawBucketPrefix}/{schemaName}/{tableName}'], 
            'groupFiles': 'none', 
            'recurse':True
        }, 
        format = 'parquet',
        transformation_ctx = tableName
    ).toDF()
    return inputDf

## Apply De-duplication logic on input data, to pickup latest record based on timestamp and operation 
def dedupCDCRecords(inputDf, keylist):
    IDWindowDF = Window.partitionBy(*keylist).orderBy(inputDf.last_update_time).rangeBetween(-sys.maxsize, sys.maxsize)
    inputDFWithTS = inputDf.withColumn('max_op_date', max(inputDf.last_update_time).over(IDWindowDF))
    
    NewInsertsDF = inputDFWithTS.filter('last_update_time=max_op_date').filter("op='I'")
    UpdateDeleteDf = inputDFWithTS.filter('last_update_time=max_op_date').filter("op IN ('U','D')")
    finalInputDF = NewInsertsDF.unionAll(UpdateDeleteDf)

    return finalInputDF

# Create database on the AWS Glue Data Catalog
def createDatabaseSparkSQL(dbName, stageS3BucketName):
    sqltemp = Template("""
        CREATE DATABASE IF NOT EXISTS $dbName LOCATION 's3://$stageS3BucketName/$dbName'
    """)
    SQLQUERY = sqltemp.substitute(
        dbName = dbName, 
        stageS3BucketName = stageS3BucketName)
    logger.info(f'****SQL QUERY IS : {SQLQUERY}')
    spark.sql(SQLQUERY)

# Create table on the AWS Glue Data Catalog
def createTableSparkSQL(stageS3BucketName, dbName, tableName, tableColumns, partitionStrSQL):
    targetPath = f's3://{stageS3BucketName}/{dbName}/{tableName}'
    inputDfWithoutControlColumns.createOrReplaceTempView('appendTable')
    logger.info('***** Creating table and inserting initial data')
    sqltemp = Template("""
        CREATE TABLE $catalog_name.$dbName.$tableName $partitionStr LOCATION '$targetPath' as SELECT $tableColumns FROM appendTable where 1=0 $partitionStrSQL
    """)
    SQLQUERY = sqltemp.substitute(
        catalog_name = catalog_name, 
        dbName = dbName, 
        tableName = tableName, 
        partitionStr = partitionStr,
        targetPath = targetPath,
        tableColumns = tableColumns, 
        partitionStrSQL = partitionStrSQL)
    logger.info(f'****SQL QUERY IS : {SQLQUERY}')
    spark.sql(SQLQUERY)

# Merge incoming changes into the Iceberg table
def upsertRecordsSparkSQL(finalInputDF, inputDfWithoutControlColumns_columns):
    finalInputDF.createOrReplaceTempView('upsertTable')
    updateTableColumnList = ''
    insertTableColumnList = ''
    for column in inputDfWithoutControlColumns_columns:
        updateTableColumnList += f" target.{column} = source.{column},"
        insertTableColumnList += f" source.{column},"

    logger.info('***** Upserting data')
    sqltemp = Template("""
        MERGE INTO $catalog_name.$dbName.$tableName target
        USING (SELECT * FROM upsertTable $partitionStrSQL) source
        ON $primary_key_condition
        WHEN MATCHED AND source.Op = 'D' THEN DELETE
        WHEN MATCHED AND source.Op = 'U' THEN UPDATE SET $updateTableColumnList
        WHEN NOT MATCHED and source.Op = 'I' THEN INSERT ($tableColumns) values ($insertTableColumnList)
    """)
    SQLQUERY = sqltemp.substitute(
        catalog_name = catalog_name, 
        dbName = dbName, 
        tableName = tableName, 
        partitionStrSQL = partitionStrSQL, 
        primary_key_condition = primary_key_condition, 
        updateTableColumnList = updateTableColumnList[ : -1], 
        tableColumns = tableColumns, 
        insertTableColumnList = insertTableColumnList[ : -1])

    logger.info(f'****SQL QUERY IS : {SQLQUERY}')
    spark.sql(SQLQUERY)

# Perform initial data loading into an empty Iceberg table
def initialLoadRecordsSparkSQL(finalInputDF, inputDfWithoutControlColumns_columns):
    finalInputDF.createOrReplaceTempView('insertTable')
    insertTableColumnList = ''
    for column in inputDfWithoutControlColumns_columns:
        insertTableColumnList += f" {column},"

    logger.info('***** Inserting initial data')
    sqltemp = Template("""
        INSERT INTO $catalog_name.$dbName.$tableName  ($insertTableColumnList)
        SELECT $insertTableColumnList FROM insertTable $partitionStrSQL
    """)
    SQLQUERY = sqltemp.substitute(
        catalog_name = catalog_name, 
        dbName = dbName, 
        tableName = tableName,
        insertTableColumnList = insertTableColumnList[ : -1],
        partitionStrSQL = partitionStrSQL)

    logger.info(f'****SQL QUERY IS : {SQLQUERY}')
    spark.sql(SQLQUERY)

# Main application
import sys
import os
import json
from pyspark.sql.session import SparkSession
from pyspark.sql.functions import max
from pyspark.sql.window import Window
from awsglue.utils import getResolvedOptions
from awsglue.context import GlueContext
from awsglue.job import Job
from pyspark.sql import SparkSession
import boto3
from botocore.exceptions import ClientError
from string import Template
glueClient = boto3.client('glue')
ssmClient = boto3.client('ssm')

## Parameters for job
args = getResolvedOptions(sys.argv, ['JOB_NAME', 'stackName'])
SSM_PARAMETER_NAME = f"{args['stackName']}-iceberg-config"
SSM_TABLE_PARAMETER_NAME = f"{args['stackName']}-iceberg-tables"
rawS3BucketName, rawBucketPrefix, stageS3BucketName, warehouse_path = getConfigFromSSM(SSM_PARAMETER_NAME)
ssm_param_table_values = json.loads(ssmClient.get_parameter(Name = SSM_TABLE_PARAMETER_NAME)['Parameter']['Value'])
dropColumnList = ['db','table_name', 'schema_name','Op', 'last_update_time', 'max_op_date']
catalog_name = 'my_catalog'
dynamodb_table = f"iceberg_table_lock_{args['stackName']}"
errored_table_list = []

## Iceberg configuration
spark = SparkSession.builder \
    .config('spark.sql.warehouse.dir', warehouse_path) \
    .config(f'spark.sql.catalog.{catalog_name}', 'org.apache.iceberg.spark.SparkCatalog') \
    .config(f'spark.sql.catalog.{catalog_name}.warehouse', warehouse_path) \
    .config(f'spark.sql.catalog.{catalog_name}.catalog-impl', 'org.apache.iceberg.aws.glue.GlueCatalog') \
    .config(f'spark.sql.catalog.{catalog_name}.io-impl', 'org.apache.iceberg.aws.s3.S3FileIO') \
    .config(f'spark.sql.catalog.{catalog_name}.lock-impl', 'org.apache.iceberg.aws.glue.DynamoLockManager') \
    .config(f'spark.sql.catalog.{catalog_name}.lock.table', dynamodb_table) \
    .config('spark.sql.extensions', 'org.apache.iceberg.spark.extensions.IcebergSparkSessionExtensions') \
    .getOrCreate()
glueContext = GlueContext(spark.sparkContext)
job = Job(glueContext)
job.init(args['JOB_NAME'], args)
logger = glueContext.get_logger()

# Iteration over tables stored on Parameter Store
for key in ssm_param_table_values:
    # Get table data
    isTableExists = False
    schemaName, tableName = key.split('.')
    logger.info(f'Processing table : {tableName}')
    try:
        primaryKey, partitionCols, dbName, keylist, primary_key_condition, partitionStr, partitionStrSQL = getTableInfoFromSSM(ssm_param_table_values[key])
    except KeyError as e:
        raise Exception(f'***** Primary key for {tableName} not found in parameter, it is required.')
    
    # Create database if not exists
    createDatabaseSparkSQL(dbName, stageS3BucketName)

    try:
        glueClient.get_table(DatabaseName = dbName, Name = tableName)
        isTableExists = True
    except ClientError as e:
        if e.response['Error']['Code'] == 'EntityNotFoundException':
            logger.info(f'***** {dbName}.{tableName} does not exist. Table will be created.')

    ## Read changes from raw bucket
    inputDf = readS3DF(rawS3BucketName, rawBucketPrefix, schemaName, tableName)
    
    if(inputDf.first() == None):
        logger.info('Dataframe is empty')
        continue;
    else:
        if('Op' in inputDf.columns):
            # Dedup incoming changes
            finalInputDF = dedupCDCRecords(inputDf, keylist)
        else:
            finalInputDF = inputDf

        inputDfWithoutControlColumns = finalInputDF.drop(*dropColumnList)
        tableColumns = ','.join(inputDfWithoutControlColumns.columns)
        try:
            if(not isTableExists):
                ## Create table if not exists
                createTableSparkSQL(stageS3BucketName, dbName, tableName, tableColumns, partitionStrSQL)

            if('Op' in inputDf.columns):
                # Upsert changes
                upsertRecordsSparkSQL(finalInputDF, inputDfWithoutControlColumns.columns)
            else:  
                # Perform initial data loading
                initialLoadRecordsSparkSQL(finalInputDF, inputDfWithoutControlColumns.columns)
        
        except Exception as e:
            # Log errors and save table into error array            
            logger.info(f'There is an issue with table: {tableName}')
            logger.info(f'The exception is : {e}')
            errored_table_list.append(tableName)
            continue
job.commit()        

# Verify if errors exists, log them and fail the job
if (len(errored_table_list)):
    logger.info('Total number of errored tables are ',len(errored_table_list))
    logger.info('Tables that failed during processing are ', *errored_table_list, sep=', ')
    raise Exception(f'***** Some tables failed to process.')
