"""
ETL mensal: lê os CSVs extraídos da RFB via DuckDB (processamento fora da
memória, sem precisar carregar os ~20GB inteiros em RAM), filtra por
UF = MG + CNAE de alimentação + situação ATIVA, e faz upsert em Postgres,
marcando `primeira_deteccao` só na primeira vez que o CNPJ aparece.

Também lê o arquivo de Sócios (QSA) para os cnpj_basico monitorados nesta
rodada e faz upsert em `socios_administradores`, marcando `eh_administrador`
para as qualificações em QUALIFICACOES_ADMINISTRADOR.

Uso:
    python -m cnpj_monitor.etl_load 2026-07
"""
import sys
import tempfile
from datetime import date
from pathlib import Path

import duckdb
import psycopg2
from psycopg2.extras import execute_values

from .config import (
    CNAES_ALIMENTACAO,
    SITUACAO_ATIVA,
    QUALIFICACOES_ADMINISTRADOR,
    QUALIFICACOES_MAP,
    POSTGRES_DSN,
    load_ufs_alvo,
)

REPO_ROOT = Path(__file__).resolve().parents[2]
EXTRACTED_DIR = REPO_ROOT / "data" / "extracted"
DUCKDB_TEMP = Path(tempfile.gettempdir()) / "duckdb_cnpj"
DUCKDB_TEMP.mkdir(exist_ok=True)


def _con() -> duckdb.DuckDBPyConnection:
    """Cria conexão DuckDB com spill-to-disk habilitado para arquivos grandes."""
    con = duckdb.connect()
    con.execute(f"SET temp_directory='{DUCKDB_TEMP}'")
    con.execute("SET memory_limit='1GB'")
    return con

# Colunas oficiais do layout RFB (sem header no CSV original)
COLS_ESTABELECIMENTOS = [
    "cnpj_basico", "cnpj_ordem", "cnpj_dv", "identificador_matriz_filial",
    "nome_fantasia", "situacao_cadastral", "data_situacao_cadastral",
    "motivo_situacao_cadastral", "nome_cidade_exterior", "pais",
    "data_inicio_atividade", "cnae_fiscal_principal", "cnae_fiscal_secundaria",
    "tipo_logradouro", "logradouro", "numero", "complemento", "bairro",
    "cep", "uf", "municipio", "ddd1", "telefone1", "ddd2", "telefone2",
    "ddd_fax", "fax", "correio_eletronico", "situacao_especial",
    "data_situacao_especial",
]

COLS_EMPRESAS = [
    "cnpj_basico", "razao_social", "natureza_juridica",
    "qualificacao_responsavel", "capital_social", "porte_empresa",
    "ente_federativo_responsavel",
]

COLS_SOCIOS = [
    # Ordem oficial do layout RFB (arquivo SociosN.zip — Dados Abertos CNPJ)
    # Posição 1: identificador_de_socio — 1=PF, 2=PJ, 3=Estrangeiro/residente exterior
    # Posição 4: cpf_cnpj_socio — para estrangeiros (id=3) contém o código do país, não CPF
    "cnpj_basico",                       # 0
    "identificador_de_socio",            # 1
    "nome_socio_razao_social",           # 2
    "cpf_cnpj_socio",                    # 3
    "qualificacao_socio",                # 4
    "data_entrada_sociedade",            # 5
    "pais",                              # 6
    "representante_legal",               # 7  ← CPF do representante legal
    "nome_do_representante",             # 8
    "qualificacao_representante_legal",  # 9
    "faixa_etaria",                      # 10
]

COLS_SIMPLES = [
    "cnpj_basico", "opcao_simples", "data_opcao_simples",
    "data_exclusao_simples", "opcao_mei", "data_opcao_mei",
    "data_exclusao_mei",
]


def _cols_sql(cols: list) -> str:
    """DuckDB zero-pads auto column names when the file has ≥ 10 columns."""
    if len(cols) >= 10:
        return ", ".join(f"column{i:02d} AS {c}" for i, c in enumerate(cols))
    return ", ".join(f"column{i} AS {c}" for i, c in enumerate(cols))


def _listar(extracted_mes: Path, prefixo: str) -> list:
    arquivos = sorted(extracted_mes.glob(f"*{prefixo}*"))
    if not arquivos:
        arquivos = sorted(extracted_mes.glob("K*"))
    return [str(f) for f in arquivos]


def _lista_sql(arquivos: list) -> str:
    """Formata lista de arquivos para uso no read_csv do DuckDB."""
    itens = ", ".join(f"'{f}'" for f in arquivos)
    return f"[{itens}]"


def _municipios_dict(extracted_mes: Path) -> dict:
    """Carrega tabela de municípios em memória como {codigo: nome}."""
    muni_files = _listar(extracted_mes, "MUNIC")
    con = _con()
    rows = con.execute(
        f"SELECT column0, column1 FROM read_csv({_lista_sql(muni_files)}, "
        f"delim=';', header=false, encoding='cp1252', all_varchar=true, ignore_errors=true)"
    ).fetchall()
    return {cod: nome for cod, nome in rows}


def _filtrar_estab_basicos(estab_path: str, cols_estab_sql: str, cnaes_in: str, ufs_in: str) -> list:
    """Passo 1: filtra Estabelecimentos por UF/CNAE/situação — seleciona todas as 30 colunas
    com aliases para evitar bug de mapeamento não-contíguo do DuckDB 1.5."""
    con = _con()
    rows = con.execute(f"""
        SELECT * FROM (
            SELECT {cols_estab_sql}
            FROM read_csv('{estab_path}', delim=';', header=false, quote='"',
                          encoding='cp1252', all_varchar=true, ignore_errors=true)
        )
        WHERE uf IN ({ufs_in})
          AND cnae_fiscal_principal IN ({cnaes_in})
          AND situacao_cadastral = '{SITUACAO_ATIVA}'
    """).fetchall()
    con.close()
    return rows


def _ler_empresas_filtrado(emp_path: str, cnpj_basicos_in: str) -> dict:
    """Passo 2: lê só as linhas de Empresas cujo cnpj_basico está no conjunto filtrado."""
    con = _con()
    rows = con.execute(f"""
        SELECT column0, column1, column2, column4, column5
        FROM read_csv('{emp_path}', delim=';', header=false, quote='"',
                      encoding='cp1252', all_varchar=true, ignore_errors=true)
        WHERE column0 IN ({cnpj_basicos_in})
    """).fetchall()
    con.close()
    # {cnpj_basico: (razao_social, natureza_juridica, capital_social, porte_empresa)}
    return {r[0]: (r[1], r[2], r[3], r[4]) for r in rows}


def rodar_filtro_duckdb(mes_ano: str):
    extracted_mes = EXTRACTED_DIR / mes_ano
    estab_files = sorted(extracted_mes.glob("*ESTABELE*"))
    emp_files   = sorted(extracted_mes.glob("*EMPRE*"))

    if len(estab_files) != len(emp_files):
        raise ValueError(
            f"Número de arquivos não coincide: {len(estab_files)} ESTABELE vs {len(emp_files)} EMPRE"
        )

    ufs_alvo = load_ufs_alvo()
    ufs_in   = ", ".join(f"'{u}'" for u in ufs_alvo)

    cols_estab_sql = _cols_sql(COLS_ESTABELECIMENTOS)
    cnaes_in = ", ".join(f"'{c}'" for c in CNAES_ALIMENTACAO)

    print("[duckdb] carregando tabela de municípios...")
    municipios = _municipios_dict(extracted_mes)

    ufs_desc = ", ".join(ufs_alvo)
    print(f"[duckdb] filtrando estabelecimentos de alimentação em {ufs_desc} ({len(estab_files)} pares)...")
    resultado = []
    for i, (estab, emp) in enumerate(zip(estab_files, emp_files)):
        print(f"  par {i}: {estab.name}")

        # passo 1: filtra estabelecimentos — retorna poucos registros (UF + CNAE)
        # r tem 30 campos na ordem de COLS_ESTABELECIMENTOS:
        # 0=cnpj_basico 1=cnpj_ordem 2=cnpj_dv 3=ident_mat_fil 4=nome_fantasia
        # 5=situacao_cadastral 6=data_situacao 7=motivo 8=cidade_exterior 9=pais
        # 10=data_inicio 11=cnae_principal 12=cnae_secundaria
        # 13=tipo_logr 14=logr 15=numero 16=complemento 17=bairro 18=cep
        # 19=uf 20=municipio 21=ddd1 22=tel1 23=ddd2 24=tel2
        # 25=ddd_fax 26=fax 27=email 28=situacao_especial 29=data_situacao_especial
        estab_rows = _filtrar_estab_basicos(str(estab), cols_estab_sql, cnaes_in, ufs_in)
        if not estab_rows:
            continue

        # passo 2: busca dados de empresa só para os cnpj_basico encontrados
        basicos_in = ", ".join(f"'{r[0]}'" for r in estab_rows)
        empresas   = _ler_empresas_filtrado(str(emp), basicos_in)

        # passo 3: join em Python + resolve município
        for r in estab_rows:
            emp_data = empresas.get(r[0])
            if emp_data is None:
                continue
            razao_social, natureza_juridica, capital_social, porte_empresa = emp_data
            resultado.append((
                r[0] + r[1] + r[2],                    # cnpj
                r[0],                                   # cnpj_basico
                r[3],                                   # identificador_matriz_filial
                razao_social,
                r[4],                                   # nome_fantasia
                natureza_juridica,
                porte_empresa,
                capital_social,
                r[11],                                  # cnae_principal
                r[12],                                  # cnae_secundaria
                r[10],                                  # data_abertura
                r[13],                                  # tipo_logradouro
                r[14],                                  # logradouro
                r[15],                                  # numero
                r[16],                                  # complemento
                r[17],                                  # bairro
                r[18],                                  # cep
                municipios.get(r[20], r[20]),           # municipio
                r[19],                                  # uf
                r[21],                                  # ddd1
                r[22],                                  # telefone1
                r[23],                                  # ddd2
                r[24],                                  # telefone2
                r[27],                                  # email
                r[5],                                   # situacao_cadastral
                r[28],                                  # situacao_especial
                r[29],                                  # data_situacao_especial
            ))

    print(f"[duckdb] {len(resultado)} estabelecimentos encontrados")
    return resultado


def rodar_filtro_socios(mes_ano: str, cnpj_basicos: set):
    """Lê o arquivo de Sócios (QSA) e retorna só as linhas cujo cnpj_basico
    está entre as empresas já filtradas nesta rodada."""
    if not cnpj_basicos:
        return []

    extracted_mes = EXTRACTED_DIR / mes_ano
    socios_files  = _listar(extracted_mes, "SOCIO")

    print(f"[duckdb] {len(socios_files)} arquivo(s) de sócios: {[Path(f).name for f in socios_files[:3]]}")

    con = _con()
    cols_socios_sql = _cols_sql(COLS_SOCIOS)

    # diagnóstico: amostra sem filtro para verificar formato
    sample = con.execute(f"""
        SELECT column00 FROM read_csv({_lista_sql(socios_files[:1])},
            delim=';', header=false, encoding='cp1252', all_varchar=true, ignore_errors=true)
        LIMIT 5
    """).fetchall()
    basicos_sample = list(cnpj_basicos)[:3]
    print(f"[debug] amostra cnpj_basico no CSV de sócios: {[r[0] for r in sample]}")
    print(f"[debug] exemplo de cnpj_basico buscado:       {basicos_sample}")

    # Usa tabela temporária em vez de IN (...) gigante para 64k+ valores
    con.execute("CREATE TEMP TABLE _filtro_basicos (b VARCHAR)")
    con.executemany("INSERT INTO _filtro_basicos VALUES (?)", [(b,) for b in cnpj_basicos])

    print("[duckdb] filtrando sócios/administradores das empresas monitoradas...")
    resultado = con.execute(f"""
        SELECT s.* FROM (
            SELECT {cols_socios_sql}
            FROM read_csv({_lista_sql(socios_files)}, delim=';', header=false, quote='"',
                           encoding='cp1252', all_varchar=true, ignore_errors=true)
        ) s
        JOIN _filtro_basicos f ON s.cnpj_basico = f.b
    """).fetchall()
    con.close()
    print(f"[duckdb] {len(resultado)} registros de sócios encontrados")
    return resultado


def upsert_postgres(registros, mes_ano: str):
    if not registros:
        print("Nenhum registro para inserir.")
        return

    hoje = date.today()
    ano, mes = map(int, mes_ano.split("-"))
    rodada = date(ano, mes, 1)

    conn = psycopg2.connect(POSTGRES_DSN)
    cur = conn.cursor()

    cnpjs = [r[0] for r in registros]
    cur.execute(
        "SELECT cnpj FROM empresas_monitoradas WHERE cnpj = ANY(%s)",
        (cnpjs,),
    )
    ja_existiam = {row[0] for row in cur.fetchall()}

    valores = []
    novos_flags = []
    for row in registros:
        (cnpj, cnpj_basico, ident_matriz_filial, razao_social, nome_fantasia,
         natureza_juridica, porte_empresa, capital_social, cnae_principal,
         cnae_secundaria, data_abertura, tipo_logradouro, logradouro, numero,
         complemento, bairro, cep, municipio, uf, ddd1, telefone1, ddd2,
         telefone2, email, situacao_cadastral, situacao_especial,
         data_situacao_especial) = row

        novos_flags.append(cnpj not in ja_existiam)
        cap = capital_social.replace(',', '.') if capital_social else None
        valores.append((
            cnpj, cnpj_basico, ident_matriz_filial, razao_social, nome_fantasia,
            natureza_juridica, porte_empresa, cap, cnae_principal,
            cnae_secundaria, data_abertura, tipo_logradouro, logradouro, numero,
            complemento, bairro, cep, municipio, uf, ddd1, telefone1, ddd2,
            telefone2, email, situacao_cadastral, situacao_especial,
            data_situacao_especial, rodada, hoje,
        ))

    # Upsert: para CNPJs já existentes, primeira_deteccao NÃO é atualizada
    # (não está no SET), preservando a data original de detecção.
    execute_values(cur, """
        INSERT INTO empresas_monitoradas (
            cnpj, cnpj_basico, identificador_matriz_filial, razao_social,
            nome_fantasia, natureza_juridica, porte_empresa, capital_social,
            cnae_principal, cnae_secundaria, data_abertura, tipo_logradouro,
            logradouro, numero, complemento, bairro, cep, municipio, uf,
            ddd1, telefone1, ddd2, telefone2, email, situacao_cadastral,
            situacao_especial, data_situacao_especial,
            primeira_deteccao, ultima_atualizacao
        ) VALUES %s
        ON CONFLICT (cnpj) DO UPDATE SET
            razao_social = EXCLUDED.razao_social,
            nome_fantasia = EXCLUDED.nome_fantasia,
            natureza_juridica = EXCLUDED.natureza_juridica,
            porte_empresa = EXCLUDED.porte_empresa,
            capital_social = EXCLUDED.capital_social,
            cnae_secundaria = EXCLUDED.cnae_secundaria,
            tipo_logradouro = EXCLUDED.tipo_logradouro,
            logradouro = EXCLUDED.logradouro,
            numero = EXCLUDED.numero,
            complemento = EXCLUDED.complemento,
            bairro = EXCLUDED.bairro,
            cep = EXCLUDED.cep,
            ddd1 = EXCLUDED.ddd1,
            telefone1 = EXCLUDED.telefone1,
            ddd2 = EXCLUDED.ddd2,
            telefone2 = EXCLUDED.telefone2,
            email = EXCLUDED.email,
            situacao_cadastral = EXCLUDED.situacao_cadastral,
            situacao_especial = EXCLUDED.situacao_especial,
            data_situacao_especial = EXCLUDED.data_situacao_especial,
            ultima_atualizacao = EXCLUDED.ultima_atualizacao
    """, valores, template="""(
        %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
        TO_DATE(%s, 'YYYYMMDD'), %s, %s, %s, %s, %s, %s, %s, %s,
        %s, %s, %s, %s, %s, %s, %s,
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'), %s, %s
    )""")

    conn.commit()

    novos = sum(novos_flags)
    print(f"[postgres] upsert de empresas concluído. Novos CNPJs nesta rodada: {novos}")

    cur.close()
    conn.close()


def rodar_filtro_simples(mes_ano: str, cnpj_basicos: set):
    """Lê o arquivo Simples.zip e retorna só as linhas cujo cnpj_basico
    está entre as empresas monitoradas nesta rodada."""
    if not cnpj_basicos:
        return []

    extracted_mes = EXTRACTED_DIR / mes_ano
    simples_files = _listar(extracted_mes, "SIMPLES")

    print(f"[duckdb] {len(simples_files)} arquivo(s) de simples: {[Path(f).name for f in simples_files[:3]]}")

    con = _con()
    cols_simples_sql = _cols_sql(COLS_SIMPLES)

    # Usa tabela temporária em vez de IN (...) gigante
    con.execute("CREATE TEMP TABLE _filtro_simples (b VARCHAR)")
    con.executemany("INSERT INTO _filtro_simples VALUES (?)", [(b,) for b in cnpj_basicos])

    print("[duckdb] filtrando dados do Simples Nacional / MEI...")
    resultado = con.execute(f"""
        SELECT s.* FROM (
            SELECT {cols_simples_sql}
            FROM read_csv({_lista_sql(simples_files)}, delim=';', header=false, quote='"',
                           encoding='cp1252', all_varchar=true, ignore_errors=true)
        ) s
        JOIN _filtro_simples f ON s.cnpj_basico = f.b
    """).fetchall()
    con.close()
    print(f"[duckdb] {len(resultado)} registros do Simples/MEI encontrados")
    return resultado


def upsert_simples(simples, mes_ano: str):
    if not simples:
        print("Nenhum dado do Simples/MEI para inserir.")
        return

    hoje = date.today()

    conn = psycopg2.connect(POSTGRES_DSN)
    cur = conn.cursor()

    valores = []
    for (cnpj_basico, opcao_simples, data_opcao_simples,
         data_exclusao_simples, opcao_mei, data_opcao_mei,
         data_exclusao_mei) in simples:

        valores.append((
            cnpj_basico,
            opcao_simples == "S",
            data_opcao_simples,
            data_exclusao_simples,
            opcao_mei == "S",
            data_opcao_mei,
            data_exclusao_mei,
            hoje,
        ))

    execute_values(cur, """
        INSERT INTO simples_nacional (
            cnpj_basico, opcao_simples, data_opcao_simples,
            data_exclusao_simples, opcao_mei, data_opcao_mei,
            data_exclusao_mei, ultima_atualizacao
        ) VALUES %s
        ON CONFLICT (cnpj_basico) DO UPDATE SET
            opcao_simples         = EXCLUDED.opcao_simples,
            data_opcao_simples    = EXCLUDED.data_opcao_simples,
            data_exclusao_simples = EXCLUDED.data_exclusao_simples,
            opcao_mei             = EXCLUDED.opcao_mei,
            data_opcao_mei        = EXCLUDED.data_opcao_mei,
            data_exclusao_mei     = EXCLUDED.data_exclusao_mei,
            ultima_atualizacao    = EXCLUDED.ultima_atualizacao
    """, valores, template="""(
        %s, %s,
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'),
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'),
        %s,
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'),
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'),
        %s
    )""")

    conn.commit()
    print(f"[postgres] upsert do Simples/MEI concluído. {len(valores)} registros processados.")
    cur.close()
    conn.close()


def valores_para_mapa(registros):
    # gera pares (cnpj_basico, cnpj) — chave=basico para lookup em upsert_socios
    for row in registros:
        yield (row[1], row[0])


def upsert_socios(socios, cnpj_por_basico: dict, mes_ano: str):
    if not socios:
        print("Nenhum sócio/administrador para inserir.")
        return

    hoje = date.today()
    ano, mes = map(int, mes_ano.split("-"))
    rodada = date(ano, mes, 1)

    conn = psycopg2.connect(POSTGRES_DSN)
    cur = conn.cursor()

    valores = []
    avisos_exterior = 0
    for (cnpj_basico, identificador_socio, nome_socio, cpf_cnpj_socio,
         qualificacao_socio, data_entrada_sociedade, pais,
         representante_legal, nome_representante,
         qualificacao_representante, faixa_etaria) in socios:

        cnpj_empresa = cnpj_por_basico.get(cnpj_basico)
        if not cnpj_empresa:
            continue  # sócio de um cnpj_basico fora do filtro desta rodada

        # Identificador 3 = Estrangeiro (residente no exterior).
        # Para eles, cpf_cnpj_socio contém o código do país (não CPF/CNPJ).
        # O UNIQUE (cnpj_basico, cpf_cnpj_socio, qualificacao) pode colidir se
        # dois estrangeiros do mesmo país tiverem a mesma qualificação na empresa;
        # nesse caso o ON CONFLICT atualiza o registro existente (comportamento correto).
        if identificador_socio == "3":
            avisos_exterior += 1

        eh_administrador = qualificacao_socio in QUALIFICACOES_ADMINISTRADOR
        cod_socio = (qualificacao_socio or '').strip()
        cod_repr  = (qualificacao_representante or '').strip()

        valores.append((
            cnpj_basico, cnpj_empresa, nome_socio, cpf_cnpj_socio,
            cod_socio, QUALIFICACOES_MAP.get(cod_socio),
            representante_legal, nome_representante,
            cod_repr, QUALIFICACOES_MAP.get(cod_repr),
            data_entrada_sociedade, eh_administrador, rodada, hoje,
        ))

    if avisos_exterior:
        print(f"[socios] {avisos_exterior} sócio(s) residentes no exterior detectados "
              f"(identificador=3; cpf_cnpj_socio contém código do país).")

    if not valores:
        print("Nenhum sócio associado a empresas monitoradas nesta rodada.")
        return

    execute_values(cur, """
        INSERT INTO socios_administradores (
            cnpj_basico, cnpj_empresa, nome_socio, cpf_cnpj_socio,
            qualificacao_socio_codigo, qualificacao_socio_desc,
            representante_legal_cpf, nome_representante,
            qualificacao_representante_codigo, qualificacao_representante_desc,
            data_entrada_sociedade, eh_administrador,
            primeira_deteccao, ultima_atualizacao
        ) VALUES %s
        ON CONFLICT (cnpj_basico, cpf_cnpj_socio, qualificacao_socio_codigo) DO UPDATE SET
            nome_socio = EXCLUDED.nome_socio,
            representante_legal_cpf = EXCLUDED.representante_legal_cpf,
            nome_representante = EXCLUDED.nome_representante,
            qualificacao_representante_codigo = EXCLUDED.qualificacao_representante_codigo,
            eh_administrador = EXCLUDED.eh_administrador,
            ultima_atualizacao = EXCLUDED.ultima_atualizacao
    """, valores, template="""(
        %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
        TO_DATE(NULLIF(%s, ''), 'YYYYMMDD'), %s, %s, %s
    )""")

    conn.commit()
    print(f"[postgres] upsert de sócios/administradores concluído. {len(valores)} registros processados.")
    cur.close()
    conn.close()


def main(mes_ano: str):
    registros = rodar_filtro_duckdb(mes_ano)
    cnpj_por_basico = dict(valores_para_mapa(registros))
    upsert_postgres(registros, mes_ano)

    cnpj_basicos = set(cnpj_por_basico.keys())
    socios = rodar_filtro_socios(mes_ano, cnpj_basicos)
    upsert_socios(socios, cnpj_por_basico, mes_ano)

    simples = rodar_filtro_simples(mes_ano, cnpj_basicos)
    upsert_simples(simples, mes_ano)


if __name__ == "__main__":
    if len(sys.argv) != 2:
        print("Uso: python -m cnpj_monitor.etl_load AAAA-MM   (ex: 2026-07)")
        sys.exit(1)
    main(sys.argv[1])
