from __future__ import annotations

import base64
from pathlib import Path

import requests
from openai import OpenAI
from PIL import Image, ImageFilter, ImageStat
from tenacity import retry, stop_after_attempt, wait_exponential

CATEGORY_VISUALS = {
    "ai": "futuristic neural interfaces, holographic model cards, luminous compute clusters",
    "hardware": "cinematic product render, exploded chip view, studio lighting, macro detail",
    "startup": "funding chart holograph, modern founder office, growth dashboards, premium pitch energy",
    "cyber": "dark SOC room, encrypted code rain, abstract lock geometry, red alert UI accents",
    "research": "futuristic lab, scientific diagrams as 3D objects, clean research aesthetic",
    "programming": "realistic IDE on ultrawide monitor, developer hands, syntax highlighting glow",
    "robotics": "cinematic humanoid robot, industrial atelier, volumetric light",
    "healthcare": "medical imaging holograms, clean clinical tech, soft cyan lighting",
    "finance": "fintech dashboard glass panels, trading floor blur, premium navy and gold",
    "cloud": "isometric cloud topology, glowing nodes, data center bokeh",
    "opensource": "collaborative coding space, terminal windows, community energy",
    "bigtech": "keynote stage lighting, product silhouette, premium brand-safe abstract",
}

# Preference order: DALL·E 3 when available, else OpenAI Images (gpt-image-1)
IMAGE_MODEL_FALLBACKS = ("dall-e-3", "gpt-image-1")


class ImageGenerator:
    """Marketing-quality stills via OpenAI Images (DALL·E 3 preferred)."""

    def __init__(self, api_key: str, model: str = "dall-e-3"):
        self.client = OpenAI(api_key=api_key)
        self.preferred_model = model
        self.model = model  # may switch after first successful generate

    def build_prompt(
        self,
        *,
        base_prompt: str,
        category: str = "ai",
        illustration_style: str = "editorial_3d",
        background_style: str = "dark_premium",
        palette: list[str] | None = None,
        hook: str = "",
    ) -> str:
        cat = CATEGORY_VISUALS.get(category, CATEGORY_VISUALS["ai"])
        colors = ", ".join(palette[:4]) if palette else "premium brand colors"
        return (
            "Generate HERO ARTWORK ONLY (not a final social post). "
            "Ultra-premium social media campaign visual for a top AI/tech media brand. "
            "Quality bar: Apple keynote / Stripe / Linear / Nvidia campaign — NOT Canva, NOT poster, NOT meme. "
            "NO readable text, NO logos of real companies, NO watermark, NO clipart icons floating in space. "
            "Scene must have strong focal subject + foreground + middle ground + background depth. "
            "Use volumetric lighting, atmospheric depth/fog, rim lighting, reflections, cinematic shadows, HDR feel. "
            "Rich environmental storytelling with realistic textures and layered objects. "
            "Reserve natural 20-30% clean composition space for future text overlay (not a giant empty gradient). "
            "The clean zone should be subtle and integrated into composition."
            " If human faces appear, keep them realistic and non-distorted. "
            f"Illustration style: {illustration_style}. Background: {background_style}. "
            f"Palette: {colors}. Visual language: {cat}. "
            f"Story: {base_prompt or hook}. "
            "Include concrete subjects: devices, UI mockups, 3D objects, environments, data viz metaphors. "
            "Avoid flat framing; compose like a premium campaign hero shot."
        )[:3900]

    def aspect_to_size(self, aspect: str = "square") -> str:
        """Map creative aspect → model size."""
        if self.model.startswith("dall-e"):
            return {"portrait": "1024x1792", "landscape": "1792x1024"}.get(aspect, "1024x1024")
        # gpt-image-1
        return {"portrait": "1024x1536", "landscape": "1536x1024"}.get(aspect, "1024x1024")

    @retry(stop=stop_after_attempt(2), wait=wait_exponential(min=2, max=30), reraise=True)
    def generate(self, prompt: str, out_path: Path, size: str | None = None, aspect: str = "square") -> Path:
        out_path.parent.mkdir(parents=True, exist_ok=True)
        models = []
        for m in (self.preferred_model, *IMAGE_MODEL_FALLBACKS):
            if m not in models:
                models.append(m)

        last_err: Exception | None = None
        for model in models:
            try:
                self.model = model
                use_size = size or self.aspect_to_size(aspect)
                if model.startswith("dall-e"):
                    allowed = {"1024x1024", "1792x1024", "1024x1792"}
                else:
                    allowed = {"1024x1024", "1536x1024", "1024x1536", "auto"}
                use_size = use_size if use_size in allowed else "1024x1024"

                kwargs: dict = {
                    "model": model,
                    "prompt": prompt[:3900],
                    "size": use_size,
                    "n": 1,
                }
                if model.startswith("dall-e"):
                    kwargs["quality"] = "hd"

                out_path = self._generate_once(kwargs, out_path.with_suffix(".png"))
                if self._passes_quality_gate(out_path):
                    return out_path

                # One stronger regeneration pass for consistency.
                stricter = dict(kwargs)
                stricter["prompt"] = (
                    kwargs["prompt"]
                    + " Enforce premium cinematic composition, no flat poster framing, "
                    + "clear focal subject, layered depth, rich textures, and subtle negative space only."
                )[:3900]
                out_path = self._generate_once(stricter, out_path)
                if self._passes_quality_gate(out_path):
                    return out_path
                raise RuntimeError("Generated image failed quality gate twice")
            except Exception as exc:
                last_err = exc
                continue
        raise RuntimeError(f"Image generation failed for models {models}: {last_err}")

    def _generate_once(self, kwargs: dict, out_path: Path) -> Path:
        result = self.client.images.generate(**kwargs)
        item = result.data[0]
        b64 = getattr(item, "b64_json", None)
        url = getattr(item, "url", None)
        if b64:
            out_path.write_bytes(base64.b64decode(b64))
        elif url:
            resp = requests.get(url, timeout=90)
            resp.raise_for_status()
            out_path.write_bytes(resp.content)
        else:
            raise RuntimeError("Image API returned neither url nor b64_json")
        return out_path

    def _passes_quality_gate(self, image_path: Path) -> bool:
        """
        Lightweight quality gate to reduce flat/poster-like generations.
        Reject if image is too low-detail, low-contrast, or overly uniform.
        """
        try:
            img = Image.open(image_path).convert("RGB")
        except Exception:
            return False

        gray = img.convert("L")
        stat = ImageStat.Stat(gray)
        std = stat.stddev[0] if stat.stddev else 0.0

        # Edge density proxy for visual richness
        edges = gray.filter(ImageFilter.FIND_EDGES)
        est = ImageStat.Stat(edges)
        edge_mean = est.mean[0] if est.mean else 0.0

        # Uniformity check on thirds (avoid giant blank gradients)
        w, h = img.size
        thirds = [
            gray.crop((0, 0, w, h // 3)),
            gray.crop((0, h // 3, w, (2 * h) // 3)),
            gray.crop((0, (2 * h) // 3, w, h)),
        ]
        third_stds = [ImageStat.Stat(t).stddev[0] for t in thirds]
        too_flat_regions = sum(1 for s in third_stds if s < 16)

        # Center activity (strong focal subject tends to create detail near center)
        cx0, cy0, cx1, cy1 = w // 4, h // 4, (3 * w) // 4, (3 * h) // 4
        center = gray.crop((cx0, cy0, cx1, cy1))
        center_std = ImageStat.Stat(center).stddev[0]

        return (
            std >= 28
            and edge_mean >= 18
            and center_std >= 24
            and too_flat_regions <= 1
        )
