"""
Hand gesture recognition using MediaPipe Hands.
Detects: open palm (STOP), pointing direction, thumbs up (RESUME),
and waving (RETURN to owner).
"""

import mediapipe as mp
import numpy as np
import math
from enum import Enum, auto


class Gesture(Enum):
    NONE = auto()
    STOP = auto()
    POINT_LEFT = auto()
    POINT_RIGHT = auto()
    THUMBS_UP = auto()
    WAVE = auto()


class GestureRecognizer:
    def __init__(self, min_detection_confidence=0.6, min_tracking_confidence=0.5):
        self.mp_hands = mp.solutions.hands
        self.hands = self.mp_hands.Hands(
            static_image_mode=False,
            max_num_hands=2,
            min_detection_confidence=min_detection_confidence,
            min_tracking_confidence=min_tracking_confidence,
        )
        self._wave_history = []
        self._wave_window = 10

    def recognize(self, frame_rgb) -> tuple:
        results = self.hands.process(frame_rgb)

        if not results.multi_hand_landmarks:
            self._wave_history.clear()
            return Gesture.NONE, None

        hand_landmarks = results.multi_hand_landmarks[0]
        h, w, _ = frame_rgb.shape

        lm = hand_landmarks.landmark
        wrist = lm[self.mp_hands.HandLandmark.WRIST]
        thumb_tip = lm[self.mp_hands.HandLandmark.THUMB_TIP]
        thumb_ip = lm[self.mp_hands.HandLandmark.THUMB_IP]
        index_tip = lm[self.mp_hands.HandLandmark.INDEX_FINGER_TIP]
        index_mcp = lm[self.mp_hands.HandLandmark.INDEX_FINGER_MCP]
        middle_tip = lm[self.mp_hands.HandLandmark.MIDDLE_FINGER_TIP]
        middle_mcp = lm[self.mp_hands.HandLandmark.MIDDLE_FINGER_MCP]
        ring_tip = lm[self.mp_hands.HandLandmark.RING_FINGER_TIP]
        ring_mcp = lm[self.mp_hands.HandLandmark.RING_FINGER_MCP]
        pinky_tip = lm[self.mp_hands.HandLandmark.PINKY_TIP]
        pinky_mcp = lm[self.mp_hands.HandLandmark.PINKY_MCP]

        hand_info = {
            "cx": int(wrist.x * w),
            "cy": int(wrist.y * h),
        }

        def dist(a, b):
            return math.sqrt((a.x - b.x) ** 2 + (a.y - b.y) ** 2)

        index_extended = dist(index_tip, wrist) > dist(index_mcp, wrist) * 1.2
        middle_extended = dist(middle_tip, wrist) > dist(middle_mcp, wrist) * 1.2
        ring_extended = dist(ring_tip, wrist) > dist(ring_mcp, wrist) * 1.2
        pinky_extended = dist(pinky_tip, wrist) > dist(pinky_mcp, wrist) * 1.2
        thumb_extended = dist(thumb_tip, wrist) > dist(thumb_ip, wrist) * 1.3

        extended_count = sum(
            [index_extended, middle_extended, ring_extended, pinky_extended, thumb_extended]
        )

        self._wave_history.append(wrist.x)
        if len(self._wave_history) > self._wave_window:
            self._wave_history.pop(0)

        is_waving = False
        if len(self._wave_history) >= self._wave_window:
            diffs = np.diff(self._wave_history)
            sign_changes = np.sum(np.abs(np.diff(np.sign(diffs))) > 0)
            is_waving = sign_changes >= 4

        if is_waving and extended_count >= 3:
            return Gesture.WAVE, hand_info

        if extended_count >= 4:
            return Gesture.STOP, hand_info

        if thumb_extended and not index_extended and not middle_extended and not ring_extended:
            if thumb_tip.y < thumb_ip.y:
                return Gesture.THUMBS_UP, hand_info

        if index_extended and not middle_extended and not ring_extended and not pinky_extended:
            dx = index_tip.x - index_mcp.x
            if dx > 0.05:
                return Gesture.POINT_RIGHT, hand_info
            elif dx < -0.05:
                return Gesture.POINT_LEFT, hand_info

        return Gesture.NONE, hand_info

    def close(self):
        self.hands.close()
