#!/usr/bin/env python3
"""
ocv_encode.py — конвертер видео в формат OCV5 для OpenComputers (Tier 3 GPU + Screen).

Вся тяжёлая работа выполняется на ПК: плеер получает готовый поток
GPU-команд и просто исполняет его, без попиксельной обработки в Lua.

Пайплайн кадра:
  1. ресэмплинг FPS (OC рисует по тикам: потолок 20 fps);
  2. letterbox + downscale в линейном свете (INTER_AREA, без ringing);
  3. unsharp mask + bilateral (выравнивает плоские области, края не трогает);
  4. временная стабилизация (гасит шум исходника -> меньше мерцания и команд);
  5. двухцветный ordered dithering (Bayer 8x8) прямо в РОДНУЮ палитру Tier 3:
     240 цветов куба 6x8x5 + 16 серых палитры (18 ровных уровней серого);
  6. diff с тем, что реально на экране -> минимальный набор команд:
     одна gpu.set() рисует любые ячейки, чьи 2 цвета ⊆ {bg, fg}
     (' ', '▄', '▀', '█'); команды сгруппированы по парам цветов;
     при бюджете --max-sets первыми рисуются самые заметные изменения.

Формат OCV5 (big-endian):
  index.ocv, 64 байта:
    [4]  "OCV5"
    [2]  width, символы
    [2]  height, символы
    [2]  fps * 100
    [4]  frames
    [2]  parts
    [48] 16 x RGB — палитра GPU 0..15
  partNNNN.ocv — кадры подряд, кадр не пересекает границу файла:
    [4]  длина payload
    команды payload:
      0x01 c              setBackground(PAL[c])
      0x02 c              setForeground(PAL[c])
      0x03 bg fg          оба
      0x04 x y n codes    gpu.set(x, y, <n символов>), коды по 2 бита, старшие биты первыми:
                          0=' '  1='▄'  2='▀'  3='█'

Зависимости: pip install opencv-python numpy
"""

from __future__ import annotations

import argparse
import os
import struct
import sys
import time
from collections import defaultdict
from dataclasses import dataclass

import cv2
import numpy as np

# ---------------------------------------------------------------- палитра OC

CUBE_R, CUBE_G, CUBE_B = 6, 8, 5
PALETTE_SIZE = 16
GRAY_LEVELS = PALETTE_SIZE + 2  # чёрный и белый берутся из куба
GRAY_STEP = 255 // (GRAY_LEVELS - 1)  # = 15


def _cube_rgb(i: int) -> tuple[int, int, int]:
    b, g, r = i % CUBE_B, (i // CUBE_B) % CUBE_G, i // (CUBE_B * CUBE_G)
    return (
        int(r * 255 / (CUBE_R - 1) + 0.5),
        int(g * 255 / (CUBE_G - 1) + 0.5),
        int(b * 255 / (CUBE_B - 1) + 0.5),
    )


GRAY_PALETTE = [(GRAY_STEP * (i + 1),) * 3 for i in range(PALETTE_SIZE)]
FULL_PALETTE = np.array(GRAY_PALETTE + [_cube_rgb(i) for i in range(240)], dtype=np.uint8)

IDX_BLACK = PALETTE_SIZE
IDX_WHITE = PALETTE_SIZE + 239

# индекс в 256-цветной палитре для каждого уровня серого 0..17
GRAY_LUT = np.array(
    [IDX_BLACK] + list(range(PALETTE_SIZE)) + [IDX_WHITE], dtype=np.uint8
)

# ---------------------------------------------------------------- дизеринг

def _bayer(n: int) -> np.ndarray:
    m = np.array([[0]], dtype=np.float32)
    while m.shape[0] < n:
        m = np.block([[4 * m, 4 * m + 2], [4 * m + 3, 4 * m + 1]])
    return (m + 0.5) / m.size  # (0, 1)


BAYER8 = _bayer(8)

LUT_BITS = 6            # входной цвет квантуется до 64 уровней на канал для LUT
MIX_CANDIDATES = 8      # сколько ближайших цветов палитры перебирать парами
LUMA_W = np.array([0.299, 0.587, 0.114], dtype=np.float32) * 3

# ---------------------------------------------------------------- команды

# Коды символов (2 бита на ячейку). bg/fg — текущие цвета GPU.
CODE_SPACE = 0   # ' '  обе половины = bg
CODE_LOWER = 1   # '▄'  верх = bg, низ = fg
CODE_UPPER = 2   # '▀'  верх = fg, низ = bg
CODE_FULL = 3    # '█'  обе половины = fg

OP_BG, OP_FG, OP_BOTH, OP_TEXT = 1, 2, 3, 4

MAX_RUN = 160      # ячеек в одной команде (ширина экрана)
MAX_BRIDGE = 12    # сколько совместимых ячеек можно перекрыть, чтобы склеить два run

# ================================================================= обработка кадра

def _srgb_to_lin(x: np.ndarray) -> np.ndarray:
    x = x / 255.0
    return np.where(x <= 0.04045, x / 12.92, ((x + 0.055) / 1.055) ** 2.4)


def _lin_to_srgb(x: np.ndarray) -> np.ndarray:
    x = np.clip(x, 0.0, 1.0)
    return np.where(x <= 0.0031308, x * 12.92, 1.055 * np.power(x, 1 / 2.4) - 0.055) * 255.0


@dataclass
class Geometry:
    width: int       # символы
    height: int      # символы

    @property
    def px_h(self) -> int:
        return self.height * 2


class FramePreprocessor:
    """letterbox + resize в линейном свете + sharpen + edge-preserving сглаживание.

    Сглаживание (bilateral) выравнивает плоские области, не трогая контуры:
    меньше мелких участков -> заметно меньше gpu.set на кадр при той же чёткости.
    """

    def __init__(self, geo: Geometry, sharpen: float, smooth: float):
        self.geo = geo
        self.sharpen = sharpen
        self.smooth = smooth
        self._to_lin = _srgb_to_lin(np.arange(256, dtype=np.float32)).astype(np.float32)

    def __call__(self, bgr: np.ndarray) -> np.ndarray:
        w, h = self.geo.width, self.geo.px_h
        sh, sw = bgr.shape[:2]
        scale = min(w / sw, h / sh)
        nw, nh = max(1, round(sw * scale)), max(1, round(sh * scale))

        lin = self._to_lin[cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)]
        small = _lin_to_srgb(cv2.resize(lin, (nw, nh), interpolation=cv2.INTER_AREA))

        if self.sharpen > 0:
            blur = cv2.GaussianBlur(small, (0, 0), 0.8)
            small = small + self.sharpen * (small - blur)
        if self.smooth > 0:
            small = cv2.bilateralFilter(np.clip(small, 0, 255).astype(np.float32), 5, self.smooth, 3)

        canvas = np.zeros((h, w, 3), dtype=np.float32)
        x0, y0 = (w - nw) // 2, (h - nh) // 2
        canvas[y0:y0 + nh, x0:x0 + nw] = np.clip(small, 0, 255)
        return canvas


class Stabilizer:
    """Держит пиксель неизменным, пока исходник не ушёл дальше порога.
    Убирает мерцание от шума/артефактов сжатия на статичных участках."""

    def __init__(self, threshold: float):
        self.threshold = threshold
        self.anchor: np.ndarray | None = None

    def __call__(self, img: np.ndarray) -> np.ndarray:
        if self.threshold <= 0:
            return img
        if self.anchor is None:
            self.anchor = img.copy()
            return img
        moved = np.abs(img - self.anchor).max(axis=2) > self.threshold
        self.anchor[moved] = img[moved]
        return self.anchor.copy()


class Quantizer:
    """Двухцветный ordered dithering в палитру Tier 3.

    Для каждого цвета заранее (LUT 64^3) находится пара цветов палитры и доля
    смешивания, минимизирующие ошибку + штраф за шум. Смешивание считается в
    линейном свете (так глаз усредняет соседние пиксели). Плоская область
    получает ровно 2 цвета -> ячейки идеально ложатся в одну пару bg/fg.
    """

    def __init__(self, geo: Geometry, gray_only: bool, noise_penalty: float):
        h, w = geo.px_h, geo.width
        self.noise_penalty = noise_penalty
        self.t = np.tile(BAYER8, (h // 8 + 1, w // 8 + 1))[:h, :w].astype(np.float32)
        allowed = GRAY_LUT if gray_only else np.arange(256, dtype=np.uint8)
        self.c1, self.c2, self.ratio = self._build_lut(np.unique(allowed), noise_penalty)

    @staticmethod
    def _build_lut(allowed: np.ndarray, noise_penalty: float):
        n = 1 << LUT_BITS
        pal = FULL_PALETTE[allowed].astype(np.float32)
        pal_lin = _srgb_to_lin(pal)
        k = min(MIX_CANDIDATES, len(pal))

        axis = np.arange(n, dtype=np.float32) * 255.0 / (n - 1)
        grid = np.stack(np.meshgrid(axis, axis, axis, indexing="ij"), -1).reshape(-1, 3)

        c1 = np.empty(len(grid), np.uint8)
        c2 = np.empty(len(grid), np.uint8)
        ratio = np.empty(len(grid), np.float32)
        ii, jj = np.triu_indices(k)  # включая i == j (чистый цвет)

        for s in range(0, len(grid), 8192):
            p = grid[s:s + 8192]
            d = (((p[:, None, :] - pal[None]) ** 2) * LUMA_W).sum(-1)
            near = np.argpartition(d, k - 1, axis=1)[:, :k]            # (B, k)
            a, b = near[:, ii], near[:, jj]                             # (B, P)
            la, lb = pal_lin[a], pal_lin[b]                             # (B, P, 3)
            pl = _srgb_to_lin(p)[:, None, :]
            diff = lb - la
            den = (diff ** 2).sum(-1)
            r = np.where(den > 0, ((pl - la) * diff).sum(-1) / np.maximum(den, 1e-12), 0.0)
            r = np.clip(np.round(r * 64) / 64, 0, 1)
            mix = _lin_to_srgb(la + r[..., None] * diff)
            err = (((mix - p[:, None, :]) ** 2) * LUMA_W).sum(-1)
            spread = (((pal[b] - pal[a]) ** 2) * LUMA_W).sum(-1)
            err += noise_penalty * r * (1 - r) * spread
            best = err.argmin(1)
            rows = np.arange(len(p))
            c1[s:s + 8192] = allowed[a[rows, best]]
            c2[s:s + 8192] = allowed[b[rows, best]]
            ratio[s:s + 8192] = r[rows, best]
        return c1, c2, ratio

    def __call__(self, img: np.ndarray) -> np.ndarray:
        n = 1 << LUT_BITS
        q = np.clip(np.rint(img * ((n - 1) / 255.0)), 0, n - 1).astype(np.int32)
        key = (q[..., 0] * n + q[..., 1]) * n + q[..., 2]
        return np.where(self.t < self.ratio[key], self.c2[key], self.c1[key])


# ================================================================= генерация команд

@dataclass
class Run:
    key: tuple[int, ...]   # 1 или 2 цвета, которыми рисуется run
    y: int
    x0: int
    x1: int
    score: float = 0.0     # визуальная ошибка, которую run исправляет


class CommandEncoder:
    """Diff отображаемого экрана с целевым кадром -> поток GPU-команд.

    Энкодер хранит то, что РЕАЛЬНО нарисовано на экране. При бюджете max_sets
    рисуются run'ы с наибольшей визуальной ошибкой, остальное доезжает позже.
    """

    def __init__(self, geo: Geometry, max_sets: int):
        self.top = np.full((geo.height, geo.width), IDX_BLACK, dtype=np.uint8)
        self.bot = self.top.copy()
        self.max_sets = max_sets
        pal = FULL_PALETTE.astype(np.float32)
        self.dist = np.sqrt((((pal[:, None] - pal[None]) ** 2) * LUMA_W).sum(-1)).astype(np.float32)
        self.age = np.zeros(self.top.shape, dtype=np.float32)  # сколько кадров ячейка ждёт отрисовки
        self.deferred = 0

    # ---------------------------------------------------------------- run'ы

    def _build_runs(self, top: np.ndarray, bot: np.ndarray, changed: np.ndarray) -> list[Run]:
        tl, bl = top.tolist(), bot.tolist()
        ys, xs = np.nonzero(changed)

        pairs: dict[tuple[int, int], dict[int, list[int]]] = defaultdict(lambda: defaultdict(list))
        singles: dict[int, dict[int, list[int]]] = defaultdict(lambda: defaultdict(list))
        for y, x in zip(ys.tolist(), xs.tolist()):
            a, b = tl[y][x], bl[y][x]
            if a == b:
                singles[a][y].append(x)
            else:
                pairs[(a, b) if a < b else (b, a)][y].append(x)

        runs: list[Run] = []

        def split(key: tuple[int, ...], y: int, xs_: list[int]) -> None:
            row_t, row_b = tl[y], bl[y]
            ok = set(key)
            xs_ = sorted(xs_)
            x0 = x1 = xs_[0]
            for x in xs_[1:]:
                if (x - x1 - 1 <= MAX_BRIDGE and x - x0 < MAX_RUN
                        and all(row_t[g] in ok and row_b[g] in ok for g in range(x1 + 1, x))):
                    x1 = x
                else:
                    runs.append(Run(key, y, x0, x1))
                    x0 = x1 = x
            runs.append(Run(key, y, x0, x1))

        # двухцветные пары, в них же вливаем одноцветные ячейки тех же цветов
        for key in sorted(pairs, key=lambda k: -sum(map(len, pairs[k].values()))):
            rows = pairs[key]
            for c in key:
                for y, xs_ in singles.pop(c, {}).items():
                    rows[y].extend(xs_)
            for y, xs_ in rows.items():
                split(key, y, xs_)
        for c, rows in singles.items():
            for y, xs_ in rows.items():
                split((c,), y, xs_)

        # оценка: суммарная ошибка исправляемых ячеек
        # отложенные ячейки дорожают с каждым кадром -> ничего не «застревает»
        err = (self.dist[self.top, top] + self.dist[self.bot, bot]) * (1.0 + self.age)
        for r in runs:
            r.score = float(err[r.y, r.x0:r.x1 + 1].sum())
        return runs

    # ---------------------------------------------------------------- кодирование

    def encode(self, pixels: np.ndarray) -> bytes:
        top, bot = pixels[0::2], pixels[1::2]
        changed = (top != self.top) | (bot != self.bot)
        if not changed.any():
            return b""

        runs = self._build_runs(top, bot, changed)
        if self.max_sets and len(runs) > self.max_sets:
            runs.sort(key=lambda r: -r.score)
            self.deferred = len(runs) - self.max_sets
            runs = runs[:self.max_sets]
        else:
            self.deferred = 0

        by_key: dict[tuple[int, ...], list[Run]] = defaultdict(list)
        for r in runs:
            by_key[r.key].append(r)

        out = bytearray()
        bg = fg = None
        tl, bl = top.tolist(), bot.tolist()

        while by_key:
            # следующая группа: та, что переиспользует текущие регистры
            key = max(by_key, key=lambda k: ((bg in k) + (fg in k), len(by_key[k])))
            group = by_key.pop(key)

            if len(key) == 1:
                c = key[0]
                if c != bg and c != fg:
                    out += bytes((OP_BG, c)); bg = c
            elif bg not in key and fg not in key:
                bg, fg = key
                out += bytes((OP_BOTH, bg, fg))
            elif bg in key and fg in key:
                pass
            elif bg in key:
                fg = key[1] if key[0] == bg else key[0]
                out += bytes((OP_FG, fg))
            else:
                bg = key[1] if key[0] == fg else key[0]
                out += bytes((OP_BG, bg))

            for r in sorted(group, key=lambda r: (r.y, r.x0)):
                row_t, row_b = tl[r.y], bl[r.y]
                n = r.x1 - r.x0 + 1
                packed = bytearray((n + 3) // 4)
                for i in range(n):
                    a, b = row_t[r.x0 + i], row_b[r.x0 + i]
                    code = (CODE_SPACE if b == bg else CODE_LOWER) if a == bg else (CODE_UPPER if b == bg else CODE_FULL)
                    packed[i >> 2] |= code << (6 - 2 * (i & 3))
                out += bytes((OP_TEXT, r.x0 + 1, r.y + 1, n))
                out += packed
                self.top[r.y, r.x0:r.x1 + 1] = top[r.y, r.x0:r.x1 + 1]
                self.bot[r.y, r.x0:r.x1 + 1] = bot[r.y, r.x0:r.x1 + 1]

        still = (top != self.top) | (bot != self.bot)
        self.age = np.where(still, self.age + 1.0, 0.0)
        return bytes(out)


# ================================================================= запись

class ChunkWriter:
    def __init__(self, out_dir: str, chunk_bytes: int):
        self.dir = out_dir
        self.limit = chunk_bytes
        self.index = 0
        self.size = 0
        self.file = None
        os.makedirs(out_dir, exist_ok=True)
        for name in os.listdir(out_dir):
            if name.endswith(".ocv"):
                os.remove(os.path.join(out_dir, name))

    def write(self, payload: bytes) -> None:
        record = struct.pack(">I", len(payload)) + payload
        if self.file is None or (self.size > 0 and self.size + len(record) > self.limit):
            self._next()
        self.file.write(record)
        self.size += len(record)

    def _next(self) -> None:
        if self.file:
            self.file.close()
        self.index += 1
        self.size = 0
        self.file = open(os.path.join(self.dir, f"part{self.index:04d}.ocv"), "wb")

    def close(self, geo: Geometry, fps: float, frames: int) -> None:
        if self.file:
            self.file.close()
        with open(os.path.join(self.dir, "index.ocv"), "wb") as f:
            f.write(b"OCV5")
            f.write(struct.pack(">HHHIH", geo.width, geo.height, round(fps * 100), frames, self.index))
            for rgb in GRAY_PALETTE:
                f.write(bytes(rgb))


# ================================================================= main

def resampled_frames(cap: cv2.VideoCapture, src_fps: float, dst_fps: float):
    """Выдаёт кадры с частотой dst_fps (ближайший исходный кадр по времени)."""
    k = 0
    i = 0
    while True:
        ok, frame = cap.read()
        if not ok:
            return
        while k * src_fps < (i + 1) * dst_fps:
            yield frame
            k += 1
        i += 1


def save_preview(enc: CommandEncoder, path: str, scale: int = 5) -> None:
    h, w = enc.top.shape
    img = np.empty((h * 2, w, 3), dtype=np.uint8)
    img[0::2], img[1::2] = FULL_PALETTE[enc.top], FULL_PALETTE[enc.bot]
    img = cv2.resize(img, (w * scale, h * 2 * scale), interpolation=cv2.INTER_NEAREST)
    cv2.imwrite(path, cv2.cvtColor(img, cv2.COLOR_RGB2BGR))


def main() -> None:
    p = argparse.ArgumentParser(description="Video -> OpenComputers OCV5 (Tier 3)")
    p.add_argument("input")
    p.add_argument("output", nargs="?", default="video")
    p.add_argument("--width", type=int, default=160, help="ширина в символах (max 160)")
    p.add_argument("--height", type=int, default=50, help="высота в символах (max 50)")
    p.add_argument("--fps", type=float, default=20.0, help="целевой FPS (OC рендерит максимум 20)")
    p.add_argument("--sharpen", type=float, default=0.5, help="сила резкости, 0 = выкл")
    p.add_argument("--smooth", type=float, default=25.0,
                   help="сглаживание с сохранением краёв, 0 = выкл. Больше -> меньше gpu.set, но меньше текстур")
    p.add_argument("--noise", type=float, default=0.6,
                   help="штраф за шумный дизеринг: больше -> чище и дешевле, меньше -> точнее цвет")
    p.add_argument("--stability", type=float, default=10.0, help="порог стабилизации шума (0..255), 0 = выкл")
    p.add_argument("--max-sets", type=int, default=0,
                   help="бюджет gpu.set на кадр, 0 = без лимита. Если плеер пишет late — ставь 600..1200")
    p.add_argument("--gray", action="store_true", help="чёрно-белое видео (18 уровней серого)")
    p.add_argument("--chunk-mb", type=float, default=1.0, help="макс. размер part-файла")
    p.add_argument("--max-frames", type=int, default=0, help="ограничить число кадров (для теста)")
    p.add_argument("--preview", type=int, nargs="*", default=[], metavar="N",
                   help="сохранить preview_N.png — ровно то, что покажет экран OC на кадре N")
    a = p.parse_args()

    if not (1 <= a.width <= 160 and 1 <= a.height <= 50):
        sys.exit("error: Tier 3 максимум 160x50")

    cap = cv2.VideoCapture(a.input)
    if not cap.isOpened():
        sys.exit(f"error: cannot open {a.input}")
    src_fps = cap.get(cv2.CAP_PROP_FPS) or 20.0
    fps = min(a.fps, src_fps, 20.0)

    geo = Geometry(a.width, a.height)
    pre = FramePreprocessor(geo, a.sharpen, a.smooth)
    stab = Stabilizer(a.stability)
    quant = Quantizer(geo, a.gray, a.noise)
    enc = CommandEncoder(geo, a.max_sets)
    writer = ChunkWriter(a.output, int(a.chunk_mb * 1024 * 1024))

    t0 = time.time()
    frames = 0
    total = 0
    try:
        for frame in resampled_frames(cap, src_fps, fps):
            payload = enc.encode(quant(stab(pre(frame))))
            writer.write(payload)
            if frames in a.preview:
                save_preview(enc, os.path.join(a.output, f"preview_{frames}.png"))
            frames += 1
            total += len(payload) + 4
            if frames % 100 == 0:
                el = time.time() - t0
                print(f"\r  {frames} frames  {total / frames / 1024:.1f} KB/frame  {frames / el:.1f} fps", end="", flush=True)
            if a.max_frames and frames >= a.max_frames:
                break
    finally:
        cap.release()
        writer.close(geo, fps, frames)

    mb = total / 1048576
    print(f"\n---- DONE ----\n"
          f"Screen   : {geo.width}x{geo.height} chars ({geo.width}x{geo.px_h} px){' gray' if a.gray else ''}\n"
          f"Frames   : {frames} @ {fps:g} fps  ({frames / fps:.1f} s)\n"
          f"Size     : {mb:.2f} MB in {writer.index} parts  ({mb * 1024 / max(frames / fps, 1e-9):.0f} KB/s)\n"
          f"Elapsed  : {time.time() - t0:.1f} s")


if __name__ == "__main__":
    main()
