"""地图合成核心:Wang 角点自动拼接渲染地形 + 物件按 y 序摆放 + 夜景光晕。

设计:
- 地形用「顶点等级网格」表达(0=最低地形,逐级 +1),相邻 tileset 链式过渡;
  每个格子的 4 个角采样顶点等级,选出对应 Wang tile(index = NW*8+NE*4+SW*2+SE)。
- 约束:相邻顶点等级差 ≤1(跨级无过渡 tile),`smooth_levels` 自动压平。
- 物件层按"脚底 y"排序后依次 alpha 合成,天然形成前后遮挡。
"""

from __future__ import annotations

import json
import logging
from dataclasses import dataclass
from pathlib import Path

from PIL import Image, ImageDraw, ImageFilter

logger = logging.getLogger(__name__)


class WangTileSheet:
    """一套 16 片的 Wang 过渡 tileset(lower→upper),按角点 index 取 tile。"""

    def __init__(self, tiles: dict[int, Image.Image], tile_px: int, lower: str, upper: str):
        self.tiles = tiles
        self.tile_px = tile_px
        self.lower = lower
        self.upper = upper

    @classmethod
    def from_files(cls, metadata_path: Path, image_path: Path) -> "WangTileSheet":
        meta = json.loads(Path(metadata_path).read_text())
        sheet = Image.open(image_path).convert("RGBA")
        tile_size = meta["tileset_data"]["tile_size"]
        tiles: dict[int, Image.Image] = {}
        for t in meta["tileset_data"]["tiles"]:
            c = t["corners"]
            idx = (
                (c["NW"] == "upper") * 8
                + (c["NE"] == "upper") * 4
                + (c["SW"] == "upper") * 2
                + (c["SE"] == "upper") * 1
            )
            bb = t["bounding_box"]
            tiles[idx] = sheet.crop((bb["x"], bb["y"], bb["x"] + bb["width"], bb["y"] + bb["height"]))
        missing = [i for i in range(16) if i not in tiles]
        if missing:
            raise ValueError(f"tileset 缺少 Wang tile {missing}({metadata_path})")
        prompts = meta.get("metadata", {}).get("terrain_prompts") or {
            "lower": meta.get("lower_description", "lower"),
            "upper": meta.get("upper_description", "upper"),
        }
        return cls(tiles, tile_size["width"], prompts["lower"], prompts["upper"])


def wang_index(nw: int, ne: int, sw: int, se: int, hi: int) -> int:
    """4 个角点等级 → Wang tile index(等级 == hi 的角视为 upper)。"""
    return (nw == hi) * 8 + (ne == hi) * 4 + (sw == hi) * 2 + (se == hi) * 1


def smooth_levels(levels: list[list[int]]) -> list[list[int]]:
    """迭代压平,使相邻(含对角)顶点等级差 ≤1——跨级相邻没有过渡 tile 可用。"""
    h, w = len(levels), len(levels[0])
    changed = True
    while changed:
        changed = False
        for y in range(h):
            for x in range(w):
                for dy in (-1, 0, 1):
                    for dx in (-1, 0, 1):
                        ny, nx = y + dy, x + dx
                        if 0 <= ny < h and 0 <= nx < w:
                            if levels[y][x] > levels[ny][nx] + 1:
                                levels[y][x] = levels[ny][nx] + 1
                                changed = True
    return levels


def render_terrain(levels: list[list[int]], sheets: list[WangTileSheet]) -> Image.Image:
    """顶点等级网格 → 地形大图。sheets[k] 负责等级 k→k+1 的过渡。"""
    vh, vw = len(levels), len(levels[0])
    ch, cw = vh - 1, vw - 1
    px = sheets[0].tile_px
    canvas = Image.new("RGBA", (cw * px, ch * px))
    for cy in range(ch):
        for cx in range(cw):
            corners = (
                levels[cy][cx], levels[cy][cx + 1],
                levels[cy + 1][cx], levels[cy + 1][cx + 1],
            )
            lo, hi = min(corners), max(corners)
            if hi == lo:  # 纯地形:等级 t 的"全 upper"tile(t=0 用第一套的全 lower)
                tile = sheets[0].tiles[0] if lo == 0 else sheets[lo - 1].tiles[15]
            elif hi - lo == 1:
                tile = sheets[lo].tiles[wang_index(*corners, hi)]
            else:
                raise ValueError(f"格子({cx},{cy})角点跨级 {corners},先用 smooth_levels 压平")
            canvas.paste(tile, (cx * px, cy * px))
    return canvas


@dataclass
class Placement:
    """一次物件摆放:图片 + 锚点像素坐标(锚点=物件底边中心,贴近'脚底'遮挡直觉)。"""

    image: Image.Image
    x: int
    y: int

    @property
    def top_left(self) -> tuple[int, int]:
        return self.x - self.image.width // 2, self.y - self.image.height


def paste_objects(canvas: Image.Image, placements: list[Placement]) -> Image.Image:
    """按脚底 y 升序合成物件(下方物件覆盖上方),返回原 canvas。"""
    for p in sorted(placements, key=lambda p: p.y):
        canvas.alpha_composite(p.image, p.top_left)
    return canvas


def add_glow(canvas: Image.Image, x: int, y: int, radius: int,
             color: tuple[int, int, int] = (255, 190, 90), strength: float = 0.55) -> None:
    """在 (x,y) 加一团柔光(screen 叠加),用于窗光/灯光氛围。"""
    glow = Image.new("RGBA", canvas.size, (0, 0, 0, 0))
    draw = ImageDraw.Draw(glow)
    alpha = int(255 * strength)
    draw.ellipse((x - radius, y - radius, x + radius, y + radius), fill=(*color, alpha))
    glow = glow.filter(ImageFilter.GaussianBlur(radius * 0.6))
    base = canvas.convert("RGB")
    lit = Image.composite(
        Image.blend(base, Image.new("RGB", canvas.size, color), 0.35), base, glow.split()[3]
    )
    canvas.paste(lit, (0, 0))


def suppress_warm(img: Image.Image, hue_range: tuple[int, int] = (10, 68),
                  sat_floor: int = 100, val_floor: int = 120) -> Image.Image:
    """压制地形里高饱和的暖亮色(橙/黄绿过渡边在夜景中过亮),压暗保留纹理。"""
    import numpy as np

    hsv = np.array(img.convert("HSV"), dtype=np.int16)
    h, s, v = hsv[..., 0], hsv[..., 1], hsv[..., 2]
    # PIL HSV 的 hue 0-255 对应 0-360°
    lo, hi = int(hue_range[0] * 255 / 360), int(hue_range[1] * 255 / 360)
    mask = (h >= lo) & (h <= hi) & (s >= sat_floor) & (v >= val_floor)
    s[mask] = (s[mask] * 0.45).astype(np.int16)
    v[mask] = (v[mask] * 0.62).astype(np.int16)
    # 低饱和的亮米黄高光边(石滩过渡 tile)同样压暗,避免夜景里像描了荧光边
    pale = (h >= 20) & (h <= 50) & (s >= 30) & (s < sat_floor) & (v >= 150)
    v[pale] = (v[pale] * 0.72).astype(np.int16)
    out = Image.fromarray(np.stack([h, s, v], axis=-1).astype("uint8"), "HSV").convert("RGBA")
    if img.mode == "RGBA":
        out.putalpha(img.split()[3])
    return out


def grade_night(img: Image.Image, darken: float = 0.12,
                tint: tuple[int, int, int] = (24, 34, 64), tint_strength: float = 0.10) -> Image.Image:
    """轻度夜景调色:整体压暗 + 冷色罩(素材本身已按夜景生成,这里只统一氛围)。"""
    base = img.convert("RGB")
    dark = Image.blend(base, Image.new("RGB", img.size, (0, 0, 0)), darken)
    cool = Image.blend(dark, Image.new("RGB", img.size, tint), tint_strength)
    return cool.convert("RGBA")
