"""
Pipeline de transações de cartão — bronze → prata → ouro (PySpark).

Implementação de referência do laboratório em simulador.html: mesmas bases,
mesmos de-paras e mesmas regras de negócio. Dados, contas e taxas são 100%
sintéticos e fictícios.

    bronze  transacoes_raw        registro como chega da origem, particionado por data
    prata   transacoes            de-paras → dimensões/câmbio → promoções → cálculos
    ouro    visao_cliente         agregado por conta
            visao_segmento        agregado por segmento × origem × canal
            tb_segmento_cliente   gasto de cada cliente em cada segmento

Uso:
    spark-submit pipeline_transacoes.py --n 200000 --out /tmp/lab
    spark-submit pipeline_transacoes.py --spread 0.05 --aliquota 0.0925 --sem-promocoes

Requer Spark 3.1+.
"""
import argparse
import math
import random
from datetime import date, datetime, timedelta, timezone

from pyspark.sql import DataFrame, SparkSession, Window
from pyspark.sql import functions as F


# ============================================================
# Bases de referência (fictícias)
# ============================================================

# cd_mcc, cd_segmento, ds_segmento | perfil do gerador: peso, vlr_min, vlr_max, prob_intl
MCC = [
    ("5411", "SUPERMERCADO", "Supermercado", 24, 40.0, 650.0, 0.02),
    ("5812", "RESTAURANTE", "Restaurantes", 18, 30.0, 380.0, 0.06),
    ("5541", "COMBUSTIVEL", "Combustível", 12, 80.0, 420.0, 0.01),
    ("5912", "FARMACIA", "Farmácia", 10, 15.0, 320.0, 0.01),
    ("5651", "VESTUARIO", "Vestuário", 10, 60.0, 950.0, 0.08),
    ("5732", "ELETRONICOS", "Eletrônicos", 8, 150.0, 6200.0, 0.15),
    ("5815", "DIGITAL", "Digital", 10, 10.0, 140.0, 0.35),
    ("4511", "VIAGEM", "Viagem", 8, 300.0, 4800.0, 0.45),
]

DE_PARA_CODIGOS = [
    ("cd_canal", "APP", "CELULAR"),
    ("cd_canal", "WEB", "COMPUTADOR"),
    ("fl_internacional", "S", "true"),
    ("fl_internacional", "N", "false"),
    ("cd_moeda", "BRL", "REAL"),
    ("cd_moeda", "USD", "DOLAR"),
]

# cd_produto, ds_produto, publico | gerador: peso, limite min/max, fator de gasto, faixas de renda
PRODUTOS = [
    ("P01", "Classic", "Varejo", 40, 1500, 5000, 0.7, ("até 3k", "3k–6k")),
    ("P02", "Gold", "Varejo", 30, 4000, 12000, 1.0, ("3k–6k", "6k–12k")),
    ("P03", "Platinum", "Alta renda", 20, 10000, 35000, 1.5, ("12k–25k", "6k–12k")),
    ("P04", "Black", "Private", 10, 30000, 120000, 2.4, ("25k+", "12k–25k")),
]

TB_INTERCHANGE = [
    # cd_produto, tx_nacional, tx_internacional
    ("P01", 0.0090, 0.0140),
    ("P02", 0.0105, 0.0155),
    ("P03", 0.0125, 0.0180),
    ("P04", 0.0145, 0.0210),
]

TB_PROMOCOES = [
    # id, cd_segmento, cd_canal_req, fl_intl_req, produtos, pct_cashback, vlr_teto
    ("PRM01", "VIAGEM", None, "S", None, 0.03, 150.0),
    ("PRM02", "SUPERMERCADO", "APP", None, None, 0.02, 30.0),
    ("PRM03", "DIGITAL", None, None, None, 0.10, 15.0),
    ("PRM04", "RESTAURANTE", None, None, "P03,P04", 0.05, 40.0),
    ("PRM05", "COMBUSTIVEL", None, None, None, 0.015, 10.0),
    ("PRM06", "ELETRONICOS", "WEB", "N", "P02,P03,P04", 0.04, 120.0),
]

CUSTO_BANDEIRA_NAC = 0.0010
CUSTO_BANDEIRA_INT = 0.0035

UFS = ["SP", "SP", "SP", "RJ", "MG", "PR", "RS", "SC", "BA", "PE", "DF", "GO"]
N_CONTAS = 40
INICIO = datetime(2026, 10, 1, 8, 0, tzinfo=timezone.utc)
SEGUNDOS_ENTRE_EVENTOS = 1980  # ~33 min de relógio simulado por evento


def build_referencias(spark: SparkSession, seed: int) -> dict:
    """Cria as bases de referência e as tabelas de perfil usadas só pelo gerador."""
    rnd = random.Random(seed)

    de_para_mcc = spark.createDataFrame(
        [m[:3] for m in MCC], "cd_mcc string, cd_segmento string, ds_segmento string"
    )
    perfil_mcc = spark.createDataFrame(
        [(m[0], m[4], m[5], m[6]) for m in MCC],
        "cd_mcc string, vlr_min double, vlr_max double, prob_intl double",
    )
    de_para_codigos = spark.createDataFrame(
        DE_PARA_CODIGOS, "campo string, valor_origem string, valor_destino string"
    )
    de_para_produto = spark.createDataFrame(
        [p[:3] for p in PRODUTOS], "cd_produto string, ds_produto string, publico string"
    )
    tb_interchange = spark.createDataFrame(
        TB_INTERCHANGE, "cd_produto string, tx_nacional double, tx_internacional double"
    )
    tb_promocoes = spark.createDataFrame(
        TB_PROMOCOES,
        "id_promocao string, cd_segmento string, cd_canal_req string, fl_intl_req string, "
        "produtos string, pct_cashback double, vlr_teto double",
    )

    # base de contas estendida + perfil de gasto (só para o gerador)
    pesos_prod = [p[3] for p in PRODUTOS]
    contas, perfil_conta = [], []
    for _ in range(N_CONTAS):
        p = rnd.choices(PRODUTOS, weights=pesos_prod)[0]
        nr = f"{10000 + rnd.randrange(89999)}-{rnd.randrange(10)}"
        limite = round((p[4] + rnd.random() * (p[5] - p[4])) / 500) * 500
        abertura = f"{2012 + rnd.randrange(14)}-{1 + rnd.randrange(12):02d}"
        renda = p[7][0] if rnd.random() < 0.7 else p[7][1]
        contas.append((nr, p[0], rnd.choice(UFS), renda, float(limite), abertura))
        perfil_conta.append((nr, 0.4 + rnd.random() * 1.6, p[6], rnd.choice(MCC)[0]))

    dim_conta_estendida = spark.createDataFrame(
        contas,
        "nr_conta string, cd_produto string, uf string, faixa_renda string, "
        "vlr_limite double, dt_abertura string",
    )
    perfil_conta = spark.createDataFrame(
        perfil_conta, "nr_conta string, peso double, fator_gasto double, cd_mcc_afinidade string"
    )

    return {
        "de_para_mcc": de_para_mcc,
        "de_para_codigos": de_para_codigos,
        "de_para_produto": de_para_produto,
        "dim_conta_estendida": dim_conta_estendida,
        "tb_interchange": tb_interchange,
        "tb_promocoes": tb_promocoes,
        "_perfil_mcc": perfil_mcc,
        "_perfil_conta": perfil_conta,
        "_pesos_mcc": [(m[0], m[3]) for m in MCC],
        "_pesos_conta": [(c[0], c[1]) for c in perfil_conta.collect()],
    }


def build_cambio(spark: SparkSession, n_eventos: int, seed: int) -> DataFrame:
    """PTAX diária fictícia: passeio aleatório suave em torno de 5,40."""
    rnd = random.Random(seed)
    dias = math.ceil(n_eventos * SEGUNDOS_ENTRE_EVENTOS / 86400) + 2
    ptax, linhas = 5.38, []
    for i in range(dias):
        ptax = max(4.6, min(6.4, ptax + (rnd.random() - 0.5) * 0.09 + (5.4 - ptax) * 0.05))
        linhas.append((INICIO.date() + timedelta(days=i), round(ptax, 4)))
    return spark.createDataFrame(linhas, "dt_cotacao date, vlr_ptax double")


# ============================================================
# Motor sintético (fonte)
# ============================================================

def faixas_de_peso(spark: SparkSession, pesos, nome: str) -> DataFrame:
    """[(valor, peso)] → intervalos acumulados [lo, hi) em 0..1 para sorteio ponderado."""
    total = float(sum(p for _, p in pesos))
    linhas, acc = [], 0.0
    for valor, p in pesos:
        linhas.append((valor, acc / total, (acc + p) / total))
        acc += p
    return spark.createDataFrame(linhas, f"{nome} string, lo double, hi double")


def escolher(df: DataFrame, col_rand: str, faixas: DataFrame) -> DataFrame:
    cond = (F.col(col_rand) >= F.col("lo")) & (F.col(col_rand) < F.col("hi"))
    return df.join(F.broadcast(faixas), cond, "left").drop("lo", "hi", col_rand)


def gerar_eventos(spark: SparkSession, ref: dict, tb_cambio: DataFrame, n: int, seed: int) -> DataFrame:
    inicio = int(INICIO.timestamp())
    ev = (
        spark.range(n)
        .withColumn("dt_evento", F.timestamp_seconds(
            F.lit(inicio) + F.col("id") * SEGUNDOS_ENTRE_EVENTOS + (F.rand(seed) * 600).cast("long")))
        .withColumn("r_conta", F.rand(seed + 1))
        .withColumn("r_mcc", F.rand(seed + 2))
    )
    ev = escolher(ev, "r_conta", faixas_de_peso(spark, ref["_pesos_conta"], "nr_conta"))
    ev = escolher(ev, "r_mcc", faixas_de_peso(spark, ref["_pesos_mcc"], "cd_mcc_sorteado"))

    ev = (
        ev.join(F.broadcast(ref["_perfil_conta"]), "nr_conta")
        .join(F.broadcast(ref["dim_conta_estendida"].select("nr_conta", "cd_produto")), "nr_conta")
        # 35% das compras caem no segmento de afinidade do cliente
        .withColumn("cd_mcc", F.when(F.rand(seed + 3) < 0.35, F.col("cd_mcc_afinidade"))
                    .otherwise(F.col("cd_mcc_sorteado")))
        .join(F.broadcast(ref["_perfil_mcc"]), "cd_mcc")
        .join(F.broadcast(tb_cambio), F.to_date("dt_evento") == F.col("dt_cotacao"), "left")
    )

    intl = F.rand(seed + 4) < F.col("prob_intl")
    total_brl = (F.col("vlr_min") + F.pow(F.rand(seed + 5), 2.2) * (F.col("vlr_max") - F.col("vlr_min"))) \
        * F.col("fator_gasto")

    ev = (
        ev.withColumn("fl_internacional", F.when(intl, "S").otherwise("N"))
        .withColumn("cd_moeda", F.when(F.col("fl_internacional") == "S", "USD").otherwise("BRL"))
        .withColumn("cd_canal", F.when(F.rand(seed + 6) < 0.62, "APP").otherwise("WEB"))
        .withColumn("qtd_itens", F.when(F.rand(seed + 7) < 0.6, 1)
                    .otherwise((F.rand(seed + 8) * 4).cast("int") + 1))
        .withColumn("_total_brl", total_brl)
        .withColumn("vlr_compra", F.round(
            F.when(F.col("fl_internacional") == "S", F.col("_total_brl") / F.col("vlr_ptax"))
            .otherwise(F.col("_total_brl")) / F.col("qtd_itens"), 2))
        .withColumn("qtd_parcelas", F.when(
            (F.col("_total_brl") > 400) & (F.rand(seed + 9) < 0.55),
            F.element_at(F.array(*[F.lit(x) for x in (2, 3, 4, 5, 6, 10, 12)]),
                         (F.rand(seed + 10) * 7).cast("int") + 1)).otherwise(1))
    )

    # Saldos que o sistema de origem manda junto: acumulados por conta no tempo
    w = Window.partitionBy("nr_conta").orderBy("dt_evento").rowsBetween(Window.unboundedPreceding, 0)
    ev = (
        ev.withColumn("vlr_total_parcelado", F.round(F.sum(
            F.when(F.col("qtd_parcelas") > 1, F.col("_total_brl")).otherwise(0.0)).over(w), 2))
        .withColumn("_n", F.count("*").over(w))
        .withColumn("vlr_parcelas_pagas", F.round(
            F.col("vlr_total_parcelado") * F.least(F.lit(0.95), F.col("_n") * 0.05), 2))
        .withColumn("saldo_credito_site", F.round(F.sum(F.col("_total_brl") * 0.01).over(w), 2))
        .withColumn("id_transacao", F.format_string("TX%08d", F.col("id")))
    )

    return ev.select(
        "id_transacao", "dt_evento", "nr_conta", "cd_mcc", "vlr_compra", "cd_moeda", "qtd_itens",
        "fl_internacional", "cd_canal", "qtd_parcelas", "saldo_credito_site",
        "vlr_total_parcelado", "vlr_parcelas_pagas",
    )


# ============================================================
# Bronze
# ============================================================

def camada_bronze(raw: DataFrame, out: str) -> None:
    bronze = (
        raw.withColumn("dt_ingestao", F.current_timestamp())
        .withColumn("dt_particao", F.to_date("dt_evento"))
    )
    (bronze.write
        .mode("overwrite")
        .partitionBy("dt_particao")
        .parquet(f"{out}/bronze/transacoes_raw"))


# ============================================================
# Prata
# ============================================================

def camada_prata(bronze: DataFrame, ref: dict, tb_cambio: DataFrame,
                 spread: float, mult_ich: float, aliquota: float, promocoes_ativas: bool) -> DataFrame:
    codigos = ref["de_para_codigos"]
    de_para_canal = (codigos.filter(F.col("campo") == "cd_canal")
                     .select(F.col("valor_origem").alias("cd_canal"), F.col("valor_destino").alias("ds_canal")))

    # passo 1 — de-paras
    prata = (
        bronze
        .join(F.broadcast(ref["de_para_mcc"]), "cd_mcc", "left")
        .join(F.broadcast(de_para_canal), "cd_canal", "left")
        .withColumn("fl_internacional", F.col("fl_internacional") == "S")
        .fillna({"cd_segmento": "OUTROS", "ds_canal": "DESCONHECIDO"})
    )

    # passo 2 — dimensões e câmbio
    prata = (
        prata
        .join(ref["dim_conta_estendida"], "nr_conta", "left")
        .join(F.broadcast(ref["de_para_produto"]), "cd_produto", "left")
        .join(F.broadcast(ref["tb_interchange"]), "cd_produto", "left")
        .join(F.broadcast(tb_cambio), F.to_date("dt_evento") == F.col("dt_cotacao"), "left")
    )

    # passo 3 — promoções: fica a elegível de maior cashback
    promos = ref["tb_promocoes"].filter(F.lit(promocoes_ativas))
    flag_sn = F.when(F.col("t.fl_internacional"), "S").otherwise("N")
    elegivel = (
        (F.col("p.cd_segmento") == F.col("t.cd_segmento"))
        & (F.col("p.cd_canal_req").isNull() | (F.col("p.cd_canal_req") == F.col("t.cd_canal")))
        & (F.col("p.fl_intl_req").isNull() | (F.col("p.fl_intl_req") == flag_sn))
        & (F.col("p.produtos").isNull()
           | F.array_contains(F.split(F.col("p.produtos"), ","), F.col("t.cd_produto")))
    )
    w = Window.partitionBy("id_transacao").orderBy(F.desc_nulls_last("pct_cashback"))
    prata = (
        prata.alias("t")
        .join(F.broadcast(promos).alias("p"), elegivel, "left")
        .select("t.*", "p.id_promocao", "p.pct_cashback", "p.vlr_teto")
        .withColumn("rk", F.row_number().over(w))
        .filter("rk = 1").drop("rk")
    )

    # passo 4 — cálculos
    intl = F.col("fl_internacional")
    prata = (
        prata
        .withColumn("tx_cambio", F.when(intl, F.col("vlr_ptax") * (1 + spread)).otherwise(1.0))
        .withColumn("vlr_faturamento_brl",
                    F.round(F.col("vlr_compra") * F.col("qtd_itens") * F.col("tx_cambio"), 2))
        .withColumn("tx_interchange",
                    F.when(intl, F.col("tx_internacional")).otherwise(F.col("tx_nacional")) * mult_ich)
        .withColumn("vlr_receita_interchange", F.col("vlr_faturamento_brl") * F.col("tx_interchange"))
        .withColumn("vlr_receita_spread", F.when(
            intl, F.col("vlr_compra") * F.col("qtd_itens") * F.col("vlr_ptax") * spread).otherwise(0.0))
        .withColumn("vlr_custo_bandeira", F.col("vlr_faturamento_brl")
                    * F.when(intl, CUSTO_BANDEIRA_INT).otherwise(CUSTO_BANDEIRA_NAC))
        .withColumn("vlr_cashback", F.coalesce(F.least(
            F.col("vlr_faturamento_brl") * F.col("pct_cashback"), F.col("vlr_teto")), F.lit(0.0)))
        .withColumn("vlr_margem", F.col("vlr_receita_interchange") + F.col("vlr_receita_spread")
                    - F.col("vlr_custo_bandeira") - F.col("vlr_cashback"))
        .withColumn("vlr_imposto", F.greatest(F.col("vlr_margem"), F.lit(0.0)) * aliquota)
        .withColumn("vlr_resultado", F.col("vlr_margem") - F.col("vlr_imposto"))
        # crédito e parcelamento
        .withColumn("vlr_saldo_aberto", F.col("vlr_total_parcelado") - F.col("vlr_parcelas_pagas"))
        .withColumn("pct_quitado", F.when(F.col("vlr_total_parcelado") > 0,
                    F.col("vlr_parcelas_pagas") / F.col("vlr_total_parcelado")).otherwise(1.0))
        .withColumn("pct_limite_usado", F.col("vlr_saldo_aberto") / F.col("vlr_limite"))
        .withColumn("dt_particao", F.to_date("dt_evento"))
    )
    return prata


# ============================================================
# Ouro
# ============================================================

def metricas():
    return [
        F.count("*").alias("qtd_transacoes"),
        F.round(F.sum("vlr_faturamento_brl"), 2).alias("vlr_faturamento"),
        F.round(F.sum("vlr_receita_interchange"), 2).alias("vlr_interchange"),
        F.round(F.sum("vlr_cashback"), 2).alias("vlr_cashback"),
        F.round(F.sum("vlr_margem"), 2).alias("vlr_margem"),
        F.round(F.sum("vlr_imposto"), 2).alias("vlr_imposto"),
        F.round(F.sum("vlr_resultado"), 2).alias("vlr_resultado"),
        F.round(F.sum(F.when(F.col("fl_internacional"), F.col("vlr_faturamento_brl")).otherwise(0.0))
                / F.sum("vlr_faturamento_brl"), 4).alias("pct_internacional"),
    ]


def camada_ouro(prata: DataFrame, out: str) -> dict:
    tb_segmento_cliente = (
        prata.groupBy("nr_conta", "cd_segmento")
        .agg(F.count("*").alias("qtd_transacoes"),
             F.round(F.sum("vlr_faturamento_brl"), 2).alias("vlr_faturamento"))
        .withColumn("pct_participacao", F.round(
            F.col("vlr_faturamento") / F.sum("vlr_faturamento").over(Window.partitionBy("nr_conta")), 4))
    )

    principal = (
        tb_segmento_cliente
        .withColumn("rk", F.row_number().over(
            Window.partitionBy("nr_conta").orderBy(F.desc("vlr_faturamento"))))
        .filter("rk = 1")
        .select("nr_conta", F.col("cd_segmento").alias("cd_segmento_principal"))
    )

    ultimo = Window.partitionBy("nr_conta").orderBy(F.desc("dt_evento"))
    limite_atual = (prata.withColumn("rk", F.row_number().over(ultimo)).filter("rk = 1")
                    .select("nr_conta", F.round("pct_limite_usado", 4).alias("pct_limite_usado")))

    visao_cliente = (
        prata.groupBy("nr_conta", "ds_produto", "uf", "faixa_renda")
        .agg(*metricas())
        .withColumn("vlr_ticket_medio", F.round(F.col("vlr_faturamento") / F.col("qtd_transacoes"), 2))
        .join(principal, "nr_conta", "left")
        .join(limite_atual, "nr_conta", "left")
    )

    visao_segmento = (
        prata.groupBy("cd_segmento", "ds_segmento", "fl_internacional", "ds_canal")
        .agg(*metricas())
        .withColumn("vlr_ticket_medio", F.round(F.col("vlr_faturamento") / F.col("qtd_transacoes"), 2))
    )

    visao_cliente.write.mode("overwrite").parquet(f"{out}/gold/visao_cliente")
    visao_segmento.write.mode("overwrite").parquet(f"{out}/gold/visao_segmento")
    tb_segmento_cliente.write.mode("overwrite").parquet(f"{out}/gold/tb_segmento_cliente")
    return {"visao_cliente": visao_cliente, "visao_segmento": visao_segmento,
            "tb_segmento_cliente": tb_segmento_cliente}


# ============================================================
# Main
# ============================================================

def main() -> None:
    ap = argparse.ArgumentParser(description="Pipeline sintético de transações de cartão.")
    ap.add_argument("--n", type=int, default=100_000, help="quantidade de eventos gerados")
    ap.add_argument("--out", default="/tmp/lab_transacoes", help="diretório de saída")
    ap.add_argument("--seed", type=int, default=42)
    ap.add_argument("--spread", type=float, default=0.04, help="spread de câmbio (0.04 = 4%%)")
    ap.add_argument("--mult-interchange", type=float, default=1.0)
    ap.add_argument("--aliquota", type=float, default=0.0925, help="alíquota sobre a margem")
    ap.add_argument("--sem-promocoes", action="store_true")
    args = ap.parse_args()

    spark = (SparkSession.builder
             .appName("lab-pipeline-transacoes")
             .config("spark.sql.session.timeZone", "UTC")
             .config("spark.sql.shuffle.partitions", "8")
             .getOrCreate())

    ref = build_referencias(spark, args.seed)
    tb_cambio = build_cambio(spark, args.n, args.seed)

    raw = gerar_eventos(spark, ref, tb_cambio, args.n, args.seed)
    camada_bronze(raw, args.out)

    bronze = spark.read.parquet(f"{args.out}/bronze/transacoes_raw").drop("dt_particao", "dt_ingestao")
    prata = camada_prata(bronze, ref, tb_cambio, args.spread, args.mult_interchange,
                         args.aliquota, not args.sem_promocoes)
    prata.write.mode("overwrite").partitionBy("dt_particao").parquet(f"{args.out}/silver/transacoes")

    prata = spark.read.parquet(f"{args.out}/silver/transacoes")
    ouro = camada_ouro(prata, args.out)

    ouro["visao_segmento"].orderBy(F.desc("vlr_faturamento")).show(20, truncate=False)
    ouro["visao_cliente"].orderBy(F.desc("vlr_faturamento")).show(10, truncate=False)

    spark.stop()


if __name__ == "__main__":
    main()
