import boto3
import logging
import os
import json
import time
import re
import subprocess
from datetime import datetime
from typing import Optional
from botocore.exceptions import ClientError

# Update imports to use dataset_utils for all dataset operations
# and only use rekognition_utils for manifest validation
from utils.rekognition_utils import validate_manifest_file
from utils.logging_utils import setup_logger, reduce_logging_verbosity
from utils.dataset_utils import (
    create_dataset_from_manifest,
    delete_dataset,
    wait_for_dataset_creation,
    list_datasets,
    check_delete_create_dataset,
    force_delete_project_datasets,
    get_correct_dataset_arn
)

# Initialize logger
logger = setup_logger()
reduce_logging_verbosity()

# Initialize clients
ssm_client = boto3.client('ssm')
s3_client = boto3.client('s3')

def get_ssm_parameter(param_name):
    """Get a parameter from SSM Parameter Store"""
    try:
        response = ssm_client.get_parameter(Name=param_name, WithDecryption=True)
        return response['Parameter']['Value']
    except Exception as e:
        logger.warning(f"Failed to get SSM parameter {param_name}: {str(e)}")
        return None

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 force_delete_dataset_using_cli(project_arn, dataset_type):
    """
    Force delete a dataset using AWS CLI directly, which sometimes works when the SDK methods fail.
    
    Args:
        project_arn: ARN of the project
        dataset_type: 'TRAIN' or 'TEST'
    
    Returns:
        bool: True if deletion was successful, False otherwise
    """
    try:
        # Use the proper dataset ARN format
        dataset_arn = get_correct_dataset_arn(project_arn, dataset_type)
        
        if not dataset_arn:
            logger.error(f"Invalid project ARN format for CLI deletion: {project_arn}")
            return False
            
        # Extract region for CLI command
        match = re.match(r'arn:aws:rekognition:([^:]+):', dataset_arn)
        if not match:
            logger.error(f"Could not extract region from dataset ARN: {dataset_arn}")
            return False
            
        region = match.group(1)
        
        logger.info(f"Attempting to delete dataset using AWS CLI: {dataset_arn}")
        
        # Set AWS CLI command
        cmd = [
            "aws", "rekognition", "delete-dataset",
            "--dataset-arn", dataset_arn,
            "--region", region
        ]
        
        logger.info(f"Running AWS CLI command: {' '.join(cmd)}")
        process = subprocess.run(cmd, capture_output=True, text=True)
        
        if process.returncode == 0:
            logger.info(f"Successfully deleted dataset using AWS CLI: {dataset_arn}")
            return True
        else:
            logger.error(f"AWS CLI deletion failed: {process.stderr}")
            return False
    except Exception as e:
        logger.error(f"Error in CLI dataset deletion: {str(e)}")
        return False

def delete_dataset_by_type(rekognition_client, project_arn, dataset_type):
    """
    Deletes a dataset by type (TRAIN or TEST) following Grok's simplified approach.
    
    Args:
        rekognition_client: Boto3 Rekognition client
        project_arn: ARN of the Rekognition project
        dataset_type: Type of dataset (TRAIN or TEST)
        
    Returns:
        dict: Result with status and message
    """
    try:
        logger.info(f"Deleting {dataset_type} dataset for project: {project_arn}")
        
        # Make sure dataset_type is upper case for API call
        dataset_type = dataset_type.upper()
        
        # Use the proper dataset ARN format
        dataset_arn = get_correct_dataset_arn(project_arn, dataset_type)
        
        if not dataset_arn:
            return {
                "success": False,
                "message": f"Could not generate a valid dataset ARN for {dataset_type}",
                "error_code": "InvalidARNFormat" 
            }
        
        logger.info(f"Using properly formatted dataset ARN: {dataset_arn}")
        
        # Use the DatasetArn parameter instead of ProjectArn and DatasetType
        try:
            response = rekognition_client.delete_dataset(
                DatasetArn=dataset_arn
            )
            logger.info(f"Successfully deleted {dataset_type} dataset using SDK")
            return {
                "success": True,
                "message": f"Successfully deleted {dataset_type} dataset",
                "response": response
            }
        except ClientError as e:
            error_code = e.response['Error']['Code']
            error_message = e.response['Error']['Message']
            logger.warning(f"SDK deletion failed: {error_code} - {error_message}")
            
            # ResourceNotFoundException is fine - dataset already gone
            if error_code == 'ResourceNotFoundException':
                return {
                    "success": True,
                    "message": f"{dataset_type} dataset does not exist or is already deleted"
                }
                
            # Try using our CLI method as fallback
            cli_success = force_delete_dataset_using_cli(project_arn, dataset_type)
            if cli_success:
                return {
                    "success": True,
                    "message": f"Successfully deleted {dataset_type} dataset using CLI fallback"
                }
                
            # If both approaches failed
            return {
                "success": False,
                "message": f"Failed to delete {dataset_type} dataset: {error_message}",
                "error_code": error_code
            }
    except Exception as e:
        logger.error(f"Error deleting dataset by type: {str(e)}", exc_info=True)
        return {
            "success": False,
            "message": f"Error: {str(e)}"
        }

def lambda_handler(event, context):
    """
    Lambda function to manage datasets for AWS Rekognition Custom Labels.
    Now focused only on dataset creation - deletion is handled by a dedicated function.
    
    Args:
        event (dict): Lambda event object containing optional parameters
        context (object): Lambda context
        
    Returns:
        dict: Operation result with status and details
    """
    try:
        logger.info(f"Received event: {json.dumps(event)}")
        
        # Get environment variables
        training_bucket = os.environ.get('TRAINING_BUCKET')
        project_arn = os.environ.get('PROJECT_ARN')
        
        # Try to get from SSM if not in environment
        if not training_bucket:
            training_bucket = get_ssm_parameter('/datafy-rekognition-stack/training-bucket')
            if not training_bucket:
                return {
                    'statusCode': 500,
                    'body': json.dumps({
                        'message': 'TRAINING_BUCKET not found in environment or SSM',
                        'error': 'Missing required configuration'
                    })
                }
            logger.info(f"Got training bucket from SSM: {training_bucket}")
        else:
            logger.info(f"Using TRAINING_BUCKET from environment: {training_bucket}")
        
        if not project_arn:
            project_arn = get_ssm_parameter('/datafy-rekognition-stack/project-name')
            if project_arn:
                # Construct the ARN using the project name
                aws_region = os.environ.get('AWS_REGION') or 'eu-west-1'
                aws_account_id = os.environ.get('AWS_ACCOUNT_ID') or '587594388832'
                project_arn = f"arn:aws:rekognition:{aws_region}:{aws_account_id}:project/{project_arn}"
                logger.info(f"Constructed project ARN from SSM: {project_arn}")
            else:
                return {
                    'statusCode': 500,
                    'body': json.dumps({
                        'message': 'PROJECT_ARN not found in environment or SSM',
                        'error': 'Missing required configuration'
                    })
                }
        
        # Fix project ARN format
        project_arn = fix_project_arn(project_arn)
        if not project_arn:
            return {
                'statusCode': 500,
                'body': json.dumps({
                    'message': 'Invalid project ARN format',
                    'error': 'Configuration error'
                })
            }
        logger.info(f"Using project ARN: {project_arn}")
        
        # Initialize AWS clients
        rekognition = boto3.client('rekognition')
        
        # For deletion operations, redirect to the specialized dataset deletion function
        if event.get('operation') == 'delete_dataset' or event.get('operation') == 'force_delete_dataset':
            # Remove API Gateway dependency and call Lambda directly
            lambda_client = boto3.client('lambda')
            
            logger.info(f"Forwarding deletion request to dedicated deletion function")
            
            # Call dedicated deletion Lambda directly
            try:
                deletion_payload = {
                    'project_arn': project_arn,
                    'dataset_type': event.get('dataset_type'),
                    'force': event.get('operation') == 'force_delete_dataset'
                }
                
                # Get function name from environment or use default
                delete_function_name = os.environ.get('DELETE_DATASETS_FUNCTION') or f"{os.environ.get('PROJECT_NAME', 'datafy-rekognition')}-delete-datasets"
                logger.info(f"Invoking deletion function: {delete_function_name}")
                
                response = lambda_client.invoke(
                    FunctionName=delete_function_name,
                    InvocationType='RequestResponse',
                    Payload=json.dumps(deletion_payload)
                )
                
                # Read and parse the response
                payload = json.loads(response['Payload'].read())
                logger.info(f"Deletion function response: {payload}")
                
                return {
                    'statusCode': 200,
                    'body': json.dumps({
                        'message': 'Dataset deletion request processed',
                        'deletion_result': payload
                    })
                }
            except Exception as e:
                logger.error(f"Error forwarding to deletion function: {str(e)}")
                return {
                    'statusCode': 500,
                    'body': json.dumps({
                        'message': f"Error forwarding deletion request: {str(e)}",
                        'error': 'Internal server error'
                    })
                }
        
        # Use default manifest path for creation
        manifest_key = "manifest.jsonl"
        manifest_s3_uri = f"s3://{training_bucket}/{manifest_key}"
        logger.info(f"Using manifest path: {manifest_s3_uri}")
        
        # Parse S3 URI
        bucket_name = manifest_s3_uri.split('/')[2]
        key = '/'.join(manifest_s3_uri.split('/')[3:])
        
        # Validate the manifest file
        logger.info(f"Validating manifest file: {manifest_s3_uri}")
        if not validate_manifest_file(bucket_name, key):
            return {
                'statusCode': 400,
                'body': json.dumps({
                    'message': 'Manifest validation failed',
                    'error': 'Invalid manifest format or content',
                    'manifestFile': manifest_s3_uri
                })
            }
        
        logger.info(f"Manifest file validated successfully: {manifest_s3_uri}")
        
        # Delete existing datasets first (using the dedicated function)
        lambda_client = boto3.client('lambda')
        
        try:
            logger.info("Deleting existing datasets before creation")
            deletion_payload = {
                'project_arn': project_arn,
                'dataset_types': ['TRAIN', 'TEST']
            }
            
            # Get function name from environment or use default
            delete_function_name = os.environ.get('DELETE_DATASETS_FUNCTION', 'DeleteDatasetsFunction')
            
            delete_response = lambda_client.invoke(
                FunctionName=delete_function_name,
                InvocationType='RequestResponse',
                Payload=json.dumps(deletion_payload)
            )
            
            payload = json.loads(delete_response['Payload'].read())
            logger.info(f"Dataset deletion result: {payload}")
            
            # Wait a moment to ensure deletion has been processed
            time.sleep(5)
        except Exception as e:
            logger.warning(f"Could not delete existing datasets: {str(e)}")
            # Continue with creation even if deletion fails
        
        # Create both datasets fresh
        train_result = create_dataset_from_manifest(
            rekognition,
            project_arn,
            'TRAIN',
            manifest_s3_uri
        )
        
        if not train_result.get('success'):
            return {
                'statusCode': 500,
                'body': json.dumps({
                    'message': 'Failed to create TRAIN dataset',
                    'datasetArn': train_result.get('dataset_arn'),
                    'error': train_result.get('error', 'Unknown error'),
                    'manifestFile': manifest_s3_uri
                })
            }
        
        # Create test dataset
        test_result = create_dataset_from_manifest(
            rekognition,
            project_arn,
            'TEST',
            manifest_s3_uri
        )
        
        # Even if test dataset creation fails, we can still return a success for the train dataset
        test_success = test_result.get('success', False)
        
        return {
            'statusCode': 200,
            'body': json.dumps({
                'message': 'Dataset creation completed',
                'trainDatasetArn': train_result['dataset_arn'],
                'testDatasetArn': test_result.get('dataset_arn'),
                'testSuccess': test_success,
                'manifestFile': manifest_s3_uri
            })
        }
        
    except Exception as e:
        logger.error(f"Error in lambda_handler: {str(e)}", exc_info=True)
        return {
            'statusCode': 500,
            'body': json.dumps({
                'message': f'Internal error: {str(e)}',
                'error': 'Unexpected exception occurred'
            })
        } 