import boto3
import logging
import os
import re
from typing import Dict, Any, Optional
from supabase import create_client, Client
from botocore.exceptions import ClientError
import json
import time
import botocore

# Configure logging
logger = logging.getLogger()
logger.setLevel(logging.INFO)

# Initialize clients
ssm = boto3.client('ssm')
rekognition_client = boto3.client('rekognition')

# Initialize global variables
SUPABASE_URL: Optional[str] = None
SUPABASE_KEY: Optional[str] = None
supabase: Optional[Client] = None
PROJECT_ARN: Optional[str] = None

def fix_project_arn(project_arn):
    """
    Ensure the project ARN ends with the correct project ID.
    The project ID should be '1725357683732'.
    
    Args:
        project_arn (str): The project ARN to check and fix if needed.
        
    Returns:
        str: The corrected project ARN.
    """
    try:
        logger.info(f"Checking project ARN format: {project_arn}")
        
        # Project ID should be '1725357683732'
        PROJECT_ID = '1725357683732'
        
        # Check if the ARN is valid
        if not project_arn or not isinstance(project_arn, str):
            logger.error("Invalid project ARN format. Using default project ID.")
            # Create a default ARN format if the input is invalid
            return f"arn:aws:rekognition:us-east-1:{os.environ.get('AWS_ACCOUNT_ID', '')}:project/{PROJECT_ID}"
            
        # Check if the ARN already ends with the correct project ID
        if project_arn.endswith(PROJECT_ID):
            logger.info(f"Project ARN already has correct format: {project_arn}")
            return project_arn
            
        # Use regex to extract parts of the ARN
        arn_pattern = r"arn:aws:rekognition:([^:]+):([^:]+):project/([^/]+)"
        match = re.match(arn_pattern, project_arn)
        
        if match:
            region, account_id, _ = match.groups()
            corrected_arn = f"arn:aws:rekognition:{region}:{account_id}:project/{PROJECT_ID}"
            logger.info(f"Corrected project ARN: {corrected_arn}")
            return corrected_arn
        else:
            # If the ARN doesn't match the expected pattern, construct a new one
            logger.warning(f"ARN pattern doesn't match expected format: {project_arn}")
            region = os.environ.get('AWS_REGION', 'us-east-1')
            account_id = os.environ.get('AWS_ACCOUNT_ID', '')
            corrected_arn = f"arn:aws:rekognition:{region}:{account_id}:project/{PROJECT_ID}"
            logger.info(f"Using default project ARN: {corrected_arn}")
            return corrected_arn
    except Exception as e:
        logger.error(f"Error fixing project ARN: {str(e)}")
        # Return the original ARN if there's an error
        return project_arn

def get_parameter(name: str) -> str:
    """
    Retrieve a parameter from AWS Parameter Store.
    
    Args:
        name: The parameter name to retrieve
        
    Returns:
        The parameter value
        
    Raises:
        ClientError: If parameter retrieval fails
    """
    try:
        response = ssm.get_parameter(Name=name, WithDecryption=True)
        return response['Parameter']['Value']
    except ClientError as e:
        logger.error(f"Failed to get parameter {name}: {str(e)}")
        raise

def initialize_parameters() -> None:
    """Initialize global parameters and clients."""
    global SUPABASE_URL, SUPABASE_KEY, supabase, PROJECT_ARN
    
    if SUPABASE_URL is None or SUPABASE_KEY is None or PROJECT_ARN is None:
        try:
            SUPABASE_URL = get_parameter('/supabase/url')
            SUPABASE_KEY = get_parameter('/supabase/anon')
            project_arn_raw = get_parameter('/datafy-rekognition-stack/project-arn')
            
            # Fix the project ARN
            PROJECT_ARN = fix_project_arn(project_arn_raw)
            if not PROJECT_ARN:
                # Try to build from project name
                project_name = get_parameter('/datafy-rekognition-stack/project-name')
                if project_name:
                    aws_region = boto3.session.Session().region_name
                    aws_account_id = boto3.client('sts').get_caller_identity().get('Account')
                    project_arn_constructed = f"arn:aws:rekognition:{aws_region}:{aws_account_id}:project/{project_name}"
                    PROJECT_ARN = fix_project_arn(project_arn_constructed)
            
            if not PROJECT_ARN:
                raise ValueError("Could not determine a valid project ARN")
                
            logger.info(f"Using project ARN: {PROJECT_ARN}")
            supabase = create_client(SUPABASE_URL, SUPABASE_KEY)
            logger.info("Successfully initialized parameters and created Supabase client")
        except Exception as e:
            logger.error(f"Failed to initialize parameters: {str(e)}")
            raise

def get_active_model_version() -> Optional[Dict[str, Any]]:
    """
    Get the latest active model version from Supabase.
    First checks for RUNNING models, then for STARTING models.
    
    Returns:
        Dictionary containing model version information or None if not found
    """
    try:
        logger.info("Looking for active models (RUNNING or STARTING)")
        
        # First try using model_status field (new approach)
        try:
            # First look for models with RUNNING status
            running_response = supabase.table('model_versions') \
                .select('*') \
                .eq('model_status', 'RUNNING') \
                .order('training_timestamp', desc=True) \
                .limit(1) \
                .execute()
                
            if running_response.data and len(running_response.data) > 0:
                model = running_response.data[0]
                logger.info(f"Found RUNNING model using model_status field: {model.get('version_name')}")
                return model
            
            # Then look for models with STARTING status
            starting_response = supabase.table('model_versions') \
                .select('*') \
                .eq('model_status', 'STARTING') \
                .order('training_timestamp', desc=True) \
                .limit(1) \
                .execute()
                
            if starting_response.data and len(starting_response.data) > 0:
                model = starting_response.data[0]
                logger.info(f"Found STARTING model using model_status field: {model.get('version_name')}")
                return model
        except Exception as inner_e:
            logger.warning(f"Error querying with model_status field: {str(inner_e)}")
        
        # Fall back to status field (legacy approach)
        logger.info("Falling back to legacy status field")
        
        # First look for models with RUNNING status
        running_response = supabase.table('model_versions') \
            .select('*') \
            .eq('status', 'RUNNING') \
            .order('training_timestamp', desc=True) \
            .limit(1) \
            .execute()
            
        if running_response.data and len(running_response.data) > 0:
            model = running_response.data[0]
            logger.info(f"Found RUNNING model using legacy status field: {model.get('version_name')}")
            return model
        
        # Then look for models with STARTING status
        starting_response = supabase.table('model_versions') \
            .select('*') \
            .eq('status', 'STARTING') \
            .order('training_timestamp', desc=True) \
            .limit(1) \
            .execute()
            
        if starting_response.data and len(starting_response.data) > 0:
            model = starting_response.data[0]
            logger.info(f"Found STARTING model using legacy status field: {model.get('version_name')}")
            return model
            
        # If no active models, check for TRAINING_COMPLETED
        logger.info("No active models found, checking for TRAINING_COMPLETED models")
        
        # Try both model_status and status fields for TRAINING_COMPLETED
        try:
            completed_response = supabase.table('model_versions') \
                .select('*') \
                .eq('model_status', 'TRAINING_COMPLETED') \
                .order('training_timestamp', desc=True) \
                .limit(1) \
                .execute()
                
            if completed_response.data and len(completed_response.data) > 0:
                model = completed_response.data[0]
                logger.info(f"Found latest completed model using model_status: {model.get('version_name')}")
                return model
        except Exception as inner_e:
            logger.warning(f"Error querying with model_status field for TRAINING_COMPLETED: {str(inner_e)}")
        
        # Fallback to legacy status field for TRAINING_COMPLETED
        completed_response = supabase.table('model_versions') \
            .select('*') \
            .eq('status', 'TRAINING_COMPLETED') \
            .order('training_timestamp', desc=True) \
            .limit(1) \
            .execute()
            
        if completed_response.data and len(completed_response.data) > 0:
            model = completed_response.data[0]
            logger.info(f"Found latest completed model using legacy status: {model.get('version_name')}")
            return model
            
        logger.warning("No active or completed models found")
        return None
        
    except Exception as e:
        logger.error(f"Error getting active model version: {str(e)}")
        return None

def check_model_status(project_version_arn: str) -> Dict[str, Any]:
    """
    Check the status of a model in AWS Rekognition.
    
    Args:
        project_version_arn: The ARN of the model version to check
        
    Returns:
        Dictionary with model status information
    """
    try:
        logger.info(f"Checking model status for: {project_version_arn}")
        
        # Extract project ARN and version name from the version ARN
        if '/version/' not in project_version_arn:
            logger.error(f"Invalid project version ARN format: {project_version_arn}")
            return {'is_running': False, 'status': 'UNKNOWN', 'error': 'Invalid ARN format'}
            
        # Split into project ARN and version info
        parts = project_version_arn.split('/version/')
        base_project_arn = parts[0]
        version_info = parts[1].split('/')
        version_name = version_info[0]
        
        # Ensure project ARN has correct project ID format
        project_id = "1725357683732"
        project_arn = f"{base_project_arn}/{project_id}"
        
        logger.info(f"Extracted project ARN: {project_arn}")
        logger.info(f"Extracted version name: {version_name}")
        
        # Use the same approach as in the check_training_job_status function
        try:
            # Call describe_project_versions with ProjectArn and VersionNames
            response = rekognition_client.describe_project_versions(
                ProjectArn=project_arn,
                VersionNames=[version_name]
            )
            
            # Check response for valid data
            if not response.get('ProjectVersionDescriptions'):
                logger.warning(f"No version descriptions found for {version_name}")
                return {
                    'is_running': False,
                    'status': 'NOT_FOUND',
                    'error': f"No version found with name {version_name}",
                    'version_name': version_name,
                    'project_arn': project_arn
                }
            
            # Extract status information
            version_info = response['ProjectVersionDescriptions'][0]
            aws_status = version_info.get('Status', 'UNKNOWN')
            
            # Log full response for troubleshooting
            logger.info(f"AWS Rekognition status for {version_name}: {aws_status}")
            logger.info(f"Status details: {json.dumps(version_info, default=str)[:200]}...")
            
            # Determine if model is running
            is_running = aws_status == 'RUNNING'
            
            # Return comprehensive status information
            return {
                'is_running': is_running,
                'status': aws_status,
                'status_message': version_info.get('StatusMessage'),
                'version_name': version_name,
                'project_arn': project_arn,
                'min_inference_units': version_info.get('MinInferenceUnits'),
                'created': version_info.get('CreationTimestamp')
            }
        
        except Exception as e:
            logger.error(f"Error calling describe_project_versions: {str(e)}")
            # Check if this is an ARN formatting issue
            if 'ValidationException' in str(e) and 'ProjectArn' in str(e):
                logger.warning("ARN format validation error. Trying alternative format.")
                # Try a different approach to formatting project ARN
                try:
                    # Extract project name from ARN 
                    project_name_match = re.search(r'project/([^/]+)', base_project_arn)
                    if project_name_match:
                        project_name = project_name_match.group(1)
                        # List all projects to find the correct ARN
                        list_response = rekognition_client.describe_projects()
                        for project in list_response.get('ProjectDescriptions', []):
                            if project_name in project.get('ProjectArn', ''):
                                correct_project_arn = project.get('ProjectArn')
                                logger.info(f"Found correct project ARN: {correct_project_arn}")
                                
                                # Try again with the correct project ARN
                                retry_response = rekognition_client.describe_project_versions(
                                    ProjectArn=correct_project_arn,
                                    VersionNames=[version_name]
                                )
                                
                                if retry_response.get('ProjectVersionDescriptions'):
                                    version_info = retry_response['ProjectVersionDescriptions'][0]
                                    aws_status = version_info.get('Status', 'UNKNOWN')
                                    logger.info(f"Successfully retrieved status with corrected ARN: {aws_status}")
                                    
                                    return {
                                        'is_running': aws_status == 'RUNNING',
                                        'status': aws_status,
                                        'status_message': version_info.get('StatusMessage'),
                                        'version_name': version_name,
                                        'project_arn': correct_project_arn
                                    }
                except Exception as inner_e:
                    logger.error(f"Error in alternate ARN approach: {str(inner_e)}")
            
            # Return error status if all approaches fail
            return {
                'is_running': False,
                'status': 'ERROR',
                'error': str(e),
                'version_name': version_name,
                'project_arn': project_arn
            }
            
    except Exception as e:
        logger.error(f"Error in check_model_status: {str(e)}")
        return {'is_running': False, 'status': 'ERROR', 'error': str(e)}

def fix_model_version_arn(model_version_arn):
    """
    Fix the model version ARN format by ensuring it has the correct project ID.
    
    Args:
        model_version_arn: The model version ARN to fix
        
    Returns:
        The fixed model version ARN
    """
    if not model_version_arn:
        logger.error("No model version ARN provided")
        return None
        
    logger.info(f"Checking model version ARN format: {model_version_arn}")
    
    # Use the known working project ID 
    project_id = "1725357683732"
    
    # Check if the ARN already has the correct format with project ID at the end
    pattern = r'project\/[^\/]+\/version\/[^\/]+\/[0-9]+$'
    if re.search(pattern, model_version_arn):
        # If it matches the pattern but has a different ID, replace it
        parts = model_version_arn.split('/')
        if parts[-1] != project_id:
            parts[-1] = project_id
            fixed_arn = '/'.join(parts)
            logger.info(f"Updated model version ARN project ID: {model_version_arn} -> {fixed_arn}")
            return fixed_arn
        return model_version_arn
        
    # If the ARN doesn't match expected pattern, attempt to fix it
    if '/project/' in model_version_arn and '/version/' in model_version_arn:
        # Extract the base part and version name
        match = re.match(r'(.*\/project\/[^\/]+)\/version\/([^\/]+)(?:\/.*)?', model_version_arn)
        if match:
            base_part = match.group(1)
            version_name = match.group(2)
            fixed_arn = f"{base_part}/version/{version_name}/{project_id}"
            logger.info(f"Reconstructed model version ARN: {model_version_arn} -> {fixed_arn}")
            return fixed_arn
    
    logger.warning(f"Could not fix model version ARN format: {model_version_arn}")
    return model_version_arn

def update_model_status(version_name: str, status: str, status_message: str = None) -> None:
    """
    Update model version status in Supabase.
    
    Args:
        version_name: Name of the model version
        status: New status to set
        status_message: Optional status message
    """
    try:
        # Use model_status instead of status to avoid breaking existing functionality
        update_data = {'model_status': status}
        if status_message:
            update_data['status_message'] = status_message
            
        logger.info(f"Updating model_status for {version_name} to {status}")
        
        supabase.table('model_versions') \
            .update(update_data) \
            .eq('version_name', version_name) \
            .execute()
            
    except Exception as e:
        logger.error(f"Failed to update model status: {str(e)}")

def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
    """
    Main Lambda handler function.
    
    Steps:
    1. Get the latest active model (RUNNING, STARTING, or TRAINING_COMPLETED)
    2. Check if the model is running in AWS
    3. Return model_available status for Step Function to proceed or wait
    """
    try:
        initialize_parameters()
        
        # Step 1: Get the latest active model
        model = get_active_model_version()
        if not model:
            logger.warning("No active or completed models found")
            response = {
                'model_available': False,
                'message': 'No active or completed models found'
            }
            logger.info(f"RESPONSE: {json.dumps(response)}")
            return response
        
        # Step 2: Get the project version ARN
        project_version_arn = model.get('project_version_arn')
        if not project_version_arn:
            logger.error("Model is missing project_version_arn")
            response = {
                'model_available': False,
                'message': 'Model is missing required ARN information'
            }
            logger.info(f"RESPONSE: {json.dumps(response)}")
            return response
            
        # Fix the project version ARN if needed
        project_version_arn = fix_model_version_arn(project_version_arn)
        
        # Step 3: Check the model status
        status_info = check_model_status(project_version_arn)
        is_running = status_info.get('is_running', False)
        aws_status = status_info.get('status', 'UNKNOWN')
        
        logger.info(f"Model {model.get('version_name')} status: {aws_status}, is_running: {is_running}")
        
        # If the model is running in AWS but not marked as running in the database, update it
        if is_running and model.get('model_status') != 'RUNNING' and model.get('status') != 'RUNNING':
            try:
                logger.info(f"Updating model status in database to RUNNING for {model.get('version_name')}")
                update_model_status(model.get('version_name'), 'RUNNING')
            except Exception as update_error:
                logger.warning(f"Failed to update model status: {str(update_error)}")
        
        # Return result for Step Function - IMPORTANT: no nested "body" object
        response = {
            'model_available': bool(is_running),
            'message': f'Model is {aws_status}',
            'version_name': model.get('version_name'),
            'aws_status': aws_status,
            'project_version_arn': project_version_arn
        }
        logger.info(f"RESPONSE: {json.dumps(response)}")
        return response
        
    except Exception as e:
        error_msg = str(e)
        logger.error(f"Error checking model availability: {error_msg}")
        response = {
            'model_available': False,
            'message': f'Error checking model availability: {error_msg}'
        }
        logger.info(f"ERROR RESPONSE: {json.dumps(response)}")
        return response 