#!/usr/bin/env python3
"""
atlas-scan -- read-only AWS network scanner -> Atlas IR.

Runs locally with the caller's own AWS credentials (env vars, ~/.aws, or --profile).
Makes ONLY describe/list/get/search calls -- no create, modify, or delete. Nothing
leaves the machine; the output IR JSON is the only artifact.

Pulls everything the Atlas engine reasons about: VPCs, subnets, route tables (+ BGP
propagations), NACLs, security groups, gateways (IGW/VGW/DX), VPC peering, Transit
Gateways (+ attachments + TGW route tables), managed prefix lists, interface endpoints
(+ private IPs + policy), ALBs (+ targets/listeners), RDS, ElastiCache, OpenSearch,
EFS, VPC-attached Lambda, and Network Firewall endpoints. VPN/DX scaffolding seeds the
onPrem block (static-route VPNs are auto-filled; BGP/DX CIDRs are flagged for review).

Usage:
  python atlas_scan.py [--regions us-east-1,us-west-2] [--profile NAME] -o inventory.ir.json
  (no --regions = all enabled regions)

Requires: boto3   (pip install boto3)
"""
import argparse, json, sys, datetime, ipaddress

try:
    import boto3
except ImportError:
    sys.exit("boto3 required:  pip install boto3")

VERSION = "atlas-scan/2.0"


def tag_name(obj, default=""):
    for t in obj.get("Tags", []) or []:
        if t.get("Key") == "Name":
            return t["Value"]
    return default


def enabled_regions(session):
    ec2 = session.client("ec2", region_name="us-east-1")
    return [r["RegionName"] for r in ec2.describe_regions(
        Filters=[{"Name": "opt-in-status", "Values": ["opt-in-not-required", "opted-in"]}]
    )["Regions"]]


# ---------------------------------------------------------------- security groups
def sg_rules(ec2, vpc_id):
    out = []
    for sg in ec2.describe_security_groups(Filters=[{"Name": "vpc-id", "Values": [vpc_id]}])["SecurityGroups"]:
        def perms(plist, key):
            rules = []
            for p in plist:
                proto = p.get("IpProtocol", "-1")
                lo, hi = p.get("FromPort"), p.get("ToPort")
                ports = "all" if lo is None else (str(lo) if lo == hi else f"{lo}-{hi}")
                for rng in p.get("IpRanges", []):
                    rules.append({"proto": proto, "ports": ports, key: rng["CidrIp"]})
                for grp in p.get("UserIdGroupPairs", []):
                    rules.append({"proto": proto, "ports": ports, key: grp["GroupId"]})
                for pl in p.get("PrefixListIds", []):
                    rules.append({"proto": proto, "ports": ports, key: pl["PrefixListId"]})
            return rules
        out.append({"id": sg["GroupId"], "name": sg.get("GroupName", sg["GroupId"]),
                    "ingress": perms(sg.get("IpPermissions", []), "source"),
                    "egress": perms(sg.get("IpPermissionsEgress", []), "dest")})
    return out


# ---------------------------------------------------------------- NACLs
def nacls(ec2, vpc_id):
    out = []
    for acl in ec2.describe_network_acls(Filters=[{"Name": "vpc-id", "Values": [vpc_id]}])["NetworkAcls"]:
        def entries(egress):
            r = []
            for e in acl["Entries"]:
                if e["Egress"] != egress:
                    continue
                pr = e.get("PortRange", {})
                ports = "0-65535" if not pr else f"{pr.get('From', 0)}-{pr.get('To', 65535)}"
                r.append({"rule": e["RuleNumber"], "action": e["RuleAction"],
                          "proto": e["Protocol"], "cidr": e.get("CidrBlock", "0.0.0.0/0"), "ports": ports})
            return sorted(r, key=lambda x: x["rule"])
        out.append({"id": acl["NetworkAclId"], "name": tag_name(acl, acl["NetworkAclId"]),
                    "associations": [a["SubnetId"] for a in acl.get("Associations", [])],
                    "inbound": entries(False), "outbound": entries(True)})
    return out


# ---------------------------------------------------------------- route tables (+ propagations)
def route_target(r):
    return (r.get("GatewayId") or r.get("NatGatewayId") or r.get("TransitGatewayId")
            or r.get("VpcPeeringConnectionId") or r.get("EgressOnlyInternetGatewayId")
            or r.get("LocalGatewayId") or r.get("CarrierGatewayId")
            or r.get("NetworkInterfaceId") or r.get("InstanceId") or "local")


def route_tables(ec2, vpc_id):
    out = []
    for rt in ec2.describe_route_tables(Filters=[{"Name": "vpc-id", "Values": [vpc_id]}])["RouteTables"]:
        routes = []
        for r in rt["Routes"]:
            if r.get("State") == "blackhole":
                continue
            dest = r.get("DestinationCidrBlock") or r.get("DestinationPrefixListId") or r.get("DestinationIpv6CidrBlock") or "?"
            routes.append({"dest": dest, "target": route_target(r)})
        out.append({"id": rt["RouteTableId"], "name": tag_name(rt, rt["RouteTableId"]),
                    "main": any(a.get("Main") for a in rt.get("Associations", [])),
                    "associations": [a["SubnetId"] for a in rt.get("Associations", []) if a.get("SubnetId")],
                    "propagations": [p["GatewayId"] for p in rt.get("PropagatingVgws", [])],
                    "routes": routes})
    return out


def subnet_is_public(subnet_id, rts):
    for rt in rts:
        if subnet_id in rt["associations"]:
            return any(str(r["target"]).startswith("igw-") for r in rt["routes"])
    # fall back to the main route table for unassociated subnets
    for rt in rts:
        if rt.get("main"):
            return any(str(r["target"]).startswith("igw-") for r in rt["routes"])
    return False


# ---------------------------------------------------------------- gateways (IGW / VGW / DX)
def gateways(ec2, vpc_id, dx_by_vpc):
    g = []
    for igw in ec2.describe_internet_gateways(Filters=[{"Name": "attachment.vpc-id", "Values": [vpc_id]}])["InternetGateways"]:
        g.append({"id": igw["InternetGatewayId"], "type": "internet-gateway", "name": tag_name(igw, "igw")})
    for vgw in ec2.describe_vpn_gateways(Filters=[{"Name": "attachment.vpc-id", "Values": [vpc_id]}])["VpnGateways"]:
        g.append({"id": vgw["VpnGatewayId"], "type": "vpn-gateway", "name": tag_name(vgw, "vgw")})
    for dx in dx_by_vpc.get(vpc_id, []):
        g.append(dx)
    return g


# ---------------------------------------------------------------- managed prefix lists
def prefix_lists(ec2):
    out = []
    try:
        for pl in ec2.describe_managed_prefix_lists()["PrefixLists"]:
            plid = pl["PrefixListId"]
            try:
                ents = ec2.get_managed_prefix_list_entries(PrefixListId=plid)["Entries"]
                out.append({"id": plid, "name": pl.get("PrefixListName", plid),
                            "entries": [e["Cidr"] for e in ents]})
            except Exception:
                pass
    except Exception:
        pass
    return out


# ---------------------------------------------------------------- transit gateways
def transit_gateways(ec2, account):
    tgws = {}
    try:
        owned = ec2.describe_transit_gateways()["TransitGateways"]
    except Exception:
        return []
    for t in owned:
        if t.get("State") in ("deleting", "deleted"):
            continue
        tgws[t["TransitGatewayId"]] = {"id": t["TransitGatewayId"], "name": tag_name(t, t["TransitGatewayId"]),
                                       "attachments": [], "routeTables": []}
    if not tgws:
        return []
    # attachment -> associated route table
    att_rtb = {}
    rtbs_by_tgw = {}
    try:
        for rtb in ec2.describe_transit_gateway_route_tables()["TransitGatewayRouteTables"]:
            if rtb.get("State") in ("deleting", "deleted"):
                continue
            rtbs_by_tgw.setdefault(rtb["TransitGatewayId"], []).append(rtb["TransitGatewayRouteTableId"])
            try:
                for a in ec2.get_transit_gateway_route_table_associations(
                        TransitGatewayRouteTableId=rtb["TransitGatewayRouteTableId"])["Associations"]:
                    att_rtb[a["TransitGatewayAttachmentId"]] = rtb["TransitGatewayRouteTableId"]
            except Exception:
                pass
    except Exception:
        pass
    # vpc attachments
    try:
        for a in ec2.describe_transit_gateway_vpc_attachments()["TransitGatewayVpcAttachments"]:
            if a.get("State") in ("deleting", "deleted"):
                continue
            tg = tgws.get(a["TransitGatewayId"])
            if not tg:
                continue
            tg["attachments"].append({"id": a["TransitGatewayAttachmentId"], "vpc": a["VpcId"],
                                      "rtb": att_rtb.get(a["TransitGatewayAttachmentId"], "")})
    except Exception:
        pass
    # route tables + routes (search, active only)
    for tgwid, rtbids in rtbs_by_tgw.items():
        tg = tgws.get(tgwid)
        if not tg:
            continue
        for rtbid in rtbids:
            routes = []
            try:
                res = ec2.search_transit_gateway_routes(
                    TransitGatewayRouteTableId=rtbid,
                    Filters=[{"Name": "state", "Values": ["active"]}])
                for r in res.get("Routes", []):
                    atts = r.get("TransitGatewayAttachments", [])
                    if not atts:
                        continue
                    routes.append({"dest": r["DestinationCidrBlock"], "target": atts[0]["TransitGatewayAttachmentId"]})
            except Exception:
                pass
            tg["routeTables"].append({"id": rtbid, "routes": routes})
    return list(tgws.values())


# ---------------------------------------------------------------- Network Firewall (best-effort)
def firewalls(session, region):
    out = []
    try:
        nfw = session.client("network-firewall", region_name=region)
    except Exception:
        return out
    try:
        names = nfw.list_firewalls().get("Firewalls", [])
    except Exception:
        return out
    for f in names:
        try:
            d = nfw.describe_firewall(FirewallArn=f["FirewallArn"])
            fw = d["Firewall"]
            sync = d.get("FirewallStatus", {}).get("SyncStates", {})
            endpoints = [v.get("Attachment", {}).get("EndpointId") for v in sync.values() if v.get("Attachment", {}).get("EndpointId")]
            entry = {"id": fw["FirewallName"], "name": fw["FirewallName"], "type": "network-firewall",
                     "endpoints": endpoints, "defaultAction": "allow", "rules": [],
                     "note": "stateful (Suricata) rules are NOT translated; only stateless 5-tuple drops/allows are modeled"}
            # best-effort: stateless default actions from the policy
            try:
                pol = nfw.describe_firewall_policy(FirewallPolicyArn=fw["FirewallPolicyArn"])["FirewallPolicy"]
                defaults = pol.get("StatelessDefaultActions", [])
                if any("drop" in a.lower() or "forward" not in a.lower() and "pass" not in a.lower() for a in defaults) and "aws:forward_to_sfe" not in defaults:
                    entry["defaultAction"] = "drop" if any("drop" in a.lower() for a in defaults) else "allow"
            except Exception:
                pass
            out.append(entry)
        except Exception:
            pass
    return out


# ---------------------------------------------------------------- on-prem (auto-derived from routes)
def collect_onprem(ec2, out_vpcs, dx_by_vpc):
    """Remote networks reachable from the VPC, split by kind:
         kind 'vpn'    -> a Site-to-Site VPN connection (a remote site/partner), CIDRs from its static routes
         kind 'onprem' -> your own datacenter reached over Direct Connect (or a BGP-only VGW)
       Link-local (169.254.x) and /32 host routes are BGP peering addresses, not networks, so dropped.
       Everything is verified:false for your review."""
    remotes = {}
    dxgw_ids = {g["id"] for lst in dx_by_vpc.values() for g in lst}
    vpn_route_cidrs = {}   # attachment -> set(cidrs claimed by a VPN connection)

    vgw_with_dx = set()
    for v in out_vpcs:
        gws = v.get("gateways", [])
        if any(g.get("type") == "direct-connect-gateway" for g in gws):
            for g in gws:
                if g.get("type") == "vpn-gateway":
                    vgw_with_dx.add(g["id"])

    def is_onprem_prefix(dest):
        d = str(dest)
        if d in ("local", "0.0.0.0/0", "::/0", "?", ""):
            return False
        if d.startswith("pl-"):
            return False
        if d.startswith("169.254."):   # link-local: VPN tunnel BGP peering address, never a real network
            return False
        return True                    # keep /32s -- they may be SNAT IPs or specific remote hosts

    # 1) Site-to-Site VPN connections -> their own entities (kind 'vpn')
    try:
        cgws = {cg["CustomerGatewayId"]: cg for cg in ec2.describe_customer_gateways().get("CustomerGateways", [])}
    except Exception:
        cgws = {}
    try:
        for vc in ec2.describe_vpn_connections()["VpnConnections"]:
            if vc.get("State") in ("deleting", "deleted"):
                continue
            att = vc.get("VpnGatewayId") or vc.get("TransitGatewayId") or ""
            vid = vc["VpnConnectionId"]
            static_only = bool((vc.get("Options") or {}).get("StaticRoutesOnly"))
            cidrs = []
            for r in vc.get("Routes", []):
                d = r.get("DestinationCidrBlock")
                if d and is_onprem_prefix(d):
                    cidrs.append(d)
                    vpn_route_cidrs.setdefault(att, set()).add(d)
            cg = cgws.get(vc.get("CustomerGatewayId"), {})
            meta = {"vpnType": vc.get("Type"), "category": vc.get("Category"),
                    "customerGatewayId": vc.get("CustomerGatewayId"),
                    "peerIp": cg.get("IpAddress"), "peerAsn": cg.get("BgpAsn"),
                    "routing": "static" if static_only else "bgp",
                    "tunnels": [t.get("OutsideIpAddress") for t in vc.get("VgwTelemetry", []) if t.get("OutsideIpAddress")]}
            name = tag_name(vc, vid)
            if name == vid and meta["peerIp"]:
                name = "VPN to " + meta["peerIp"]
            remotes[vid] = {"id": vid, "name": name, "kind": "vpn",
                            "via": "Site-to-Site VPN", "cidrs": list(dict.fromkeys(cidrs)),
                            "attachment": att, "advertisedRoutes": [], "verified": False, "meta": meta}
    except Exception:
        pass

    def in_own(dest, own):
        try:
            n = ipaddress.ip_network(dest, strict=False)
        except Exception:
            return False
        for c in own:
            try:
                if n.subnet_of(ipaddress.ip_network(c, strict=False)):
                    return True
            except Exception:
                pass
        return False

    # 2) route-table prefixes to a VGW/DX that AREN'T claimed by a VPN site -> on-prem (kind 'onprem')
    onprem_by_att = {}
    for v in out_vpcs:
        own = [c for c in (v.get("cidrs") or [v.get("cidr", "")]) if c]
        for rt in v.get("routeTables", []):
            for r in rt.get("routes", []):
                t, dest = r.get("target", ""), r.get("dest", "")
                if not (str(t).startswith("vgw-") or t in dxgw_ids):
                    continue
                if not is_onprem_prefix(dest) or in_own(dest, own):
                    continue
                if dest in vpn_route_cidrs.get(t, set()):
                    continue
                onprem_by_att.setdefault(t, []).append(dest)
    for att, cidrs in onprem_by_att.items():
        is_dx = att in dxgw_ids or att in vgw_with_dx
        remotes["onprem:" + att] = {"id": att, "name": "on-prem via " + att, "kind": "onprem",
                                    "via": "Direct Connect" if is_dx else "VPN",
                                    "cidrs": list(dict.fromkeys(cidrs)), "attachment": att,
                                    "advertisedRoutes": [], "verified": False}
    return list(remotes.values())


# ---------------------------------------------------------------- endpoint policy (best-effort, conservative)
def endpoint_allowed_sources(policy_doc):
    """Only extract source restrictions we can map CLEANLY to CIDRs (IpAddress conditions).
       Anything else is left for manual review -- we never fabricate a verdict."""
    try:
        pol = json.loads(policy_doc) if isinstance(policy_doc, str) else policy_doc
    except Exception:
        return None
    allowed = []
    for st in pol.get("Statement", []):
        if st.get("Effect") != "Allow":
            continue
        cond = st.get("Condition", {}) or {}
        for op in ("IpAddress",):
            for key, vals in (cond.get(op, {}) or {}).items():
                if "sourceip" in key.lower():
                    allowed += vals if isinstance(vals, list) else [vals]
    return allowed or None


# ---------------------------------------------------------------- ENI IP resolution
def eni_ip_by_id(ec2, eni_ids):
    out = {}
    if not eni_ids:
        return out
    try:
        for ni in ec2.describe_network_interfaces(NetworkInterfaceIds=eni_ids)["NetworkInterfaces"]:
            out[ni["NetworkInterfaceId"]] = ni.get("PrivateIpAddress", "")
    except Exception:
        pass
    return out


def scan_region(session, region, account):
    ec2 = session.client("ec2", region_name=region)
    elbv2 = session.client("elbv2", region_name=region)
    rds = session.client("rds", region_name=region)

    vpcs_raw = ec2.describe_vpcs()["Vpcs"]
    if not vpcs_raw:
        return None, []

    notes = []
    peerings = ec2.describe_vpc_peering_connections(
        Filters=[{"Name": "status-code", "Values": ["active"]}])["VpcPeeringConnections"]

    # Direct Connect gateways are global; associate to VPCs via their gateway associations
    dx_by_vpc = {}
    try:
        dc = session.client("directconnect", region_name=region)
        dxgws = dc.describe_direct_connect_gateways().get("directConnectGateways", [])
        for dxg in dxgws:
            assocs = dc.describe_direct_connect_gateway_associations(
                directConnectGatewayId=dxg["directConnectGatewayId"]).get("directConnectGatewayAssociations", [])
            for a in assocs:
                # associatedGateway can be a VGW or TGW; we surface the DX gateway on the VPC the VGW attaches to (best-effort)
                gid = a.get("associatedGateway", {}).get("id", "")
                if gid.startswith("vgw-"):
                    for att in ec2.describe_vpn_gateways(VpnGatewayIds=[gid]).get("VpnGateways", []):
                        for va in att.get("VpcAttachments", []):
                            if va.get("State") == "attached":
                                dx_by_vpc.setdefault(va["VpcId"], []).append(
                                    {"id": dxg["directConnectGatewayId"], "type": "direct-connect-gateway",
                                     "name": dxg.get("directConnectGatewayName", "Direct Connect")})
    except Exception:
        pass

    # pre-index resources by subnet
    nat_by, vpce_by, alb_by, rds_by, cache_by, os_by, efs_by, lam_by = ({} for _ in range(8))

    for nat in ec2.describe_nat_gateways()["NatGateways"]:
        if nat["State"] != "available":
            continue
        ip = (nat.get("NatGatewayAddresses") or [{}])[0].get("PrivateIp", "")
        nat_by.setdefault(nat["SubnetId"], []).append(
            {"id": nat["NatGatewayId"], "type": "nat-gateway", "name": tag_name(nat, "nat"), "ip": ip, "sgs": [], "meta": {}})

    for ep in ec2.describe_vpc_endpoints()["VpcEndpoints"]:
        if ep["VpcEndpointType"] != "Interface":
            continue
        ip_map = eni_ip_by_id(ec2, ep.get("NetworkInterfaceIds", []))
        ip = next(iter(ip_map.values()), "")
        meta = {"endpointType": "Interface", "service": ep["ServiceName"]}
        pdoc = ep.get("PolicyDocument")
        if pdoc:
            meta["policyDocument"] = pdoc  # raw, for transparency
            allowed = endpoint_allowed_sources(pdoc)
            if allowed:
                meta["policy"] = {"allowedSources": allowed}
        for sn in ep.get("SubnetIds", []):
            vpce_by.setdefault(sn, []).append(
                {"id": ep["VpcEndpointId"], "type": "vpc-endpoint", "name": ep["ServiceName"].split(".")[-1],
                 "ip": ip, "sgs": [g["GroupId"] for g in ep.get("Groups", [])], "meta": dict(meta)})

    # ALBs: resolve ENI IP + listeners + targets
    for lb in elbv2.describe_load_balancers()["LoadBalancers"]:
        arn = lb["LoadBalancerArn"]
        listeners = []
        try:
            listeners = [str(l["Port"]) for l in elbv2.describe_listeners(LoadBalancerArn=arn).get("Listeners", [])]
        except Exception:
            pass
        targets = []
        try:
            for tg in elbv2.describe_target_groups(LoadBalancerArn=arn).get("TargetGroups", []):
                for th in elbv2.describe_target_health(TargetGroupArn=tg["TargetGroupArn"]).get("TargetHealthDescriptions", []):
                    tid = th.get("Target", {}).get("Id", "")
                    if tid.startswith("i-"):
                        targets.append(tid)
        except Exception:
            pass
        ip_map = eni_ip_by_id(ec2, [])  # ALB ENIs resolved below by description filter
        try:
            nis = ec2.describe_network_interfaces(
                Filters=[{"Name": "description", "Values": [f"ELB {lb['LoadBalancerName']}", f"ELB {arn.split('/')[-3]}/{arn.split('/')[-2]}/{arn.split('/')[-1]}"]}]).get("NetworkInterfaces", [])
            alb_ip = next((ni.get("PrivateIpAddress", "") for ni in nis), "")
        except Exception:
            alb_ip = ""
        for az in lb.get("AvailabilityZones", []):
            alb_by.setdefault(az["SubnetId"], []).append(
                {"id": arn.split("/")[-1], "type": "alb", "name": lb["LoadBalancerName"], "ip": alb_ip,
                 "sgs": lb.get("SecurityGroups", []),
                 "meta": {"scheme": lb["Scheme"], "listeners": ",".join(listeners)},
                 "targets": sorted(set(targets))})
            break

    for db in rds.describe_db_instances()["DBInstances"]:
        subs = [s["SubnetIdentifier"] for s in db.get("DBSubnetGroup", {}).get("Subnets", [])][:1]
        for sn in subs:
            rds_by.setdefault(sn, []).append(
                {"id": db["DBInstanceIdentifier"], "type": "rds", "name": db["DBInstanceIdentifier"], "ip": "",
                 "sgs": [g["VpcSecurityGroupId"] for g in db.get("VpcSecurityGroups", [])],
                 "meta": {"engine": db["Engine"], "publiclyAccessible": db.get("PubliclyAccessible", False)}})

    # ElastiCache (resolve node ENI IPs via the cache SGs is unreliable; place by subnet group, ip best-effort)
    try:
        ecc = session.client("elasticache", region_name=region)
        sng = {g["CacheSubnetGroupName"]: [s["SubnetIdentifier"] for s in g.get("Subnets", [])]
               for g in ecc.describe_cache_subnet_groups().get("CacheSubnetGroups", [])}
        for c in ecc.describe_cache_clusters(ShowCacheNodeInfo=True).get("CacheClusters", []):
            subs = sng.get(c.get("CacheSubnetGroupName", ""), [])
            sn = subs[0] if subs else None
            if not sn:
                continue
            sgs = [g["SecurityGroupId"] for g in c.get("SecurityGroups", [])]
            ip = ""
            cache_by.setdefault(sn, []).append(
                {"id": c["CacheClusterId"], "type": "elasticache", "name": c["CacheClusterId"], "ip": ip,
                 "sgs": sgs, "meta": {"engine": c.get("Engine", "")}})
    except Exception:
        pass

    # OpenSearch
    try:
        oss = session.client("opensearch", region_name=region)
        for dn in oss.list_domain_names().get("DomainNames", []):
            try:
                d = oss.describe_domain(DomainName=dn["DomainName"])["DomainStatus"]
            except Exception:
                continue
            vo = d.get("VPCOptions", {})
            subs = vo.get("SubnetIds", [])
            if not subs:
                continue
            os_by.setdefault(subs[0], []).append(
                {"id": d["DomainName"], "type": "opensearch", "name": d["DomainName"], "ip": "",
                 "sgs": vo.get("SecurityGroupIds", []), "meta": {"engine": "opensearch"}})
    except Exception:
        pass

    flow_log_vpcs = {fl["ResourceId"] for fl in ec2.describe_flow_logs()["FlowLogs"]}

    try:
        efs = session.client("efs", region_name=region)
        for fsd in efs.describe_file_systems().get("FileSystems", []):
            fsid = fsd["FileSystemId"]
            for mt in efs.describe_mount_targets(FileSystemId=fsid).get("MountTargets", []):
                try:
                    sgs = efs.describe_mount_target_security_groups(MountTargetId=mt["MountTargetId"]).get("SecurityGroups", [])
                except Exception:
                    sgs = []
                efs_by.setdefault(mt["SubnetId"], []).append(
                    {"id": mt["MountTargetId"], "type": "efs", "name": fsd.get("Name", fsid),
                     "ip": mt.get("IpAddress", ""), "sgs": sgs, "meta": {"fileSystemId": fsid}})
    except Exception:
        pass
    try:
        lam = session.client("lambda", region_name=region)
        for page in lam.get_paginator("list_functions").paginate():
            for fn in page.get("Functions", []):
                vc = fn.get("VpcConfig") or {}
                subs = vc.get("SubnetIds") or []
                if not subs:
                    continue
                lam_by.setdefault(subs[0], []).append(
                    {"id": fn["FunctionName"], "type": "lambda", "name": fn["FunctionName"], "ip": "",
                     "sgs": vc.get("SecurityGroupIds", []), "meta": {"runtime": fn.get("Runtime", ""), "vpcAttached": True}})
    except Exception:
        pass

    def res_in(sn):
        out = []
        for r in ec2.describe_instances(Filters=[{"Name": "subnet-id", "Values": [sn]}])["Reservations"]:
            for inst in r["Instances"]:
                if inst["State"]["Name"] == "terminated":
                    continue
                out.append({"id": inst["InstanceId"], "type": "ec2", "name": tag_name(inst, inst["InstanceId"]),
                            "ip": inst.get("PrivateIpAddress", ""), "sgs": [s["GroupId"] for s in inst.get("SecurityGroups", [])],
                            "meta": {"instanceType": inst.get("InstanceType", ""), "state": inst["State"]["Name"]}})
        for idx in (nat_by, vpce_by, alb_by, rds_by, cache_by, os_by, efs_by, lam_by):
            out += idx.get(sn, [])
        return out

    out_vpcs, edges = [], []
    for v in vpcs_raw:
        vid = v["VpcId"]
        rts = route_tables(ec2, vid)
        subnets = ec2.describe_subnets(Filters=[{"Name": "vpc-id", "Values": [vid]}])["Subnets"]
        az_map = {}
        for sn in subnets:
            rt_id = next((rt["id"] for rt in rts if sn["SubnetId"] in rt["associations"]),
                         next((rt["id"] for rt in rts if rt.get("main")), ""))
            az_map.setdefault(sn["AvailabilityZone"], []).append({
                "id": sn["SubnetId"], "name": tag_name(sn, sn["SubnetId"]), "cidr": sn["CidrBlock"],
                "public": subnet_is_public(sn["SubnetId"], rts), "routeTable": rt_id, "resources": res_in(sn["SubnetId"])})
        cidrs = [a.get("CidrBlock") for a in (v.get("CidrBlockAssociationSet") or [])
                 if a.get("CidrBlock") and (a.get("CidrBlockState", {}) or {}).get("State", "associated") == "associated"]
        cidr = v.get("CidrBlock") or (cidrs[0] if cidrs else "")
        if cidr and cidr not in cidrs:
            cidrs = [cidr] + cidrs
        out_vpcs.append({
            "id": vid, "name": tag_name(v, vid), "cidr": cidr, "cidrs": cidrs,
            "flowLogs": vid in flow_log_vpcs, "gateways": gateways(ec2, vid, dx_by_vpc),
            "routeTables": rts, "nacls": nacls(ec2, vid), "securityGroups": sg_rules(ec2, vid),
            "azs": [{"az": az, "subnets": s} for az, s in sorted(az_map.items())]})
    for px in peerings:
        edges.append({"type": "peering", "id": px["VpcPeeringConnectionId"],
                      "from": px["RequesterVpcInfo"]["VpcId"], "to": px["AccepterVpcInfo"]["VpcId"],
                      "meta": {"accepted": True}})

    region_obj = {"region": region, "vpcs": out_vpcs, "edges": edges,
                  "transitGateways": transit_gateways(ec2, account),
                  "prefixLists": prefix_lists(ec2),
                  "firewalls": firewalls(session, region)}
    onprem = collect_onprem(ec2, out_vpcs, dx_by_vpc)
    return region_obj, onprem


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--regions", help="comma list; default = all enabled")
    ap.add_argument("--profile")
    ap.add_argument("-o", "--out", default="inventory.ir.json")
    args = ap.parse_args()

    session = boto3.Session(profile_name=args.profile) if args.profile else boto3.Session()
    acct = session.client("sts").get_caller_identity()["Account"]
    regions = args.regions.split(",") if args.regions else enabled_regions(session)

    ir = {"irVersion": "1.0", "generatedAt": datetime.datetime.utcnow().isoformat() + "Z",
          "generatedBy": VERSION, "account": acct, "onPrem": [], "regions": []}
    for rg in regions:
        sys.stderr.write(f"scanning {rg} ...\n")
        try:
            r, onprem = scan_region(session, rg, acct)
            if r and r["vpcs"]:
                ir["regions"].append(r)
                ir["onPrem"].extend(onprem)
        except Exception as e:
            sys.stderr.write(f"  skipped {rg}: {e}\n")

    json.dump(ir, open(args.out, "w"), indent=2)
    tgw = sum(len(r.get("transitGateways", [])) for r in ir["regions"])
    print(f"wrote {args.out}  ({len(ir['regions'])} regions, {tgw} transit gateways, "
          f"{len(ir['onPrem'])} on-prem links, account {acct})")
    if ir["onPrem"]:
        print("on-prem CIDRs were derived from routes to VGW/DX gateways (verified:false). Review them; "
              "if route propagation is disabled on a route table, its prefixes won't appear and need adding by hand.")


if __name__ == "__main__":
    main()
