from pathlib import Path
import argparse
import json
import os
from datetime import datetime

from PIL import Image, ImageOps

# Run from the project root, examples:
#   python scripts/create_thumbnails.py
#   python scripts/create_thumbnails.py --regenerate
#   python scripts/create_thumbnails.py --state-file instance/image_variant_job.json
#   python scripts/create_thumbnails.py --root static/uploads/listings/ST25695
#
# It creates two lightweight WebP variants next to original listing/admin images:
#   image.jpg -> image_thumb.webp   (cards/admin)
#   image.jpg -> image_gallery.webp (public gallery/lightbox)
#
# Originals remain untouched for export/admin download.

ROOTS = [
    Path("static/uploads/listings"),
    Path("static/uploads/generated_visuals"),
    Path("static/uploads/ai_edits"),
]

THUMB_SUFFIX = "_thumb"
GALLERY_SUFFIX = "_gallery"
THUMB_SIZE = (700, 525)
GALLERY_SIZE = (1800, 1350)
THUMB_QUALITY = 72
GALLERY_QUALITY = 82
EXTS = {".jpg", ".jpeg", ".png", ".webp", ".gif"}


def now_utc():
    return datetime.utcnow().replace(microsecond=0).isoformat() + "Z"


def resample_filter():
    try:
        return Image.Resampling.LANCZOS
    except AttributeError:
        return getattr(Image, "LANCZOS", getattr(Image, "ANTIALIAS", None))


def variant_path(img_path, suffix):
    if img_path.stem.endswith((THUMB_SUFFIX, GALLERY_SUFFIX)):
        return None
    return img_path.with_name(f"{img_path.stem}{suffix}.webp")


def create_variant(img_path, suffix, max_size, quality, regenerate):
    out = variant_path(img_path, suffix)
    if out is None:
        return False, "derived-source"
    if out.exists() and not regenerate:
        return False, "exists"
    try:
        with Image.open(img_path) as im:
            im = ImageOps.exif_transpose(im)
            if im.mode not in ("RGB", "L"):
                im = im.convert("RGB")
            elif im.mode == "L":
                im = im.convert("RGB")
            filt = resample_filter()
            if filt is not None:
                im.thumbnail(max_size, filt)
            else:
                im.thumbnail(max_size)
            im.save(out, "WEBP", quality=quality, method=6)
        return True, str(out)
    except Exception as exc:
        return False, str(exc)


def candidate_images(roots):
    images = []
    for root in roots:
        if not root.exists():
            continue
        for img in root.rglob("*"):
            if not img.is_file():
                continue
            if img.suffix.lower() not in EXTS:
                continue
            if img.stem.endswith((THUMB_SUFFIX, GALLERY_SUFFIX)):
                continue
            images.append(img)
    return images


def needs_work(img_path):
    thumb = variant_path(img_path, THUMB_SUFFIX)
    gallery = variant_path(img_path, GALLERY_SUFFIX)
    return bool((thumb and not thumb.exists()) or (gallery and not gallery.exists()))


def write_state(path, payload):
    if not path:
        return
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    payload = dict(payload)
    payload["updated_at"] = now_utc()
    path.write_text(json.dumps(payload, indent=2), encoding="utf-8")


def process_image(img_path, regenerate):
    written = 0
    skipped = 0
    errors = []
    for suffix, max_size, quality in (
        (THUMB_SUFFIX, THUMB_SIZE, THUMB_QUALITY),
        (GALLERY_SUFFIX, GALLERY_SIZE, GALLERY_QUALITY),
    ):
        ok, message = create_variant(img_path, suffix, max_size, quality, regenerate)
        if ok:
            written += 1
            print(f"OK: {message}")
        elif message == "exists":
            skipped += 1
        elif message != "derived-source":
            errors.append(f"{img_path}: {message}")
            print(f"SKIP: {img_path} -> {message}")
    return written, skipped, errors


def main():
    parser = argparse.ArgumentParser(description="Generate _thumb.webp and _gallery.webp files.")
    parser.add_argument("--regenerate", action="store_true", help="Rebuild files even if they already exist.")
    parser.add_argument("--root", action="append", default=[], help="Optional custom root folder(s) to scan.")
    parser.add_argument("--state-file", default="", help="Optional JSON state file updated during background runs.")
    args = parser.parse_args()

    roots = [Path(item) for item in args.root] if args.root else ROOTS
    images = candidate_images(roots)
    if not args.regenerate:
        images = [img for img in images if needs_work(img)]

    created = 0
    skipped = 0
    processed = 0
    errors = []
    state = {
        "status": "running",
        "started_at": now_utc(),
        "pid": os.getpid(),
        "regenerate": bool(args.regenerate),
        "processed_images": 0,
        "total_images": len(images),
        "created_variants": 0,
        "skipped_variants": 0,
        "error_count": 0,
        "current_image": "",
    }
    write_state(args.state_file, state)

    try:
        for img in images:
            state["current_image"] = str(img)
            write_state(args.state_file, state)
            written, skipped_local, errs = process_image(img, regenerate=args.regenerate)
            processed += 1
            created += written
            skipped += skipped_local
            errors.extend(errs)
            state.update(
                {
                    "processed_images": processed,
                    "created_variants": created,
                    "skipped_variants": skipped,
                    "error_count": len(errors),
                }
            )
            if errs:
                state["last_error"] = errs[-1]
            if processed % 5 == 0:
                write_state(args.state_file, state)

        state.update(
            {
                "status": "completed",
                "finished_at": now_utc(),
                "processed_images": processed,
                "created_variants": created,
                "skipped_variants": skipped,
                "error_count": len(errors),
                "current_image": "",
            }
        )
        if errors:
            state["last_error"] = errors[-1]
        write_state(args.state_file, state)
        print(f"Done. Processed={processed}, written={created}, skipped={skipped}, errors={len(errors)}")
        if errors:
            print("First errors:")
            for error in errors[:12]:
                print(f" - {error}")
    except Exception as exc:
        state.update(
            {
                "status": "failed",
                "finished_at": now_utc(),
                "processed_images": processed,
                "created_variants": created,
                "skipped_variants": skipped,
                "error_count": len(errors) + 1,
                "last_error": str(exc),
            }
        )
        write_state(args.state_file, state)
        raise


if __name__ == "__main__":
    main()
