#!/usr/bin/env python3

import hashlib
import fcntl
import json
import os
import re
import sys
import time
import tempfile
import urllib.error
import urllib.request
from collections import Counter
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from ipaddress import ip_address, ip_network
from pathlib import Path
from typing import Any

import boto3
from botocore.config import Config

# Regions whose CloudTrail Event History you want to scan.
REGIONS_TO_SCAN = [
    r.strip()
    for r in os.getenv("AWS_REGIONS_TO_SCAN", "eu-west-2").split(",")
    if r.strip()
]

# AWS source IP regions that are acceptable.
# Example: eu-west-2 means AWS IPs from us-east-1, eu-west-1, GLOBAL, etc. alert.
ALLOWED_AWS_SOURCE_REGIONS = {
    r.strip()
    for r in os.getenv("ALLOWED_AWS_SOURCE_REGIONS", "eu-west-2").split(",")
    if r.strip()
}

# Your expected non-AWS source IP ranges, e.g. office/VPN/NAT.
EXPECTED_CIDRS = [
    ip_network(c.strip())
    for c in os.getenv("EXPECTED_CIDRS", "").split(",")
    if c.strip()
]

PROFILE = os.getenv("AWS_PROFILE", "default").strip() or "default"
STATE_FILE = Path(
    os.getenv(
        "STATE_FILE", f"./{PROFILE}-cloudtrail-event-watch-state.json"
    )
)

LOOKBACK_MINUTES = int(os.getenv("LOOKBACK_MINUTES", "90"))

ALLOW_AWS_SERVICE_ORIGINATED = (
    os.getenv("ALLOW_AWS_SERVICE_ORIGINATED", "true").lower() == "true"
)

IGNORE_READ_ONLY = os.getenv("IGNORE_READ_ONLY", "false").lower() == "true"

# One JSON endpoint under your control. It can route to Google Chat, Slack, email, etc.
WEBHOOK_URL = os.getenv("WEBHOOK_URL", "").strip()
WEBHOOK_TOKEN = os.getenv("WEBHOOK_TOKEN", "").strip()
WEBHOOK_TIMEOUT_SECONDS = int(os.getenv("WEBHOOK_TIMEOUT_SECONDS", "15"))

# CI-friendly default:
#   0 = alert sent successfully, build remains green
#   2 = alert sent successfully, build fails unless handled by CI
ALERT_EXIT_CODE = int(os.getenv("ALERT_EXIT_CODE", "0"))

# Useful for local testing.
PRINT_ALERT_TO_STDOUT = os.getenv("PRINT_ALERT_TO_STDOUT", "false").lower() == "true"

ACCOUNT_LABEL = os.getenv("ACCOUNT_LABEL", PROFILE)

# Optional exact allowlist for noisy known-good role/API combinations.
#
# Format:
#   role_arn,eventSource,eventName;role_arn,eventSource,eventName
#
# Example:
#   arn:aws:iam::123456789012:role/ecs-instance-role-prod,ssm.amazonaws.com,UpdateInstanceInformation
EXPECTED_ROLE_EVENTS = {
    tuple(x.strip() for x in row.split(",", 2))
    for row in os.getenv("EXPECTED_ROLE_EVENTS", "").split(";")
    if row.strip() and len(row.split(",", 2)) == 3
}

BOTO_CONFIG = Config(
    retries={"max_attempts": 10, "mode": "standard"},
)

AWS_ACCOUNT_ID_RE = re.compile(r"(?<!\d)(\d{5})\d{7}(?!\d)")
AWS_ACCESS_KEY_ID_RE = re.compile(r"\b((?:AKIA|ASIA)[A-Z0-9])[A-Z0-9]{15}\b")


def mask_aws_identifiers(value):
    """Truncate AWS account IDs and access key IDs in actor data.

    Account IDs are 12 digits and are kept as their first 5 digits.
    Access key IDs are 20 characters and are kept as their first 5 characters.
    """
    if isinstance(value, str):
        value = AWS_ACCOUNT_ID_RE.sub(r"\1", value)
        return AWS_ACCESS_KEY_ID_RE.sub(r"\1", value)

    if isinstance(value, dict):
        return {key: mask_aws_identifiers(item) for key, item in value.items()}

    if isinstance(value, list):
        return [mask_aws_identifiers(item) for item in value]

    if isinstance(value, tuple):
        return tuple(mask_aws_identifiers(item) for item in value)

    return value


@dataclass(frozen=True)
class AwsPrefix:
    network: Any
    region: str
    service: str
    network_border_group: str


def load_state() -> dict[str, Any]:
    if STATE_FILE.exists():
        return json.loads(STATE_FILE.read_text())

    return {"seen_event_ids": []}


def save_state(state: dict[str, Any]) -> None:
    # Keep enough IDs for overlapping hourly runs.
    state["seen_event_ids"] = list(dict.fromkeys(state["seen_event_ids"]))[-10000:]

    STATE_FILE.parent.mkdir(parents=True, exist_ok=True)
    fd, temporary = tempfile.mkstemp(prefix=f".{STATE_FILE.name}.", dir=STATE_FILE.parent)
    try:
        with os.fdopen(fd, "w") as stream:
            json.dump(state, stream, indent=2, sort_keys=True)
            stream.flush()
            os.fsync(stream.fileno())
        os.replace(temporary, STATE_FILE)
    finally:
        if os.path.exists(temporary):
            os.unlink(temporary)


def fetch_aws_prefixes() -> list[AwsPrefix]:
    with urllib.request.urlopen(
        "https://ip-ranges.amazonaws.com/ip-ranges.json",
        timeout=20,
    ) as resp:
        data = json.loads(resp.read().decode("utf-8"))

    prefixes: list[AwsPrefix] = []

    for item in data.get("prefixes", []):
        prefixes.append(
            AwsPrefix(
                network=ip_network(item["ip_prefix"]),
                region=item.get("region", ""),
                service=item.get("service", ""),
                network_border_group=item.get("network_border_group", ""),
            )
        )

    for item in data.get("ipv6_prefixes", []):
        prefixes.append(
            AwsPrefix(
                network=ip_network(item["ipv6_prefix"]),
                region=item.get("region", ""),
                service=item.get("service", ""),
                network_border_group=item.get("network_border_group", ""),
            )
        )

    return prefixes


def source_ip_as_ip(value: str):
    if not value:
        return None

    try:
        return ip_address(value)
    except ValueError:
        return None


def matching_aws_prefixes(ip, aws_prefixes: list[AwsPrefix]) -> list[AwsPrefix]:
    return [prefix for prefix in aws_prefixes if ip in prefix.network]


def ip_in_expected_cidrs(ip) -> bool:
    return any(ip in network for network in EXPECTED_CIDRS)


def is_aws_service_originated(event: dict[str, Any]) -> bool:
    identity = event.get("userIdentity", {}) or {}
    source_ip = str(event.get("sourceIPAddress", "") or "")

    # eventSource is the API service being called, not the caller.
    # These are better signals that AWS itself initiated the call.
    if identity.get("type") == "AWSService":
        return True

    if identity.get("invokedBy"):
        return True

    if source_ip.startswith("AWS Internal"):
        return True

    if event.get("eventType") == "AwsServiceEvent":
        return True

    return False


def issuer_role_arn(event: dict[str, Any]) -> str:
    identity = event.get("userIdentity", {}) or {}
    session_context = identity.get("sessionContext", {}) or {}
    issuer = session_context.get("sessionIssuer", {}) or {}
    return issuer.get("arn", "")


def is_expected_role_event(event: dict[str, Any]) -> bool:
    key = (
        issuer_role_arn(event),
        event.get("eventSource", ""),
        event.get("eventName", ""),
    )

    return key in EXPECTED_ROLE_EVENTS


def actor_from_event(event: dict[str, Any]) -> dict[str, Any]:
    identity = event.get("userIdentity", {}) or {}
    session_context = identity.get("sessionContext", {}) or {}
    issuer = session_context.get("sessionIssuer", {}) or {}
    attributes = session_context.get("attributes", {}) or {}
    in_scope = identity.get("inScopeOf", {}) or {}

    identity_type = identity.get("type", "")
    arn = identity.get("arn", "")
    user_name = identity.get("userName", "")
    principal_id = identity.get("principalId", "")
    access_key_id = identity.get("accessKeyId", "")
    account_id = identity.get("accountId", "")

    actor: dict[str, Any] = {
        "identityType": identity_type,
        "accountId": account_id,
        "arn": arn,
        "userName": user_name,
        "principalId": principal_id,
        "accessKeyId": access_key_id,
    }

    if identity_type == "IAMUser":
        actor["summary"] = f"IAM user {user_name or arn}"

    elif identity_type == "Root":
        actor["summary"] = f"Root user for account {account_id}"

    elif identity_type == "AssumedRole":
        session_name = ""

        if ":assumed-role/" in arn:
            # arn:aws:sts::123456789012:assumed-role/RoleName/SessionName
            parts = arn.split(":assumed-role/", 1)[1].split("/", 1)
            if len(parts) == 2:
                session_name = parts[1]

        actor["issuerRoleArn"] = issuer.get("arn", "")
        actor["issuerRoleName"] = issuer.get("userName", "")
        actor["sessionName"] = session_name
        actor["mfaAuthenticated"] = attributes.get("mfaAuthenticated", "")
        actor["creationDate"] = attributes.get("creationDate", "")
        actor["sourceIdentity"] = session_context.get("sourceIdentity", "")

        if in_scope:
            actor["inScopeOf"] = in_scope

        actor["summary"] = f"assumed role {issuer.get('arn', '') or arn}" + (
            f" session {session_name}" if session_name else ""
        )

        if in_scope.get("credentialsIssuedTo"):
            actor["summary"] += f" issued to {in_scope['credentialsIssuedTo']}"

    elif identity.get("invokedBy"):
        actor["summary"] = f"AWS service {identity.get('invokedBy')}"

    else:
        actor["summary"] = (
            arn or user_name or principal_id or identity_type or "unknown"
        )

    return mask_aws_identifiers(actor)


def classify_event(
    event: dict[str, Any],
    aws_prefixes: list[AwsPrefix],
) -> tuple[bool, str, list[AwsPrefix]]:
    """
    Return:
      should_alert, reason, matched_aws_prefixes
    """

    if (
        IGNORE_READ_ONLY
        and event.get("readOnly") is True
    ):
        return False, "readOnly event ignored", []

    if ALLOW_AWS_SERVICE_ORIGINATED and is_aws_service_originated(event):
        return False, "AWS-service-originated event", []

    if is_expected_role_event(event):
        return False, "expected role/event combination", []

    raw_source_ip = str(event.get("sourceIPAddress", "") or "")
    source_ip = source_ip_as_ip(raw_source_ip)

    # Policy:
    # if it is not an IP address at all, do not alert.
    if source_ip is None:
        return (
            False,
            (
                "sourceIPAddress is not an IP address; treating as internal/service-mediated: "
                f"{raw_source_ip!r}"
            ),
            [],
        )

    if EXPECTED_CIDRS and ip_in_expected_cidrs(source_ip):
        return False, "source IP is explicitly expected", []

    matches = matching_aws_prefixes(source_ip, aws_prefixes)

    if not matches:
        return True, "source IP is not in AWS published IP ranges", []

    matched_regions = {m.region for m in matches}

    if matched_regions & ALLOWED_AWS_SOURCE_REGIONS:
        return (
            False,
            (
                "source IP is in an allowed AWS source region: "
                + ", ".join(sorted(matched_regions & ALLOWED_AWS_SOURCE_REGIONS))
            ),
            matches,
        )

    return (
        True,
        (
            "source IP is AWS, but not in allowed AWS source regions; "
            f"matched regions: {', '.join(sorted(matched_regions))}"
        ),
        matches,
    )


def compact_prefixes(prefixes: list[AwsPrefix]) -> list[dict[str, str]]:
    result = []
    seen = set()

    for prefix in prefixes:
        key = (
            str(prefix.network),
            prefix.region,
            prefix.service,
            prefix.network_border_group,
        )

        if key in seen:
            continue

        seen.add(key)

        result.append(
            {
                "prefix": str(prefix.network),
                "region": prefix.region,
                "service": prefix.service,
                "networkBorderGroup": prefix.network_border_group,
            }
        )

    return result


def compact_event(
    event: dict[str, Any],
    scan_region: str,
    reason: str,
    prefixes: list[AwsPrefix],
) -> dict[str, Any]:
    return {
        "reason": reason,
        "scanRegion": scan_region,
        "eventTime": event.get("eventTime"),
        "eventSource": event.get("eventSource"),
        "eventName": event.get("eventName"),
        "sourceIPAddress": event.get("sourceIPAddress"),
        "actor": actor_from_event(event),
        "awsIpMatches": compact_prefixes(prefixes),
        "recipientAccountId": event.get("recipientAccountId"),
        "eventID": event.get("eventID"),
        "readOnly": event.get("readOnly"),
        "errorCode": event.get("errorCode") if "errorCode" in event else None,
        "errorMessage": event.get("errorMessage"),
        "managementEvent": event.get("managementEvent"),
        "eventType": event.get("eventType"),
        "eventCategory": event.get("eventCategory"),
    }


def lookup_region_events(region: str, start_time: datetime, end_time: datetime):
    client = boto3.client("cloudtrail", region_name=region, config=BOTO_CONFIG)

    kwargs = {
        "StartTime": start_time,
        "EndTime": end_time,
        "MaxResults": 50,
    }

    while True:
        response = client.lookup_events(**kwargs)

        for item in response.get("Events", []):
            raw = item.get("CloudTrailEvent")
            if raw:
                yield json.loads(raw)

        token = response.get("NextToken")
        if not token:
            break

        kwargs["NextToken"] = token

        # Stay below CloudTrail LookupEvents throttle.
        time.sleep(0.6)


def actor_group_key(finding: dict[str, Any]) -> str:
    actor = finding.get("actor", {}) or {}

    # Prefer stable identifiers over display names.
    return (
        actor.get("arn")
        or actor.get("issuerRoleArn")
        or actor.get("principalId")
        or actor.get("summary")
        or "unknown"
    )


def actor_display_name(actor: dict[str, Any]) -> str:
    identity_type = actor.get("identityType", "")
    summary = actor.get("summary", "")
    arn = actor.get("arn", "")
    access_key_id = actor.get("accessKeyId", "")

    parts = []

    if summary:
        parts.append(summary)
    elif arn:
        parts.append(arn)
    else:
        parts.append("unknown actor")

    if identity_type:
        parts.append(f"type={identity_type}")

    if arn and arn not in parts[0]:
        parts.append(f"arn={arn}")

    # Access key ID is not the secret key, but it is still useful context.
    if access_key_id:
        parts.append(f"accessKeyId={access_key_id}")

    return " | ".join(parts)


def alert_id_from_findings(findings: list[dict[str, Any]]) -> str:
    event_ids = sorted(
        str(finding.get("eventID")) for finding in findings if finding.get("eventID")
    )
    material = "\n".join(event_ids).encode("utf-8")
    return hashlib.sha256(material).hexdigest()[:32]


def build_rollup_payload(findings: list[dict[str, Any]]) -> dict[str, Any]:
    grouped: dict[str, dict[str, Any]] = {}

    for finding in findings:
        actor = finding.get("actor", {}) or {}
        key = actor_group_key(finding)

        if key not in grouped:
            grouped[key] = {
                "actor": actor,
                "actor_display": actor_display_name(actor),
                "source_ips": set(),
                "reasons": set(),
                "scan_regions": set(),
                "recipient_accounts": set(),
                "events": Counter(),
                "event_times": [],
                "event_ids": [],
                "count": 0,
            }

        group = grouped[key]
        group["count"] += 1

        source_ip = finding.get("sourceIPAddress")
        if source_ip:
            group["source_ips"].add(str(source_ip))

        reason = finding.get("reason")
        if reason:
            group["reasons"].add(str(reason))

        scan_region = finding.get("scanRegion")
        if scan_region:
            group["scan_regions"].add(str(scan_region))

        recipient_account = finding.get("recipientAccountId")
        if recipient_account:
            group["recipient_accounts"].add(str(recipient_account))

        event_source = finding.get("eventSource", "unknown")
        event_name = finding.get("eventName", "unknown")
        group["events"][f"{event_source}:{event_name}"] += 1

        event_time = finding.get("eventTime")
        if event_time:
            group["event_times"].append(str(event_time))

        event_id = finding.get("eventID")
        if event_id:
            group["event_ids"].append(str(event_id))

    total_events = sum(group["count"] for group in grouped.values())
    all_times = sorted(
        str(finding.get("eventTime"))
        for finding in findings
        if finding.get("eventTime")
    )

    actor_rows = []
    for group in sorted(grouped.values(), key=lambda g: g["actor_display"]):
        event_times = sorted(group["event_times"])

        actor_rows.append(
            {
                "actor": group["actor"],
                "actor_display": group["actor_display"],
                "source_ips": sorted(group["source_ips"]),
                "events": [
                    {"event": event, "count": count}
                    for event, count in group["events"].most_common()
                ],
                "count": group["count"],
                "first_seen": event_times[0] if event_times else None,
                "last_seen": event_times[-1] if event_times else None,
                "recipient_accounts": sorted(group["recipient_accounts"]),
                "scan_regions": sorted(group["scan_regions"]),
                "reasons": sorted(group["reasons"]),
                "event_ids": sorted(group["event_ids"]),
            }
        )

    alert_id = alert_id_from_findings(findings)

    return {
        "source_title": "CloudTrail unexpected source IP alert",
        "alert_id": alert_id,
        "account_label": ACCOUNT_LABEL,
        "total_events": total_events,
        "actor_count": len(grouped),
        "first_seen": all_times[0] if all_times else None,
        "last_seen": all_times[-1] if all_times else None,
        "scanned_cloudtrail_regions": REGIONS_TO_SCAN,
        "allowed_aws_source_regions": sorted(ALLOWED_AWS_SOURCE_REGIONS),
        "expected_cidrs_configured": len(EXPECTED_CIDRS),
        "lookback_minutes": LOOKBACK_MINUTES,
        "actors": actor_rows,
    }


def post_json_webhook(
    url: str,
    payload: dict[str, Any],
    *,
    headers: dict[str, str] | None = None,
    timeout_seconds: int = 15,
) -> None:
    body = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")

    request_headers = {
        "Content-Type": "application/json; charset=UTF-8",
        "X-Webhook-Timestamp": str(int(time.time())),
    }

    if headers:
        request_headers.update(headers)

    request = urllib.request.Request(
        url,
        data=body,
        headers=request_headers,
        method="POST",
    )

    try:
        with urllib.request.urlopen(request, timeout=timeout_seconds) as response:
            status = getattr(response, "status", 200)
            if status < 200 or status >= 300:
                raise RuntimeError(f"webhook returned HTTP {status}")

    except urllib.error.HTTPError as exc:
        raise RuntimeError(f"webhook returned HTTP {exc.code}") from exc

    except urllib.error.URLError as exc:
        raise RuntimeError(f"webhook request failed: {exc}") from exc


def send_webhook(payload: dict[str, Any]) -> None:
    headers = {"Authorization": f"Bearer {WEBHOOK_TOKEN}"} if WEBHOOK_TOKEN else None
    post_json_webhook(WEBHOOK_URL, payload, headers=headers,
                      timeout_seconds=WEBHOOK_TIMEOUT_SECONDS)


def update_seen_events(state: dict[str, Any], event_ids: list[str]) -> None:
    state["seen_event_ids"] = list(state.get("seen_event_ids", [])) + event_ids
    save_state(state)


def main() -> int:
    state = load_state()
    seen = set(state.get("seen_event_ids", []))

    newly_seen_event_ids: list[str] = []
    newly_seen_non_alert_event_ids: list[str] = []

    aws_prefixes = fetch_aws_prefixes()

    end_time = datetime.now(timezone.utc)
    start_time = end_time - timedelta(minutes=LOOKBACK_MINUTES)

    findings = []

    for scan_region in REGIONS_TO_SCAN:
        for event in lookup_region_events(scan_region, start_time, end_time):
            event_id = event.get("eventID")

            if not event_id:
                continue

            if event_id in seen:
                continue

            newly_seen_event_ids.append(event_id)
            seen.add(event_id)

            should_alert, reason, prefixes = classify_event(event, aws_prefixes)

            if should_alert:
                finding = compact_event(
                    event=event,
                    scan_region=scan_region,
                    reason=reason,
                    prefixes=prefixes,
                )
                findings.append(finding)

            else:
                newly_seen_non_alert_event_ids.append(event_id)

    if not findings:
        # Clean run: save state and intentionally print nothing.
        update_seen_events(state, newly_seen_event_ids)
        return 0

    payload = build_rollup_payload(findings)
    if PRINT_ALERT_TO_STDOUT:
        print(json.dumps(payload, indent=2, sort_keys=True))

    if not WEBHOOK_URL:
        if newly_seen_non_alert_event_ids:
            update_seen_events(state, newly_seen_non_alert_event_ids)
        print("ERROR: findings exist but WEBHOOK_URL is not set", file=sys.stderr)
        return 1

    try:
        send_webhook(payload)
    except Exception as exc:
        if newly_seen_non_alert_event_ids:
            update_seen_events(state, newly_seen_non_alert_event_ids)
        print(f"ERROR: webhook delivery failed: {exc}", file=sys.stderr)
        return 1

    update_seen_events(state, newly_seen_event_ids)
    return ALERT_EXIT_CODE


if __name__ == "__main__":
    STATE_FILE.parent.mkdir(parents=True, exist_ok=True)
    with open(str(STATE_FILE) + ".lock", "w") as lock:
        try:
            fcntl.flock(lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
        except BlockingIOError:
            print("ERROR: a scan is already running", file=sys.stderr)
            sys.exit(1)
        sys.exit(main())
