import sys
from awsglue.utils import getResolvedOptions  # type: ignore
from awsglue.dynamicframe import DynamicFrame # type: ignore
from awsglue.context import GlueContext # type: ignore
from pyspark.sql import SparkSession # type: ignore
from pyspark.context import SparkContext  # type: ignore
from awsglue.context import GlueContext  # type: ignore
from awsglue.job import Job  # type: ignore
import boto3

def sparkSqlQuery(glueContext, query, mapping, transformation_ctx, spark) -> tuple[DynamicFrame, int]:
    for alias, frame in mapping.items():
        frame.toDF().createOrReplaceTempView(alias)
    result = spark.sql(query)
    return DynamicFrame.fromDF(result, glueContext, transformation_ctx), result.count()

def get_data_from_postgresql_and_save_to_s3(
        glueContext: GlueContext, 
        spark: SparkSession, 
        transform_sql: str,
        db_connection_options: dict, 
        s3_connection_options: dict):
    
    print(f"db_connection_options: {db_connection_options}")
    print(f"s3_connection_options: {s3_connection_options}")

    database_query_result = glueContext.create_dynamic_frame.from_options(
        connection_type="postgresql",
        connection_options=db_connection_options,
        transformation_ctx=f"database_query_result"
    )
    
    print(f"Running SQL transformation: {transform_sql}")
    final_result, row_count = sparkSqlQuery(
        glueContext, 
        query=transform_sql, 
        mapping={"result": database_query_result}, 
        transformation_ctx=f"final_result",
        spark=spark
    )

    AmazonS3_node = glueContext.write_dynamic_frame.from_options(
        frame=final_result, 
        connection_type="s3", 
        format="glueparquet", 
        connection_options=s3_connection_options, 
        format_options={"compression": "uncompressed"}, 
        transformation_ctx=f"AmazonS3_node"
    )

    return row_count

def init_spark_context():
    args = getResolvedOptions(sys.argv, ["JOB_NAME"])
    sc = SparkContext()
    glueContext = GlueContext(sc)
    spark = glueContext.spark_session
    job = Job(glueContext)
    job.init(args["JOB_NAME"], args)

    return sc, glueContext, spark, job

def get_environment():
    sts = boto3.client("sts")
    account_id = sts.get_caller_identity()["Account"]
    PROD_ACCOUNT_ID = "734185074900"
    STAGING_ACCOUNT_ID = "590183763515"
    if account_id == PROD_ACCOUNT_ID:
        return "PROD"
    elif account_id == STAGING_ACCOUNT_ID:
        return "STAGING"
    else:
        raise Exception("Invalid account ID.")
