import sys
import random
import re
import os
import time
import urllib.parse
import requests
from collections import defaultdict
from cw_rpa import Logger, Input, HttpClient, ResultLevel

sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))

log = Logger()
http_client = HttpClient()
input = Input()
log.info("Imports completed successfully")

cwpsa_base_url = "https://aus.myconnectwise.net"
cwpsa_base_url_path = "/v4_6_release/apis/3.0"
msgraph_base_url_base = "https://graph.microsoft.com"
msgraph_base_url_path = "/v1.0"
msgraph_base_url_beta_base = "https://graph.microsoft.com"
msgraph_base_url_beta_path = "/beta"
sender_email = "support@example.com"

graph_integration_name = "azure_o365"
psa_integration_name = "cw_psa"
extra_integration_name = ""

data_to_log = {}
bot_name = "REPORT - M365 License User Report"
log.info("Static variables set")

def record_result(log, level, message):
    log.result_message(level, f"[{bot_name}]: {message}")
    if level in (ResultLevel.WARNING, ResultLevel.ERROR):
        data_to_log["status_result"] = "Fail"
    elif level == ResultLevel.SUCCESS:
        if "status_result" not in data_to_log or data_to_log["status_result"] != "Fail":
            data_to_log["status_result"] = "Success"

def execute_api_call(log, http_client, method, endpoint, data=None, retries=5, integration_name=None, headers=None, params=None):
    base_delay = 5
    log.info(f"Executing API call: {method.upper()} {endpoint}")
    for attempt in range(retries):
        try:
            if integration_name:
                response = (
                    getattr(http_client.third_party_integration(integration_name), method)(url=endpoint, json=data)
                    if data else getattr(http_client.third_party_integration(integration_name), method)(url=endpoint)
                )
            else:
                request_args = {"url": endpoint}
                if params:
                    request_args["params"] = params
                if headers:
                    request_args["headers"] = headers
                if data:
                    if (headers and headers.get("Content-Type") == "application/x-www-form-urlencoded"):
                        request_args["data"] = data
                    else:
                        request_args["json"] = data
                response = getattr(requests, method)(**request_args)

            if 200 <= response.status_code < 300:
                return response
            elif response.status_code in [429, 503]:
                retry_after = response.headers.get("Retry-After")
                wait_time = int(retry_after) if retry_after else base_delay * (2 ** attempt) + random.uniform(0, 3)
                log.warning(f"Rate limit exceeded. Retrying in {wait_time:.2f} seconds")
                time.sleep(wait_time)
            elif 400 <= response.status_code < 500:
                if response.status_code == 404:
                    log.warning(f"Skipping non-existent resource [{endpoint}]")
                    return None
                log.error(f"Client error Status: {response.status_code}, Response: {response.text}")
                return response
            elif 500 <= response.status_code < 600:
                log.warning(f"Server error Status: {response.status_code}, attempt {attempt + 1} of {retries}, Response: {getattr(response, 'text', '')[:1000]}")
                time.sleep(base_delay * (2 ** attempt) + random.uniform(0, 3))
            else:
                log.error(f"Unexpected response Status: {response.status_code}, Response: {response.text}")
                return response

        except Exception as e:
            log.exception(e, f"Exception during API call to {endpoint}")
            return None
    return None

def get_company_data_from_ticket(log, http_client, cwpsa_base_url, cwpsa_base_url_path, ticket_number):
    log.info(f"Retrieving company details for ticket [{ticket_number}]")
    ticket_endpoint = f"{cwpsa_base_url}{cwpsa_base_url_path}/service/tickets/{ticket_number}"
    ticket_response = execute_api_call(log, http_client, "get", ticket_endpoint, integration_name=psa_integration_name)
    if ticket_response:
        ticket_data = ticket_response.json()
        company = ticket_data.get("company", {})
        company_id = company["id"]
        company_identifier = company["identifier"]
        company_name = company["name"]
        log.info(f"Company ID: [{company_id}], Identifier: [{company_identifier}], Name: [{company_name}]")
        company_endpoint = f"{cwpsa_base_url}{cwpsa_base_url_path}/company/companies/{company_id}"
        company_response = execute_api_call(log, http_client, "get", company_endpoint, integration_name=psa_integration_name)
        company_types = []
        if company_response:
            company_data = company_response.json()
            types = company_data.get("types", [])
            company_types = [t.get("name", "") for t in types if "name" in t]
            log.info(f"Company types for ID [{company_id}]: {company_types}")
        else:
            log.warning(f"Unable to retrieve company types for ID [{company_id}]")
        return company_identifier, company_name, company_id, company_types
    return "", "", 0, []

def get_aad_user_data(log, http_client, msgraph_base_url_base, msgraph_base_url_path, user_identifier):
    log.info(f"Resolving user ID and email for [{user_identifier}]")
    if re.fullmatch(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}", user_identifier):
        endpoint = f"{msgraph_base_url_base}{msgraph_base_url_path}/users/{user_identifier}"
        response = execute_api_call(log, http_client, "get", endpoint, integration_name=graph_integration_name)
        if response:
            user = response.json()
            return user.get("id", ""), user.get("userPrincipalName", ""), user.get("onPremisesSamAccountName", ""), user.get("onPremisesSyncEnabled", False)
        return "", "", "", False
    filters = [
        f"startswith(displayName,'{user_identifier}')",
        f"startswith(userPrincipalName,'{user_identifier}')",
        f"startswith(mail,'{user_identifier}')"
    ]
    filter_query = " or ".join(filters)
    endpoint = f"{msgraph_base_url_base}{msgraph_base_url_path}/users?$filter={urllib.parse.quote(filter_query)}"
    response = execute_api_call(log, http_client, "get", endpoint, integration_name=graph_integration_name)
    if response:
        users = response.json().get("value", [])
        if len(users) > 1:
            log.error(f"Multiple users found for [{user_identifier}]")
            return users
        if users:
            user = users[0]
            return user.get("id", ""), user.get("userPrincipalName", ""), user.get("onPremisesSamAccountName", ""), user.get("onPremisesSyncEnabled", False)
    return "", "", "", False

def get_subscribed_skus(log, http_client, msgraph_base_url_base, msgraph_base_url_path):
    log.info("Fetching subscribed SKUs")
    endpoint = f"{msgraph_base_url_base}{msgraph_base_url_path}/subscribedSkus"
    response = execute_api_call(log, http_client, "get", endpoint, integration_name=graph_integration_name)
    sku_info = {}
    if not response:
        log.error("Failed to fetch subscribed SKUs")
        return sku_info
    for sku in response.json().get("value", []):
        sku_id = sku.get("skuId")
        if not sku_id:
            continue
        sku_info[sku_id] = {
            "name": sku.get("skuPartNumber", sku_id),
            "consumed": sku.get("consumedUnits", 0),
            "enabled": sku.get("prepaidUnits", {}).get("enabled", 0),
            "status": sku.get("capabilityStatus", "")
        }
    log.info(f"Retrieved [{len(sku_info)}] subscribed SKUs")
    return sku_info

def get_users_by_license(log, http_client, msgraph_base_url_base, msgraph_base_url_path):
    log.info("Fetching users with assigned licenses")
    users_by_sku = defaultdict(list)
    next_link = f"{msgraph_base_url_base}{msgraph_base_url_path}/users?$select=userPrincipalName,assignedLicenses&$top=999"
    while next_link:
        response = execute_api_call(log, http_client, "get", next_link, integration_name=graph_integration_name)
        if not response:
            log.error("Failed to fetch users with assigned licenses")
            break
        data = response.json()
        for user in data.get("value", []):
            upn = user.get("userPrincipalName", "")
            if not upn:
                continue
            for license_entry in user.get("assignedLicenses", []):
                sku_id = license_entry.get("skuId")
                if sku_id:
                    users_by_sku[sku_id].append(upn)
        next_link = data.get("@odata.nextLink")
    log.info(f"Mapped licenses for [{sum(len(v) for v in users_by_sku.values())}] user assignments")
    return users_by_sku

def format_license_user_report(sku_info, users_by_sku):
    active_skus = [
        (sku_id, info)
        for sku_id, info in sku_info.items()
        if info.get("consumed", 0) > 0
    ]
    active_skus.sort(key=lambda item: item[1].get("name", ""))
    if not active_skus:
        return "No assigned licenses found in tenant"
    sections = []
    for sku_id, info in active_skus:
        license_name = info.get("name", sku_id)
        count = info.get("consumed", 0)
        users = sorted(set(users_by_sku.get(sku_id, [])))
        user_lines = "\n".join(users) if users else "No users found"
        sections.append(f"{license_name}: {count}\n\n{user_lines}")
    return "\n\n".join(sections)

def generate_license_user_report(log, http_client, msgraph_base_url_base, msgraph_base_url_path):
    sku_info = get_subscribed_skus(log, http_client, msgraph_base_url_base, msgraph_base_url_path)
    if not sku_info:
        return ""
    users_by_sku = get_users_by_license(log, http_client, msgraph_base_url_base, msgraph_base_url_path)
    return format_license_user_report(sku_info, users_by_sku)

def main():
    try:
        try:
            ticket_number = input.get_value("TicketNumber_xxxxxxxxxxxxx")
        except Exception:
            record_result(log, ResultLevel.ERROR, "Failed to fetch input values")
            return

        ticket_number = ticket_number.strip() if ticket_number else ""

        log.info(f"Ticket Number = [{ticket_number}]")

        if not ticket_number:
            record_result(log, ResultLevel.WARNING, "Ticket number is required but missing")
            return

        log.info(f"Retrieving company data for ticket [{ticket_number}]")
        company_identifier, company_name, company_id, company_type = get_company_data_from_ticket(log, http_client, cwpsa_base_url, cwpsa_base_url_path, ticket_number)
        if not company_identifier:
            record_result(log, ResultLevel.ERROR, f"Failed to retrieve company identifier from ticket [{ticket_number}]")
            return

        log.info(f"Generating license user report for [{company_name}]")
        report = generate_license_user_report(log, http_client, msgraph_base_url_base, msgraph_base_url_path)
        if not report:
            record_result(log, ResultLevel.ERROR, f"Failed to generate license user report for [{company_name}]")
            return

        record_result(log, ResultLevel.SUCCESS, f"\n\n{report}")

    except Exception as e:
        log.error(f"Unhandled error in main: {str(e)}")
        record_result(log, ResultLevel.ERROR, "Unhandled exception occurred during execution")
    finally:
        log.result_data(data_to_log)

if __name__ == "__main__":
    main()