import cv2
import mediapipe as mp
import numpy as np

# ----- Helper Functions -----
def compute_distance(p1, p2):
    return np.hypot(p2[0] - p1[0], p2[1] - p1[1])

# ----- Configuration -----
CAMERA_INDICES = [0]          # List of camera device indices
FRAME_WIDTH = 1280
FRAME_HEIGHT = 720
PIXELS_PER_INCH = 20          # will be set by calibration
PINCH_THRESHOLD = 40          # px to start touch
RELEASE_THRESHOLD = 60        # px to end touch
CIRCLE_TOUCH_THRESHOLD = 20   # px tolerance for circle touch

# ----- Initialize Hand Detector -----
mp_hands = mp.solutions.hands
mp_draw  = mp.solutions.drawing_utils
hands = mp_hands.Hands(
    static_image_mode=False,
    max_num_hands=1,
    min_detection_confidence=0.7,
    min_tracking_confidence=0.5
)

# ----- Global State -----
measuring = False             # measurement in progress
start_pt = None               # measurement start point
calibrating = False           # pinch-based calibration flag
cal_start = None              # pinch calibration start point
object_calibrating = False    # object calibration flag
cal_circle = None             # reference circle (x,y,r)
fingertip_idx_global = None   # last detected fingertip position

# ----- Per-camera Processing -----
def process_frame(frame):
    global measuring, start_pt, calibrating, cal_start, object_calibrating, cal_circle, PIXELS_PER_INCH, fingertip_idx_global
    frame_out = cv2.flip(frame, 1)
    h, w, _ = frame_out.shape

    # Hand detection
    rgb = cv2.cvtColor(frame_out, cv2.COLOR_BGR2RGB)
    rgb.flags.writeable = False
    results = hands.process(rgb)
    rgb.flags.writeable = True
    frame_vis = cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR)

    fingertip_idx = None
    fingertip_mid = None

    if results.multi_hand_landmarks:
        hand = results.multi_hand_landmarks[0]
        mp_draw.draw_landmarks(frame_vis, hand, mp_hands.HAND_CONNECTIONS)
        # get index and middle finger tips and PIP joints
        idx_tip = hand.landmark[mp_hands.HandLandmark.INDEX_FINGER_TIP]
        idx_pip = hand.landmark[mp_hands.HandLandmark.INDEX_FINGER_PIP]
        mid_tip = hand.landmark[mp_hands.HandLandmark.MIDDLE_FINGER_TIP]
        mid_pip = hand.landmark[mp_hands.HandLandmark.MIDDLE_FINGER_PIP]
        ix, iy = int(idx_tip.x * w), int(idx_tip.y * h)
        mx, my = int(mid_tip.x * w), int(mid_tip.y * h)
        fingertip_idx = (ix, iy)
        fingertip_mid = (mx, my)
        fingertip_idx_global = fingertip_idx
        # draw fingertips
        cv2.circle(frame_vis, fingertip_idx, 8, (0,255,0), -1)
        cv2.circle(frame_vis, fingertip_mid,  8, (0,255,0), -1)
        # check extension
        index_ext = idx_tip.y < idx_pip.y
        middle_ext = mid_tip.y < mid_pip.y
        # if measurement in progress but fingers no longer both extended, stop measuring
        if measuring and not (index_ext and middle_ext):
            measuring = False
        # pinch distance
        pinch = compute_distance(fingertip_idx, fingertip_mid)
        cv2.putText(frame_vis, f"Pinch: {int(pinch)} px", (10,30), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0,255,0),2)

        # pinch calibration
        if calibrating:
            if pinch < PINCH_THRESHOLD and cal_start is None:
                cal_start = fingertip_idx
                print("Pinch calibration start set")
            elif pinch > RELEASE_THRESHOLD and cal_start is not None:
                cal_end = fingertip_idx
                px = compute_distance(cal_start, cal_end)
                inches = float(input("Enter actual distance between points (inches): "))
                PIXELS_PER_INCH = px / inches
                print(f"Calibrated: {PIXELS_PER_INCH:.2f} px/inch")
                calibrating = False
                cal_start = None
        # measurement gesture
        elif not object_calibrating and index_ext and middle_ext:
            if pinch < PINCH_THRESHOLD and not measuring:
                measuring = True
                start_pt = fingertip_idx
            elif pinch > RELEASE_THRESHOLD and measuring:
                measuring = False

    # object calibration
    if object_calibrating and fingertip_idx is not None:
        hsv = cv2.cvtColor(frame_vis, cv2.COLOR_BGR2HSV)
        mask = cv2.inRange(hsv, np.array([10,100,100]), np.array([25,255,255]))
        masked = cv2.bitwise_and(frame_vis, frame_vis, mask=mask)
        gray = cv2.cvtColor(masked, cv2.COLOR_BGR2GRAY)
        gray = cv2.medianBlur(gray,5)
        circles = cv2.HoughCircles(gray, cv2.HOUGH_GRADIENT, 1.2, 100,
                                   param1=50, param2=30, minRadius=10, maxRadius=300)
        if circles is not None:
            circles = np.round(circles[0]).astype(int)
            touched = [(x,y,r) for x,y,r in circles
                       if abs(compute_distance((x,y), fingertip_idx)-r) < CIRCLE_TOUCH_THRESHOLD]
            if touched:
                touched.sort(key=lambda c: abs(compute_distance((c[0],c[1]), fingertip_idx)-c[2]))
                x,y,r = touched[0]
                cal_circle = (x,y,r)
                PIXELS_PER_INCH = 2 * r
                print(f"Circle calib: {PIXELS_PER_INCH:.2f} px/inch")
                object_calibrating = False

    # permanent reference circle
    if cal_circle:
        cx,cy,cr = cal_circle
        cv2.circle(frame_vis,(cx,cy),cr,(0,0,255),2)
        cv2.drawMarker(frame_vis,(cx,cy),(0,0,255),cv2.MARKER_TILTED_CROSS,15,1)
        cv2.putText(frame_vis,f"Ref r={cr} px",(cx-cr,cy+cr+20),cv2.FONT_HERSHEY_SIMPLEX,0.5,(0,0,255),1)

    return frame_vis

# ----- Main -----
def main():
    caps = []
    for idx in CAMERA_INDICES:
        cap = cv2.VideoCapture(idx)
        cap.set(cv2.CAP_PROP_FRAME_WIDTH, FRAME_WIDTH)
        cap.set(cv2.CAP_PROP_FRAME_HEIGHT, FRAME_HEIGHT)
        caps.append(cap)
    if not all(cap.isOpened() for cap in caps):
        print("Error: could not open all cameras")
        return

    print("Press 'c' for pinch calib, 'o' for circle calib, 'q' to quit.")
    while True:
        frames = [cap.read()[1] for cap in caps]
        frame = next((f for f in frames if f is not None), None)
        if frame is None:
            break
        # process
        full_view = process_frame(frame)
        # create proj output
        proj = np.zeros_like(full_view)
        if cal_circle:
            cx,cy,cr = cal_circle
            cv2.circle(proj,(cx,cy),cr,(0,0,255),2)
        if measuring and start_pt and fingertip_idx_global:
            cv2.line(proj, start_pt, fingertip_idx_global, (255,0,0),2)
            px = compute_distance(start_pt, fingertip_idx_global)
            inch = px/PIXELS_PER_INCH; cm = inch*2.54
            mid = ((start_pt[0]+fingertip_idx_global[0])//2,(start_pt[1]+fingertip_idx_global[1])//2)
            cv2.putText(proj,f"{inch:.2f}in/{cm:.1f}cm",(mid[0]+10,mid[1]-10),
                        cv2.FONT_HERSHEY_SIMPLEX,0.7,(255,0,0),2)
        # overlay proj onto debug
        debug = full_view.copy()
        # overlay ref circle
        if cal_circle:
            cx,cy,cr = cal_circle
            cv2.circle(debug,(cx,cy),cr,(0,0,255),2)
        # overlay measurement
        if measuring and start_pt and fingertip_idx_global:
            cv2.line(debug, start_pt, fingertip_idx_global, (255,0,0),2)
            cv2.putText(debug,f"{inch:.2f}in/{cm:.1f}cm",(mid[0]+10,mid[1]-10),
                        cv2.FONT_HERSHEY_SIMPLEX,0.7,(255,0,0),2)

        # show windows
        cv2.imshow('Hand Measure', debug)
        cv2.imshow('Projector Output', proj)

        key = cv2.waitKey(1) & 0xFF
        if key == ord('q'):
            break
        elif key == ord('c'):
            global calibrating
            calibrating = True
            cal_start = None
            print("Entered pinch calibration mode.")
        elif key == ord('o'):
            global object_calibrating
            object_calibrating = True
            print("Entered circle calibration mode.")

    for cap in caps:
        cap.release()
    cv2.destroyAllWindows()

if __name__ == '__main__':
    main()