import serial
import mss
import numpy as np
import time

PORT = '/dev/ttyUSB0'

# --- Бюджет пропускной способности ---
# Кадр = 3 байта заголовка ('A','d','a') + LEDS*4 байта (G,R,B,W) + 1 байт checksum
# 60 светодиодов: 3 + 240 + 1 = 244 байта/кадр
# На 460800 бод (8N1, ~10 бит/байт) канал даёт ~46080 байт/с -> до ~188 кадров/с с запасом.
# ВАЖНО: значение ниже должно СОВПАДАТЬ с Serial.begin() на ESP32!
BAUDRATE = 460800

# --- Выбор монитора ---
# 0 = "виртуальный" экран (объединение всех мониторов в mss), 1 = первый физический,
# 2 = второй физический и т.д. Если индекс недоступен - скрипт сам откатится на monitors[0].
MONITOR_INDEX = 1

# --- Конфигурация ленты по сторонам экрана ---
# Сейчас физически подключены только LEFT и TOP (50 из 60 светодиодов).
# RIGHT и BOTTOM - ЗАГЛУШКИ для расширения до полного периметра в будущем.
TOTAL_LEDS = 60
LEDS_LEFT = 16
LEDS_TOP = 34
LEDS_RIGHT = 0    # ЗАГЛУШКА: поставьте реальное число светодиодов правой стороны
LEDS_BOTTOM = 0   # ЗАГЛУШКА: поставьте реальное число светодиодов нижней стороны

# Глубина зоны сэмплирования задана в % от меньшей стороны экрана, а не в пикселях.
# Это и есть "универсальность для разных экранов": на 1280x720 и на 3840x2160
# зона будет занимать одну и ту же ОТНОСИТЕЛЬНУЮ долю картинки.
ZONE_DEPTH_PERCENT = 0.05

# --- Опциональная линеаризация (гамма-коррекция) ---
# Физика: PWM-скважность светодиода линейна по электрической мощности, а восприятие
# яркости человеком - нелинейно (степенной закон, гамма ~1.8-2.4). Если в тёмных
# участках картинки лента выглядит "пустой" - поставьте GAMMA, например, 2.0.
GAMMA = 1.0
_gamma_lut = None
if GAMMA != 1.0:
    _gamma_lut = np.array(
        [round(((i / 255) ** (1 / GAMMA)) * 255) for i in range(256)],
        dtype=np.uint8,
    )


def apply_gamma(rgb_uint8):
    if _gamma_lut is None:
        return rgb_uint8
    return _gamma_lut[rgb_uint8]


def sample_edge(screenshot, screen_w, screen_h, count, edge, depth_px, reverse=False):
    """
    Делит указанную сторону экрана на `count` равных секторов и возвращает
    список средних цветов [B,G,R,A] (формат mss) для каждого сектора.

    edge: 'left' | 'right' | 'top' | 'bottom'
    reverse=True - обходить сектора в обратном порядке (например, "снизу вверх")

    Это ОДНА функция для всех четырёх сторон - именно поэтому добавить правую
    или нижнюю сторону = раскомментировать вызов ниже, а не писать новый код.
    """
    if count <= 0:
        return []

    colors = []
    indices = reversed(range(count)) if reverse else range(count)

    if edge in ('left', 'right'):
        sector_h = screen_h / count
        x_start, x_end = (0, depth_px) if edge == 'left' else (screen_w - depth_px, screen_w)
        for i in indices:
            y_start = max(0, int(i * sector_h))
            y_end = min(screen_h, int((i + 1) * sector_h))
            zone = screenshot[y_start:y_end:4, x_start:x_end:4]
            colors.append(zone.mean(axis=(0, 1)).astype(np.uint8) if zone.size > 0 else np.zeros(4, dtype=np.uint8))
    else:  # 'top' or 'bottom'
        sector_w = screen_w / count
        y_start, y_end = (0, depth_px) if edge == 'top' else (screen_h - depth_px, screen_h)
        for i in indices:
            x_start = max(0, int(i * sector_w))
            x_end = min(screen_w, int((i + 1) * sector_w))
            zone = screenshot[y_start:y_end:4, x_start:x_end:4]
            colors.append(zone.mean(axis=(0, 1)).astype(np.uint8) if zone.size > 0 else np.zeros(4, dtype=np.uint8))

    return colors


def append_edge_to_packet(led_packet, screenshot, screen_w, screen_h, count, edge, depth_px, reverse=False):
    for color in sample_edge(screenshot, screen_w, screen_h, count, edge, depth_px, reverse):
        color = apply_gamma(color)
        led_packet.extend([int(color[1]), int(color[2]), int(color[0]), 0])  # G,R,B,W


try:
    ser = serial.Serial(PORT, BAUDRATE, timeout=1)
    print(f"[УСПЕХ] Подключено к ESP32 на порту {PORT} @ {BAUDRATE} бод")
except Exception as e:
    print(f"[ПРЕДУПРЕЖДЕНИЕ] Не удалось открыть порт {PORT}: {e}")
    ser = None

with mss.MSS() as sct:
    target_monitor = sct.monitors[MONITOR_INDEX] if len(sct.monitors) > MONITOR_INDEX else sct.monitors[0]
    screen_w = target_monitor["width"]
    screen_h = target_monitor["height"]

zone_depth_px = max(1, int(min(screen_w, screen_h) * ZONE_DEPTH_PERCENT))

print(f"[ИНФО] Разрешение экрана: {screen_w}x{screen_h}, глубина зоны: {zone_depth_px}px")
print(f"[ИНФО] Начинаем захват экрана. Для выхода нажмите Ctrl+C\n")

while True:
    start_time = time.time()

    with mss.MSS() as sct:
        target_monitor = sct.monitors[MONITOR_INDEX] if len(sct.monitors) > MONITOR_INDEX else sct.monitors[0]
        screenshot = np.array(sct.grab(target_monitor))

    led_packet = []

    # 1. ЛЕВАЯ СТОРОНА (снизу вверх) - физически подключена
    append_edge_to_packet(led_packet, screenshot, screen_w, screen_h, LEDS_LEFT, 'left', zone_depth_px, reverse=True)

    # 2. ВЕРХНЯЯ СТОРОНА (слева направо) - физически подключена
    append_edge_to_packet(led_packet, screenshot, screen_w, screen_h, LEDS_TOP, 'top', zone_depth_px, reverse=False)

    # 3. ПРАВАЯ СТОРОНА - ЗАГЛУШКА.
    #    Раскомментируйте, когда физически добавите светодиоды на правую сторону рамки.
    #    Направление "сверху вниз" - логичное продолжение от верхней стороны по кругу.
    # append_edge_to_packet(led_packet, screenshot, screen_w, screen_h, LEDS_RIGHT, 'right', zone_depth_px, reverse=False)

    # 4. НИЖНЯЯ СТОРОНА - ЗАГЛУШКА.
    #    Раскомментируйте, когда физически добавите светодиоды на нижнюю сторону рамки.
    #    Направление "справа налево" - замыкает кольцо обратно к левой стороне.
    # append_edge_to_packet(led_packet, screenshot, screen_w, screen_h, LEDS_BOTTOM, 'bottom', zone_depth_px, reverse=True)

    # 5. ДОПОЛНЕНИЕ ДО TOTAL_LEDS
    current_leds_in_packet = len(led_packet) // 4
    unused_leds = TOTAL_LEDS - current_leds_in_packet
    if unused_leds > 0:
        led_packet.extend([0, 0, 0, 0] * unused_leds)
    led_packet = led_packet[:TOTAL_LEDS * 4]

    # Вывод отладки
    first_led_g, first_led_r, first_led_b = led_packet[0], led_packet[1], led_packet[2]
    last_active_index = (LEDS_LEFT + LEDS_TOP - 1) * 4
    last_led_g = led_packet[last_active_index]
    last_led_r = led_packet[last_active_index + 1]
    last_led_b = led_packet[last_active_index + 2]

    print(f"\r[ТЕСТ] Низ-Лево: R:{first_led_r:<3} G:{first_led_g:<3} B:{first_led_b:<3} | "
          f"Верх-Право: R:{last_led_r:<3} G:{last_led_g:<3} B:{last_led_b:<3}", end="", flush=True)

    if ser is not None:
        try:
            packet_bytes = bytes(led_packet)
            checksum = 0
            for b in packet_bytes:
                checksum ^= b

            ser.write(b'Ada')
            ser.write(packet_bytes)
            ser.write(bytes([checksum]))
        except serial.SerialException:
            print("\n[ОШИБКА] Соединение разорвано.")
            ser = None

    delay = (1 / 60) - (time.time() - start_time)
    if delay > 0:
        time.sleep(delay)
