diff --git a/traceroute6/endpoint.go b/traceroute6/endpoint.go new file mode 100644 index 0000000..3879d34 --- /dev/null +++ b/traceroute6/endpoint.go @@ -0,0 +1,78 @@ +package traceroute6 + +import ( + "fmt" + "net" + "os" +) + +// localEndpoint resolves the source IPv6 and outbound interface for a +// destination. If conf.LocalAddr / conf.Interface are set they are validated +// and returned; otherwise they are auto-detected by consulting the routing +// table via a connect(2) on a UDP6 socket (no packets are sent). +func localEndpoint(conf *Config, dst net.IP) (srcIP net.IP, iface string, err error) { + if conf.LocalAddr != "" { + ip := net.ParseIP(conf.LocalAddr) + if ip == nil || ip.To4() != nil || ip.To16() == nil { + return nil, "", fmt.Errorf("invalid local IPv6 address: %q", conf.LocalAddr) + } + srcIP = ip + } else { + srcIP, err = outboundIP(dst) + if err != nil { + return nil, "", err + } + } + + iface = conf.Interface + if iface == "" { + iface = interfaceForIP(srcIP) + if iface == "" { + return nil, "", fmt.Errorf("cannot determine outbound interface for %s, use --interface/-I", srcIP) + } + } + return srcIP, iface, nil +} + +// outboundIP returns the local IPv6 the kernel would use to reach dst. +func outboundIP(dst net.IP) (net.IP, error) { + conn, err := net.Dial("udp6", net.JoinHostPort(dst.String(), "33434")) + if err != nil { + return nil, fmt.Errorf("resolve local IP for %s: %w", dst, err) + } + defer func() { _ = conn.Close() }() + if ua, ok := conn.LocalAddr().(*net.UDPAddr); ok { + if ip := ua.IP; ip != nil && ip.To16() != nil { + return ip, nil + } + } + return nil, fmt.Errorf("no IPv6 source address for %s", dst) +} + +// interfaceForIP returns the name of the interface that owns ip. +func interfaceForIP(ip net.IP) string { + ifaces, err := net.Interfaces() + if err != nil { + return "" + } + for _, iface := range ifaces { + if iface.Flags&net.FlagUp == 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + continue + } + for _, addr := range addrs { + if ipnet, ok := addr.(*net.IPNet); ok && ipnet.IP.Equal(ip) { + return iface.Name + } + } + } + return "" +} + +// pid returns the lower 16 bits of the process id, used to tag probes. +func pid() uint16 { + return uint16(os.Getpid() & 0xFFFF) +} diff --git a/traceroute6/endpoint_test.go b/traceroute6/endpoint_test.go new file mode 100644 index 0000000..9c550bd --- /dev/null +++ b/traceroute6/endpoint_test.go @@ -0,0 +1,69 @@ +package traceroute6 + +import ( + "net" + "testing" +) + +func TestLocalEndpointRejectsNonV6(t *testing.T) { + conf := &Config{LocalAddr: "10.0.0.1"} + if _, _, err := localEndpoint(conf, net.ParseIP("2001:4860:4860::8888")); err == nil { + t.Errorf("expected error for IPv4 local addr") + } + conf = &Config{LocalAddr: "garbage"} + if _, _, err := localEndpoint(conf, net.ParseIP("2001:4860:4860::8888")); err == nil { + t.Errorf("expected error for garbage local addr") + } +} + +// TestInterfaceForIP exercises the IP-to-interface mapping using whatever IPv6 +// address the loopback interface owns (::1 on most systems). +func TestInterfaceForIP(t *testing.T) { + ifaces, err := net.Interfaces() + if err != nil { + t.Skip("cannot enumerate interfaces") + } + var v6 net.IP + var wantIface string + for _, iface := range ifaces { + if iface.Flags&net.FlagUp == 0 { + continue + } + addrs, _ := iface.Addrs() + for _, addr := range addrs { + if ipnet, ok := addr.(*net.IPNet); ok && ipnet.IP.To4() == nil && ipnet.IP.To16() != nil { + v6 = ipnet.IP + wantIface = iface.Name + break + } + } + if v6 != nil { + break + } + } + if v6 == nil { + t.Skip("no IPv6 address on any interface") + } + if got := interfaceForIP(v6); got != wantIface { + t.Errorf("interfaceForIP(%s) = %q, want %q", v6, got, wantIface) + } + + // An address owned by no interface yields "". + if got := interfaceForIP(net.ParseIP("2001:db8::dead:beef")); got != "" { + t.Errorf("interfaceForIP of unowned addr = %q, want empty", got) + } +} + +func TestLocalEndpointExplicitInterface(t *testing.T) { + conf := &Config{LocalAddr: "2001:db8::1", Interface: "eth-test"} + src, iface, err := localEndpoint(conf, net.ParseIP("2001:4860:4860::8888")) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !src.Equal(net.ParseIP("2001:db8::1")) { + t.Errorf("src = %v, want 2001:db8::1", src) + } + if iface != "eth-test" { + t.Errorf("iface = %q, want eth-test", iface) + } +}