import cv2
import numpy as np
import os
import json
import logging
from typing import List, Dict, Any, Tuple, Optional
import pytesseract
from PIL import Image

logger = logging.getLogger(__name__)

# Constants
DAYS_OF_WEEK = ['Day', 'Monday', 'Tuesday', 'Wednesday', 'Thursday', 'Friday', 'Saturday', 'Sunday']
DEBUG_DIR = "test_data"

class SlotCountService:
    """
    Service for extracting cells from routine images and calculating slot counts.
    Uses grid-based table reconstruction to accurately determine column spans.
    Includes OCR text extraction and day detection.
    """

    @staticmethod
    def extract_cells_with_slot_count(file_path: str) -> Dict[str, Any]:
        """
        Extract cells from an image/PDF and calculate slot counts based on column span.
        
        Args:
            file_path (str): Path to the image or PDF file
            
        Returns:
            Dict containing days, slots, and backward compatible fields.
        """
        try:
            # 1. Convert to image with high DPI for consistency
            img = SlotCountService._convert_file_to_opencv_image(file_path)
            if img is None:
                return {'error': 'Failed to load/convert image', 'days': [], 'slots': [], 'structured_data': [], 'cells': []}
            
            # 2. Detect grid lines
            h_lines, v_lines = SlotCountService._detect_grid_lines(img)
            
            # 3. Extract x and y boundaries
            row_boundaries = SlotCountService._get_line_boundaries(h_lines, is_horizontal=True)
            col_boundaries = SlotCountService._get_line_boundaries(v_lines, is_horizontal=False)
            
            # Save debug grid visualization
            SlotCountService._save_debug_grid(img, row_boundaries, col_boundaries)

            if len(row_boundaries) < 2 or len(col_boundaries) < 2:
                return {'error': 'Could not detect a valid grid', 'days': [], 'slots': [], 'structured_data': [], 'cells': []}

            # 4. Extract cells using the grid
            cells = SlotCountService._extract_grid_cells(img, h_lines, v_lines, row_boundaries, col_boundaries)
            SlotCountService._save_debug_data("extract_grid_cells.json", cells)
            # 5. Detect days from header
            days_map, is_day_in_col = SlotCountService._detect_days_from_grid(cells)
            # 6. Process the content into timetable slots
            slots = SlotCountService._process_timetable_cells(cells, days_map, is_day_in_col)
            SlotCountService._save_debug_data("slot_data.json", slots)
            # Extract unique days maintaining order of DAYS_OF_WEEK
            days_list = sorted(list(set(days_map.values())), key=lambda d: DAYS_OF_WEEK.index(d) if d in DAYS_OF_WEEK else 99)
            # Debug save
            SlotCountService._save_debug_data("extracted_slots.json", slots)

            # Format result with backward compatibility
            # Remove cell_img (ndarray) to prevent JSON serialization errors upstream
            clean_cells = []
            for cell in cells:
                clean_cell = cell.copy()
                if 'cell_img' in clean_cell:
                    del clean_cell['cell_img']
                clean_cells.append(clean_cell)

            return {
                "days": days_list,
                "slots": slots,
                "structured_data": slots, # Alias for backward compatibility
                "cells": clean_cells,     # For debugging/backward compat
                "base_width": 1.0,        # Dummy value for old logic
                "error": None
            }
        except Exception as e:
            logger.error(f"Error in extract_cells_with_slot_count: {str(e)}")
            return {'error': str(e), 'days': [], 'slots': [], 'structured_data': [], 'cells': []}

    @staticmethod
    def _convert_file_to_opencv_image(file_path: str) -> Optional[np.ndarray]:
        """Convert PDF or image file to OpenCV format (BGR numpy array)."""
        file_ext = os.path.splitext(file_path)[1].lower()
        
        try:
            if file_ext == '.pdf':
                from pdf2image import convert_from_path
                # Convert PDF using high DPI (300) for better OCR
                images = convert_from_path(file_path, first_page=1, last_page=1, dpi=300)
                if not images:
                    logger.error("Could not convert PDF to image")
                    return None
                img = cv2.cvtColor(np.array(images[0]), cv2.COLOR_RGB2BGR)
                logger.info(f"Successfully converted PDF to image: {img.shape}")
                return img
                
            elif file_ext in ['.jpg', '.jpeg', '.png']:
                img = cv2.imread(file_path)
                if img is None:
                    logger.error(f"Failed to read image file: {file_path}")
                    return None
                logger.info(f"Successfully read image: {img.shape}")
                return img
                
            else:
                logger.error(f"Unsupported file format: {file_ext}")
                return None
                
        except Exception as e:
            logger.error(f"Error converting file: {str(e)}")
            return None

    @staticmethod
    def _detect_grid_lines(img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
        """Detect all vertical and horizontal grid lines using morphology."""
        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
        
        # Adaptive thresholding handles varying illumination
        thresh = cv2.adaptiveThreshold(
            cv2.bitwise_not(gray), 255,
            cv2.ADAPTIVE_THRESH_GAUSSIAN_C, cv2.THRESH_BINARY, 15, -2
        )

        # Morphological operations to find horizontal lines
        h_length = max(40, img.shape[1] // 40)
        h_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (h_length, 1))
        h_lines = cv2.morphologyEx(thresh, cv2.MORPH_OPEN, h_kernel, iterations=2)

        # Morphological operations to find vertical lines
        v_length = max(40, img.shape[0] // 40)
        v_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (1, v_length))
        v_lines = cv2.morphologyEx(thresh, cv2.MORPH_OPEN, v_kernel, iterations=2)

        return h_lines, v_lines

    @staticmethod
    def _get_line_boundaries(line_img: np.ndarray, is_horizontal: bool) -> List[int]:
        """Extract boundaries (x or y coordinates) from grid lines."""
        axis = 1 if is_horizontal else 0
        proj = np.sum(line_img, axis=axis)
        
        # Filter noise
        threshold = np.max(proj) * 0.1
        
        boundaries = []
        in_line = False
        start = 0
        
        for i, val in enumerate(proj):
            if val > threshold and not in_line:
                in_line = True
                start = i
            elif val <= threshold and in_line:
                in_line = False
                boundaries.append((start + i) // 2)
        
        if in_line:
            boundaries.append((start + len(proj)) // 2)
            
        # Filter out boundaries that are too close (e.g., thick lines)
        filtered = []
        min_dist = 20
        for b in boundaries:
            if not filtered or b - filtered[-1] > min_dist:
                filtered.append(b)
                
        return filtered

    @staticmethod
    def _extract_grid_cells(img: np.ndarray, h_lines: np.ndarray, v_lines: np.ndarray, 
                            row_bounds: List[int], col_bounds: List[int]) -> List[Dict[str, Any]]:
        """Extract cells and calculate span properties using the detected boundaries."""
        # Combine horizontal and vertical lines
        grid = cv2.add(v_lines, h_lines)
        
        # Dilate to close small gaps in the grid
        kernel = np.ones((3,3), np.uint8)
        grid = cv2.dilate(grid, kernel, iterations=1)
        
        contours, _ = cv2.findContours(grid, cv2.RETR_TREE, cv2.CHAIN_APPROX_SIMPLE)
        
        cells = []
        # Exclude the outermost frame contour by checking size
        max_w = img.shape[1] * 0.95
        max_h = img.shape[0] * 0.95
        
        for cnt in contours:
            x, y, w, h = cv2.boundingRect(cnt)
            
            # Filter very small noise and the outer page boundary
            if w > 30 and h > 20 and w < max_w and h < max_h:
                # Map cell corners to the nearest grid boundaries
                c_start = SlotCountService._get_closest_index(x, col_bounds)
                c_end = SlotCountService._get_closest_index(x + w, col_bounds)
                
                r_start = SlotCountService._get_closest_index(y, row_bounds)
                r_end = SlotCountService._get_closest_index(y + h, row_bounds)
                
                # Verify valid mapping
                if c_start < c_end and r_start < r_end:
                    cell_img = img[y:y+h, x:x+w]
                    # Calculate slot count using column span! (CRITICAL improvement)
                    slot_count = c_end - c_start
                    
                    # Store cell info without running OCR immediately to save time
                    cells.append({
                        'x': x, 'y': y, 'width': w, 'height': h,
                        'row': r_start,
                        'row_span': r_end - r_start,
                        'col_start': c_start,
                        'col_end': c_end - 1,
                        'slot_count': slot_count,
                        'cell_img': cell_img
                    })
        
        # Sort cells by row then column
        cells.sort(key=lambda c: (c['row'], c['col_start']))
        return cells

    @staticmethod
    def _get_closest_index(val: int, bounds: List[int]) -> int:
        """Find the index of the boundary closest to the given value."""
        return min(range(len(bounds)), key=lambda i: abs(bounds[i] - val))

    @staticmethod
    def _ocr_cell(cell_img: np.ndarray) -> str:
        """Improve OCR accuracy via preprocessing (Grayscale, Resize, Otsu)."""
        if cell_img.size == 0:
            return ""
            
        # 1. Grayscale
        if len(cell_img.shape) == 3:
            gray = cv2.cvtColor(cell_img, cv2.COLOR_BGR2GRAY)
        else:
            gray = cell_img
            
        # 2. Resize (3x scaling)
        gray = cv2.resize(gray, None, fx=3, fy=3, interpolation=cv2.INTER_CUBIC)
        
        # 3. Otsu threshold (yields black text on white bg for tesseract)
        _, thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)
        
        # 4. Use pytesseract with config '--psm 6' (Assume a single uniform block of text)
        custom_config = r'--oem 3 --psm 6'
        text = pytesseract.image_to_string(thresh, config=custom_config)
        
        # Extract only the first meaningful line
        lines = [line.strip() for line in text.split('\n') if line.strip()]
        if lines:
            return lines[0]
        return ""

    @staticmethod
    def _get_matched_day(text: str) -> Optional[str]:
        """Return day string if matched."""
        text_lower = text.lower()
        for d in DAYS_OF_WEEK:
            if d.lower() == text_lower:
                return d
        return None

    @staticmethod
    def _detect_days_from_grid(cells: List[Dict[str, Any]]) -> Tuple[Dict[int, str], bool]:
        """
        Detect days using OCR on the header row/col.
        Returns a mapping from index to Day string, and boolean if days are columns.
        """
        days_map = {}
        
        # Priority 1: Days are in header row (columns = days)
        header_cells = [c for c in cells if c['row'] == 0]
        day_in_row = False
        for cell in header_cells:
            text = SlotCountService._ocr_cell(cell['cell_img'])
            day = SlotCountService._get_matched_day(text)
            if day:
                day_in_row = True
                for col in range(cell['col_start'], cell['col_end'] + 1):
                    days_map[col] = day
                    
        if day_in_row:
            return days_map, True
            
        # Priority 2: Days are in first column (rows = days)
        first_col_cells = [c for c in cells if c['col_start'] == 0]
        for cell in first_col_cells:
            text = SlotCountService._ocr_cell(cell['cell_img'])
            day = SlotCountService._get_matched_day(text)
            if day:
                for row in range(cell['row'], cell['row'] + cell['row_span']):
                    days_map[row] = day         
        return days_map, False

    @staticmethod
    def _is_header_or_empty(text: str) -> bool:
        if not text:
            return True
        text_lower = text.lower()
        keywords = ["course code", "course title", "room", "instructor", "teacher", "class", "time"]
        return any(k in text_lower for k in keywords)

    @staticmethod
    def _process_timetable_cells(cells: List[Dict[str, Any]], days_map: Dict[int, str], is_day_in_col: bool) -> List[Dict[str, Any]]:
        """Extract final slot information for valid content cells."""
        slots = []
        for cell in cells:
            # Skip header column/row
            if is_day_in_col and cell['row'] == 0:
                continue
            if not is_day_in_col and cell['col_start'] == 0:
                continue
                
            text = SlotCountService._ocr_cell(cell['cell_img'])
            if SlotCountService._get_matched_day(text) or SlotCountService._is_header_or_empty(text):
                continue
                
            # Populate text field to be backwards compatible with old code
            cell['text'] = text
                
            idx = cell['col_start'] if is_day_in_col else cell['row']
            day = days_map.get(idx)
            
            if not day:
                continue
                
            time_slot = cell['row'] if is_day_in_col else cell['col_start']
            
            # Map column to day
            cell['day'] = day
            
            slots.append({
                "day": day,
                "time_slot": time_slot,
                "course": text,
                "text": text, # alias for backward compatibility
                "slot_count": cell['slot_count'],
                "position": {
                    "row": cell['row'],
                    "col_start": cell['col_start'],
                    "col_end": cell['col_end']
                }
            })
            
        return slots

    @staticmethod
    def _save_debug_grid(img: np.ndarray, row_bounds: List[int], col_bounds: List[int]) -> None:
        """Save an image with drawn grid lines for debugging."""
        try:
            os.makedirs(DEBUG_DIR, exist_ok=True)
            debug_img = img.copy()
            
            for y in row_bounds:
                cv2.line(debug_img, (0, y), (img.shape[1], y), (0, 0, 255), 2)
            for x in col_bounds:
                cv2.line(debug_img, (x, 0), (x, img.shape[0]), (0, 255, 0), 2)
                
            cv2.imwrite(os.path.join(DEBUG_DIR, 'debug_grid.jpg'), debug_img)
        except Exception as e:
            logger.error(f"Error saving debug grid: {str(e)}")

    @staticmethod
    def _save_debug_data(filename: str, data: Any) -> None:
        try:
            os.makedirs(DEBUG_DIR, exist_ok=True)
            filepath = os.path.join(DEBUG_DIR, filename)
            with open(filepath, 'w', encoding='utf-8') as f:
                json.dump(data, f, indent=2, default=str)
        except Exception as e:
            logger.error(f"Error saving debug data: {str(e)}")

    @staticmethod
    def merge_slot_data_with_parsed_data(
        parsed_data: List[Dict[str, Any]],
        slot_data: Dict[str, Any]
    ) -> List[Dict[str, Any]]:
        """
        Merge slot count information with parsed routine data.
        Maintains backward compatibility.
        """
        if slot_data.get('error'):
            logger.warning(f"Slot data has error: {slot_data.get('error')}")
            return SlotCountService._apply_default_slots(parsed_data)
        
        try:
            slots = slot_data.get('slots', [])
            if not slots:
                slots = slot_data.get('structured_data', [])
                if not slots:
                    return SlotCountService._apply_default_slots(parsed_data)
            
            # Group by day string first
            day_groups = {}
            for struct_item in slots:
                if not isinstance(struct_item, dict): continue
                day = struct_item.get('day')
                if day:
                    if day not in day_groups:
                        day_groups[day] = []
                    day_groups[day].append(struct_item)
            
            # If OCR failed to extract distinct day names (e.g., only "Day" was found)
            # fallback to grouping by grid row to maintain day separation
            use_row_grouping = len(day_groups) <= 1
            row_groups = {}
            sorted_rows = []
            if use_row_grouping:
                for struct_item in slots:
                    if not isinstance(struct_item, dict): continue
                    pos = struct_item.get('position', {})
                    row = pos.get('row', pos.get('col_start')) # handle row or col depending on orientation
                    if row is not None:
                        if row not in row_groups:
                            row_groups[row] = []
                        row_groups[row].append(struct_item)
                sorted_rows = sorted(row_groups.keys())
            if isinstance(parsed_data, dict):
                if 'routine_data' in parsed_data:
                    routine_data = parsed_data.get('routine_data', {})
                    for day_idx, (day_name, entries) in enumerate(routine_data.items()):
                        if not isinstance(entries, list): continue
                        
                        ocr_entries = []
                        if use_row_grouping and day_idx < len(sorted_rows):
                            ocr_entries = row_groups[sorted_rows[day_idx]]
                        else:
                            ocr_entries = day_groups.get(day_name, [])
                            
                        fallback_entries = ocr_entries if ocr_entries else slots
                        
                        for idx, entry in enumerate(entries):
                            if not isinstance(entry, dict): continue
                            
                            # Try index matching if we have valid grouped entries
                            ocr_match = SlotCountService._find_match_entry(entry, fallback_entries)
                            if ocr_match:
                                entry['slot_count'] = ocr_match.get('slot_count', 1)
                                entry['slot'] = ocr_match.get('time_slot', 1) - 1
                                if 'position' in ocr_match:
                                    entry['position'] = ocr_match['position']
                            else:
                                entry['slot_count'] = entry.get('slot_count', 1)
                                entry['slot'] = entry.get('time_slot', 1) - 1
                                
                return parsed_data
            
            if isinstance(parsed_data, list):
                for entry in parsed_data:
                    if not isinstance(entry, dict): continue
                    day = entry.get('day')
                    if not day:
                        entry['slot_count'] = entry.get('slot_count', 1)
                        continue
                    
                    ocr_entries = day_groups.get(day, slots)
                    ocr_match = SlotCountService._find_match_entry(entry, ocr_entries)
                    
                    if ocr_match:
                        entry['slot_count'] = ocr_match.get('slot_count', 1)
                        entry['slot'] = ocr_match.get('time_slot', 1) - 1
                        if 'position' in ocr_match:
                            entry['position'] = ocr_match['position']
                    else:
                        entry['slot_count'] = entry.get('slot_count', 1)
                        entry['slot'] = entry.get('time_slot', 1) - 1
                return parsed_data
            
            return parsed_data
            
        except Exception as e:
            logger.error(f"Error in merge_slot_data_with_parsed_data: {str(e)}")
            return SlotCountService._apply_default_slots(parsed_data)
            
    @staticmethod
    def _apply_default_slots(parsed_data: Any) -> Any:
        if isinstance(parsed_data, dict):
            return parsed_data
        for entry in parsed_data:
            if isinstance(entry, dict) and 'slot_count' not in entry:
                entry['slot_count'] = 1
        return parsed_data

    @staticmethod
    def _find_match_entry(entry: Dict[str, Any], ocr_entries: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
        course = str(entry.get('course', '')).strip().upper()
        teacher = str(entry.get('teacher', '')).strip().upper()
        
        best_match = None
        best_score = 0
        
        for ocr_entry in ocr_entries:
            if not isinstance(ocr_entry, dict): continue
            
            ocr_text = str(ocr_entry.get('course', ocr_entry.get('text', ''))).strip().upper()
            score = 0
            
            if course and course in ocr_text:
                score += 10
            if teacher and teacher in ocr_text:
                score += 5
            if course and len(course) > 2 and course in ocr_text:
                score += 3
                
            if score > best_score:
                best_score = score
                best_match = ocr_entry
                
        return best_match
