import cv2
import numpy as np
import torch
from sam2.build_sam import build_sam2
from sam2.sam2_image_predictor import SAM2ImagePredictor

def interactive_sam2_extractor(image_path, checkpoint_path, model_cfg):
    print("جاري تهيئة SAM 2 والبحث عن مسرعات الرسوميات...")
    
    if torch.cuda.is_available():
        device = torch.device("cuda")
        torch.autocast("cuda", dtype=torch.bfloat16).__enter__()
        print("تم العثور على NVIDIA GPU!")
    elif torch.backends.mps.is_available():
        device = torch.device("mps")
        print("تم العثور على Apple MPS!")
    else:
        device = torch.device("cpu")
        print("سيتم استخدام المعالج المركزي CPU.")

    print("جاري تحميل أوزان SAM 2...")
    sam2_model = build_sam2(model_cfg, checkpoint_path, device=device)
    predictor = SAM2ImagePredictor(sam2_model)

    img = cv2.imread(image_path)
    if img is None:
        print("خطأ: لم يتم العثور على الصورة.")
        return

    h, w = img.shape[:2]
    max_w = 1000
    if w > max_w:
        img = cv2.resize(img, (max_w, int(h * (max_w / w))))

    print("جاري تحليل البنية البصرية للصخرة...")
    predictor.set_image(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))

    white_paper = np.ones_like(img) * 255

    print("\n✅ جاهز! انقر بالماوس الأيسر على أي نقش لاستخراجه.")
    print("اضغط 'q' للخروج، 'c' لمسح الورقة، أو 's' للحفظ.")

    def get_click(event, x, y, flags, param):
        nonlocal white_paper
        if event == cv2.EVENT_LBUTTONDOWN:
            print(f"تم النقر على: ({x}, {y})")
            input_point = np.array([[x, y]])
            input_label = np.array([1])
            masks, scores, logits = predictor.predict(
                point_coords=input_point,
                point_labels=input_label,
                multimask_output=False,
            )
            mask = masks[0]
            mask_uint8 = (mask * 255).astype(np.uint8)
            contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
            cv2.drawContours(white_paper, contours, -1, (0, 0, 0), 2)
            cv2.imshow("SAM 2 Extracted (White Paper)", white_paper)

    cv2.namedWindow("Original Image - Click Here")
    cv2.setMouseCallback("Original Image - Click Here", get_click)
    cv2.imshow("Original Image - Click Here", img)
    cv2.imshow("SAM 2 Extracted (White Paper)", white_paper)

    while True:
        key = cv2.waitKey(1) & 0xFF
        if key == ord('q'):
            break
        elif key == ord('c'):
            white_paper = np.ones_like(img) * 255
            cv2.imshow("SAM 2 Extracted (White Paper)", white_paper)
            print("تم مسح الورقة.")
        elif key == ord('s'):
            cv2.imwrite("sam2_extracted_symbols.png", white_paper)
            print("تم الحفظ كـ: sam2_extracted_symbols.png")

    cv2.destroyAllWindows()

interactive_sam2_extractor(
    image_path="rock.jpg",
    checkpoint_path="sam2_hiera_large.pt",
    model_cfg="sam2_hiera_l.yaml"
)
