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

# 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 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 fix_project_arn(project_arn):
    """
    Fix the project ARN format by ensuring it ends with the correct project ID.
    AWS Rekognition project ARNs must use the correct project ID: 1725357683732
    """
    if not project_arn:
        logger.error("No project ARN provided")
        return None
        
    logger.info(f"Checking project ARN format: {project_arn}")
    
    # Use the known working project ID 
    project_id = "1725357683732"
        
    # Check if the ARN already has a version number
    pattern = r'project\/[a-zA-Z0-9_\.-]{1,255}\/[0-9]+$'
    if re.search(pattern, project_arn):
        # If it matches the pattern but has a different ID, replace it
        if not project_arn.endswith(f"/{project_id}"):
            # Extract everything up to the last slash
            base_arn = '/'.join(project_arn.split('/')[:-1])
            fixed_arn = f"{base_arn}/{project_id}"
            logger.info(f"Updated project ID: {project_arn} -> {fixed_arn}")
            return fixed_arn
        else:
            logger.info(f"Project ARN format is already correct: {project_arn}")
            return project_arn
    
    # Get just the project name without any trailing slashes
    if '/project/' in project_arn:
        base_arn = project_arn.rstrip('/')
        # Standard case: append the correct project ID
        fixed_arn = f"{base_arn}/{project_id}"
        logger.info(f"Fixed ARN by adding correct project ID: {project_arn} -> {fixed_arn}")
        return fixed_arn
    
    logger.warning(f"Cannot determine correct ARN format from: {project_arn}")
    
    # Fall back to simple approach if more sophisticated extraction fails
    if project_arn.endswith('/'):
        fixed_arn = f"{project_arn}{project_id}"
    else:
        fixed_arn = f"{project_arn}/{project_id}"
        
    logger.info(f"Applied fallback ARN fix: {project_arn} -> {fixed_arn}")
    return fixed_arn

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 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)
        except Exception as e:
            logger.error(f"Failed to initialize parameters: {str(e)}")
            raise

def get_latest_model_version() -> Optional[Dict[str, Any]]:
    """
    Get the latest trained model version from Supabase.
    
    Returns:
        Dictionary containing model version information or None if not found
    """
    try:
        # First try to find models with TRAINING_COMPLETED status (original field)
        response = supabase.table('model_versions') \
            .select('*') \
            .eq('status', 'TRAINING_COMPLETED') \
            .order('training_timestamp', desc=True) \
            .limit(1) \
            .execute()
            
        if response.data and len(response.data) > 0:
            logger.info("Found latest trained model using status field")
            return response.data[0]
        
        # If no models found, try using model_status field if it exists
        try:
            alt_response = supabase.table('model_versions') \
                .select('*') \
                .eq('model_status', 'TRAINING_COMPLETED') \
                .order('training_timestamp', desc=True) \
                .limit(1) \
                .execute()
                
            if alt_response.data and len(alt_response.data) > 0:
                logger.info("Found latest trained model using model_status field")
                return alt_response.data[0]
        except Exception as inner_e:
            logger.warning(f"Error querying with model_status field: {str(inner_e)}")
        
        logger.info("No trained models found")
        return None
        
    except Exception as e:
        logger.error(f"Failed to get latest model version: {str(e)}")
        return None

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 start_model(model_version: Dict[str, Any]) -> str:
    """
    Start the Rekognition model without waiting for completion.
    
    Args:
        model_version: Dictionary containing model version information
        
    Returns:
        Status message
    """
    min_inference_units = 1

    try:
        # Get the model version ARN directly from the database record
        # We assume the ARN is already correctly formatted when stored in the database
        project_version_arn = model_version['project_version_arn']
        if not project_version_arn:
            raise ValueError("Missing project_version_arn in model record")
            
        # Start the model
        logger.info(f"Starting model: {project_version_arn}")
        response = rekognition_client.start_project_version(
            ProjectVersionArn=project_version_arn,
            MinInferenceUnits=min_inference_units
        )
        
        # Update status to STARTING
        update_model_status(model_version['version_name'], 'STARTING')
        
        return "Starting"
        
    except Exception as e:
        error_msg = str(e)
        logger.error(f"Error starting model: {error_msg}")
        # Update status to ERROR
        update_model_status(
            model_version['version_name'],
            'ERROR',
            error_msg
        )
        return error_msg

def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
    """
    Main Lambda handler function.
    
    Args:
        event: Lambda event
        context: Lambda context
        
    Returns:
        Dictionary containing result information
    """
    try:
        initialize_parameters()
        
        # Get latest trained model version
        model_version = get_latest_model_version()
        if not model_version:
            error_msg = "No trained model version found"
            logger.error(error_msg)
            return {
                'statusCode': 404,
                'body': error_msg
            }
        
        # Start the model without waiting for it to be running
        result = start_model(model_version)
        
        return {
            'statusCode': 200 if result == "Starting" else 500,
            'body': {
                'message': result,
                'version_name': model_version['version_name'],
                'status': 'STARTING',
                'project_version_arn': model_version['project_version_arn']
            }
        }
        
    except Exception as e:
        error_msg = str(e)
        logger.error(f"Lambda execution failed: {error_msg}")
        return {
            'statusCode': 500,
            'body': f"Internal server error: {error_msg}"
        }
        