#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
carta_pdf.py — carta náutica raster (BSB/KAP ou GeoTIFF georreferenciado) → PDF de
impressão em tamanho real, com moldura de coordenadas (grade de latitude/longitude)
desenhada em vetor. Usa só o GDAL (osgeo), nada de QGIS/ArcGIS.

Uso:
  python3 carta_pdf.py entrada.KAP saida.pdf
  python3 carta_pdf.py entrada.KAP saida.pdf --dpi 254 --margem 15 --fonte 7
  python3 carta_pdf.py entrada.KAP saida.pdf --sem-linhas --sem-titulo

O DPI padrão vem do cabeçalho do KAP (campo DU; nas cartas da DHN é 254), o que
faz a página do PDF ter exatamente o tamanho físico da carta na escala nominal.
"""
import argparse
import math
import os
import re
import shutil
import sqlite3
import sys
import tempfile

import numpy as np
from osgeo import gdal, ogr, osr

gdal.UseExceptions()

PT_POR_MM = 72.0 / 25.4  # pontos tipográficos por milímetro

# --- Geometria da moldura, em milímetros a partir da borda do raster (neatline) ---
# Padrão das cartas INT/DHN: faixa interna com tracinhos (décimos de minuto) e faixa
# externa com barras alternadas preto/branco (minutos); números do lado de fora.
BANDA_INT = 1.2    # faixa interna (tracinhos)
BANDA_EXT = 1.2    # faixa externa (barras preto/branco)
LINHA_GROSSA = 3.2 # distância da linha externa grossa
ROTULO = 4.6       # distância onde começa o texto dos rótulos
TITULO = 9.5       # distância do título (acima da carta)
ESP_FINA = 0.25    # espessura das linhas finas da moldura
ESP_MARCA = 0.2    # espessura dos tracinhos
ESP_GROSSA = 0.6   # espessura da linha externa
ESP_GRADE = 0.12   # espessura das linhas de grade dentro da carta


def passos_automaticos(mm_por_minuto):
    """Passo das barras, dos tracinhos e dos números (em minutos) conforme a escala.
    Segue o costume das cartas: barras de 1' sempre que couberem (>= 5 mm), tracinhos
    nos décimos quando legíveis (>= 2 mm), números a 1', 5', 10', 30' ou 1° (>= 50 mm)."""
    if mm_por_minuto >= 200:          # escala muito grande: barra de meio minuto ou menos
        barra = next((s for s in (1, 0.5, 0.2, 0.1) if s * mm_por_minuto <= 200), 0.1)
    elif mm_por_minuto >= 5:
        barra = 1
    else:                              # escala pequena: barras de 2', 5', 10'...
        barra = next((s for s in (2, 5, 10, 30) if s * mm_por_minuto >= 5), 30)
    marca = next((s for s in (0.1, 0.2, 0.5, 1, 2, 5, 10)
                  if s * mm_por_minuto >= 2.0 and s <= barra and abs(barra / s - round(barra / s)) < 1e-6), barra)
    rotulo = next((s for s in (0.1, 0.2, 0.5, 1, 5, 10, 30, 60)
                   if s >= barra and s * mm_por_minuto >= 50 and abs(s / barra - round(s / barra)) < 1e-6), barra)
    return barra, marca, rotulo


def ler_cabecalho_kap(caminho):
    """Lê a parte texto do KAP (até o byte 0x1A) e devolve os campos de BSB/ e KNP/."""
    info = {}
    try:
        with open(caminho, 'rb') as f:
            bruto = f.read(65536)
    except OSError:
        return info
    texto = bruto.split(b'\x1a')[0].decode('latin-1', 'replace')
    texto = re.sub(r'\r?\n[ \t]+', ',', texto)  # linhas de continuação
    for linha in texto.splitlines():
        chave, sep, resto = linha.partition('/')
        if sep and chave in ('BSB', 'KNP'):
            for par in resto.split(','):
                k, s, v = par.partition('=')
                if s:
                    info[k.strip()] = v.strip()
    return info


def escolhe_passo(mm_por_minuto, minimo_mm, candidatos):
    for p in candidatos:
        if p * mm_por_minuto >= minimo_mm:
            return p
    return candidatos[-1]


def fmt_minutos(centesimos, passo_c):
    """centésimos de minuto (0..5999) → texto tipo 20' ou 20,5'."""
    m = centesimos // 100
    frac = centesimos % 100
    if passo_c % 100 == 0 or frac == 0:
        return f"{m}'"
    if passo_c % 10 == 0:
        return f"{m},{frac // 10}'"
    return f"{m},{frac:02d}'"


def rotulo_coord(valor_c, passo_c, eixo, com_grau):
    """valor_c em centésimos de minuto (com sinal). eixo 'lon' ou 'lat'."""
    neg = valor_c < 0
    v = abs(valor_c)
    graus, resto = divmod(v, 6000)
    hemi = ('W' if neg else 'E') if eixo == 'lon' else ('S' if neg else 'N')
    if resto == 0:
        return f"{graus}°{hemi}"
    minutos = fmt_minutos(resto, passo_c)
    if com_grau:
        return f"{graus}°{minutos}{hemi}"
    return minutos


class Camada:
    """Acumula feições (WKT + estilo OGR) para gravar num GeoPackage."""

    def __init__(self, nome):
        self.nome = nome
        self.feicoes = []

    def linha(self, pontos, esp_mm, cor='#000000'):
        wkt = 'LINESTRING(' + ','.join(f'{x:.6f} {y:.6f}' for x, y in pontos) + ')'
        self.feicoes.append((wkt, f'PEN(c:{cor},w:{esp_mm * PT_POR_MM:.4f}mm)'))

    def texto(self, x, y, txt, tam_pt, ancora, fonte='Helvetica', cor='#000000'):
        txt = txt.replace('\\', '\\\\').replace('"', '\\"')
        self.feicoes.append((f'POINT({x:.6f} {y:.6f})',
                             f'LABEL(f:"{fonte}",s:{tam_pt}mm,t:"{txt}",c:{cor},p:{ancora})'))


def gravar_gpkg(caminho, srs_wkt, camadas):
    srs = osr.SpatialReference()
    srs.ImportFromWkt(srs_wkt)
    ds = ogr.GetDriverByName('GPKG').CreateDataSource(caminho)
    for cam in camadas:
        ly = ds.CreateLayer(cam.nome, srs, ogr.wkbUnknown, options=['SPATIAL_INDEX=NO'])
        ly.CreateField(ogr.FieldDefn('OGR_STYLE', ogr.OFTString))
        defn = ly.GetLayerDefn()
        for wkt, estilo in cam.feicoes:
            f = ogr.Feature(defn)
            f.SetGeometry(ogr.CreateGeometryFromWkt(wkt))
            f.SetField('OGR_STYLE', estilo)
            ly.CreateFeature(f)
    ds = None
    # A fonte Helvetica do PDF usa WinAnsi (1 byte por caractere). O GDAL copia os
    # bytes do texto como estão, então gravamos os rótulos com acento/grau em Latin-1.
    con = sqlite3.connect(caminho)
    con.text_factory = bytes
    for cam in camadas:
        linhas = con.execute(f'SELECT fid, OGR_STYLE FROM "{cam.nome}"').fetchall()
        for fid, estilo in linhas:
            s = estilo.decode('utf-8')
            if any(ord(ch) > 127 for ch in s):
                con.execute(f'UPDATE "{cam.nome}" SET OGR_STYLE=? WHERE fid=?',
                            (s.encode('latin-1', 'replace'), fid))
    con.commit()
    con.close()


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument('entrada', help='arquivo .KAP (BSB) ou GeoTIFF georreferenciado')
    ap.add_argument('saida', help='PDF de saída')
    ap.add_argument('--dpi', type=float, default=None,
                    help='DPI do raster no papel (padrão: DU do cabeçalho do KAP, senão 254)')
    ap.add_argument('--margem', type=float, default=15.0, help='margem branca em mm ao redor da carta (padrão 15)')
    ap.add_argument('--fonte', type=float, default=7.0, help='tamanho dos rótulos em pt (padrão 7)')
    ap.add_argument('--barra', type=float, default=None, help='minutos por barra preto/branco (padrão: automático)')
    ap.add_argument('--marca', type=float, default=None, help='minutos entre tracinhos da faixa interna (padrão: automático)')
    ap.add_argument('--rotulo', type=float, default=None, help='minutos entre números (padrão: automático)')
    ap.add_argument('--sem-linhas', action='store_true', help='não desenha as linhas de grade dentro da carta')
    ap.add_argument('--sem-titulo', action='store_true', help='não escreve o título acima da carta')
    ap.add_argument('--titulo', default=None, help='texto do título (padrão: número, nome e escala do cabeçalho)')
    args = ap.parse_args()

    cab = ler_cabecalho_kap(args.entrada)
    dpi = args.dpi or float(cab.get('DU', 254))
    if args.margem < ROTULO + 6:
        sys.exit(f'--margem precisa ser >= {ROTULO + 6:.0f} mm para caber a moldura e os rótulos')

    src = gdal.Open(args.entrada)
    W, H = src.RasterXSize, src.RasterYSize
    gt = src.GetGeoTransform(can_return_null=True)
    if gt is None and src.GetGCPCount() >= 3:
        gt = gdal.GCPsToGeoTransform(src.GetGCPs())
    if gt is None:
        sys.exit('raster sem georreferenciamento')
    if abs(gt[2]) > 1e-9 or abs(gt[4]) > 1e-9:
        sys.exit('raster rotacionado; reprojete antes (gdalwarp) para uma grade alinhada')
    srs_wkt = src.GetProjection() or src.GetGCPProjection()
    if not srs_wkt:
        sys.exit('raster sem sistema de coordenadas')

    u = gt[1]                      # unidades projetadas por pixel
    px_por_mm = dpi / 25.4
    def mm2u(mm):                  # milímetros no papel → unidades projetadas
        return mm * px_por_mm * u

    # --- raster preenchido com margem branca (cópia pixel a pixel, sem reamostrar) ---
    margem_px = int(round(args.margem * px_por_mm))
    ct = src.GetRasterBand(1).GetRasterColorTable()
    tmpdir = tempfile.mkdtemp(prefix='carta_pdf_')
    try:
        preenchido = os.path.join(tmpdir, 'preenchido.tif')
        if ct is not None and src.RasterCount == 1:
            brancos = [i for i in range(ct.GetCount()) if ct.GetColorEntry(i)[:3] == (255, 255, 255)]
            if brancos:
                dados = src.GetRasterBand(1).ReadAsArray()
                fundo, bandas = brancos[0], [dados]
            else:
                ct = None
        if ct is None:
            if src.RasterCount == 1:
                src = gdal.Translate(os.path.join(tmpdir, 'rgb.tif'), src, rgbExpand='rgb')
            bandas = [src.GetRasterBand(i + 1).ReadAsArray() for i in range(3)]
            fundo = 255
        drv = gdal.GetDriverByName('GTiff')
        out = drv.Create(preenchido, W + 2 * margem_px, H + 2 * margem_px, len(bandas), gdal.GDT_Byte,
                         options=['COMPRESS=DEFLATE', 'BIGTIFF=IF_SAFER'])
        out.SetGeoTransform((gt[0] - margem_px * u, u, 0.0, gt[3] + margem_px * u, 0.0, -u))
        out.SetProjection(srs_wkt)
        for i, b in enumerate(bandas):
            arr = np.full((H + 2 * margem_px, W + 2 * margem_px), fundo, np.uint8)
            arr[margem_px:margem_px + H, margem_px:margem_px + W] = b
            banda = out.GetRasterBand(i + 1)
            if ct is not None:  # paleta antes dos pixels (exigência do GTiff)
                banda.SetRasterColorTable(ct)
                banda.SetColorInterpretation(gdal.GCI_PaletteIndex)
            banda.WriteArray(arr)
        out.FlushCache()
        out = None
        del bandas

        # --- transformações geográficas ---
        srs = osr.SpatialReference(); srs.ImportFromWkt(srs_wkt)
        srs.SetAxisMappingStrategy(osr.OAMS_TRADITIONAL_GIS_ORDER)
        ll = osr.SpatialReference(); ll.ImportFromEPSG(4326)
        ll.SetAxisMappingStrategy(osr.OAMS_TRADITIONAL_GIS_ORDER)
        para_proj = osr.CoordinateTransformation(ll, srs)
        para_ll = osr.CoordinateTransformation(srs, ll)

        x0, x1 = gt[0], gt[0] + W * u          # bordas oeste/leste do raster
        y1, y0 = gt[3], gt[3] - H * u          # bordas norte (y1) / sul (y0)
        xm, ym = (x0 + x1) / 2, (y0 + y1) / 2
        lon_w, lat_m = para_ll.TransformPoint(x0, ym)[:2]
        lon_e = para_ll.TransformPoint(x1, ym)[0]
        lon_m = (lon_w + lon_e) / 2
        lat_s = para_ll.TransformPoint(xm, y0)[1]
        lat_n = para_ll.TransformPoint(xm, y1)[1]
        def x_de_lon(lon): return para_proj.TransformPoint(lon, lat_m)[0]
        def y_de_lat(lat): return para_proj.TransformPoint(lon_m, lat)[1]

        # mm de papel por minuto de arco em cada eixo
        mm_min_lon = abs(x_de_lon(lon_m + 1 / 60) - x_de_lon(lon_m)) / u / px_por_mm
        mm_min_lat = abs(y_de_lat(lat_m + 1 / 60) - y_de_lat(lat_m)) / u / px_por_mm
        mm_min = min(mm_min_lon, mm_min_lat)
        passo_barra, passo_marca, passo_rotulo = passos_automaticos(mm_min)
        passo_barra = args.barra or passo_barra
        passo_marca = args.marca or passo_marca
        passo_rotulo = args.rotulo or passo_rotulo
        pb_c, pm_c, pr_c = (int(round(p * 100)) for p in (passo_barra, passo_marca, passo_rotulo))

        def marcas(v_ini, v_fim, passo_c):
            """valores (centésimos de minuto) múltiplos de passo_c dentro de [v_ini, v_fim]."""
            a, b = sorted((v_ini * 6000, v_fim * 6000))
            i0 = math.ceil(a / passo_c - 1e-6)
            i1 = math.floor(b / passo_c + 1e-6)
            return [i * passo_c for i in range(i0, i1 + 1)]

        moldura = Camada('Moldura')
        grade = Camada('Grade')

        def faixa(eixo):
            """Faixas, tracinhos e números ao longo de um eixo ('lon' → bordas sup/inf, 'lat' → esq/dir)."""
            if eixo == 'lon':
                ini, fim = lon_w, lon_e
                pos = lambda c: x_de_lon(c / 6000)
                p_ini, p_fim = x0, x1
                valor_em = lambda p: para_ll.TransformPoint(p, ym)[0]
            else:
                ini, fim = lat_s, lat_n
                pos = lambda c: y_de_lat(c / 6000)
                p_ini, p_fim = y0, y1
                valor_em = lambda p: para_ll.TransformPoint(xm, p)[1]

            def transversal(p, d1, d2, esp):
                """Segmento perpendicular à borda, na posição p, entre as distâncias d1 e d2 (mm)."""
                a, b = mm2u(d1), mm2u(d2)
                if eixo == 'lon':
                    moldura.linha([(p, y0 - a), (p, y0 - b)], esp)
                    moldura.linha([(p, y1 + a), (p, y1 + b)], esp)
                else:
                    moldura.linha([(x0 - a, p), (x0 - b, p)], esp)
                    moldura.linha([(x1 + a, p), (x1 + b, p)], esp)

            dentro = lambda p: p_ini - 1e-6 <= p <= p_fim + 1e-6

            # faixa interna: tracinhos (meia altura; altura inteira nas divisões da barra)
            for c in marcas(ini, fim, pm_c):
                p = pos(c)
                if not dentro(p):
                    continue
                inteiro = (c % pb_c == 0) or (pb_c % 2 == 0 and c % (pb_c // 2) == 0)
                transversal(p, 0.0, BANDA_INT if inteiro else BANDA_INT / 2, ESP_MARCA)

            # faixa externa: barras alternadas preto/branco a cada passo_barra
            limites = [p_ini] + [pos(c) for c in marcas(ini, fim, pb_c) if p_ini < pos(c) < p_fim] + [p_fim]
            for a, b in zip(limites[:-1], limites[1:]):
                if b - a <= 0:
                    continue
                v = valor_em((a + b) / 2)
                if int(math.floor(v * 6000 / pb_c + 1e-6)) % 2 != 0:
                    continue
                d = mm2u(BANDA_INT + BANDA_EXT / 2)
                if eixo == 'lon':
                    moldura.linha([(a, y0 - d), (b, y0 - d)], BANDA_EXT)
                    moldura.linha([(a, y1 + d), (b, y1 + d)], BANDA_EXT)
                else:
                    moldura.linha([(x0 - d, a), (x0 - d, b)], BANDA_EXT)
                    moldura.linha([(x1 + d, a), (x1 + d, b)], BANDA_EXT)
            for c in marcas(ini, fim, pb_c):       # divisórias das barras, através das duas faixas
                if dentro(pos(c)):
                    transversal(pos(c), 0.0, BANDA_INT + BANDA_EXT, ESP_FINA)

            # números: marca até a linha grossa e texto do lado de fora
            rot = [c for c in marcas(ini, fim, pr_c) if dentro(pos(c))]
            primeiro = max(rot) if eixo == 'lat' else min(rot)   # primeiro lido: norte / oeste
            for c in rot:
                p = pos(c)
                txt = rotulo_coord(c, pr_c, eixo, com_grau=(c == primeiro))
                transversal(p, BANDA_INT + BANDA_EXT, LINHA_GROSSA, ESP_FINA)
                r = mm2u(ROTULO)
                if eixo == 'lon':
                    moldura.texto(p, y0 - r, txt, args.fonte, 8)    # abaixo: âncora topo-centro
                    moldura.texto(p, y1 + r, txt, args.fonte, 11)   # acima: âncora base-centro
                    if not args.sem_linhas and p_ini < p < p_fim:
                        grade.linha([(p, y0), (p, y1)], ESP_GRADE)
                else:
                    moldura.texto(x0 - r, p, txt, args.fonte, 6)    # esquerda: âncora centro-direita
                    moldura.texto(x1 + r, p, txt, args.fonte, 4)    # direita: âncora centro-esquerda
                    if not args.sem_linhas and p_ini < p < p_fim:
                        grade.linha([(x0, p), (x1, p)], ESP_GRADE)

        faixa('lon')
        faixa('lat')

        # molduras retangulares por cima das barras
        for dist, esp in ((0.0, ESP_FINA), (BANDA_INT, ESP_FINA), (BANDA_INT + BANDA_EXT, ESP_FINA), (LINHA_GROSSA, ESP_GROSSA)):
            d = mm2u(dist)
            moldura.linha([(x0 - d, y0 - d), (x1 + d, y0 - d), (x1 + d, y1 + d), (x0 - d, y1 + d), (x0 - d, y0 - d)], esp)

        # título
        titulo = args.titulo
        if titulo is None:
            partes = [cab.get('NU'), cab.get('NA')]
            if cab.get('SC'):
                partes.append(f"Escala 1:{int(float(cab['SC'])):,}".replace(',', '.'))
            if cab.get('PR'):
                partes.append(f"Projeção {cab['PR']}" + (f" ({cab['GD']})" if cab.get('GD') else ''))
            titulo = '   ·   '.join(p for p in partes if p)
        if titulo and not args.sem_titulo:
            moldura.texto(x0, y1 + mm2u(TITULO), titulo, args.fonte + 1, 10, fonte='Helvetica-Bold')

        camadas = [moldura] + ([grade] if grade.feicoes else [])
        gpkg = os.path.join(tmpdir, 'moldura.gpkg')
        gravar_gpkg(gpkg, srs_wkt, camadas)

        opcoes = [
            f'DPI={dpi:g}',
            'WRITE_USERUNIT=NO',       # MediaBox em pontos: tamanho físico inequívoco na gráfica
            'COMPRESS=DEFLATE',        # sem perdas
            'LAYER_NAME=Carta',
            f'OGR_DATASOURCE={gpkg}',
            'OGR_DISPLAY_LAYER_NAMES=' + ','.join(c.nome for c in camadas),
            'OGR_WRITE_ATTRIBUTES=NO',
            'GEO_ENCODING=ISO32000',
        ]
        if titulo:
            opcoes.append(f'TITLE={titulo}')
        gdal.Translate(args.saida, preenchido, format='PDF', creationOptions=opcoes)
    finally:
        shutil.rmtree(tmpdir, ignore_errors=True)

    larg_mm = (W + 2 * margem_px) / px_por_mm
    alt_mm = (H + 2 * margem_px) / px_por_mm
    print(f'PDF gravado: {args.saida}')
    print(f'  página: {larg_mm:.1f} x {alt_mm:.1f} mm  (carta {W / px_por_mm:.1f} x {H / px_por_mm:.1f} mm + margem {args.margem:g} mm)')
    print(f'  DPI: {dpi:g}  |  tracinhos a cada {passo_marca:g}\'  |  barras a cada {passo_barra:g}\'  |  números e grade a cada {passo_rotulo:g}\'')
    if cab.get('SC') and cab.get('DU'):
        esc = float(cab['SC']) * dpi / float(cab['DU'])  # pixel maior no papel => escala maior
        print(f'  escala no papel: 1:{esc:,.0f}'.replace(',', '.') + f'  (nominal da carta 1:{int(float(cab["SC"])):,})'.replace(',', '.'))
    print('  na gráfica: imprimir em "tamanho real" / 100%, sem "ajustar à página"')


if __name__ == '__main__':
    main()
