#!/usr/bin/env python3
"""
OpenCV Image Matching Script for Dijital Rehber Application
This script performs feature-based image matching using OpenCV to identify exhibits
"""

import cv2
import numpy as np
import argparse
import sys
import json
from typing import List, Tuple, Dict, Optional


def log_debug(message: str):
    print(message, file=sys.stderr)


def detect_and_compute_features(image_path: str, detector_type: str = 'sift'):
    """
    Detect and compute features from an image using specified detector
    """
    # Read the image
    img = cv2.imread(image_path, cv2.IMREAD_COLOR)
    if img is None:
        raise ValueError(f"Could not load image: {image_path}")

    # Convert to grayscale for feature detection
    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)

    # Initialize the detector
    if detector_type.lower() == 'sift':
        detector = cv2.SIFT_create()
    elif detector_type.lower() == 'orb':
        detector = cv2.ORB_create(
            nfeatures=800,
            scaleFactor=1.2,
            edgeThreshold=31,
            patchSize=31,
            WTA_K=2,
            scoreType=cv2.ORB_HARRIS_SCORE
        )
    elif detector_type.lower() == 'akaze':
        detector = cv2.AKAZE_create()
    else:
        raise ValueError(f"Unsupported detector type: {detector_type}")

    # Detect keypoints and compute descriptors
    keypoints, descriptors = detector.detectAndCompute(gray, None)

    log_debug(f"[FEATURES] image={image_path} detector={detector_type} kpts={len(keypoints)} desc_shape={None if descriptors is None else descriptors.shape} dtype={None if descriptors is None else descriptors.dtype}")

    return keypoints, descriptors, img


def match_images(target_desc: np.ndarray,
                 reference_desc: np.ndarray,
                 matcher_type: str = 'bf',
                 ratio_threshold: float = 0.75) -> Tuple[float, int]:
    """
    Match descriptors between target and reference images
    Returns a tuple of (similarity score, good match count)
    """
    if target_desc is None or reference_desc is None:
        log_debug(f"[MATCH] skipped because target_desc or reference_desc is None target={target_desc is None} ref={reference_desc is None}")
        return 0.0, 0

    # Initialize matcher with a norm appropriate for descriptor type
    if matcher_type.lower() == 'bf':
        if target_desc.dtype == np.uint8 or reference_desc.dtype == np.uint8:
            matcher = cv2.BFMatcher(cv2.NORM_HAMMING, crossCheck=False)
        else:
            matcher = cv2.BFMatcher(crossCheck=False)
    elif matcher_type.lower() == 'flann':
        # FLANN expects float descriptors for most algorithms
        if target_desc.dtype != np.float32:
            target_desc = np.float32(target_desc)
        if reference_desc.dtype != np.float32:
            reference_desc = np.float32(reference_desc)
        matcher = cv2.FlannBasedMatcher(indexParams=dict(algorithm=1, trees=5), searchParams=dict(checks=50))
    else:
        raise ValueError(f"Unsupported matcher type: {matcher_type}")

    # Perform matching
    good_matches = []
    try:
        matches = matcher.knnMatch(target_desc, reference_desc, k=2)
        raw_match_count = len(matches)
        sample_distances = []
        for i, match_pair in enumerate(matches[:5]):
            if len(match_pair) == 2:
                m, n = match_pair
                sample_distances.append(f"{i}:({m.distance:.2f},{n.distance:.2f})")
            else:
                sample_distances.append(f"{i}:(len={len(match_pair)})")

        log_debug(f"[MATCH] raw_matches={raw_match_count} sample_distances={' | '.join(sample_distances)}")

        # Apply Lowe's ratio test to filter good matches
        for match_pair in matches:
            if len(match_pair) == 2:
                m, n = match_pair
                if m.distance < ratio_threshold * n.distance:
                    good_matches.append(m)
    except cv2.error as e:
        log_debug(f"[MATCH] knnMatch failed, falling back to match: {e}")
        try:
            direct_matches = matcher.match(target_desc, reference_desc)
            direct_matches = sorted(direct_matches, key=lambda x: x.distance)
            raw_match_count = len(direct_matches)
            sample_distances = [f"{i}:({m.distance:.2f})" for i, m in enumerate(direct_matches[:5])]
            log_debug(f"[MATCH] fallback raw_matches={raw_match_count} sample_distances={' | '.join(sample_distances)}")
            for m in direct_matches:
                if m.distance < 60:
                    good_matches.append(m)
        except cv2.error as e2:
            log_debug(f"[MATCH] fallback match failed too: {e2}")
            return 0.0, 0

    if not good_matches:
        log_debug(f"[MATCH] no good matches target_desc={target_desc.shape if hasattr(target_desc, 'shape') else target_desc} ref_desc={reference_desc.shape if hasattr(reference_desc, 'shape') else reference_desc} ratio_threshold={ratio_threshold}")
        return 0.0, 0

    # Calculate similarity score based on number of good matches
    denominator = max(10, int(min(len(target_desc), len(reference_desc)) * 0.10))
    score = len(good_matches) / denominator

    # Cap the score between 0 and 1
    score = min(score, 1.0)
    log_debug(f"[MATCH] score={score} good_matches={len(good_matches)} target_desc_len={len(target_desc)} reference_desc_len={len(reference_desc)}")
    return score, len(good_matches)


def match_target_to_references(target_path: str,
                              reference_paths: List[str],
                              exhibit_ids: List[str],
                              exhibit_titles: List[str],
                              detector_type: str = 'sift',
                              matcher_type: str = 'bf') -> List[Dict]:
    """
    Match a target image against multiple reference images
    """
    try:
        log_debug(f"[INPUT] target={target_path} refs={len(reference_paths)} ids={len(exhibit_ids)} titles={len(exhibit_titles)} detector={detector_type} matcher={matcher_type}")

        # Extract features from target image
        target_kp, target_desc, target_img = detect_and_compute_features(target_path, detector_type)

        if target_desc is None:
            log_debug("[ERROR] target descriptor extraction failed")
            print(json.dumps([]))
            return []

        results = []

        for ref_path, exhibit_id, exhibit_title in zip(reference_paths, exhibit_ids, exhibit_titles):
            try:
                # Extract features from reference image
                ref_kp, ref_desc, ref_img = detect_and_compute_features(ref_path, detector_type)

                if ref_desc is None:
                    log_debug(f"[ERROR] reference descriptor extraction failed ref={ref_path}")
                    continue

                # Calculate similarity score
                similarity_score, good_match_count = match_images(target_desc, ref_desc, matcher_type)

                if similarity_score >= 0.01:  # Debug threshold to expose weak candidate matches
                    results.append({
                        'exhibit_id': exhibit_id,
                        'exhibit_title': exhibit_title,
                        'confidence': float(similarity_score),
                        'matches_count': int(good_match_count)
                    })
                    log_debug(f"[REF] matched ref={ref_path} id={exhibit_id} title={exhibit_title} score={similarity_score}")
                else:
                    log_debug(f"[REF] skipped ref={ref_path} id={exhibit_id} title={exhibit_title} score={similarity_score}")
            except Exception as e:
                print(f"Error processing reference image {ref_path}: {str(e)}", file=sys.stderr)
                continue

        # Sort results by confidence score (descending)
        results.sort(key=lambda x: x['confidence'], reverse=True)
        log_debug(f"[RESULTS] total_matches={len(results)}")

        return results

    except Exception as e:
        print(f"Error in image matching: {str(e)}", file=sys.stderr)
        return []


def main():
    parser = argparse.ArgumentParser(description='OpenCV Image Matching for Exhibit Recognition')
    parser.add_argument('--target', required=True, help='Path to target image')
    parser.add_argument('--refs', help='Comma-separated paths to reference images')
    parser.add_argument('--ids', help='Comma-separated exhibit IDs')
    parser.add_argument('--titles', help='Comma-separated exhibit titles')
    parser.add_argument('--refs-json', help='JSON array of reference image paths')
    parser.add_argument('--ids-json', help='JSON array of exhibit IDs')
    parser.add_argument('--titles-json', help='JSON array of exhibit titles')
    parser.add_argument('--detector', default='sift', choices=['sift', 'orb', 'akaze'],
                       help='Feature detector to use (default: sift)')
    parser.add_argument('--matcher', default='bf', choices=['bf', 'flann'],
                       help='Descriptor matcher to use (default: bf)')

    args = parser.parse_args()

    # Parse values from JSON arrays if provided, otherwise fall back to comma-separated values
    if args.refs_json:
        reference_paths = json.loads(args.refs_json)
    elif args.refs:
        reference_paths = args.refs.split(',')
    else:
        reference_paths = []

    if args.ids_json:
        exhibit_ids = json.loads(args.ids_json)
    elif args.ids:
        exhibit_ids = args.ids.split(',')
    else:
        exhibit_ids = []

    if args.titles_json:
        exhibit_titles = json.loads(args.titles_json)
    elif args.titles:
        exhibit_titles = args.titles.split(',')
    else:
        exhibit_titles = []

    # Validate inputs
    if len(reference_paths) != len(exhibit_ids) or len(reference_paths) != len(exhibit_titles):
        print("Error: Number of reference paths, IDs, and titles must match", file=sys.stderr)
        sys.exit(1)

    # Perform matching
    results = match_target_to_references(
        args.target,
        reference_paths,
        exhibit_ids,
        exhibit_titles,
        args.detector,
        args.matcher
    )

    # Output results as JSON
    print(json.dumps(results))


if __name__ == "__main__":
    main()
