import datetime
import json
import re
import traceback
import requests

from Crypto.Cipher import AES
from Crypto.Util.Padding import pad, unpad
import base64
from django.conf import settings
from django.db import connection, transaction, models
from django.db.models.functions import Cast, Replace
from django.forms.models import model_to_dict
from rest_framework.exceptions import APIException

from api.models.error_log import ErrorLog
from django.apps import apps
from django.db.models import ForeignKey, Value
import hashlib

def encrypt_aes_128_cbc(text: str) -> str:
    cipher = AES.new(settings.PWD_ENCRYPTION_KEY, AES.MODE_CBC, settings.PWD_ENCRYPTION_IV)
    padded = pad(text.encode(), AES.block_size)
    return base64.b64encode(cipher.encrypt(padded)).decode()

def decrypt_aes_128_cbc(enc_text: str) -> str:
    cipher = AES.new(settings.PWD_ENCRYPTION_KEY, AES.MODE_CBC, settings.PWD_ENCRYPTION_IV)
    decrypted = unpad(cipher.decrypt(base64.b64decode(enc_text)), AES.block_size)
    return decrypted.decode()

def get_ip(request):
    x_forwarded_for = request.META.get('HTTP_X_FORWARDED_FOR')
    if x_forwarded_for:
        # Take first IP from comma-separated list
        return x_forwarded_for.split(',')[0].strip()
    return request.META.get('REMOTE_ADDR')

def throw_validation_error(code=None, error=None):
    e = APIException()
    e.status_code = code
    e.detail = {
        "code": code,
        "error": error
    }
    raise e

def log_critical_error(user_id=None, descr="", url="", trace: list = None, data: list = None):
    if isinstance(trace, (list, dict)):
        trace = json.dumps(trace)
    if isinstance(data, (list, dict)):
        data = json.dumps(data)

    ErrorLog.insert_log(user_id, descr, trace, data,url)


def get_exception_detail(e):
    descr = str(e)
    url = e.__traceback__.tb_frame.f_code.co_filename
    trace = "".join(traceback.format_exception(type(e), e, e.__traceback__))

    return {
        "descr": descr,
        "url": url,
        "trace": trace
    }

def get_old_new_data(model, old_data, ignore_fields=None):
    """
    Compare current instance with old_data snapshot.
    old_data must use model_to_dict be4 pass in
    """
    ignore_fields = ignore_fields or []
    changes = {"old_data": {}, "new_data": {}}

    # Get current state
    new_data = model_to_dict(model)

    # combine both old and new data key, so even old data is not set, but new data have set, it also can store the value
    for field in set(old_data.keys()) | set(new_data.keys()):
        if field in ignore_fields:
            continue

        old_value = old_data.get(field, "")
        new_value = new_data.get(field, "")

        if old_value != new_value:
            # make sure is not in datetime/date/time format in the value, if not when store to json field, will have error
            if isinstance(old_value, (datetime.date, datetime.time, datetime.datetime)):
                old_value = old_value.isoformat()
            if isinstance(new_value, (datetime.date, datetime.time, datetime.datetime)):
                new_value = new_value.isoformat()
            changes["old_data"][field] = old_value
            changes["new_data"][field] = new_value

    return changes

def get_multiple_old_new_data(models, request_data, ignore_field, valid_fields):
    old_new_data = {
        "old_data": {},
        "new_data": {}
    }

    ignore_field.append("password")

    for model in models:
        changes_new = {}
        changes_old = {}
        for key, new_val in request_data:
            if key not in ignore_field and key in valid_fields:
                old_val = getattr(model, key)
                if old_val != new_val:  # only log changed fields
                    # make sure is not in datetime/date/time format in the value, if not when store to json field, will have error
                    if isinstance(old_val, (datetime.date, datetime.time, datetime.datetime)):
                        old_val = old_val.isoformat()
                    if isinstance(new_val, (datetime.date, datetime.time, datetime.datetime)):
                        new_val = new_val.isoformat()
                    changes_old[key] = old_val
                    changes_new[key] = new_val

        if changes_new:
            old_new_data["old_data"][model.id] = changes_old
            old_new_data["new_data"][model.id] = changes_new

    return old_new_data

def get_add_new_data(model,skip_empty=True):
    """
    Return audit data when inserting a new record
    """

    #exclude many to many field, cause i need get, will run extra query to perform it
    exclude_fields = [f.name for f in model._meta.many_to_many]

    data = model_to_dict(model, exclude=exclude_fields)

    if skip_empty:
        data = {
            key: val for key, val in data.items()
            if val not in (None, "", [], {})
        }

    return {
        "table": model._meta.db_table,
        "data": data,
    }

def search_old_new_data(old_new_data, search_key):
    """
    mostly used for find out the id need update to redis
    :param old_new_data: data form by get_multiple_old_new_data or get_old_new_data
    :param search_key: the key want to get
    """
    # using set() can get unique id, if direct use list, need extra handle to filter for unique id
    ids = set()
    # check old and new data, combine the suitability id and get relation from db
    if old_new_data['old_data']:
        for key, obj in old_new_data['old_data'].items():
            if search_key in obj:
                val = obj.get(search_key, [])
                if isinstance(val, (list, tuple, set)):
                    ids.update(val)
                elif val is not None:  # int, str, etc.
                    ids.add(val)

    if old_new_data['new_data']:
        for key, obj in old_new_data['new_data'].items():
            if search_key in obj:
                val = obj.get(search_key, [])
                if isinstance(val, (list, tuple, set)):
                    ids.update(val)
                elif val is not None:  # int, str, etc.
                    ids.add(val)

    return list(ids)

def bulk_sync_to_redis(redis_client, data, chunk_size=500):
    """
    :param redis_client: Redis client instance
    :param data: dict with keys "insert", "update", "delete"
    :param chunk_size: number of commands per pipeline execution
    """
    try:
        # record down all redis command need to run
        commands = []
        # Prepare insert commands
        for model_name, objects in data.get("insert", {}).items():
            for obj in objects:
                key = f"{model_name}:{obj['id']}"
                commands.append(("json_set", key, "$", obj, {}))

        # Prepare update commands
        for model_name, objects in data.get("update", {}).items():
            for obj in objects:
                key = f"{model_name}:{obj['id']}"
                for field, value in obj.items():
                    if field != "id":
                        commands.append(("json_set", key, f"$.{field}", value, {"xx": True}))

        # Prepare delete commands
        for model_name, ids in data.get("delete", {}).items():
            for obj_id in ids:
                key = f"{model_name}:{obj_id}"
                commands.append(("delete", key))

        # Execute commands in chunks
        for i in range(0, len(commands), chunk_size):
            with redis_client.pipeline() as pipe:
                for cmd in commands[i:i+chunk_size]:
                    if cmd[0] == "json_set":
                        pipe.json().set(cmd[1], cmd[2], cmd[3], **cmd[4])
                    elif cmd[0] == "delete":
                        pipe.delete(cmd[1])
                pipe.execute()

    except Exception as e:
        # Log the error
        e_details = get_exception_detail(e)
        log_critical_error(user_id=None, descr=e_details['descr'], url=e_details['url'], trace=e_details['trace'])

def calculate_preset_pattern_time(type,pattern,minute_per_slot,slot_per_day):
    """
    curently used for preset availability, start_preference, usage_preference only
    :param type: only accept availability,start_preference,usage_preference
    :param pattern: pattern get from db
    :param minute_per_slot: get from setting table
    :return:
    """
    start_time = None
    end_time = None
    slot_per_day = int(slot_per_day)
    minute_per_slot = int(minute_per_slot)
    # seperate the pattern(this should be 1 week pattern) become daily
    daily_patterns = [pattern[i:i + slot_per_day] for i in range(0, len(pattern), slot_per_day)]
    # must correct type only will check, if not direct throw error, cause by right wont have this problem when coding
    match type:
        case "usage_preference" | "start_preference":
            value_check = "5" # in string type, so can match compare with pattern string
        case "availability":
            value_check = "1" # in string type, so can match compare with pattern string

    for pattern in daily_patterns:
        # reset current_time every day
        # only datetime value can using timedelta to do calculation, will ignore the date
        current_datetime = datetime.datetime.strptime("00:00", "%H:%M")
        delta = datetime.timedelta(minutes=minute_per_slot)
        for key, char in enumerate(pattern):
            current_time = current_datetime.time()
            if char != value_check:
                # need preset a value if is None, if not got error
                if start_time is None:
                    start_time = current_time
                if end_time is None:
                    end_time = current_time

                if current_time < start_time:
                    start_time = current_time
                if current_time > end_time:
                    end_time = current_time
            current_datetime = current_datetime + delta

    return {
        "start_time": start_time,
        "end_time": end_time,
    }

def get_auto_generate_name(name: str,prefix,leading_zero_length):
    """
    when name or code or any other column need the auto generate data, will pass in last name/code using, this function will increase 1
    Increment a code like 'A000001' or 'AB0099' → 'A000002', 'AB0100'.
    Keeps prefix and leading zeros.
    """

    if name is None:
        # first code
        return f"{prefix}{'0' * (leading_zero_length - 1)}1"
    else:
        pattern = f"^{re.escape(prefix)}(\\d+)$"
        match = re.match(pattern, name)
        if match:
            num = match.group(1)
            new_number = int(num) + 1
            return f"{prefix}{new_number:0{leading_zero_length}d}"

    # no name found return default name only, cause return false will have error
    return False


def get_related_fk_table_ids(related_model_name, ids, redis_format = False, return_column="id"):
    """
    this function used to get all having FK of related_model_name table, the main table id,
    eg. related_model_name = TtDepartment, ids = [1,2], so it will return all table record of id having FK of TtDepartment with id [1,2]
    related_model_name is the name in model like TtDepartment
    ids must be array like [1,2,3] even 1 id only also need put [1]
    redis_format: if True, the return table name will remove prefix tt_
    return_column: default will return id, some case will return by other column
    Returns a dict: {ModelName: queryset of ids}
    """
    results = {}
    return_data = {}

    for model in apps.get_models():
        for field in model._meta.get_fields():
            if isinstance(field, ForeignKey) and field.related_model.__name__ == related_model_name:
                # Use __in to filter by list of ids
                qs = model.objects.filter(**{f"{field.name}__in": ids})
                if qs.exists():
                    field_names = {f.name for f in model._meta.get_fields()}
                    # incase the custom column not found, will return by id
                    if return_column not in field_names:
                        final_column = "id"
                    else:
                        final_column = return_column
                    results[model._meta.db_table] = list(qs.values_list(final_column, flat=True))

    if results and redis_format:
        for key,val in results.items():
            key_name = key.removeprefix(settings.TIMETABLER_TABLE_PREFIX)
            return_data[key_name] = val
    else:
        return_data = results

    return return_data

def get_missing_code_from_db(prefix,table_name,leading_zero_length):
    """
    query explain
    SELECT CAST(REPLACE(code, '{prefix}', '') AS BIGINT) AS num: remove prefix and convert to int
    SELECT generate_series(1, (SELECT MAX(num) FROM codes)) AS n: from codes, form a data like [1,2,3,4,5,...all other code]
    SELECT '{prefix}' || LPAD(n::text, 10, '0') AS missing_code: result will add back the prefix like "MV0000000008"
    WHERE c.num IS NULL: when Null means is missing code

    leading_zero_length = need make number have how many character, eg if 5, when number is 1 means will be 00001, if 9999, will be 09999
    """
    query = f"""
    WITH codes AS (
        SELECT CAST(
            REPLACE(
                REGEXP_REPLACE(code, ' [0-9]{{3}}$', ''), -- Remove suffix " 001"
                '{prefix}', 
                ''
            ) AS BIGINT
        ) AS num
        FROM {table_name}
        WHERE code LIKE '{prefix}%'
    ),
    numbers AS (
        SELECT generate_series(1, (SELECT MAX(num) FROM codes)) AS n
    )
    SELECT '{prefix}' || LPAD(n::text, {leading_zero_length}, '0') AS missing_code
    FROM numbers
    LEFT JOIN codes c ON c.num = n
    WHERE c.num IS NULL
    ORDER BY n
    """
    with connection.cursor() as cursor:
        cursor.execute(query)
        rows = cursor.fetchall()

    new_codes = [row[0] for row in rows]
    return new_codes

def get_missing_name(name_list, prefix, leading_zero_length: int):
    """
    Get missing codes for a given prefix, this function wont go db, so different from get_missing_code_from_db.

    Parameters:
    - name_list: list of existing codes/names
    - prefix: prefix to filter
    - leading_zero_length: need make number have how many character, eg if 5, when number is 1 means will be 00001, if 9999, will be 09999

    Returns:
    - List of missing full name with prefix
    """
    filtered = [n for n in name_list if n.startswith(prefix)]

    # Extract numeric part
    numbers = [int(n.split('/')[-1]) for n in filtered]
    if numbers:
        max_num = max(numbers)
        # Determine padding length
        pad_len = leading_zero_length

        # Find missing numbers
        missing_numbers = [i for i in range(1, max_num) if i not in numbers]

        # Build full missing codes with prefix
        missing_name = [f"{prefix}{str(i).zfill(pad_len)}" for i in missing_numbers]

        return missing_name

    return []

def bitwise_calculator_from_string(string1,string2):
    """
    string1 = "10011"
    string2 = "01001"
    will return "11011" in string format
    """
    # convert to int base 2
    num1 = int(string1, 2)
    num2 = int(string2, 2)

    # bitwise OR
    merged_num = num1 | num2

    # convert back to string, keeping original length
    merged_str = bin(merged_num)[2:].zfill(len(string1))

    return merged_str

def convert_to_redis_week(weeks):
    """
    Convert 1-based week numbers to 0-based for Redis.
    Example: [1,2,3,5] -> [0,1,2,4]
    Reason: scheduling part to read redis is using 0-based for the week, so 0 means week 1, 1 means week 2
    """
    return [w - 1 for w in weeks if w is not None and w > 0]

def recalculate_resource_map(type,id_with_week, slot_per_week):
    """
    used pure query to recalculate the latest resource map for staff/location/student_set
    :param type: accept staff/location/student_set
    :param id_with_week: eg.
        {
            1: [1, 2],  # staff_id 1 → week 1, 2
            2: [1, 4],  # staff_id 2 → week 1, 4
        }
    :param default_pattern: when in that week no activity, will used default pattern replace to db
    """
    activity_table = "tt_activity"
    match type:
        case "staff":
            activity_relation_table = "tt_activity_staff"
            resource_map_table = "tt_staff_resource_map"
        case "location":
            activity_relation_table = "tt_activity_location"
            resource_map_table = "tt_location_resource_map"
        case "student_set":
            activity_relation_table = "tt_student_set_activity"
            resource_map_table = "tt_student_set_resource_map"
        case _:
            # invalid type just return False
            return False

    week_pairs = [
        (id, week_id)
        for id, weeks in id_with_week.items()  # <- must use .items()
        for week_id in weeks
    ]
    # safety check, if nothings return false
    if not week_pairs:
        return False

    slot_per_week = int(slot_per_week)
    default_resource_map = "0" * slot_per_week

    with transaction.atomic():
        with connection.cursor() as cursor:
            cursor.execute(f"""
                WITH targets AS (
                    -- 1. Identify ALL rows we intend to update (based on your input list)
                    SELECT {type}_id AS entity_id, week_id
                    FROM {resource_map_table}
                    WHERE ({type}_id, week_id) IN %s
                ),
                activity_masks AS (
                    -- 2. Calculate patterns ONLY for weeks that actually have activities
                    SELECT
                        ar.{type}_id AS entity_id,
                        COALESCE(wpw.week_id, aw.week_id) AS week_id,
                        bit_or(
                            (
                                overlay(
                                    repeat('0', {slot_per_week})
                                    placing repeat('1', a.slot_required)
                                    from a.scheduled_start_slot + 1
                                )::bit({slot_per_week})
                            )
                        ) AS new_pattern
                    FROM {activity_relation_table} ar
                    JOIN {activity_table} a ON a.id = ar.activity_id AND a.scheduled = 1
                    LEFT JOIN tt_week_pattern_week wpw ON a.week_pattern_id = wpw.week_pattern_id
                    LEFT JOIN tt_activity_week aw 
                        ON aw.activity_id = a.id AND a.week_pattern_id IS NULL
                    WHERE (ar.{type}_id, COALESCE(wpw.week_id, aw.week_id)) IN %s
                    GROUP BY ar.{type}_id, COALESCE(wpw.week_id, aw.week_id)
                )
                -- 3. Update using a LEFT JOIN logic
                UPDATE {resource_map_table} rm
                SET pattern = COALESCE(am.new_pattern, B'{default_resource_map}')
                FROM targets t
                LEFT JOIN activity_masks am 
                    ON t.entity_id = am.entity_id AND t.week_id = am.week_id
                WHERE rm.{type}_id = t.entity_id 
                  AND rm.week_id = t.week_id;
            """, [tuple(week_pairs), tuple(week_pairs)])

def push_websocket_notification(url,data,socket_id):
    try:
        payload = {
            "socket_id": socket_id,
            "data": data
        }
        res = requests.post(url,
                      data=json.dumps(payload),
                      headers={"Content-Type": "application/json"},
                      timeout=2
                      )
    except Exception as e:
        e_details = get_exception_detail(e)
        log_critical_error(user_id=None, descr=e_details['descr'], url=e_details['url'], trace=e_details['trace'])

def format_variant_activity_name_week_pattern_to_ranges(current_week_pattern):
    """
    :param current_week_pattern: week pattern array than have id,week key
    :return: return data become <1-4,7,10-12>
    """
    if not current_week_pattern:
        return ""

    weeks = sorted({w["week"] for w in current_week_pattern})
    ranges = []

    start = prev = weeks[0]

    for w in weeks[1:]:
        if w == prev + 1:
            prev = w
        else:
            ranges.append(
                str(start) if start == prev else f"{start}-{prev}"
            )
            start = prev = w

    # close last range
    ranges.append(
        str(start) if start == prev else f"{start}-{prev}"
    )

    return "<" + ",".join(ranges) + ">"

def remove_variant_activity_name_range(name):
    # remove last <...> at the end, will used when variant again those variant record
    return re.sub(r"\s*<[^>]+>$", "", name)


def get_next_available_code(model, field_name, prefix, leading_zero_length):
    """
    Different from get_missing_code_from_db, this will get either missing or latest code from db
    Finds the first available gap in the sequence.
    Example: If DB has [MV0000000001, MV0000000002, MV0000000004], it returns 'MV0000000003'.
    """
    # for search in db i want get exact this length of code, so can different out MV00001, MV0000000002 in db, if in this case will return MV0000000001
    total_expected_length = len(prefix) + leading_zero_length
    existing_codes = (
        model.objects.filter(
            **{f"{field_name}__startswith": prefix}
        )
        .extra(where=[f"LENGTH({field_name}) = %s"], params=[total_expected_length])  # Only look at total_expected_length count strings
        .annotate(
            num_val=Cast( # Cast is convert str to output_field (int)
                Replace(field_name, Value(prefix), Value("")),# replace the prefix "MV" so remaining the int of "0000000001"
                output_field=models.IntegerField()
            )
        )
        .values_list("num_val", flat=True)
        .order_by("num_val")
    )

    # 2. Find the first gap
    # We compare the actual list [1, 2, 4] vs a perfect sequence [1, 2, 3, 4]
    next_num = 1
    for val in existing_codes:
        if val == next_num:
            next_num += 1
        elif val > next_num:
            # We found a gap!
            break

    # 3. Format and return
    return f"{prefix}{next_num:0{leading_zero_length}d}"

def create_resource_map_single(type, obj_id, week_ids, default_pattern = None):
    inserted = 0
    week_placeholders = ",".join(["%s"] * len(week_ids))

    sql = f"""
    INSERT INTO tt_{type}_resource_map ({type}_id, week_id, pattern, created_at)
    SELECT %s, w.id, %s, NOW()
    FROM tt_week w
    WHERE w.id IN ({week_placeholders})
    """

    params = [obj_id, default_pattern] + week_ids

    with transaction.atomic(), connection.cursor() as cursor: 
        cursor.execute(sql, params)
        inserted = cursor.rowcount

    return inserted

def create_resource_map_all(type, week_ids, default_pattern = None):
    if type not in {"staff", "location", "student_set"}:
        raise ValueError("Invalid type")

    inserted = 0
    week_placeholders = ",".join(["%s"] * len(week_ids))

    sql = f"""
    INSERT INTO tt_{type}_resource_map ({type}_id, week_id, pattern, created_at)
    SELECT s.id, w.id, %s, NOW()
    FROM tt_{type} s
    JOIN tt_week w ON w.id IN ({week_placeholders})
    """
    params = [default_pattern] + week_ids

    with transaction.atomic(), connection.cursor() as cursor: 
        cursor.execute(sql, params)
        inserted = cursor.rowcount

    return inserted

def check_resource_map_available(resource_maps,start_slot,slot_required, exclude_start=None, exclude_end=None):
    """
    :param resource_maps: data format in {1:"00000000000",2:"0000111011100",3:"01010101010101"}
    :param start_slot: integer get from db column scheduled_start_slot
    :param slot_required: integer get from db column slot_required
    :param exclude_start: start of the slot need to exclude
    :param exclude_end: end of the slot need to exclude
    exclude_start and exclude_end must used tgt, so if start is 10, end is 13, means if 10-13 slot is 1, will be ignore
    use case is want move activity start from 10:00 to 10:30 in same staff and location, so need used this exclude part to exclude from current pattern
    :return: True or False
    """

    # if start_slot = 15 and slot_required = 2, the end slot will be 17
    end_slot = start_slot + slot_required

    for week_id, pattern in resource_maps.items():
        target_slice = pattern[start_slot:end_slot]

        if len(target_slice) < slot_required:
            return False

        if "1" in target_slice:
            # If no exclusion, any '1' is an instant fail
            if exclude_start is None or exclude_end is None:
                return False

            # if got exclude, will check is this "1" is in the range of exclude, if yes, will ignore it, cause might be move slot 10-13 to 11-14 only
            for i in range(start_slot, end_slot):
                # even pattern is str, can get by pattern[0] will get first character
                if i < len(pattern) and pattern[i] == "1":
                    # so if "i" not between exclude start and end, return False
                    if not (exclude_start <= i <= exclude_end):
                        return False
    return True


def check_availability_available(pattern, start_slot, slot_required, old_start=None, old_end=None):
    # Calculate the boundaries of the new requested position
    new_end = start_slot + slot_required

    # 1. Slice the pattern for the new requested range
    target_slice = pattern[start_slot:new_end]

    # 2. Safety check: ensure the pattern isn't shorter than the required time
    if len(target_slice) < slot_required:
        return False

    # 3. If there are '1's in the way, we check if they are "self-conflicts"
    if "1" in target_slice:
        # Check every individual slot in the requested range
        for i in range(start_slot, new_end):
            if i < len(pattern) and pattern[i] == "1":
                # If we have an 'old' position to ignore, check if this '1' is inside it
                if old_start is not None and old_end is not None:
                    if old_start <= i < old_end:
                        continue  # This '1' is the activity itself! Skip to next slot.

                # If it's a '1' and NOT in the ignore range, it's a real conflict
                return False

    return True

def fetch_code_id_map(model, code_field, codes):
    if not codes:
        return {}
    return dict(
        model.objects
        .filter(**{f"{code_field}__in": list(codes)})
        .values_list(code_field, "id")
    )

# for insert/creating data
def generate_kafka_data(new_data, relation_fields = []):
    if not new_data:
        return {}
    
    kafka_data = new_data.get("data", {}).copy()
    kafka_data.pop('created_by', None)

    for field in relation_fields:
        value = new_data.get(field)
        if value:
            kafka_data[field] = value
    
    #print("kafka", kafka_data)
    return kafka_data