"""KARC Live-Demo Backend (von KarcController via Symfony Process aufgerufen).

Ehrlicher KARC-vs-JPEG-Kopf-an-Kopf:
  - BD-Rate ueber die ganze Rate-Distortion-Kurve DES BILDES (die korrekte Kennzahl,
    an der KARC JPEG real schlaegt; ein Einzelpunkt waere irrefuehrend).
  - Beide RD-Kurven (bpp vs PSNR) fuer ein Diagramm.
  - Original / KARC / JPEG bei gleicher Dateigroesse zum Selber-Beurteilen.
WebP/AVIF werden NICHT als "Gewinner-Plakette" gezeigt (ehrlich auf der Testbericht-Seite).

Aufruf:  python compare.py <bildpfad> [level 1..3]
"""
import os
import sys
import json
import base64
import warnings
from io import BytesIO
import numpy as np
from PIL import Image

# Decompression-Bomben hart abwehren: Maße deckeln und PILs Bomb-Warnung in einen
# Fehler verwandeln, damit ein winziges Riesen-Bild nicht erst GB an RAM allokiert.
Image.MAX_IMAGE_PIXELS = 8000 * 8000
warnings.simplefilter('error', Image.DecompressionBombWarning)

BILD = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'codec', 'verfahren')
for _p in ('common', '09_integration_codec', '10_evaluation_tests'):
    sys.path.insert(0, os.path.join(BILD, _p))
import metrics       # noqa: E402
import baselines     # noqa: E402
import marc_codec    # noqa: E402

MAX_SIDE = 400                          # kleiner = deutlich schneller (Demo bleibt fluessig)
QSWEEP = [14, 24, 40, 60]               # 4 Punkte reichen fuer die BD-Rate (kubisch)
LEVEL_QSTEP = {1: 60, 2: 40, 3: 24}
JQ = list(range(8, 96, 6))


def b64(arr):
    bio = BytesIO()
    Image.fromarray(arr.astype(np.uint8)).save(bio, 'PNG')
    return 'data:image/png;base64,' + base64.b64encode(bio.getvalue()).decode('ascii')


def load_capped(path):
    im = Image.open(path)
    im.load()
    im = im.convert('RGB')
    w, h = im.size
    s = max(w, h)
    if s > MAX_SIDE:
        f = MAX_SIDE / s
        im = im.resize((max(1, int(w * f)), max(1, int(h * f))), Image.LANCZOS)
    return np.asarray(im)


def main():
    if len(sys.argv) < 2:
        print(json.dumps({'ok': False, 'message': 'kein Bild'}))
        return 1
    level = 2
    if len(sys.argv) >= 3:
        try:
            level = max(1, min(3, int(sys.argv[2])))
        except ValueError:
            pass
    try:
        img = load_capped(sys.argv[1])
    except Exception as e:
        print(json.dumps({'ok': False, 'message': 'Kein gueltiges Bild: %s' % e}))
        return 1
    H, W = img.shape[:2]

    karc = {}
    for qs in QSWEEP:
        blob = marc_codec.encode(img, qs)
        dec = marc_codec.decode(blob)
        karc[qs] = {'bpp': metrics.bpp(len(blob), H, W), 'psnr': metrics.psnr(img, dec),
                    'ssim': metrics.ssim(img, dec), 'dec': dec}
    jpeg = []
    for q in JQ:
        size, dec = baselines.jpeg_rd(img, q)
        jpeg.append({'bpp': metrics.bpp(size, H, W), 'psnr': metrics.psnr(img, dec),
                     'ssim': metrics.ssim(img, dec), 'q': q, 'dec': dec})

    try:
        bd = metrics.bd_rate([(p['bpp'], p['psnr']) for p in jpeg],
                             [(karc[q]['bpp'], karc[q]['psnr']) for q in QSWEEP], 'psnr')
    except Exception:
        bd = float('nan')

    vqs = LEVEL_QSTEP[level]
    kv = karc[vqs]
    tb = kv['bpp']
    jv = min(jpeg, key=lambda p: abs(p['bpp'] - tb))

    out = {
        'ok': True, 'width': int(W), 'height': int(H), 'level': level,
        'bd_jpeg': None if bd != bd else round(bd, 1),
        'target_bpp': round(tb, 3),
        'curves': {
            'karc': sorted([round(karc[q]['bpp'], 4), round(karc[q]['psnr'], 3)] for q in QSWEEP),
            'jpeg': sorted([round(p['bpp'], 4), round(p['psnr'], 3)] for p in jpeg),
        },
        'original': b64(img),
        'karc': {'bpp': round(kv['bpp'], 3), 'psnr': round(kv['psnr'], 2), 'ssim': round(kv['ssim'], 4), 'image': b64(kv['dec'])},
        'jpeg': {'bpp': round(jv['bpp'], 3), 'psnr': round(jv['psnr'], 2), 'ssim': round(jv['ssim'], 4), 'quality': int(jv['q']), 'image': b64(jv['dec'])},
    }
    sys.stdout.write(json.dumps(out))
    return 0


if __name__ == '__main__':
    raise SystemExit(main())
