#!/usr/bin/env python3
# -*- encoding: utf-8; py-indent-offset: 4 -*-
# Microsoft Defender Alerts
# +------------------------------------------------------------------+
#      _____ ____   _____    _____  ____  _               _____  
#     |_   _|  _ \ / ____|  / ____|/ __ \| |        /\   |  __ \ 
#       | | | |_) | |      | (___ | |  | | |       /  \  | |__) |
#       | | |  _ <| |       \___ \| |  | | |      / /\ \ |  _  / 
#      _| |_| |_) | |____   ____) | |__| | |____ / ____ \| | \ \ 
#     |_____|____/ \_____| |_____/ \____/|______/_/    \_\_|  \_\
#                                                                                                                        
# +------------------------------------------------------------------+
# Copyright  IBC SOLAR AG - License: GNU General Public License v2.
# Author: Marvin Sobeck
# Link: https://www.ibc-solar.com/

import argparse
import json
import sys
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, TypedDict

import requests

import cmk.utils.password_store
from cmk.utils.http_proxy_config import HTTPProxyConfig, deserialize_http_proxy_config


GRAPH_URL = "https://graph.microsoft.com/v1.0/security/alerts_v2"
RESOURCE_SCOPE = "https://graph.microsoft.com/.default"


class DefenderAlert(TypedDict):
    incident_id: str
    title: str
    severity: str
    status: str
    classification: str
    determination: str
    service_source: str
    detection_source: str
    product_name: str
    category: str
    rule_name: str
    created: str
    updated: str
    resolved: str
    web_url: str


def parse_arguments() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Checkmk special agent for Microsoft Defender alerts",
    )
    parser.add_argument(
        "--tenant-id",
        required=True,
        help="Microsoft Entra tenant ID",
    )
    parser.add_argument(
        "--app-id",
        required=True,
        help="Application (client) ID of the Microsoft Entra app registration",
    )
    parser.add_argument(
        "--app-secret",
        required=True,
        help=(
            "Password store reference of the application (client) secret "
            "(e.g., 'secret_id:/omd/sites/<site>/var/check_mk/passwords_merged')"
        ),
    )
    parser.add_argument(
        "--lookback",
        required=False,
        type=float,
        default=604800.0,
        help="Alert lookback in seconds (default: %(default)s)",
    )
    parser.add_argument(
        "--proxy",
        required=False,
        help="HTTP proxy (FROM_ENVIRONMENT, NO_PROXY, or URL) (default: environment settings)",
    )
    parser.add_argument(
        "--timeout",
        required=False,
        type=float,
        default=10.0,
        help="API request timeout in seconds (default: %(default)s)",
    )
    return parser.parse_args()


def handle_error(err: Exception, context: str, exit_code: int = 1) -> None:
    err_msg = str(err)
    response = getattr(err, "response", None)
    if response is not None:
        err_msg += f" Response: {getattr(response, 'text', 'No response text')}"
    sys.stderr.write(f"{err_msg}\n\n{context}\n")
    sys.exit(exit_code)


def get_access_token(
    tenant_id: str,
    app_id: str,
    app_secret: str,
    resource_scope: str,
    timeout: float,
    proxy: HTTPProxyConfig,
) -> str:
    token_url = f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
    headers = {"Content-Type": "application/x-www-form-urlencoded"}
    body = {
        "client_id": app_id,
        "client_secret": app_secret,
        "grant_type": "client_credentials",
        "scope": resource_scope,
    }

    try:
        token_response = requests.post(
            token_url,
            headers=headers,
            data=body,
            timeout=timeout,
            proxies=proxy.to_requests_proxies(),
        )
        token_response.raise_for_status()
    except requests.exceptions.Timeout as err:
        handle_error(err, "Timeout while getting access token.", 11)
    except requests.exceptions.RequestException as err:
        error_message = "Failed to get access token."
        error_message_details = {
            400: f"{error_message} Please check tenant ID and client ID.",
            401: f"{error_message} Please check client secret.",
            429: f"{error_message} Request has been throttled.",
        }
        status_code = getattr(getattr(err, "response", None), "status_code", 0)
        handle_error(err, error_message_details.get(status_code, error_message), 1)

    token_payload = token_response.json()
    if "access_token" not in token_payload:
        handle_error(
            RuntimeError("Microsoft Entra token response contains no access_token."),
            "Unable to authenticate against Microsoft Graph.",
            1,
        )
    return str(token_payload["access_token"])


def normalize_field_name(field_name: str) -> str:
    return (
        field_name.lower()
        .replace("_", "")
        .replace("-", "")
        .replace(" ", "")
    )


def get_rule_name(value: Any) -> str:
    rule_name_keys = {
        "rulename",
        "dlprulename",
        "policyrulename",
        "matchedrulename",
    }

    if isinstance(value, dict):
        for key, nested_value in value.items():
            if normalize_field_name(str(key)) in rule_name_keys and nested_value:
                return str(nested_value)

        for nested_value in value.values():
            rule_name = get_rule_name(nested_value)
            if rule_name:
                return rule_name

    elif isinstance(value, list):
        for item in value:
            rule_name = get_rule_name(item)
            if rule_name:
                return rule_name

    return ""


def string_or_empty(value: Any) -> str:
    if value is None:
        return ""

    return str(value)
    

def normalize_alert(
    alert: dict[str, Any],
) -> DefenderAlert:
    return {
        "incident_id": string_or_empty(
            alert.get("incidentId")
        ),
        "title": string_or_empty(
            alert.get("title")
        ) or "Unknown security alert",
        "severity": (
            string_or_empty(
                alert.get("severity")
            ) or "unknown"
        ).lower(),
        "status": (
            string_or_empty(
                alert.get("status")
            ) or "unknown"
        ).lower(),
        "classification": (
            string_or_empty(
                alert.get("classification")
            ) or "unknown"
        ),
        "determination": (
            string_or_empty(
                alert.get("determination")
            ) or "unknown"
        ),
        "service_source": string_or_empty(
            alert.get("serviceSource")
        ),
        "detection_source": string_or_empty(
            alert.get("detectionSource")
        ),
        "product_name": string_or_empty(
            alert.get("productName")
        ),
        "category": string_or_empty(
            alert.get("category")
        ),
        "rule_name": get_rule_name(alert),
        "created": string_or_empty(
            alert.get("createdDateTime")
        ),
        "updated": string_or_empty(
            alert.get("lastUpdateDateTime")
        ),
        "resolved": string_or_empty(
            alert.get("resolvedDateTime")
        ),
        "web_url": string_or_empty(
            alert.get("alertWebUrl")
        ),
    }


def get_ms_defender_alerts(
    token: str,
    lookback: float,
    timeout: float,
    proxy: HTTPProxyConfig,
) -> list[DefenderAlert]:
    created_after = datetime.now(timezone.utc) - timedelta(seconds=lookback)
    created_filter = created_after.isoformat(timespec="seconds").replace("+00:00", "Z")
    headers = {"Authorization": f"Bearer {token}", "Accept": "application/json"}
    params: dict[str, str] | None = {
        "$filter": f"createdDateTime ge {created_filter}",
        "$top": "200",
    }
    url: str | None = GRAPH_URL
    alerts: list[DefenderAlert] = []

    while url:
        try:
            response = requests.get(
                url,
                headers=headers,
                params=params,
                timeout=timeout,
                proxies=proxy.to_requests_proxies(),
            )
            response.raise_for_status()
        except requests.exceptions.Timeout as err:
            handle_error(err, "Timeout while getting Microsoft Defender alerts.", 12)
        except requests.exceptions.RequestException as err:
            error_message = "Failed to get Microsoft Defender alerts."
            error_message_details = {
                401: f"{error_message} Please check application authentication.",
                403: (
                    f"{error_message} Please check application API permissions. "
                    "At least SecurityAlert.Read.All is required."
                ),
                429: f"{error_message} Request has been throttled.",
            }
            status_code = getattr(getattr(err, "response", None), "status_code", 0)
            handle_error(err, error_message_details.get(status_code, error_message), 2)

        payload = response.json()
        alerts.extend(normalize_alert(alert) for alert in payload.get("value", []))
        url = payload.get("@odata.nextLink")
        params = None

    return alerts


def main() -> None:
    args = parse_arguments()
    tenant_id = args.tenant_id
    app_id = args.app_id
    proxy = deserialize_http_proxy_config(args.proxy)
    lookback = args.lookback
    timeout = args.timeout

    try:
        pw_id, pw_path = args.app_secret.split(":", 1)
    except ValueError as err:
        handle_error(
            err,
            "Invalid Checkmk password-store reference for the application client secret.",
            3,
        )

    app_secret = cmk.utils.password_store.lookup(Path(pw_path), pw_id)
    token = get_access_token(
        tenant_id,
        app_id,
        app_secret,
        RESOURCE_SCOPE,
        timeout,
        proxy,
    )

    print("<<<ms_defender:sep(0)>>>")
    print(
        json.dumps(
            get_ms_defender_alerts(token, lookback, timeout, proxy),
            ensure_ascii=False,
        )
    )


if __name__ == "__main__":
    main()
