"""생성 이미지 바이트 형식 정규화 (2026-08-06 사용자 지시 — PNG 로 정규화).

발단은 실측이다. 프로젝트의 생성 산출은 전부 `.png` 이름으로 저장되는데,
멀티롤 롤(`_a`/`_b`/`_c`)·선정본(`_sel`)·수정본(`_fix`) 의 **내용은 JPEG**
였다 (C 실행 1,149장 전건, A 실행 888장 전건 — 7월부터 같은 상태).
Gemini 이미지 모델이 `inlineData.mimeType: image/jpeg` 로 돌려주는데
클라이언트가 mimeType 을 읽지 않고 바이트를 그대로 `.png` 에 썼기 때문이다.

두 가지가 걸린다.

① **이름과 내용이 어긋난다.** Gemini 판정은 media type 불일치를 조용히
   받아줬지만 Anthropic 은 400 으로 거부한다("specified using the image/png
   media type, but the image appears to be a image/jpeg image"). 선정 판정을
   Claude 로 옮기는 순간 전 호출이 죽는다.

② **i2i 체인마다 재압축된다.** 롤 → 선정본 → 수정본으로 이어지는 편집
   체인이 JPEG 를 매번 다시 인코딩한다. 손실이 누적된다.

그래서 **생성 경로의 반환 지점에서 PNG 로 정규화**한다. 이미 PNG 면 바이트를
그대로 돌려준다(무손실·무변경). 이미 잃은 화질은 되돌아오지 않지만 그 뒤로
더 잃지는 않는다.
"""
from __future__ import annotations

import io
import logging

logger = logging.getLogger(__name__)

_PNG_MAGIC = b"\x89PNG\r\n\x1a\n"


def is_png(data: bytes) -> bool:
    return bool(data) and data.startswith(_PNG_MAGIC)


def ensure_png_bytes(data: bytes, *, context: str = "") -> bytes:
    """이미지 바이트를 PNG 로 정규화한다.

    이미 PNG 면 **같은 객체를 그대로** 돌려준다 — 재인코딩하지 않으므로
    기존 산출과 byte-identical 이다.

    디코드하지 못하는 바이트(빈 값·손상·미지원 포맷)는 손대지 않고 그대로
    돌려준다. 여기서 예외를 올리면 이미 유료로 받아 온 이미지를 버리게 된다 —
    형식 정규화가 생성 실패로 번지면 안 된다.
    """
    if not data or is_png(data):
        return data
    try:
        from PIL import Image

        with Image.open(io.BytesIO(data)) as im:
            # JPEG 는 알파가 없다. RGB/L 그대로 두고, 팔레트·알파 계열만
            # PNG 가 다룰 수 있는 모드로 맞춘다.
            if im.mode in ("P", "LA", "PA"):
                im = im.convert("RGBA")
            elif im.mode == "CMYK":
                im = im.convert("RGB")
            buf = io.BytesIO()
            im.save(buf, format="PNG", optimize=False)
        out = buf.getvalue()
    except Exception as exc:  # noqa: BLE001
        logger.warning(
            "ensure_png_bytes: PNG 정규화 실패 — 원본 바이트 유지 (%s): %r",
            context or "?", exc,
        )
        return data
    logger.info(
        "ensure_png_bytes: %s → PNG 정규화 (%d → %d bytes)%s",
        _sniff(data), len(data), len(out),
        f" [{context}]" if context else "",
    )
    return out


def _sniff(data: bytes) -> str:
    if data.startswith(b"\xff\xd8\xff"):
        return "JPEG"
    if data.startswith((b"GIF87a", b"GIF89a")):
        return "GIF"
    if data[:4] == b"RIFF" and data[8:12] == b"WEBP":
        return "WEBP"
    return "unknown"
