#!/usr/bin/env python3
"""hopfind - work out how many hops away the thing blocking your port is.

Walks the IPv4 TTL, or the IPv6 hop limit, up from 1 and records which router
answers at each step - using the protocol and port you actually care about
instead of traceroute's default UDP high ports.

Run it twice. Once against something that works, once against the port that
does not. The hop where the answers stop is the device dropping you, and the
run that works gives you its address.

    python3 hopfind.py example.net 445            # the port under suspicion
    python3 hopfind.py example.net 443            # the reference run
    python3 hopfind.py example.net 53 --proto udp
    python3 hopfind.py 2001:db8::1 443 -6

On Linux this needs no privileges at all: IP_RECVERR and IPV6_RECVERR hand the
ICMP errors back on the ordinary socket that caused them. On macOS, the BSDs
and Solaris the errors have to be read off a raw ICMP socket, which means root.

Written for https://blogs.damiendye.uk/networking/how-far-away-is-the-firewall/
Public domain. Do what you like with it.
"""

import argparse
import errno
import os
import select
import socket
import struct
import sys
import time

# Linux socket options. Absent from the socket module on some builds, so they
# are spelled out rather than looked up.
IP_RECVERR = 11
IPV6_RECVERR = 25

# ee_origin values from linux/errqueue.h. Anything else means the errno came
# from the local stack rather than from a router.
SO_EE_ORIGIN_ICMP = 2
SO_EE_ORIGIN_ICMP6 = 3

ICMP_V4 = {
    (11, 0): "time exceeded in-transit",
    (11, 1): "fragment reassembly time exceeded",
    (3, 0): "net unreachable",
    (3, 1): "host unreachable",
    (3, 2): "protocol unreachable",
    (3, 3): "port unreachable",
    (3, 4): "fragmentation needed",
    (3, 9): "net administratively prohibited",
    (3, 10): "host administratively prohibited",
    (3, 13): "communication administratively prohibited",
    (5, 0): "redirect",
}

ICMP_V6 = {
    (3, 0): "hop limit exceeded in-transit",
    (3, 1): "fragment reassembly time exceeded",
    (1, 0): "no route to destination",
    (1, 1): "communication administratively prohibited",
    (1, 3): "address unreachable",
    (1, 4): "port unreachable",
    (2, 0): "packet too big",
}


def describe(family, icmp_type, icmp_code):
    table = ICMP_V4 if family == socket.AF_INET else ICMP_V6
    return table.get((icmp_type, icmp_code), "unrecognised")


def is_expiry(family, icmp_type):
    """Was this the router saying 'your hop budget ran out here'?"""
    return icmp_type == (11 if family == socket.AF_INET else 3)


class ErrorQueue:
    """Linux. The kernel reports the ICMP error on the socket that provoked it."""

    def arm(self, sock, family):
        if family == socket.AF_INET:
            sock.setsockopt(socket.IPPROTO_IP, IP_RECVERR, 1)
        else:
            sock.setsockopt(socket.IPPROTO_IPV6, IPV6_RECVERR, 1)

    def extra_readers(self):
        return []

    def collect(self, sock, family):
        try:
            _, ancillary, _, _ = sock.recvmsg(0, 1024, socket.MSG_ERRQUEUE)
        except OSError:
            return None
        wanted = (socket.IPPROTO_IP, IP_RECVERR) if family == socket.AF_INET \
            else (socket.IPPROTO_IPV6, IPV6_RECVERR)
        for level, kind, data in ancillary:
            if (level, kind) != wanted or len(data) < 16:
                continue
            # struct sock_extended_err, then the sockaddr of the router that
            # sent the error - SO_EE_OFFENDER in the kernel headers.
            _, origin, icmp_type, icmp_code = struct.unpack_from("=IBBB", data, 0)
            if origin not in (SO_EE_ORIGIN_ICMP, SO_EE_ORIGIN_ICMP6):
                return None
            addr = None
            if len(data) >= 24:
                offender_family, = struct.unpack_from("=H", data, 16)
                if offender_family == socket.AF_INET:
                    addr = socket.inet_ntoa(data[20:24])
                elif offender_family == socket.AF_INET6 and len(data) >= 40:
                    addr = socket.inet_ntop(socket.AF_INET6, data[24:40])
            return addr, icmp_type, icmp_code
        return None


class RawIcmp:
    """macOS, the BSDs, illumos, Solaris. Read the ICMP off a raw socket, as root."""

    def __init__(self, family):
        proto = socket.IPPROTO_ICMP if family == socket.AF_INET else socket.IPPROTO_ICMPV6
        self.sock = socket.socket(family, socket.SOCK_RAW, proto)
        self.sock.setblocking(False)

    def arm(self, sock, family):
        pass

    def extra_readers(self):
        return [self.sock]

    def collect(self, sock, family):
        try:
            packet, peer = self.sock.recvfrom(1500)
        except OSError:
            return None
        if family == socket.AF_INET:
            # BSD raw sockets hand back the IP header too.
            header_len = (packet[0] & 0x0F) * 4
            packet = packet[header_len:]
        if len(packet) < 2:
            return None
        return peer[0], packet[0], packet[1]


def probe(dest, port, proto, family, hop_limit, timeout, listener):
    """One probe at one hop limit. Returns (icmp, socket_state, note)."""
    kind = socket.SOCK_STREAM if proto == "tcp" else socket.SOCK_DGRAM
    sock = socket.socket(family, kind)
    if family == socket.AF_INET:
        sock.setsockopt(socket.IPPROTO_IP, socket.IP_TTL, hop_limit)
    else:
        sock.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_UNICAST_HOPS, hop_limit)
    listener.arm(sock, family)
    sock.setblocking(False)

    try:
        if kind == socket.SOCK_DGRAM:
            sock.connect((dest, port))
            sock.send(b"\x00" * 32)
        else:
            try:
                sock.connect((dest, port))
            except BlockingIOError:
                pass
    except OSError as exc:
        sock.close()
        return None, None, "local error: %s" % exc.strerror

    readers = [sock] + listener.extra_readers()
    writers = [] if kind == socket.SOCK_DGRAM else [sock]
    deadline = time.time() + timeout
    icmp = state = None

    while time.time() < deadline:
        ready_r, ready_w, ready_x = select.select(
            readers, writers, [sock], max(0.01, deadline - time.time()))
        if not (ready_r or ready_w or ready_x):
            continue
        # Drain the error queue first, always. An ICMP error reaches a TCP
        # socket as a plain errno, so SO_ERROR on its own cannot tell you
        # whether a router spoke or the far end did.
        icmp = listener.collect(sock, family)
        if icmp:
            break
        if ready_w:
            err = sock.getsockopt(socket.SOL_SOCKET, socket.SO_ERROR)
            if err == 0:
                state = "connected"
            elif err == errno.ECONNREFUSED:
                state = "TCP reset"
            else:
                state = os.strerror(err)
            break

    sock.close()
    if icmp or state:
        return icmp, state, None
    return None, None, "no reply"


def walk(dest, port, proto, family, first, last, timeout, listener):
    print("walking to %s  %s/%d  hop limit %d-%d" % (dest, proto.upper(), port, first, last))
    answered = []
    for hop in range(first, last + 1):
        started = time.time()
        icmp, state, _ = probe(dest, port, proto, family, hop, timeout, listener)
        rtt = (time.time() - started) * 1000
        if icmp:
            addr, icmp_type, icmp_code = icmp
            print(" %2d  %-39s %7.1f ms  ICMP %d/%d %s"
                  % (hop, addr or "?", rtt, icmp_type, icmp_code,
                     describe(family, icmp_type, icmp_code)))
            if is_expiry(family, icmp_type):
                answered.append((hop, addr))
            else:
                return answered, hop, "icmp-reject", addr
        elif state:
            print(" %2d  %-39s %7.1f ms  %s" % (hop, dest, rtt, state))
            return answered, hop, state, dest
        else:
            print(" %2d  *" % hop)
    return answered, None, "silent", None


def main():
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("host")
    parser.add_argument("port", nargs="?", type=int, default=443)
    parser.add_argument("--proto", choices=("tcp", "udp"), default="tcp")
    parser.add_argument("--first", type=int, default=1, help="hop limit to start at")
    parser.add_argument("--max", type=int, default=20, help="hop limit to stop at")
    parser.add_argument("--wait", type=float, default=2.0, help="seconds to wait per hop")
    parser.add_argument("-6", dest="v6", action="store_true", help="force IPv6")
    parser.add_argument("-4", dest="v4", action="store_true", help="force IPv4")
    args = parser.parse_args()

    family = socket.AF_INET6 if args.v6 else socket.AF_INET
    kind = socket.SOCK_STREAM if args.proto == "tcp" else socket.SOCK_DGRAM
    dest = socket.getaddrinfo(args.host, args.port, family, kind)[0][4][0]

    if sys.platform.startswith("linux"):
        listener = ErrorQueue()
    else:
        try:
            listener = RawIcmp(family)
        except PermissionError:
            sys.exit("%s cannot report ICMP errors on a normal socket, so this "
                     "needs a raw one. Run it as root." % sys.platform)

    answered, stop, why, who = walk(dest, args.port, args.proto, family,
                                    args.first, args.max, args.wait, listener)
    what = "%s/%d" % (args.proto.upper(), args.port)
    print()

    if why == "connected":
        print("Verdict: %s is open. It answered at hop %d." % (what, stop))
    elif why == "icmp-reject":
        print("Verdict: %s at hop %d is refusing %s on policy, and is honest "
              "enough to say so." % (who, stop, what))
    elif why == "TCP reset":
        print("Verdict: a reset came back to a probe with a hop limit of %d." % stop)
        print("         Nothing more than %s away can have sent it, so check the reply"
              % ("one hop" if stop == 1 else "%d hops" % stop))
        print("         TTL before you believe the host did.")
    elif answered:
        last_hop, last_addr = answered[-1]
        print("Verdict: answers stop after hop %d (%s)." % (last_hop, last_addr))
        print("         Whatever swallows %s is hop %d." % (what, last_hop + 1))
        print("         Walk a port that works and read off the address at hop %d."
              % (last_hop + 1))
    else:
        print("Verdict: nothing answered at all, not even the first hop. Either the")
        print("         first hop is the one dropping you, or the ICMP errors are being")
        print("         filtered on the way back. Walk a port that works to tell those")
        print("         two apart.")


if __name__ == "__main__":
    main()
