import cv2
import numpy as np
from segment_anything import sam_model_registry, SamPredictor

def interactive_sam_extractor(image_path, checkpoint_path):
    print("جاري تحميل نموذج SAM إلى الذاكرة... (قد يستغرق ثوانٍ)")
    
    # 1. إعداد النموذج (نستخدم النسخة الأساسية vit_b)
    sam_checkpoint = checkpoint_path
    model_type = "vit_b"
    sam = sam_model_registry[model_type](checkpoint=sam_checkpoint)
    
    # إذا كان لديك كرت شاشة يدعم CUDA، يمكنك تفعيل السطر التالي لتسريع العملية جداً
    # sam.to(device="cuda")
    
    predictor = SamPredictor(sam)

    # 2. قراءة الصورة
    img = cv2.imread(image_path)
    if img is None:
        print("خطأ: تأكد من مسار الصورة.")
        return
        
    # تصغير الصورة لتناسب الشاشة لتسهيل النقر
    h, w = img.shape[:2]
    max_w = 800
    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("جاهز! انقر بالماوس على أي نقش في الصورة لاستخراجه.")
    print("اضغط 'q' للخروج، أو 'c' لمسح الورقة والبدء من جديد، أو 's' للحفظ.")

    # 3. دالة الاستجابة لضغطات الماوس
    def get_click(event, x, y, flags, param):
        nonlocal white_paper
        
        if event == cv2.EVENT_LBUTTONDOWN:
            print(f"تم النقر على الإحداثيات: ({x}, {y})، جاري الاستخراج...")
            
            # تحديد النقطة التي ضغطت عليها كـ "إشارة إيجابية" (1)
            input_point = np.array([[x, y]])
            input_label = np.array([1])
            
            # توليد القناع (Mask) للشكل الذي تم النقر عليه
            masks, scores, logits = predictor.predict(
                point_coords=input_point,
                point_labels=input_label,
                multimask_output=False,
            )
            
            mask = masks[0] # القناع المنطقي (True للنقش، False للخلفية)
            
            # تحويل القناع إلى صيغة مناسبة لاستخراج الحدود
            mask_uint8 = (mask * 255).astype(np.uint8)
            
            # استخراج الكفاف (الحدود الخارجية للرمز)
            contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
            
            # رسم الرمز على الورقة البيضاء باللون الأسود (سُمك القلم = 2)
            cv2.drawContours(white_paper, contours, -1, (0, 0, 0), 2)
            
            # تحديث عرض النتيجة
            cv2.imshow("Extracted Symbols (White Paper)", white_paper)

    # 4. إعداد النوافذ
    cv2.namedWindow("Original Image - Click Here")
    cv2.setMouseCallback("Original Image - Click Here", get_click)
    
    cv2.imshow("Original Image - Click Here", img)
    cv2.imshow("Extracted Symbols (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("Extracted Symbols (White Paper)", white_paper)
            print("تم مسح الورقة.")
        elif key == ord('s'): # حفظ
            cv2.imwrite("sam_extracted_symbols.png", white_paper)
            print("تم حفظ النتيجة بنجاح!")

    cv2.destroyAllWindows()

# تأكد من وضع مسار الصورة ومسار ملف الأوزان بشكل صحيح
interactive_sam_extractor("rock.jpg", "sam_vit_b_01ec64.pth")
