import base64
import json
from datetime import datetime, timezone
import redis
import os
import asyncio
from hashlib import sha256
whilte_list_event_name = [
    'Audio-Player-Event', 
    'Audio-Creation-Event', 
    'Audio-Action-Event', 
    'Playlist-Action-Event', 
    'Link-Click-Event',
    'Web-Page-Event',
    'Artist-Action-Event', 
    'General-Event', 
    'Audio-Editing-Event', 
    'Account-Action-Event',
    'Search-Event',
    'App-Audio-Player-Event', 
    'App-Audio-Creation-Event', 
    'App-Audio-Action-Event', 
    'App-Playlist-Action-Event', 
    'App-Link-Click-Event',
    'App-Artist-Action-Event', 
    'App-General-Event', 
    'App-Audio-Editing-Event', 
    'App-Account-Action-Event',
    'App-Search-Event',
    'Web-User-Event',
    'Hook-Web-Event',
    ]
required_fields = ['event', 'properties', 'timestamp', 'context']
SESSION_IDENTITY_CACHE_KEY_PREFIX = "json-session-identity-"
ANALYTICS_PIPELINE = 1
BACKEND_SESSION = 2

#TODO: adding each event fields name and fields type validation logic
#      adding datahash check to remove duplicate event

def get_session_identity(
    user_uid: str, anonymous_id: str, platform: str, session_type: int = BACKEND_SESSION
) -> str:
    session_identity = platform + "-" + str(session_type) + "-" + anonymous_id + "-" + user_uid
    hash = sha256(session_identity.encode()).hexdigest()
    return hash

def is_whitelisted_event(json_value):
    return json_value['event'] in whilte_list_event_name

def is_required_fields_missing(json_value):
    for field in required_fields:
        if field not in json_value:
            return True
    return False
    

def lambda_handler(event, context):
    output = {'records': []}
    redis_host = os.environ.get('REDIS_HOST')
    redis_port = os.environ.get('REDIS_PORT')
    redis_db = os.environ.get('REDIS_DB')
    redis_client = redis.Redis(host=redis_host, port=int(redis_port), db=int(redis_db))
    session_keys = []
    session_id_dict = {}
    for record in event['records']:
        payload = base64.b64decode(record['data'])
        json_value = json.loads(payload)
        client_ip = ''
        request_time = ''
        domain_name = ''
        if 'request-time' in json_value:
            request_time = json_value['request-time']
        if 'domain-name' in json_value:
            domain_name = json_value['domain-name']
        if 'x-client-ip' in json_value:
            client_ip = json_value['x-client-ip']
            json_value = json.loads(base64.b64decode(json_value['request_body']))

        
        # dropped the event if not pass the validation check
        # TODO: deal with the dropped messages
        if is_required_fields_missing(json_value) or is_whitelisted_event(json_value) is False:
            firehose_record_output = {
                'recordId': record['recordId'],
                'result': 'Dropped'
            }
        else:
            current_hour = datetime.now(timezone.utc).hour
            current_day = datetime.now(timezone.utc).day
            current_month = datetime.now(timezone.utc).month
            current_year = datetime.now(timezone.utc).year
            
            partition_keys = {
                "EventName": json_value['event'],
                "EventVersion": 0,
                "year": current_year,
                "month": current_month,
                "date": current_day,
                "hour": current_hour,
            }
            
            result = {}
            ct = datetime.now(timezone.utc)
            result['server_timestamp'] = str(ct)
            result['client_ip'] = client_ip
            result['request_time'] = request_time
            result['domain_name'] = domain_name
            for field in required_fields:
                result[field] =  json_value[field]
            session_id = ''
            properties = json_value.get('properties')
            if properties:
                user_uid = properties.get('userId')
                anonymous_id = json_value.get('anonymousId')
                result['anonymous_id'] = anonymous_id
                platform = "web"
                if user_uid and anonymous_id:
                    session_identity = get_session_identity(
                        user_uid=str(user_uid),
                        anonymous_id=anonymous_id,
                        platform=platform,
                        session_type=BACKEND_SESSION,
                    )
                    session_identity_cache_key = SESSION_IDENTITY_CACHE_KEY_PREFIX + session_identity
                    session_keys.append(session_identity_cache_key)
                    # session_id = redis_client.get(session_identity_cache_key)
                    # if session_id:
                    #     session_id = session_id.decode('utf-8')
                    #     result['backend_session_id'] = session_id
                    session_id_dict[record['recordId']] = {
                        'recordId': record['recordId'],
                        'result': result,
                        'metadata': { 'partitionKeys': partition_keys },
                        'session_identity_cache_key': session_identity_cache_key
                    }
                else:
                    session_id_dict[record['recordId']] = {
                        'recordId': record['recordId'],
                        'result': result,
                        'metadata': { 'partitionKeys': partition_keys },
                        'session_identity_cache_key': '' 
                    }
            else:
                session_id_dict[record['recordId']] = {
                    'recordId': record['recordId'],
                    'result': result,
                    'metadata': { 'partitionKeys': partition_keys },
                    'session_identity_cache_key': '' 
                }
                print(f"properties not found in json_value: {json_value}")
            
            # firehose_record_output = {
            #     'recordId': record['recordId'],
            #     'data': base64.b64encode(json.dumps(result).encode('utf-8')).decode('utf-8'),
            #     'result': 'Ok',
            #     'metadata': { 'partitionKeys': partition_keys }
            # }

        
        # output['records'].append(firehose_record_output)
    session_values = redis_client.mget(session_keys)
    session_result_dict = {}
    for key, value in zip(session_keys, session_values):
        if value:
            session_result_dict[key] = value.decode('utf-8')

    for key, value in session_id_dict.items():
        if value['session_identity_cache_key'] in session_result_dict:
            value['result']['backend_session_id'] = session_result_dict[value['session_identity_cache_key']]
        output['records'].append({
            'recordId': value['recordId'],
            'data': base64.b64encode(json.dumps(value['result']).encode('utf-8')).decode('utf-8'),
            'result': 'Ok',
            'metadata': value['metadata']
        })
            
    
    return output
