#!/usr/bin/env python3
import cv2
import time
import sys
import os
import signal
from edge_impulse_linux.image import ImageImpulseRunner

# ================== 配置区域 ==================
MODEL_FILE = "/home/arduino/1223-linux-aarch64-v2-impulse-#1.eim"

# 采集分辨率：若追求更丝滑，可改用 320x240
FRAME_WIDTH = 320
FRAME_HEIGHT = 240
DEFAULT_CAMERA_DEVICE = "/dev/video2"

# 阈值配置
CONFIDENCE_THRESHOLD = 0.50  # 基础过滤阈值
NMS_IOU_THRESHOLD = 0.40     # NMS 重叠抑制阈值：越小对重合框过滤越严密

# 性能优化：每隔多少帧进行一次 AI 推理 (设为 2 表示每两帧识别一次，其余帧平滑显示)
SKIP_FRAMES = 3
# =============================================

runner = None
show_camera = True


def sigint_handler(sig, frame):
    print('\nInterrupted by SIGINT')
    if runner:
        runner.stop()
    sys.exit(0)


signal.signal(signal.SIGINT, sigint_handler)


def find_camera_device():
    for cam_id in range(10):
        device = f"/dev/video{cam_id}"
        if not os.path.exists(device):
            continue
        cap = cv2.VideoCapture(cam_id, cv2.CAP_V4L2)
        if cap.isOpened():
            ret, _ = cap.read()
            cap.release()
            if ret:
                print(f"Using camera: {device}")
                return device
        else:
            cap.release()
    return DEFAULT_CAMERA_DEVICE


def open_camera():
    device = find_camera_device()
    cam_id = int(device.replace("/dev/video", ""))
    cap = cv2.VideoCapture(cam_id, cv2.CAP_V4L2)
    cap.set(cv2.CAP_PROP_FRAME_WIDTH, FRAME_WIDTH)
    cap.set(cv2.CAP_PROP_FRAME_HEIGHT, FRAME_HEIGHT)
    cap.set(cv2.CAP_PROP_BUFFERSIZE, 1)
    return cap


def main():
    global runner

    if not os.path.exists(MODEL_FILE):
        print(f"Error: Model file not found at: {MODEL_FILE}")
        sys.exit(1)

    with ImageImpulseRunner(MODEL_FILE) as runner:
        cap = None
        try:
            model_info = runner.init()
            in_w = model_info['model_parameters']['image_input_width']
            in_h = model_info['model_parameters']['image_input_height']

            cap = open_camera()
            if not cap.isOpened():
                print("Failed to open camera.")
                sys.exit(1)

            actual_w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
            actual_h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))

            # 居中映射坐标计算
            crop_size = min(actual_w, actual_h)
            scale = crop_size / in_w
            offset_x = (actual_w - crop_size) // 2
            offset_y = (actual_h - crop_size) // 2

            if show_camera:
                cv2.namedWindow('UNO Q Object Detection', cv2.WINDOW_NORMAL)
                cv2.resizeWindow('UNO Q Object Detection', actual_w, actual_h)

            frame_count = 0
            cached_boxes = []  # 缓存上一帧的检测结果，保证画面流畅

            while True:
                ret, frame = cap.read()
                if not ret or frame is None:
                    continue

                frame_count += 1

                # 隔帧推理：只有满足指定间隔才触发一次耗时的 AI 模型推理
                if frame_count % SKIP_FRAMES == 0:
                    img_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
                    features, _ = runner.get_features_from_image(img_rgb)
                    res = runner.classify(features)

                    raw_boxes = []
                    scores = []
                    labels = []

                    if "bounding_boxes" in res["result"].keys():
                        for bb in res["result"]["bounding_boxes"]:
                            if bb['value'] >= CONFIDENCE_THRESHOLD:
                                # 还原到屏幕真实坐标
                                x = int(bb['x'] * scale + offset_x)
                                y = int(bb['y'] * scale + offset_y)
                                w = int(bb['width'] * scale)
                                h = int(bb['height'] * scale)

                                raw_boxes.append([x, y, w, h])
                                scores.append(float(bb['value']))
                                labels.append(bb['label'])

                        # 核心：使用 NMS 算法合并重合框，解决同一物体出现多个框的问题
                        cached_boxes = []
                        if len(raw_boxes) > 0:
                            indices = cv2.dnn.NMSBoxes(
                                raw_boxes, scores, CONFIDENCE_THRESHOLD, NMS_IOU_THRESHOLD
                            )
                            # 兼容不同版本 OpenCV 返回的列表/数组维度
                            if len(indices) > 0:
                                for i in indices.flatten():
                                    cached_boxes.append({
                                        "box": raw_boxes[i],
                                        "score": scores[i],
                                        "label": labels[i]
                                    })

                # 每一帧都绘制缓存中的检测框（即使当前帧没做推理，也能保证流畅跟手）
                for item in cached_boxes:
                    x, y, w, h = item["box"]
                    score = item["score"]
                    label = item["label"]

                    cv2.rectangle(frame, (x, y), (x + w, y + h), (0, 255, 0), 2)
                    text = f"{label}: {score:.2f}"
                    cv2.putText(frame, text, (x, max(y - 8, 20)),
                                cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2)

                if show_camera:
                    cv2.imshow('UNO Q Object Detection', frame)
                    if cv2.waitKey(1) == ord('q'):
                        break

        except KeyboardInterrupt:
            print("\nStopped.")
        finally:
            if cap is not None:
                cap.release()
            if show_camera:
                cv2.destroyAllWindows()
            if runner:
                runner.stop()


if __name__ == "__main__":
    main()