From b2b1a10f82bf2546763b6f480189ee66cfcf0729 Mon Sep 17 00:00:00 2001 From: v-byte-cpu <65545655+v-byte-cpu@users.noreply.github.com> Date: Wed, 12 Aug 2026 03:27:46 +0400 Subject: [PATCH] feat!: add IPv6 scanning support Add NDP neighbor discovery and dual-stack support for ICMP, TCP, UDP, and SOCKS scans, including scoped link-local addresses. Replace the ARP-specific cache with a shared neighbor cache and add family-aware interface, routing, BPF, and packet handling. BREAKING CHANGE: Public scan and IP APIs now use netip.Addr and netip.Prefix. ARP cache APIs moved from pkg/scan/arp to pkg/scan/neighbor. --- README.md | 34 +- command/arp.go | 7 +- command/arp_test.go | 3 +- command/config.go | 268 ++++++++++++---- command/config_test.go | 96 ++++-- command/icmp.go | 36 ++- command/icmp_test.go | 16 +- command/log/logger_test.go | 18 +- command/ndp.go | 126 ++++++++ command/ndp_test.go | 28 ++ command/root.go | 3 +- command/tcp.go | 42 ++- command/tcp_test.go | 4 +- command/udp.go | 27 +- command/udp_test.go | 4 +- pkg/ip/ip.go | 94 ++++-- pkg/ip/ip_darwin.go | 73 ++++- pkg/ip/ip_darwin_test.go | 24 +- pkg/ip/ip_linux.go | 31 +- pkg/ip/ip_other.go | 5 +- pkg/ip/ip_test.go | 94 ++++-- pkg/packet/afpacket/readwriter.go | 8 +- pkg/packet/afpacket/readwriter_darwin.go | 36 ++- pkg/packet/afpacket/readwriter_darwin_test.go | 27 +- pkg/packet/afpacket/readwriter_other.go | 2 +- pkg/scan/arp/arp.go | 25 +- pkg/scan/arp/arp_test.go | 3 +- pkg/scan/arp/bpf.go | 4 +- pkg/scan/arp/bpf_test.go | 11 +- pkg/scan/arp/cache.go | 96 ------ pkg/scan/arp/cache_test.go | 301 ------------------ pkg/scan/arp/result_easyjson.go | 106 ------ pkg/scan/engine.go | 6 +- pkg/scan/engine_test.go | 18 +- pkg/scan/generator_test.go | 33 +- pkg/scan/icmp/bpf.go | 12 +- pkg/scan/icmp/bpf_test.go | 11 +- pkg/scan/icmp/icmp.go | 154 ++++++++- pkg/scan/icmp/icmp_test.go | 29 +- pkg/scan/icmp/ipv6_test.go | 56 ++++ pkg/scan/icmp/result_easyjson.go | 13 +- pkg/scan/ip_test.go | 7 + pkg/scan/mock_request_test.go | 8 +- pkg/scan/ndp/bpf.go | 17 + pkg/scan/ndp/bpf_test.go | 17 + pkg/scan/ndp/ndp.go | 134 ++++++++ pkg/scan/ndp/ndp_test.go | 85 +++++ pkg/scan/neighbor/cache.go | 124 ++++++++ pkg/scan/neighbor/cache_test.go | 72 +++++ .../{arp => neighbor}/mock_request_test.go | 6 +- pkg/scan/neighbor/result.go | 30 ++ pkg/scan/request.go | 106 +++--- pkg/scan/request_test.go | 268 ++++++++-------- pkg/scan/socks5/ipv6_test.go | 19 ++ pkg/scan/socks5/socks5.go | 20 +- pkg/scan/tcp/bpf.go | 13 +- pkg/scan/tcp/bpf_test.go | 20 +- pkg/scan/tcp/ipv6_test.go | 70 ++++ pkg/scan/tcp/tcp.go | 117 ++++++- pkg/scan/tcp/tcp_test.go | 29 +- pkg/scan/udp/ipv6_test.go | 31 ++ pkg/scan/udp/udp.go | 70 +++- pkg/scan/udp/udp_test.go | 29 +- 63 files changed, 2199 insertions(+), 1077 deletions(-) create mode 100644 command/ndp.go create mode 100644 command/ndp_test.go delete mode 100644 pkg/scan/arp/cache.go delete mode 100644 pkg/scan/arp/cache_test.go delete mode 100644 pkg/scan/arp/result_easyjson.go create mode 100644 pkg/scan/icmp/ipv6_test.go create mode 100644 pkg/scan/ip_test.go create mode 100644 pkg/scan/ndp/bpf.go create mode 100644 pkg/scan/ndp/bpf_test.go create mode 100644 pkg/scan/ndp/ndp.go create mode 100644 pkg/scan/ndp/ndp_test.go create mode 100644 pkg/scan/neighbor/cache.go create mode 100644 pkg/scan/neighbor/cache_test.go rename pkg/scan/{arp => neighbor}/mock_request_test.go (90%) create mode 100644 pkg/scan/neighbor/result.go create mode 100644 pkg/scan/socks5/ipv6_test.go create mode 100644 pkg/scan/tcp/ipv6_test.go create mode 100644 pkg/scan/udp/ipv6_test.go diff --git a/README.md b/README.md index ca079dd..44e122d 100644 --- a/README.md +++ b/README.md @@ -31,7 +31,8 @@ The goal of this project is to create the fastest network scanner with clean and ## ✨ Features * **⚡ 30x times faster** than nmap - * **ARP scan**: Scan your local networks to detect live devices + * **ARP and NDP scans**: Discover IPv4 and IPv6 neighbors on local networks + * **Dual-stack scanning**: ICMP, TCP, UDP, and SOCKS5 support IPv4 and IPv6 * **ICMP scan**: Use advanced ICMP scanning techniques to detect live hosts and firewall rules * **TCP SYN scan**: Traditional half-open scan to find open TCP ports * **TCP FIN / NULL / Xmas scans**: Scan techniques to bypass some firewall rules @@ -111,20 +112,28 @@ Live scan mode that rescans network every 10 seconds: sx arp 192.168.0.1/24 --live 10s ``` +ARP is an IPv4-only protocol. For IPv6, use Neighbor Discovery Protocol (NDP): + +``` +sx ndp --json 'fe80::%en0/120' | tee neighbor.cache +``` + +Scoped link-local hosts and prefixes are supported. The interface zone is retained in JSON output, for example `fe80::1%en0`. NDP also supports `--live` and `--file` in the same way as ARP and the other IP scanners. + ### TCP scan -Unlike nmap and other scanners that implicitly perform ARP requests to resolve IP addresses to MAC addresses before the actual scan, `sx` explicitly uses the **ARP cache** concept. ARP cache file is a simple text file containing JSON string on each line ([JSONL](https://jsonlines.org/) file), which has the same JSON fields as the ARP scan JSON output described above. Scans of higher-level protocols like TCP and UDP read the ARP cache file from the stdin and then start the actual scan. +Unlike nmap and other scanners that implicitly resolve link-layer addresses before the actual scan, `sx` explicitly uses a **neighbor cache**. The cache is a JSONL file with `ip`, `mac`, and optional `vendor` fields. It can contain IPv4 entries produced by `sx arp` and IPv6 entries produced by `sx ndp`. Higher-level scans read it from stdin by default. -This not only simplifies the design of the program, but also speeds up the scanning process, since it is not necessary to perform an ARP scan every time. +This also avoids repeating ARP or NDP discovery for every higher-level scan. -Let's assume that the actual ARP cache is in the `arp.cache` file. We can create it manually -or use ARP scan as shown below: +Let's assume that the current IPv4 neighbor cache is in the `arp.cache` file. We can create it manually +or populate it with an ARP scan: ``` sx arp 192.168.0.1/24 --json | tee arp.cache ``` -Once we have the ARP cache file, we can run scans of higher-level protocols like TCP SYN scan: +Once we have the neighbor cache file, we can run higher-level scans such as TCP SYN: ``` cat arp.cache | sx tcp -p 1-65535 192.168.0.171 @@ -185,7 +194,7 @@ sample input file: {"ip":"10.0.2.2","port":1081} ``` -It is possible to specify the ARP cache file using the `-a` or `--arp-cache` options: +It is possible to specify the neighbor cache file using `-a` or `--neighbor-cache`. The old `--arp-cache` name remains as a deprecated alias: ``` sx tcp -a arp.cache -p 22,443 192.168.0.171 @@ -197,6 +206,17 @@ or stdin redirect: sx tcp -p 22,443 192.168.0.171 < arp.cache ``` +IPv6 works with the same TCP, UDP, ICMP, and SOCKS commands: + +``` +sx ndp --json 'fe80::%en0/120' | sx tcp -p 22,443 'fe80::1%en0' +sx socks -p 1080 '2001:db8::/120' +``` + +For file-only raw-packet scans, pass `--ipv6` to select an IPv6 source/interface. IPv6 packet fields use `--hop-limit`, `--next-header`, and `--payload-length`; ICMPv6 additionally uses `--icmpv6-type` and `--icmpv6-code`. IPv4-only fields such as `--ttl`, `--ipproto`, `--ipflags`, and `--iplen` are rejected for IPv6 scans. + +To keep randomized target generation bounded, directly generated IPv6 ranges must be `/96` or narrower. Use `--file` for larger or explicitly selected address sets. + You can also use the `tcp syn` subcommand instead of the `tcp`: ``` diff --git a/command/arp.go b/command/arp.go index 9fdfa72..4d25808 100644 --- a/command/arp.go +++ b/command/arp.go @@ -30,16 +30,19 @@ func newARPCmd() *arpCmd { if len(args) != 1 { return errors.New("requires one ip subnet argument") } - dstSubnet, err := ip.ParseIPNet(args[0]) + dstPrefix, dstZone, err := ip.ParsePrefix(args[0]) if err != nil { return } + if !dstPrefix.Addr().Is4() { + return errors.New("ARP supports IPv4 only; use ndp for IPv6") + } if err = c.opts.parseRawOptions(); err != nil { return } var r *scan.Range - if r, err = c.opts.getScanRange(dstSubnet); err != nil { + if r, err = c.opts.getScanRange(dstPrefix, dstZone); err != nil { return err } if r.SrcMAC == nil { diff --git a/command/arp_test.go b/command/arp_test.go index 20908c6..1c87bfe 100644 --- a/command/arp_test.go +++ b/command/arp_test.go @@ -1,7 +1,6 @@ package command import ( - "net" "strings" "testing" "time" @@ -51,7 +50,7 @@ func TestARPCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) diff --git a/command/config.go b/command/config.go index 83c7d34..eccc6ee 100644 --- a/command/config.go +++ b/command/config.go @@ -6,6 +6,7 @@ import ( "errors" "io" "net" + "net/netip" "os" "strconv" "strings" @@ -16,7 +17,7 @@ import ( "github.com/v-byte-cpu/sx/command/log" "github.com/v-byte-cpu/sx/pkg/ip" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" "github.com/yl2chen/cidranger" "go.uber.org/ratelimit" ) @@ -24,23 +25,26 @@ import ( const ( defaultWorkerCount = 100 defaultExitDelay = 300 * time.Millisecond + flagHopLimit = "hop-limit" + flagNextHeader = "next-header" + flagPayloadLength = "payload-length" ) var ( - errSrcIP = errors.New("invalid source IP") - errSrcMAC = errors.New("invalid source MAC") - errSrcInterface = errors.New("invalid source interface") - errRateLimit = errors.New("invalid ratelimit") - errARPCacheStdin = errors.New("ARP cache is expected from file or stdin pipe") - errIPFlags = errors.New("invalid ip flags") - errNoDstIP = errors.New("requires one ip subnet argument or file with ip/port pairs") - errARPStdin = errors.New("ARP cache and IP file can not be read from stdin at the same time") + errSrcIP = errors.New("invalid source IP") + errSrcMAC = errors.New("invalid source MAC") + errSrcInterface = errors.New("invalid source interface") + errRateLimit = errors.New("invalid ratelimit") + errNeighborCacheStdin = errors.New("neighbor cache is expected from file or stdin pipe") + errIPFlags = errors.New("invalid ip flags") + errNoDstIP = errors.New("requires one ip subnet argument or file with ip/port pairs") + errNeighborStdin = errors.New("neighbor cache and IP file can not be read from stdin at the same time") ) type packetScanCmdOpts struct { json bool iface *net.Interface - srcIP net.IP + srcIP netip.Addr srcMAC net.HardwareAddr rateCount int rateWindow time.Duration @@ -48,6 +52,7 @@ type packetScanCmdOpts struct { excludeIPs scan.IPContainer rawInterface string + rawSrcIP string rawSrcMAC string rawRateLimit string rawExcludeFile string @@ -56,7 +61,7 @@ type packetScanCmdOpts struct { func (o *packetScanCmdOpts) initCliFlags(cmd *cobra.Command) { cmd.Flags().BoolVar(&o.json, "json", false, "enable JSON output") cmd.Flags().StringVarP(&o.rawInterface, "iface", "i", "", "set interface to send/receive packets") - cmd.Flags().IPVar(&o.srcIP, "srcip", nil, "set source IP address for generated packets") + cmd.Flags().StringVar(&o.rawSrcIP, "srcip", "", "set source IP address for generated packets") cmd.Flags().StringVar(&o.rawSrcMAC, "srcmac", "", "set source MAC address for generated packets") cmd.Flags().StringVar(&o.rawExcludeFile, "exclude", "", strings.Join([]string{ @@ -80,6 +85,12 @@ func (o *packetScanCmdOpts) parseRawOptions() (err error) { return } } + if len(o.rawSrcIP) > 0 { + if o.srcIP, err = netip.ParseAddr(o.rawSrcIP); err != nil { + return errSrcIP + } + o.srcIP = o.srcIP.Unmap() + } if len(o.rawSrcMAC) > 0 { if o.srcMAC, err = net.ParseMAC(o.rawSrcMAC); err != nil { return @@ -100,8 +111,12 @@ func (o *packetScanCmdOpts) parseRawOptions() (err error) { return } -func (o *packetScanCmdOpts) getScanRange(dstSubnet *net.IPNet) (*scan.Range, error) { - iface, srcIP, err := o.getInterface(dstSubnet) +func (o *packetScanCmdOpts) getScanRange(dstPrefix netip.Prefix, dstZone string) (*scan.Range, error) { + return o.getScanRangeForFamily(dstPrefix, dstZone, dstPrefix.IsValid() && dstPrefix.Addr().Is6()) +} + +func (o *packetScanCmdOpts) getScanRangeForFamily(dstPrefix netip.Prefix, dstZone string, ipv6 bool) (*scan.Range, error) { + iface, srcAddr, err := o.getInterface(dstPrefix, dstZone, ipv6) if err != nil { return nil, err } @@ -109,10 +124,13 @@ func (o *packetScanCmdOpts) getScanRange(dstSubnet *net.IPNet) (*scan.Range, err return nil, errSrcInterface } - if o.srcIP != nil { - srcIP = o.srcIP + if o.srcIP.IsValid() { + srcAddr = o.srcIP + } + if !srcAddr.IsValid() { + return nil, errSrcIP } - if srcIP == nil { + if dstPrefix.IsValid() && srcAddr.Is4() != dstPrefix.Addr().Is4() { return nil, errSrcIP } @@ -123,36 +141,82 @@ func (o *packetScanCmdOpts) getScanRange(dstSubnet *net.IPNet) (*scan.Range, err return &scan.Range{ Interface: iface, - DstSubnet: dstSubnet, - SrcIP: srcIP.To4(), + DstPrefix: dstPrefix, + DstZone: dstZone, + SrcIP: srcAddr, SrcMAC: srcMAC}, nil } -func (o *packetScanCmdOpts) getInterface(dstSubnet *net.IPNet) (iface *net.Interface, ifaceIP net.IP, err error) { - if dstSubnet != nil { +func (o *packetScanCmdOpts) getInterface(dstPrefix netip.Prefix, dstZone string, ipv6 bool) (iface *net.Interface, ifaceIP netip.Addr, err error) { + if scopedIface, scoped, scopedErr := o.getScopedSourceInterface(dstZone); scoped || scopedErr != nil { + return scopedIface, o.srcIP, scopedErr + } + if dstPrefix.IsValid() { // try to find directly connected interface - if iface, ifaceIP, err = o.getLocalSubnetInterface(dstSubnet); err != nil { + if iface, ifaceIP, err = o.getLocalPrefixInterface(dstPrefix, dstZone); err != nil { return } // found local interface - if iface != nil && ifaceIP != nil { + if iface != nil && ifaceIP.IsValid() { return } } + target := netip.IPv4Unspecified() + if ipv6 { + target = netip.IPv6Unspecified() + } + if dstPrefix.IsValid() { + target = dstPrefix.Addr() + if dstZone != "" { + target = target.WithZone(dstZone) + } + } if o.iface != nil { // try to get first ip address - ifaceIP, err = ip.GetInterfaceIP(o.iface) + ifaceIP, err = ip.GetInterfaceIP(o.iface, target) return o.iface, ifaceIP, err } // fallback to interface of default gateway - return ip.GetDefaultInterface() + return ip.GetDefaultInterface(target) +} + +func (o *packetScanCmdOpts) getScopedSourceInterface(dstZone string) (*net.Interface, bool, error) { + sourceZone := o.srcIP.Zone() + if sourceZone == "" { + return nil, false, nil + } + if dstZone != "" && dstZone != sourceZone { + return nil, true, errSrcInterface + } + if o.iface != nil { + if o.iface.Name != sourceZone { + return nil, true, errSrcInterface + } + return o.iface, true, nil + } + iface, err := net.InterfaceByName(sourceZone) + if err != nil { + return nil, true, err + } + o.iface = iface + return iface, true, nil } -func (o *packetScanCmdOpts) getLocalSubnetInterface(dstSubnet *net.IPNet) (iface *net.Interface, ifaceIP net.IP, err error) { +func (o *packetScanCmdOpts) getLocalPrefixInterface(dstPrefix netip.Prefix, dstZone string) (iface *net.Interface, ifaceIP netip.Addr, err error) { + if dstZone != "" { + if o.iface != nil && o.iface.Name != dstZone { + return nil, netip.Addr{}, errSrcInterface + } + if o.iface == nil { + if o.iface, err = net.InterfaceByName(dstZone); err != nil { + return nil, netip.Addr{}, err + } + } + } if o.iface == nil { - return ip.GetLocalSubnetInterface(dstSubnet) + return ip.GetLocalPrefixInterface(dstPrefix) } - ifaceIP, err = ip.GetLocalSubnetInterfaceIP(o.iface, dstSubnet) + ifaceIP, err = ip.GetLocalPrefixInterfaceIP(o.iface, dstPrefix) return o.iface, ifaceIP, err } @@ -165,16 +229,33 @@ func (o *packetScanCmdOpts) getLogger(name string, w io.Writer) (logger log.Logg return } +func validateIPVersionFlags(cmd *cobra.Command, ipv6 bool, ipv4Flags, ipv6Flags []string) error { + unsupported := ipv6Flags + family := "IPv4" + if ipv6 { + unsupported = ipv4Flags + family = "IPv6" + } + for _, name := range unsupported { + if cmd.Flags().Changed(name) { + return errors.New("--" + name + " is not supported with " + family) + } + } + return nil +} + type ipScanCmdOpts struct { packetScanCmdOpts - ipFile string - arpCacheFile string - gatewayMAC net.HardwareAddr - vpnMode bool + ipFile string + neighborCacheFile string + arpCacheFile string + gatewayMAC net.HardwareAddr + vpnMode bool + ipv6 bool logger log.Logger scanRange *scan.Range - cache *arp.Cache + cache *neighbor.Cache rawGatewayMAC string } @@ -183,8 +264,11 @@ func (o *ipScanCmdOpts) initCliFlags(cmd *cobra.Command) { o.packetScanCmdOpts.initCliFlags(cmd) cmd.Flags().StringVar(&o.rawGatewayMAC, "gwmac", "", "set gateway MAC address to send generated packets to") cmd.Flags().StringVarP(&o.ipFile, "file", "f", "", "set JSONL file with IPs to scan") - cmd.Flags().StringVarP(&o.arpCacheFile, "arp-cache", "a", "", - strings.Join([]string{"set ARP cache file", "reads from stdin by default"}, "\n")) + cmd.Flags().BoolVar(&o.ipv6, "ipv6", false, "use IPv6 for file-only scans") + cmd.Flags().StringVarP(&o.neighborCacheFile, "neighbor-cache", "a", "", + strings.Join([]string{"set neighbor cache file", "reads from stdin by default"}, "\n")) + cmd.Flags().StringVar(&o.arpCacheFile, "arp-cache", "", "deprecated alias for --neighbor-cache") + _ = cmd.Flags().MarkDeprecated("arp-cache", "use --neighbor-cache instead") } func (o *ipScanCmdOpts) parseRawOptions() (err error) { @@ -196,18 +280,26 @@ func (o *ipScanCmdOpts) parseRawOptions() (err error) { return } } + if o.neighborCacheFile != "" && o.arpCacheFile != "" { + return errors.New("neighbor-cache and arp-cache can not be used together") + } return } func (o *ipScanCmdOpts) parseOptions(scanName string, args []string) (err error) { - dstSubnet, err := o.parseDstSubnet(args) + dstPrefix, dstZone, err := o.parseDstPrefix(args) if err != nil { return } - if o.scanRange, err = o.getScanRange(dstSubnet); err != nil { + ipv6 := o.ipv6 + if dstPrefix.IsValid() { + ipv6 = dstPrefix.Addr().Is6() + } + if o.scanRange, err = o.getScanRangeForFamily(dstPrefix, dstZone, ipv6); err != nil { return } + o.ipv6 = o.scanRange.SrcIP.Is6() if o.scanRange.SrcMAC == nil { o.vpnMode = true } @@ -216,15 +308,15 @@ func (o *ipScanCmdOpts) parseOptions(scanName string, args []string) (err error) return } - // disable arp cache parsing for vpn mode + // VPN interfaces exchange raw IP packets and do not need neighbor MACs. if o.vpnMode { return } - if err = o.validateARPStdin(); err != nil { + if err = o.validateNeighborStdin(); err != nil { return } - if o.cache, err = o.parseARPCache(); err != nil { + if o.cache, err = o.parseNeighborCache(); err != nil { return } @@ -234,37 +326,48 @@ func (o *ipScanCmdOpts) parseOptions(scanName string, args []string) (err error) return } -func (o *ipScanCmdOpts) validateARPStdin() (err error) { - if o.isARPCacheFromStdin() && o.ipFile == "-" { - return errARPStdin +func (o *ipScanCmdOpts) validateNeighborStdin() (err error) { + if o.isNeighborCacheFromStdin() && o.ipFile == "-" { + return errNeighborStdin } return } -func (o *ipScanCmdOpts) parseDstSubnet(args []string) (ipnet *net.IPNet, err error) { +func (o *ipScanCmdOpts) parseDstPrefix(args []string) (netip.Prefix, string, error) { if len(args) == 0 && len(o.ipFile) == 0 { - return nil, errNoDstIP + return netip.Prefix{}, "", errNoDstIP } if len(args) == 0 { - return + return netip.Prefix{}, "", nil + } + return ip.ParsePrefix(args[0]) +} + +func (o *ipScanCmdOpts) parseDstSubnet(args []string) (*net.IPNet, error) { + prefix, _, err := o.parseDstPrefix(args) + if err != nil || !prefix.IsValid() { + return nil, err } - return ip.ParseIPNet(args[0]) + return &net.IPNet{ + IP: net.IP(prefix.Addr().AsSlice()), + Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), + }, nil } -func (o *ipScanCmdOpts) parseARPCache() (cache *arp.Cache, err error) { +func (o *ipScanCmdOpts) parseNeighborCache() (cache *neighbor.Cache, err error) { var r io.ReadCloser - if r, err = o.openARPCache(); err != nil { + if r, err = o.openNeighborCache(); err != nil { return } defer r.Close() - cache = arp.NewCache() - err = arp.FillCache(cache, r) + cache = neighbor.NewCache() + err = neighbor.FillCache(cache, r) return } -func (o *ipScanCmdOpts) openARPCache() (r io.ReadCloser, err error) { - if !o.isARPCacheFromStdin() { - return os.Open(o.arpCacheFile) +func (o *ipScanCmdOpts) openNeighborCache() (r io.ReadCloser, err error) { + if !o.isNeighborCacheFromStdin() { + return os.Open(o.cacheFile()) } // read from stdin var info os.FileInfo @@ -274,25 +377,36 @@ func (o *ipScanCmdOpts) openARPCache() (r io.ReadCloser, err error) { // only data being piped to stdin is valid if (info.Mode() & os.ModeCharDevice) != 0 { // stdin from terminal is not valid - return nil, errARPCacheStdin + return nil, errNeighborCacheStdin } r = io.NopCloser(os.Stdin) return } -func (o *ipScanCmdOpts) isARPCacheFromStdin() bool { - return len(o.arpCacheFile) == 0 || o.arpCacheFile == "-" +func (o *ipScanCmdOpts) isNeighborCacheFromStdin() bool { + cacheFile := o.cacheFile() + return len(cacheFile) == 0 || cacheFile == "-" +} + +func (o *ipScanCmdOpts) cacheFile() string { + if o.neighborCacheFile != "" { + return o.neighborCacheFile + } + return o.arpCacheFile } -func (o *ipScanCmdOpts) getGatewayMAC(iface *net.Interface, cache *arp.Cache) (mac net.HardwareAddr, err error) { +func (o *ipScanCmdOpts) getGatewayMAC(iface *net.Interface, cache *neighbor.Cache) (mac net.HardwareAddr, err error) { if o.gatewayMAC != nil { return o.gatewayMAC, nil } - var gatewayIP net.IP - if gatewayIP, err = ip.GetDefaultGatewayIP(iface); err != nil { + var gatewayIP netip.Addr + if gatewayIP, err = ip.GetDefaultGatewayIP(iface, o.scanRange.SrcIP); err != nil { return } - mac = cache.Get(gatewayIP.To4()) + if gatewayIP.IsLinkLocalUnicast() { + gatewayIP = gatewayIP.WithZone(iface.Name) + } + mac = cache.Get(gatewayIP) return } @@ -435,22 +549,23 @@ func (o *genericScanCmdOpts) parseRawOptions() (err error) { } func (o *genericScanCmdOpts) parseScanRange(args []string) (r *scan.Range, err error) { - dstSubnet, err := o.parseDstSubnet(args) + dstPrefix, dstZone, err := o.parseDstPrefix(args) r = &scan.Range{ - DstSubnet: dstSubnet, + DstPrefix: dstPrefix, + DstZone: dstZone, Ports: o.portRanges, } return } -func (o *genericScanCmdOpts) parseDstSubnet(args []string) (ipnet *net.IPNet, err error) { +func (o *genericScanCmdOpts) parseDstPrefix(args []string) (netip.Prefix, string, error) { if len(args) == 0 && len(o.ipFile) == 0 { - return nil, errNoDstIP + return netip.Prefix{}, "", errNoDstIP } if len(args) == 0 { - return + return netip.Prefix{}, "", nil } - return ip.ParseIPNet(args[0]) + return ip.ParsePrefix(args[0]) } func (o *genericScanCmdOpts) getLogger(name string, w io.Writer) (logger log.Logger, err error) { @@ -576,6 +691,14 @@ func parseIPFlags(inputFlags string) (result uint8, err error) { type openFileFunc func() (io.ReadCloser, error) +type ipContainer struct { + ranger cidranger.Ranger +} + +func (c *ipContainer) Contains(addr netip.Addr) (bool, error) { + return c.ranger.Contains(net.IP(addr.WithZone("").AsSlice())) +} + func parseExcludeFile(openFile openFileFunc) (excludeIPs scan.IPContainer, err error) { input, err := openFile() if err != nil { @@ -594,15 +717,20 @@ func parseExcludeFile(openFile openFileFunc) (excludeIPs scan.IPContainer, err e if len(line) == 0 { continue } - var ipnet *net.IPNet - if ipnet, err = ip.ParseIPNet(line); err != nil { + prefix, zone, parseErr := ip.ParsePrefix(line) + if parseErr != nil || zone != "" { + err = ip.ErrInvalidAddr return } - if err = ranger.Insert(cidranger.NewBasicRangerEntry(*ipnet)); err != nil { + ipnet := net.IPNet{ + IP: net.IP(prefix.Addr().AsSlice()), + Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), + } + if err = ranger.Insert(cidranger.NewBasicRangerEntry(ipnet)); err != nil { return } } - excludeIPs = ranger + excludeIPs = &ipContainer{ranger: ranger} return } diff --git a/command/config_test.go b/command/config_test.go index 5aad87c..fa34008 100644 --- a/command/config_test.go +++ b/command/config_test.go @@ -4,6 +4,7 @@ import ( "errors" "io" "net" + "net/netip" "strings" "testing" "time" @@ -27,7 +28,7 @@ func TestPacketScanCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) @@ -49,6 +50,13 @@ func TestPacketScanCmdOptsParseRawOptions(t *testing.T) { require.Equal(t, 7*time.Second, opts.rateWindow) } +func TestPacketScanCmdOptsParseScopedIPv6Source(t *testing.T) { + t.Parallel() + opts := &packetScanCmdOpts{rawSrcIP: "fe80::1%en0"} + require.NoError(t, opts.parseRawOptions()) + require.Equal(t, netip.MustParseAddr("fe80::1%en0"), opts.srcIP) +} + func TestIPScanCmdOptsInitCliFlags(t *testing.T) { t.Parallel() var opts ipScanCmdOpts @@ -64,7 +72,7 @@ func TestIPScanCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) @@ -72,7 +80,24 @@ func TestIPScanCmdOptsInitCliFlags(t *testing.T) { require.Equal(t, "11:22:33:44:55:66", opts.rawGatewayMAC) require.Equal(t, "ip_file.jsonl", opts.ipFile) - require.Equal(t, "arp.cache", opts.arpCacheFile) + require.Equal(t, "arp.cache", opts.neighborCacheFile) +} + +func TestIPScanCmdOptsNeighborCacheCompatibilityFlags(t *testing.T) { + t.Parallel() + + var current ipScanCmdOpts + currentCmd := &cobra.Command{} + current.initCliFlags(currentCmd) + require.NoError(t, currentCmd.ParseFlags(strings.Fields("--neighbor-cache neighbors.jsonl --ipv6"))) + require.Equal(t, "neighbors.jsonl", current.cacheFile()) + require.True(t, current.ipv6) + + var deprecated ipScanCmdOpts + deprecatedCmd := &cobra.Command{} + deprecated.initCliFlags(deprecatedCmd) + require.NoError(t, deprecatedCmd.ParseFlags(strings.Fields("--arp-cache arp.jsonl"))) + require.Equal(t, "arp.jsonl", deprecated.cacheFile()) } func TestIPScanCmdOptsParseRawOptions(t *testing.T) { @@ -111,7 +136,7 @@ func TestIPPortScanCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) @@ -119,7 +144,7 @@ func TestIPPortScanCmdOptsInitCliFlags(t *testing.T) { require.Equal(t, "11:22:33:44:55:66", opts.rawGatewayMAC) require.Equal(t, "ip_file.jsonl", opts.ipFile) - require.Equal(t, "arp.cache", opts.arpCacheFile) + require.Equal(t, "arp.cache", opts.neighborCacheFile) require.Equal(t, "23-57,71-2733", opts.rawPortRanges) require.Equal(t, "ports.txt", opts.portFile) @@ -189,7 +214,7 @@ func TestGenericScanCmdOptsParseRawOptions(t *testing.T) { require.Equal(t, 7*time.Second, opts.rateWindow) } -func TestIPScanCmdOptsIsARPCacheFromStdin(t *testing.T) { +func TestIPScanCmdOptsIsNeighborCacheFromStdin(t *testing.T) { t.Parallel() tests := []struct { name string @@ -216,12 +241,12 @@ func TestIPScanCmdOptsIsARPCacheFromStdin(t *testing.T) { tt := tt t.Run(tt.name, func(t *testing.T) { t.Parallel() - require.Equal(t, tt.expected, tt.opts.isARPCacheFromStdin()) + require.Equal(t, tt.expected, tt.opts.isNeighborCacheFromStdin()) }) } } -func TestIPScanCmdOptsValidateARPStdin(t *testing.T) { +func TestIPScanCmdOptsValidateNeighborStdin(t *testing.T) { t.Parallel() tests := []struct { name string @@ -271,7 +296,7 @@ func TestIPScanCmdOptsValidateARPStdin(t *testing.T) { tt := tt t.Run(tt.name, func(t *testing.T) { t.Parallel() - err := tt.opts.validateARPStdin() + err := tt.opts.validateNeighborStdin() if tt.shouldErr { require.Error(t, err) } else { @@ -296,6 +321,12 @@ func TestIPScanCmdOptsParseDstSubnet(t *testing.T) { args: []string{"192.168.0.1"}, expected: &net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.IPv4Mask(255, 255, 255, 255)}, }, + { + name: "ValidIPv6Host", + opts: ipScanCmdOpts{}, + args: []string{"2001:db8::1"}, + expected: &net.IPNet{IP: net.ParseIP("2001:db8::1"), Mask: net.CIDRMask(128, 128)}, + }, { name: "ValidDstSubnet", opts: ipScanCmdOpts{}, @@ -330,32 +361,42 @@ func TestIPScanCmdOptsParseDstSubnet(t *testing.T) { } } -func TestGenericScanCmdOptsParseDstSubnet(t *testing.T) { +func TestIPScanCmdOptsParseDstPrefixIPv6Scoped(t *testing.T) { + t.Parallel() + + prefix, zone, err := (&ipScanCmdOpts{}).parseDstPrefix([]string{"fe80::1%en0"}) + + require.NoError(t, err) + require.Equal(t, netip.MustParsePrefix("fe80::1/128"), prefix) + require.Equal(t, "en0", zone) +} + +func TestGenericScanCmdOptsParseDstPrefix(t *testing.T) { t.Parallel() tests := []struct { name string opts genericScanCmdOpts args []string - expected *net.IPNet + expected netip.Prefix shouldErr bool }{ { name: "ValidDstHost", opts: genericScanCmdOpts{}, args: []string{"192.168.0.1"}, - expected: &net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.IPv4Mask(255, 255, 255, 255)}, + expected: netip.MustParsePrefix("192.168.0.1/32"), }, { name: "ValidDstSubnet", opts: genericScanCmdOpts{}, args: []string{"10.0.0.1/16"}, - expected: &net.IPNet{IP: net.IPv4(10, 0, 0, 0).To4(), Mask: net.IPv4Mask(255, 255, 0, 0)}, + expected: netip.MustParsePrefix("10.0.0.0/16"), }, { name: "IPFile", opts: genericScanCmdOpts{ipFile: "ip_file"}, args: []string{}, - expected: nil, + expected: netip.Prefix{}, }, { name: "NoIPHosts", @@ -368,12 +409,13 @@ func TestGenericScanCmdOptsParseDstSubnet(t *testing.T) { tt := tt t.Run(tt.name, func(t *testing.T) { t.Parallel() - result, err := tt.opts.parseDstSubnet(tt.args) + result, zone, err := tt.opts.parseDstPrefix(tt.args) if tt.shouldErr { require.Error(t, err) } else { require.NoError(t, err) require.Equal(t, tt.expected, result) + require.Empty(t, zone) } }) } @@ -394,7 +436,7 @@ func TestGenericScanCmdOptsParseScanRange(t *testing.T) { args: []string{"192.168.0.1"}, expected: &scan.Range{ Ports: []*scan.PortRange{{StartPort: 22, EndPort: 100}}, - DstSubnet: &net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.IPv4Mask(255, 255, 255, 255)}, + DstPrefix: netip.MustParsePrefix("192.168.0.1/32"), }, }, { @@ -403,7 +445,7 @@ func TestGenericScanCmdOptsParseScanRange(t *testing.T) { args: []string{"10.0.0.1/16"}, expected: &scan.Range{ Ports: []*scan.PortRange{{StartPort: 22, EndPort: 100}}, - DstSubnet: &net.IPNet{IP: net.IPv4(10, 0, 0, 0).To4(), Mask: net.IPv4Mask(255, 255, 0, 0)}, + DstPrefix: netip.MustParsePrefix("10.0.0.0/16"), }, }, { @@ -843,6 +885,20 @@ func TestParseExcludeFileWithInvalidFile(t *testing.T) { require.Error(t, err) } +func TestParseExcludeFileIPv6(t *testing.T) { + t.Parallel() + container, err := parseExcludeFile(func() (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader("2001:db8::/126")), nil + }) + require.NoError(t, err) + contained, err := container.Contains(netip.MustParseAddr("2001:db8::2")) + require.NoError(t, err) + require.True(t, contained) + contained, err = container.Contains(netip.MustParseAddr("2001:db8::4")) + require.NoError(t, err) + require.False(t, contained) +} + func TestParseExcludeFile(t *testing.T) { t.Parallel() @@ -948,7 +1004,11 @@ func TestParseExcludeFile(t *testing.T) { return } for _, ip := range tt.contains { - ok, err := ips.Contains(ip) + addr, valid := netip.AddrFromSlice(ip) + if !assert.True(t, valid) { + return + } + ok, err := ips.Contains(addr.Unmap()) if !assert.NoError(t, err) { return } diff --git a/command/icmp.go b/command/icmp.go index 8ef49a0..8779431 100644 --- a/command/icmp.go +++ b/command/icmp.go @@ -10,8 +10,8 @@ import ( "github.com/spf13/cobra" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" "github.com/v-byte-cpu/sx/pkg/scan/icmp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" ) func newICMPCmd() *icmpCmd { @@ -32,6 +32,15 @@ func newICMPCmd() *icmpCmd { if err = c.opts.parseRawOptions(); err != nil { return } + ipv6 := c.opts.ipv6 + if prefix, _, prefixErr := c.opts.parseDstPrefix(args); prefixErr == nil && prefix.IsValid() { + ipv6 = prefix.Addr().Is6() + } + if err = validateIPVersionFlags(cmd, ipv6, + []string{"ttl", "ipproto", "ipflags", "iplen", "type", "code"}, + []string{flagHopLimit, flagNextHeader, flagPayloadLength, "icmpv6-type", "icmpv6-code"}); err != nil { + return err + } if err = c.opts.parseOptions(icmp.ScanType, args); err != nil { return } @@ -71,9 +80,14 @@ type icmpCmdOpts struct { ipProtocol uint8 ipTotalLen uint16 - icmpType uint8 - icmpCode uint8 - icmpPayload []byte + icmpType uint8 + icmpCode uint8 + icmpPayload []byte + hopLimit uint8 + nextHeader uint8 + payloadLength uint16 + icmpv6Type uint8 + icmpv6Code uint8 rawIPFlags string rawICMPPayload string @@ -93,6 +107,11 @@ func (o *icmpCmdOpts) initCliFlags(cmd *cobra.Command) { cmd.Flags().Uint8VarP(&o.icmpCode, "code", "c", 0, "set ICMP code of generated packet") cmd.Flags().StringVarP(&o.rawICMPPayload, "payload", "p", "", strings.Join([]string{"set byte payload of generated packet", "48 random bytes by default"}, "\n")) + cmd.Flags().Uint8Var(&o.hopLimit, flagHopLimit, 64, "set IPv6 Hop Limit field") + cmd.Flags().Uint8Var(&o.nextHeader, flagNextHeader, 58, "set IPv6 Next Header field") + cmd.Flags().Uint16Var(&o.payloadLength, flagPayloadLength, 0, "set IPv6 Payload Length field (calculated by default)") + cmd.Flags().Uint8Var(&o.icmpv6Type, "icmpv6-type", 128, "set ICMPv6 type") + cmd.Flags().Uint8Var(&o.icmpv6Code, "icmpv6-code", 0, "set ICMPv6 code") } func (o *icmpCmdOpts) parseRawOptions() (err error) { @@ -124,12 +143,12 @@ func (o *icmpCmdOpts) newICMPScanMethod(ctx context.Context) *icmp.ScanMethod { reqgen = scan.NewFilterIPRequestGenerator(reqgen, o.excludeIPs) } if o.cache != nil { - reqgen = arp.NewCacheRequestGenerator(reqgen, o.gatewayMAC, o.cache) + reqgen = neighbor.NewCacheRequestGenerator(reqgen, o.gatewayMAC, o.cache) } pktgen := scan.NewPacketMultiGenerator(icmp.NewPacketFiller(o.getICMPOptions()...), runtime.NumCPU()) psrc := scan.NewPacketSource(reqgen, pktgen) results := scan.NewResultChan(ctx, 1000) - return icmp.NewScanMethod(psrc, results, o.vpnMode) + return icmp.NewScanMethodForFamily(psrc, results, o.vpnMode, o.ipv6, o.scanRange.Interface.Name) } func (o *icmpCmdOpts) getICMPOptions() (opts []icmp.PacketFillerOption) { @@ -140,6 +159,11 @@ func (o *icmpCmdOpts) getICMPOptions() (opts []icmp.PacketFillerOption) { icmp.WithIPTotalLength(o.ipTotalLen), icmp.WithType(o.icmpType), icmp.WithCode(o.icmpCode), + icmp.WithHopLimit(o.hopLimit), + icmp.WithNextHeader(o.nextHeader), + icmp.WithPayloadLength(o.payloadLength), + icmp.WithICMPv6Type(o.icmpv6Type), + icmp.WithICMPv6Code(o.icmpv6Code), icmp.WithVPNmode(o.vpnMode)) if len(o.icmpPayload) > 0 { diff --git a/command/icmp_test.go b/command/icmp_test.go index cb8ba74..c53d8e8 100644 --- a/command/icmp_test.go +++ b/command/icmp_test.go @@ -37,6 +37,18 @@ func TestICMPCmdDstSubnetError(t *testing.T) { } } +func TestICMPCmdRejectsFamilySpecificFlags(t *testing.T) { + t.Parallel() + + ipv6 := newICMPCmd().cmd + require.NoError(t, ipv6.ParseFlags([]string{"--ttl", "37"})) + require.EqualError(t, ipv6.RunE(ipv6, []string{"2001:db8::1"}), "--ttl is not supported with IPv6") + + ipv4 := newICMPCmd().cmd + require.NoError(t, ipv4.ParseFlags([]string{"--hop-limit", "37"})) + require.EqualError(t, ipv4.RunE(ipv4, []string{"192.0.2.1"}), "--hop-limit is not supported with IPv4") +} + func TestICMPCmdOptsInitCliFlags(t *testing.T) { t.Parallel() var opts icmpCmdOpts @@ -53,14 +65,14 @@ func TestICMPCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) require.Equal(t, "11:22:33:44:55:66", opts.rawGatewayMAC) require.Equal(t, "ip_file.jsonl", opts.ipFile) - require.Equal(t, "arp.cache", opts.arpCacheFile) + require.Equal(t, "arp.cache", opts.neighborCacheFile) require.Equal(t, uint8(128), opts.ipTTL) require.Equal(t, uint8(6), opts.ipProtocol) diff --git a/command/log/logger_test.go b/command/log/logger_test.go index 6972988..81af76c 100644 --- a/command/log/logger_test.go +++ b/command/log/logger_test.go @@ -11,7 +11,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" ) func scanResultToJSON(t *testing.T, result scan.Result) string { @@ -27,7 +27,7 @@ func TestJSONLoggerResults(t *testing.T) { tests := []struct { name string expected []byte - results []*arp.ScanResult + results []*neighbor.ScanResult }{ { name: "emptyResults", @@ -37,7 +37,7 @@ func TestJSONLoggerResults(t *testing.T) { { name: "oneResult", expected: []byte(scanResultToJSON(t, newScanResult(net.IPv4(192, 168, 0, 3).To4())) + "\n"), - results: []*arp.ScanResult{ + results: []*neighbor.ScanResult{ newScanResult(net.IPv4(192, 168, 0, 3).To4()), }, }, @@ -47,7 +47,7 @@ func TestJSONLoggerResults(t *testing.T) { scanResultToJSON(t, newScanResult(net.IPv4(192, 168, 0, 3).To4())), scanResultToJSON(t, newScanResult(net.IPv4(192, 168, 0, 5).To4())), }, "\n") + "\n"), - results: []*arp.ScanResult{ + results: []*neighbor.ScanResult{ newScanResult(net.IPv4(192, 168, 0, 3).To4()), newScanResult(net.IPv4(192, 168, 0, 5).To4()), }, @@ -75,8 +75,8 @@ func TestJSONLoggerResults(t *testing.T) { } } -func newScanResult(ip net.IP) *arp.ScanResult { - return &arp.ScanResult{ +func newScanResult(ip net.IP) *neighbor.ScanResult { + return &neighbor.ScanResult{ IP: ip.String(), MAC: net.HardwareAddr{0x11, 0x22, 0x33, 0x44, 0x55, 0x66}.String(), Vendor: "Sunny Industries", @@ -89,7 +89,7 @@ func TestPlainLoggerResults(t *testing.T) { tests := []struct { name string expected []byte - results []*arp.ScanResult + results []*neighbor.ScanResult }{ { name: "emptyResults", @@ -99,7 +99,7 @@ func TestPlainLoggerResults(t *testing.T) { { name: "oneResult", expected: []byte(newScanResult(net.IPv4(192, 168, 0, 3).To4()).String() + "\n"), - results: []*arp.ScanResult{ + results: []*neighbor.ScanResult{ newScanResult(net.IPv4(192, 168, 0, 3).To4()), }, }, @@ -109,7 +109,7 @@ func TestPlainLoggerResults(t *testing.T) { newScanResult(net.IPv4(192, 168, 0, 3).To4()).String(), newScanResult(net.IPv4(192, 168, 0, 5).To4()).String(), }, "\n") + "\n"), - results: []*arp.ScanResult{ + results: []*neighbor.ScanResult{ newScanResult(net.IPv4(192, 168, 0, 3).To4()), newScanResult(net.IPv4(192, 168, 0, 5).To4()), }, diff --git a/command/ndp.go b/command/ndp.go new file mode 100644 index 0000000..9ba43ec --- /dev/null +++ b/command/ndp.go @@ -0,0 +1,126 @@ +package command + +import ( + "context" + "errors" + "io" + "net/netip" + "os" + "os/signal" + "runtime" + "strings" + "time" + + "github.com/spf13/cobra" + "github.com/v-byte-cpu/sx/command/log" + "github.com/v-byte-cpu/sx/pkg/ip" + "github.com/v-byte-cpu/sx/pkg/scan" + "github.com/v-byte-cpu/sx/pkg/scan/ndp" +) + +func newNDPCmd() *ndpCmd { + command := &ndpCmd{} + cmd := &cobra.Command{ + Use: "ndp [flags] [subnet]", + Short: "Perform IPv6 Neighbor Discovery scan", + Example: strings.Join([]string{"ndp fe80::/120", "ndp --file ips.jsonl", "ndp --live 5s fe80::1%en0"}, "\n"), + RunE: func(cmd *cobra.Command, args []string) (err error) { + if len(args) > 1 || (len(args) == 0 && command.opts.ipFile == "") { + return errors.New("requires one IPv6 subnet argument or file") + } + + var dstZone string + var dstPrefix netip.Prefix + if len(args) == 1 { + if dstPrefix, dstZone, err = ip.ParsePrefix(args[0]); err != nil { + return err + } + if !dstPrefix.Addr().Is6() { + return errors.New("NDP supports IPv6 only") + } + } + if err = command.opts.parseRawOptions(); err != nil { + return err + } + if command.opts.scanRange, err = command.opts.getScanRangeForFamily(dstPrefix, dstZone, true); err != nil { + return err + } + if command.opts.scanRange.SrcMAC == nil { + return errSrcMAC + } + if command.opts.logger, err = command.opts.getLogger(); err != nil { + return err + } + + ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt) + defer cancel() + method := command.opts.newScanMethod(ctx) + return startPacketScanEngine(ctx, newPacketScanConfig( + withPacketScanMethod(method), + withPacketBPFFilter(ndp.BPFFilter), + withRateCount(command.opts.rateCount), + withRateWindow(command.opts.rateWindow), + withPacketEngineConfig(newEngineConfig( + withLogger(command.opts.logger), + withScanRange(command.opts.scanRange), + withExitDelay(command.opts.exitDelay), + )), + )) + }, + } + command.opts.initCliFlags(cmd) + command.cmd = cmd + return command +} + +type ndpCmd struct { + cmd *cobra.Command + opts ndpCmdOpts +} + +type ndpCmdOpts struct { + packetScanCmdOpts + ipFile string + liveTimeout time.Duration + scanRange *scan.Range + logger log.Logger +} + +func (o *ndpCmdOpts) initCliFlags(cmd *cobra.Command) { + o.packetScanCmdOpts.initCliFlags(cmd) + cmd.Flags().StringVarP(&o.ipFile, "file", "f", "", "set JSONL file with IPv6 addresses to scan") + cmd.Flags().DurationVar(&o.liveTimeout, "live", 0, "enable live mode") +} + +func (o *ndpCmdOpts) getLogger() (log.Logger, error) { + logger, err := o.packetScanCmdOpts.getLogger(ndp.ScanType, os.Stdout) + if err == nil && o.liveTimeout > 0 { + logger = log.NewUniqueLogger(logger) + } + return logger, err +} + +func (o *ndpCmdOpts) newScanMethod(ctx context.Context) *ndp.ScanMethod { + ipGenerator := scan.NewIPGenerator() + if o.ipFile != "" { + ipGenerator = scan.NewFileIPGenerator(func() (io.ReadCloser, error) { + if o.ipFile == "-" { + return io.NopCloser(os.Stdin), nil + } + return os.Open(o.ipFile) + }) + } + requests := scan.NewIPRequestGenerator(ipGenerator) + if o.excludeIPs != nil { + requests = scan.NewFilterIPRequestGenerator(requests, o.excludeIPs) + } + if o.liveTimeout > 0 { + requests = scan.NewLiveRequestGenerator(requests, o.liveTimeout) + } + packets := scan.NewPacketMultiGenerator(ndp.NewPacketFiller(), runtime.NumCPU()) + return ndp.NewScanMethod( + scan.NewPacketSource(requests, packets), + scan.NewResultChan(ctx, 1000), + o.scanRange.Interface.Name, + ) +} diff --git a/command/ndp_test.go b/command/ndp_test.go new file mode 100644 index 0000000..e7b626d --- /dev/null +++ b/command/ndp_test.go @@ -0,0 +1,28 @@ +package command + +import ( + "strings" + "testing" + + "github.com/spf13/cobra" + "github.com/stretchr/testify/require" +) + +func TestNDPCmdRejectsIPv4(t *testing.T) { + t.Parallel() + + cmd := newNDPCmd().cmd + err := cmd.RunE(cmd, []string{"192.0.2.1"}) + require.EqualError(t, err, "NDP supports IPv6 only") +} + +func TestNDPCmdOptsInitCliFlags(t *testing.T) { + t.Parallel() + + var opts ndpCmdOpts + cmd := &cobra.Command{} + opts.initCliFlags(cmd) + require.NoError(t, cmd.ParseFlags(strings.Fields("--file ips.jsonl --live 5s --srcip fe80::1%en0"))) + require.Equal(t, "ips.jsonl", opts.ipFile) + require.Equal(t, "fe80::1%en0", opts.rawSrcIP) +} diff --git a/command/root.go b/command/root.go index 2ae3fb6..715ac6e 100644 --- a/command/root.go +++ b/command/root.go @@ -38,6 +38,7 @@ func newRootCmd(version string) *cobra.Command { cmd.AddCommand( newARPCmd().cmd, + newNDPCmd().cmd, newICMPCmd().cmd, newUDPCmd().cmd, tcpCmd, @@ -161,7 +162,7 @@ func startPacketScanEngine(ctx context.Context, conf *packetScanConfig) error { r := &conf.scanRange // setup network interface to read/write packets - ps, err := afpacket.NewPacketSource(r.Interface.Name, conf.vpnMode) + ps, err := afpacket.NewPacketSource(r.Interface.Name, conf.vpnMode, r.SrcIP.Is6()) if err != nil { return err } diff --git a/command/tcp.go b/command/tcp.go index e9e3e94..dfa9844 100644 --- a/command/tcp.go +++ b/command/tcp.go @@ -10,7 +10,7 @@ import ( "github.com/spf13/cobra" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" "github.com/v-byte-cpu/sx/pkg/scan/tcp" ) @@ -102,7 +102,7 @@ type tcpFlagsCmdOpts struct { } func (o *tcpFlagsCmdOpts) initCliFlags(cmd *cobra.Command) { - o.ipPortScanCmdOpts.initCliFlags(cmd) + o.tcpCmdOpts.initCliFlags(cmd) cmd.Flags().StringVar(&o.rawTCPFlags, "flags", "", "set TCP flags") } @@ -144,6 +144,32 @@ func parseTCPFlags(tcpFlags string) ([]string, error) { type tcpCmdOpts struct { ipPortScanCmdOpts + hopLimit uint8 + nextHeader uint8 + payloadLength uint16 + cmd *cobra.Command +} + +func (o *tcpCmdOpts) initCliFlags(cmd *cobra.Command) { + o.ipPortScanCmdOpts.initCliFlags(cmd) + o.cmd = cmd + cmd.Flags().Uint8Var(&o.hopLimit, flagHopLimit, 64, "set IPv6 Hop Limit field") + cmd.Flags().Uint8Var(&o.nextHeader, flagNextHeader, 6, "set IPv6 Next Header field") + cmd.Flags().Uint16Var(&o.payloadLength, flagPayloadLength, 0, "set IPv6 Payload Length field (calculated by default)") +} + +func (o *tcpCmdOpts) parseOptions(scanName string, args []string) error { + ipv6 := o.ipv6 + if prefix, _, err := o.parseDstPrefix(args); err == nil && prefix.IsValid() { + ipv6 = prefix.Addr().Is6() + } + if err := validateIPVersionFlags(o.cmd, ipv6, nil, []string{flagHopLimit, flagNextHeader, flagPayloadLength}); err != nil { + return err + } + if err := o.ipPortScanCmdOpts.parseOptions(scanName, args); err != nil { + return err + } + return nil } func (o *tcpCmdOpts) newTCPScanMethod(ctx context.Context, opts ...tcpScanConfigOption) *tcp.ScanMethod { @@ -153,9 +179,13 @@ func (o *tcpCmdOpts) newTCPScanMethod(ctx context.Context, opts ...tcpScanConfig } reqgen := o.newIPPortGenerator() if o.cache != nil { - reqgen = arp.NewCacheRequestGenerator(reqgen, o.gatewayMAC, o.cache) + reqgen = neighbor.NewCacheRequestGenerator(reqgen, o.gatewayMAC, o.cache) } - c.packetFillerOpts = append(c.packetFillerOpts, tcp.WithFillerVPNmode(o.vpnMode)) + c.packetFillerOpts = append(c.packetFillerOpts, + tcp.WithFillerVPNmode(o.vpnMode), + tcp.WithHopLimit(o.hopLimit), + tcp.WithNextHeader(o.nextHeader), + tcp.WithPayloadLength(o.payloadLength)) pktgen := scan.NewPacketMultiGenerator(tcp.NewPacketFiller(c.packetFillerOpts...), runtime.NumCPU()) psrc := scan.NewPacketSource(reqgen, pktgen) results := scan.NewResultChan(ctx, 1000) @@ -163,7 +193,9 @@ func (o *tcpCmdOpts) newTCPScanMethod(ctx context.Context, opts ...tcpScanConfig c.scanName, psrc, results, tcp.WithPacketFilterFunc(c.packetFilter), tcp.WithPacketFlagsFunc(c.packetFlags), - tcp.WithScanVPNmode(o.vpnMode)) + tcp.WithScanVPNmode(o.vpnMode), + tcp.WithScanIPv6(o.ipv6), + tcp.WithScanZone(o.scanRange.Interface.Name)) } type tcpScanConfig struct { diff --git a/command/tcp_test.go b/command/tcp_test.go index 1c4d100..b1e67d3 100644 --- a/command/tcp_test.go +++ b/command/tcp_test.go @@ -60,14 +60,14 @@ func TestTCPCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) require.Equal(t, "11:22:33:44:55:66", opts.rawGatewayMAC) require.Equal(t, "ip_file.jsonl", opts.ipFile) - require.Equal(t, "arp.cache", opts.arpCacheFile) + require.Equal(t, "arp.cache", opts.neighborCacheFile) require.Equal(t, "23-57,71-2733", opts.rawPortRanges) diff --git a/command/udp.go b/command/udp.go index 3307bb0..000ccdb 100644 --- a/command/udp.go +++ b/command/udp.go @@ -9,8 +9,8 @@ import ( "github.com/spf13/cobra" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" "github.com/v-byte-cpu/sx/pkg/scan/icmp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" "github.com/v-byte-cpu/sx/pkg/scan/udp" ) @@ -33,6 +33,15 @@ func newUDPCmd() *udpCmd { if err = c.opts.parseRawOptions(); err != nil { return } + ipv6 := c.opts.ipv6 + if prefix, _, prefixErr := c.opts.parseDstPrefix(args); prefixErr == nil && prefix.IsValid() { + ipv6 = prefix.Addr().Is6() + } + if err = validateIPVersionFlags(cmd, ipv6, + []string{"ttl", "ipproto", "ipflags", "iplen"}, + []string{flagHopLimit, flagNextHeader, flagPayloadLength}); err != nil { + return err + } if err = c.opts.parseOptions(udp.ScanType, args); err != nil { return } @@ -72,7 +81,10 @@ type udpCmdOpts struct { ipProtocol uint8 ipTotalLen uint16 - udpPayload []byte + udpPayload []byte + hopLimit uint8 + nextHeader uint8 + payloadLength uint16 rawIPFlags string rawUDPPayload string @@ -89,6 +101,9 @@ func (o *udpCmdOpts) initCliFlags(cmd *cobra.Command) { cmd.Flags().StringVar(&o.rawUDPPayload, "payload", "", strings.Join([]string{"set byte payload of generated packet", "0 bytes by default"}, "\n")) + cmd.Flags().Uint8Var(&o.hopLimit, flagHopLimit, 64, "set IPv6 Hop Limit field") + cmd.Flags().Uint8Var(&o.nextHeader, flagNextHeader, 17, "set IPv6 Next Header field") + cmd.Flags().Uint16Var(&o.payloadLength, flagPayloadLength, 0, "set IPv6 Payload Length field (calculated by default)") } func (o *udpCmdOpts) parseRawOptions() (err error) { @@ -111,12 +126,12 @@ func (o *udpCmdOpts) parseRawOptions() (err error) { func (o *udpCmdOpts) newUDPScanMethod(ctx context.Context) *udp.ScanMethod { reqgen := o.newIPPortGenerator() if o.cache != nil { - reqgen = arp.NewCacheRequestGenerator(o.newIPPortGenerator(), o.gatewayMAC, o.cache) + reqgen = neighbor.NewCacheRequestGenerator(reqgen, o.gatewayMAC, o.cache) } pktgen := scan.NewPacketMultiGenerator(udp.NewPacketFiller(o.getUDPOptions()...), runtime.NumCPU()) psrc := scan.NewPacketSource(reqgen, pktgen) results := scan.NewResultChan(ctx, 1000) - return udp.NewScanMethod(psrc, results, o.vpnMode) + return udp.NewScanMethodForFamily(psrc, results, o.vpnMode, o.ipv6, o.scanRange.Interface.Name) } func (o *udpCmdOpts) getUDPOptions() (opts []udp.PacketFillerOption) { @@ -126,6 +141,10 @@ func (o *udpCmdOpts) getUDPOptions() (opts []udp.PacketFillerOption) { udp.WithIPFlags(o.ipFlags), udp.WithIPTotalLength(o.ipTotalLen), udp.WithVPNmode(o.vpnMode)) + opts = append(opts, + udp.WithHopLimit(o.hopLimit), + udp.WithNextHeader(o.nextHeader), + udp.WithPayloadLength(o.payloadLength)) if len(o.udpPayload) > 0 { opts = append(opts, udp.WithPayload(o.udpPayload)) diff --git a/command/udp_test.go b/command/udp_test.go index f5d423b..2fa0c5a 100644 --- a/command/udp_test.go +++ b/command/udp_test.go @@ -55,14 +55,14 @@ func TestUDPCmdOptsInitCliFlags(t *testing.T) { require.NoError(t, err) require.True(t, opts.json) require.Equal(t, "eth0", opts.rawInterface) - require.Equal(t, net.IPv4(192, 168, 0, 1), opts.srcIP) + require.Equal(t, "192.168.0.1", opts.rawSrcIP) require.Equal(t, "00:11:22:33:44:55", opts.rawSrcMAC) require.Equal(t, "500/7s", opts.rawRateLimit) require.Equal(t, 10*time.Second, opts.exitDelay) require.Equal(t, "11:22:33:44:55:66", opts.rawGatewayMAC) require.Equal(t, "ip_file.jsonl", opts.ipFile) - require.Equal(t, "arp.cache", opts.arpCacheFile) + require.Equal(t, "arp.cache", opts.neighborCacheFile) require.Equal(t, "23-57,71-2733", opts.rawPortRanges) diff --git a/pkg/ip/ip.go b/pkg/ip/ip.go index ff0153a..c99c67a 100644 --- a/pkg/ip/ip.go +++ b/pkg/ip/ip.go @@ -4,63 +4,107 @@ import ( "errors" "fmt" "net" + "net/netip" + "strings" ) var ErrInvalidAddr = errors.New("invalid IP subnet/host") -func ParseIPNet(subnet string) (*net.IPNet, error) { - _, result, err := net.ParseCIDR(subnet) - if err == nil { - return result, err +// ParsePrefix parses an IP host or prefix. A scoped IPv6 host returns its zone +// separately because netip.Prefix deliberately does not retain zones. +func ParsePrefix(input string) (netip.Prefix, string, error) { + if percent, slash := strings.LastIndexByte(input, '%'), strings.LastIndexByte(input, '/'); percent >= 0 && slash > percent { + zone := input[percent+1 : slash] + prefix, err := netip.ParsePrefix(input[:percent] + input[slash:]) + if err != nil || zone == "" || !prefix.Addr().Is6() { + return netip.Prefix{}, "", ErrInvalidAddr + } + return prefix.Masked(), zone, nil } - // try to parse host IP address instead - ipAddr := net.ParseIP(subnet) - if ipAddr == nil { - return nil, ErrInvalidAddr + if prefix, err := netip.ParsePrefix(input); err == nil { + addr := prefix.Addr() + bits := prefix.Bits() + if addr.Is4In6() { + if bits < 96 { + return netip.Prefix{}, "", ErrInvalidAddr + } + addr = addr.Unmap() + bits -= 96 + } + return netip.PrefixFrom(addr, bits).Masked(), "", nil } - return &net.IPNet{IP: ipAddr.To4(), Mask: net.CIDRMask(32, 32)}, nil + + addr, err := netip.ParseAddr(input) + if err != nil { + return netip.Prefix{}, "", ErrInvalidAddr + } + zone := addr.Zone() + addr = addr.WithZone("").Unmap() + return netip.PrefixFrom(addr, addr.BitLen()), zone, nil } -func GetInterfaceIP(iface *net.Interface) (ifaceIP net.IP, err error) { - var addrs []net.Addr - if addrs, err = iface.Addrs(); err != nil || len(addrs) == 0 { - return +func GetInterfaceIP(iface *net.Interface, target netip.Addr) (netip.Addr, error) { + addrs, err := iface.Addrs() + if err != nil { + return netip.Addr{}, err } + prefixes := make([]netip.Prefix, 0, len(addrs)) for _, addr := range addrs { - if ipnet, ok := addr.(*net.IPNet); ok && ipnet.IP.To4() != nil { - return ipnet.IP.To4(), nil + prefix, err := netip.ParsePrefix(addr.String()) + if err == nil { + prefixes = append(prefixes, prefix) + } + } + result := selectInterfaceIP(prefixes, target) + if !result.IsValid() { + return netip.Addr{}, fmt.Errorf("interface has no matching IP address: %s", iface.Name) + } + return result, nil +} + +func selectInterfaceIP(addresses []netip.Prefix, target netip.Addr) netip.Addr { + wantIPv4 := target.Is4() + wantLinkLocal := target.Is6() && target.IsLinkLocalUnicast() + for _, prefix := range addresses { + addr := prefix.Addr().Unmap() + if addr.Is4() != wantIPv4 { + continue + } + if addr.Is6() && addr.IsLinkLocalUnicast() != wantLinkLocal { + continue } + return addr } - return nil, fmt.Errorf("interface has no IPv4 address: %s", iface.Name) + return netip.Addr{} } -func GetLocalSubnetInterface(dstSubnet *net.IPNet) (iface *net.Interface, ifaceIP net.IP, err error) { +func GetLocalPrefixInterface(dstPrefix netip.Prefix) (iface *net.Interface, ifaceIP netip.Addr, err error) { var ifaces []net.Interface if ifaces, err = net.Interfaces(); err != nil { return } for _, v := range ifaces { viface := v - if ifaceIP, err = GetLocalSubnetInterfaceIP(&viface, dstSubnet); err != nil { + if ifaceIP, err = GetLocalPrefixInterfaceIP(&viface, dstPrefix); err != nil { return } - if ifaceIP != nil { + if ifaceIP.IsValid() { return &viface, ifaceIP, nil } } return } -func GetLocalSubnetInterfaceIP(iface *net.Interface, dstSubnet *net.IPNet) (net.IP, error) { - dstSubnetIP := dstSubnet.IP.Mask(dstSubnet.Mask) +func GetLocalPrefixInterfaceIP(iface *net.Interface, dstPrefix netip.Prefix) (netip.Addr, error) { addrs, err := iface.Addrs() if err != nil { - return nil, err + return netip.Addr{}, err } for _, addr := range addrs { - if ipnet, ok := addr.(*net.IPNet); ok && ipnet.Contains(dstSubnetIP) { - return ipnet.IP, nil + prefix, err := netip.ParsePrefix(addr.String()) + if err == nil && prefix.Contains(dstPrefix.Addr()) { + return prefix.Addr().Unmap(), nil } } - return nil, nil + return netip.Addr{}, nil } diff --git a/pkg/ip/ip_darwin.go b/pkg/ip/ip_darwin.go index 42de767..1d5d037 100644 --- a/pkg/ip/ip_darwin.go +++ b/pkg/ip/ip_darwin.go @@ -2,6 +2,7 @@ package ip import ( "net" + "net/netip" "syscall" "golang.org/x/net/route" @@ -9,33 +10,33 @@ import ( type defaultRoute struct { interfaceIndex int - gatewayIP net.IP + gatewayIP netip.Addr } -func GetDefaultInterface() (iface *net.Interface, ifaceIP net.IP, err error) { - defaultRoute, err := findDefaultRoute(0) +func GetDefaultInterface(target netip.Addr) (iface *net.Interface, ifaceIP netip.Addr, err error) { + defaultRoute, err := findDefaultRoute(0, target) if err != nil || defaultRoute == nil { - return nil, nil, err + return nil, netip.Addr{}, err } if iface, err = net.InterfaceByIndex(defaultRoute.interfaceIndex); err != nil { - return nil, nil, err + return nil, netip.Addr{}, err } - if ifaceIP, err = GetInterfaceIP(iface); err != nil { - return nil, nil, err + if ifaceIP, err = GetInterfaceIP(iface, target); err != nil { + return nil, netip.Addr{}, err } return iface, ifaceIP, nil } -func GetDefaultGatewayIP(iface *net.Interface) (gatewayIP net.IP, err error) { - defaultRoute, err := findDefaultRoute(iface.Index) +func GetDefaultGatewayIP(iface *net.Interface, target netip.Addr) (gatewayIP netip.Addr, err error) { + defaultRoute, err := findDefaultRoute(iface.Index, target) if err != nil || defaultRoute == nil { - return nil, err + return netip.Addr{}, err } return defaultRoute.gatewayIP, nil } -func findDefaultRoute(interfaceIndex int) (*defaultRoute, error) { - rib, err := route.FetchRIB(syscall.AF_INET, route.RIBTypeRoute, 0) +func findDefaultRoute(interfaceIndex int, target netip.Addr) (*defaultRoute, error) { + rib, err := route.FetchRIB(routeAddressFamily(target), route.RIBTypeRoute, 0) if err != nil { return nil, err } @@ -48,7 +49,7 @@ func findDefaultRoute(interfaceIndex int) (*defaultRoute, error) { if !ok { continue } - defaultRoute, ok := parseDefaultRoute(routeMessage) + defaultRoute, ok := parseDefaultRoute(routeMessage, target) if ok && (interfaceIndex == 0 || interfaceIndex == defaultRoute.interfaceIndex) { return &defaultRoute, nil } @@ -56,10 +57,17 @@ func findDefaultRoute(interfaceIndex int) (*defaultRoute, error) { return nil, nil } -func parseDefaultRoute(message *route.RouteMessage) (defaultRoute, bool) { +func parseDefaultRoute(message *route.RouteMessage, target netip.Addr) (defaultRoute, bool) { if message.Err != nil || message.Index == 0 || message.Flags&syscall.RTF_GATEWAY == 0 { return defaultRoute{}, false } + if target.Is4() { + return parseDefaultIPv4Route(message) + } + return parseDefaultIPv6Route(message) +} + +func parseDefaultIPv4Route(message *route.RouteMessage) (defaultRoute, bool) { dst, ok := routeAddr[*route.Inet4Addr](message.Addrs, syscall.RTAX_DST) if !ok || !isZeroInet4Addr(dst) { return defaultRoute{}, false @@ -74,7 +82,26 @@ func parseDefaultRoute(message *route.RouteMessage) (defaultRoute, bool) { } return defaultRoute{ interfaceIndex: message.Index, - gatewayIP: net.IP(gateway.IP[:]).To4(), + gatewayIP: netip.AddrFrom4(gateway.IP), + }, true +} + +func parseDefaultIPv6Route(message *route.RouteMessage) (defaultRoute, bool) { + dst, ok := routeAddr[*route.Inet6Addr](message.Addrs, syscall.RTAX_DST) + if !ok || !isZeroInet6Addr(dst) { + return defaultRoute{}, false + } + netmask, ok := routeAddr[*route.Inet6Addr](message.Addrs, syscall.RTAX_NETMASK) + if ok && !isZeroInet6Addr(netmask) { + return defaultRoute{}, false + } + gateway, ok := routeAddr[*route.Inet6Addr](message.Addrs, syscall.RTAX_GATEWAY) + if !ok || isZeroInet6Addr(gateway) { + return defaultRoute{}, false + } + return defaultRoute{ + interfaceIndex: message.Index, + gatewayIP: netip.AddrFrom16(gateway.IP), }, true } @@ -95,3 +122,19 @@ func isZeroInet4Addr(addr *route.Inet4Addr) bool { } return true } + +func isZeroInet6Addr(addr *route.Inet6Addr) bool { + for _, octet := range addr.IP { + if octet != 0 { + return false + } + } + return true +} + +func routeAddressFamily(target netip.Addr) int { + if target.Is4() { + return syscall.AF_INET + } + return syscall.AF_INET6 +} diff --git a/pkg/ip/ip_darwin_test.go b/pkg/ip/ip_darwin_test.go index cafd1c6..8f0957e 100644 --- a/pkg/ip/ip_darwin_test.go +++ b/pkg/ip/ip_darwin_test.go @@ -2,7 +2,7 @@ package ip import ( "errors" - "net" + "net/netip" "syscall" "testing" @@ -22,11 +22,27 @@ func TestParseDefaultRoute(t *testing.T) { Addrs: addrs, } - result, ok := parseDefaultRoute(message) + result, ok := parseDefaultRoute(message, netip.IPv4Unspecified()) require.True(t, ok) require.Equal(t, 11, result.interfaceIndex) - require.Equal(t, net.IPv4(192, 168, 0, 1).To4(), result.gatewayIP) + require.Equal(t, netip.MustParseAddr("192.168.0.1"), result.gatewayIP) +} + +func TestParseDefaultIPv6Route(t *testing.T) { + t.Parallel() + + addrs := make([]route.Addr, syscall.RTAX_MAX) + addrs[syscall.RTAX_DST] = &route.Inet6Addr{} + addrs[syscall.RTAX_GATEWAY] = &route.Inet6Addr{IP: [16]byte{0xfe, 0x80, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}} + addrs[syscall.RTAX_NETMASK] = &route.Inet6Addr{} + message := &route.RouteMessage{Flags: syscall.RTF_GATEWAY, Index: 11, Addrs: addrs} + + result, ok := parseDefaultRoute(message, netip.IPv6Unspecified()) + + require.True(t, ok) + require.Equal(t, 11, result.interfaceIndex) + require.Equal(t, netip.MustParseAddr("fe80::1"), result.gatewayIP) } func TestParseDefaultRouteRejectsNonDefaultRoutes(t *testing.T) { @@ -105,7 +121,7 @@ func TestParseDefaultRouteRejectsNonDefaultRoutes(t *testing.T) { tt := vtt t.Run(tt.name, func(t *testing.T) { t.Parallel() - _, ok := parseDefaultRoute(tt.message) + _, ok := parseDefaultRoute(tt.message, netip.IPv4Unspecified()) require.False(t, ok) }) } diff --git a/pkg/ip/ip_linux.go b/pkg/ip/ip_linux.go index 4384011..f73ba2b 100644 --- a/pkg/ip/ip_linux.go +++ b/pkg/ip/ip_linux.go @@ -3,44 +3,57 @@ package ip import ( "math" "net" + "net/netip" "github.com/vishvananda/netlink" "github.com/vishvananda/netlink/nl" ) -func GetDefaultInterface() (iface *net.Interface, ifaceIP net.IP, err error) { +func GetDefaultInterface(target netip.Addr) (iface *net.Interface, ifaceIP netip.Addr, err error) { var routes []netlink.Route - if routes, err = netlink.RouteList(nil, nl.FAMILY_V4); err != nil { + if routes, err = netlink.RouteList(nil, routeFamily(target)); err != nil { return } priority := math.MaxInt32 for _, route := range routes { // found default gateway - if route.Dst == nil && route.Src == nil && route.Priority < priority { + if route.Dst == nil && route.Priority < priority { priority = route.Priority if iface, err = net.InterfaceByIndex(route.LinkIndex); err != nil { return } - if ifaceIP, err = GetInterfaceIP(iface); err != nil { - return + if route.Src != nil { + if ifaceIP, _ = netip.AddrFromSlice(route.Src); ifaceIP.IsValid() { + ifaceIP = ifaceIP.Unmap() + continue + } } + ifaceIP, err = GetInterfaceIP(iface, target) } } return } -func GetDefaultGatewayIP(iface *net.Interface) (gatewayIP net.IP, err error) { +func GetDefaultGatewayIP(iface *net.Interface, target netip.Addr) (gatewayIP netip.Addr, err error) { var routes []netlink.Route - if routes, err = netlink.RouteList(nil, nl.FAMILY_V4); err != nil { + if routes, err = netlink.RouteList(nil, routeFamily(target)); err != nil { return } priority := math.MaxInt32 for _, route := range routes { // found default gateway - if route.Dst == nil && route.Src == nil && route.LinkIndex == iface.Index && route.Priority < priority { + if route.Dst == nil && route.LinkIndex == iface.Index && route.Priority < priority { priority = route.Priority - gatewayIP = route.Gw + gatewayIP, _ = netip.AddrFromSlice(route.Gw) + gatewayIP = gatewayIP.Unmap() } } return } + +func routeFamily(target netip.Addr) int { + if target.Is4() { + return nl.FAMILY_V4 + } + return nl.FAMILY_V6 +} diff --git a/pkg/ip/ip_other.go b/pkg/ip/ip_other.go index 72f9544..944963b 100644 --- a/pkg/ip/ip_other.go +++ b/pkg/ip/ip_other.go @@ -5,16 +5,17 @@ package ip import ( "errors" "net" + "net/netip" ) var errOS = errors.New("OS platform is not supported") -func GetDefaultInterface() (iface *net.Interface, ifaceIP net.IP, err error) { +func GetDefaultInterface(_ netip.Addr) (iface *net.Interface, ifaceIP netip.Addr, err error) { err = errOS return } -func GetDefaultGatewayIP(_ *net.Interface) (gatewayIP net.IP, err error) { +func GetDefaultGatewayIP(_ *net.Interface, _ netip.Addr) (gatewayIP netip.Addr, err error) { err = errOS return } diff --git a/pkg/ip/ip_test.go b/pkg/ip/ip_test.go index cba834f..5847d18 100644 --- a/pkg/ip/ip_test.go +++ b/pkg/ip/ip_test.go @@ -1,51 +1,91 @@ package ip import ( - "net" + "net/netip" "testing" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func TestParseIPNetWithError(t *testing.T) { +func TestParsePrefix(t *testing.T) { t.Parallel() - _, err := ParseIPNet("") - assert.Error(t, err) -} -func TestParseIPNet(t *testing.T) { - t.Parallel() tests := []struct { - name string - in string - expected *net.IPNet + name string + input string + expectedPrefix netip.Prefix + expectedZone string }{ { - name: "subnet", - in: "192.168.0.1/24", - expected: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0).To4(), - Mask: net.CIDRMask(24, 32), - }, + name: "IPv4Host", + input: "192.0.2.1", + expectedPrefix: netip.MustParsePrefix("192.0.2.1/32"), + }, + { + name: "ScopedIPv6Prefix", + input: "fe80::1%en0/64", + expectedPrefix: netip.MustParsePrefix("fe80::/64"), + expectedZone: "en0", + }, + { + name: "IPv6Prefix", + input: "2001:db8::1/64", + expectedPrefix: netip.MustParsePrefix("2001:db8::/64"), }, { - name: "host", - in: "10.0.0.1", - expected: &net.IPNet{ - IP: net.IPv4(10, 0, 0, 1).To4(), - Mask: net.CIDRMask(32, 32), - }, + name: "ScopedIPv6Host", + input: "fe80::1%en0", + expectedPrefix: netip.MustParsePrefix("fe80::1/128"), + expectedZone: "en0", + }, + { + name: "IPv4MappedHost", + input: "::ffff:192.0.2.1", + expectedPrefix: netip.MustParsePrefix("192.0.2.1/32"), + }, + { + name: "IPv4MappedPrefix", + input: "::ffff:192.0.2.1/128", + expectedPrefix: netip.MustParsePrefix("192.0.2.1/32"), }, } - for _, vtt := range tests { - tt := vtt + + for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - result, err := ParseIPNet(tt.in) + + prefix, zone, err := ParsePrefix(tt.input) + require.NoError(t, err) - assert.Equal(t, tt.expected, result) + require.Equal(t, tt.expectedPrefix, prefix) + require.Equal(t, tt.expectedZone, zone) + }) + } +} + +func TestSelectInterfaceIPByAddressFamilyAndScope(t *testing.T) { + t.Parallel() + + addresses := []netip.Prefix{ + netip.MustParsePrefix("fe80::2/64"), + netip.MustParsePrefix("2001:db8::2/64"), + netip.MustParsePrefix("192.0.2.2/24"), + } + tests := []struct { + name string + target netip.Addr + expected netip.Addr + }{ + {name: "IPv4", target: netip.MustParseAddr("198.51.100.1"), expected: netip.MustParseAddr("192.0.2.2")}, + {name: "IPv6Global", target: netip.MustParseAddr("2001:db8:1::1"), expected: netip.MustParseAddr("2001:db8::2")}, + {name: "IPv6LinkLocal", target: netip.MustParseAddr("fe80::1"), expected: netip.MustParseAddr("fe80::2")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.expected, selectInterfaceIP(addresses, tt.target)) }) } } diff --git a/pkg/packet/afpacket/readwriter.go b/pkg/packet/afpacket/readwriter.go index 1d9c800..c7bc6ad 100644 --- a/pkg/packet/afpacket/readwriter.go +++ b/pkg/packet/afpacket/readwriter.go @@ -20,14 +20,18 @@ type Source struct { // Assert that AfPacketSource conforms to the packet.ReadWriter interface var _ packet.ReadWriter = (*Source)(nil) -func NewPacketSource(iface string, vpnMode bool) (*Source, error) { +func NewPacketSource(iface string, vpnMode bool, ipv6 bool) (*Source, error) { handle, err := afp.NewTPacket(afp.SocketRaw, afp.OptInterface(iface)) if err != nil { return nil, err } linkType := layers.LinkTypeEthernet if vpnMode { - linkType = layers.LinkTypeIPv4 + if ipv6 { + linkType = layers.LinkTypeIPv6 + } else { + linkType = layers.LinkTypeIPv4 + } } return &Source{handle, linkType}, nil } diff --git a/pkg/packet/afpacket/readwriter_darwin.go b/pkg/packet/afpacket/readwriter_darwin.go index d80efa7..9b0d3fd 100644 --- a/pkg/packet/afpacket/readwriter_darwin.go +++ b/pkg/packet/afpacket/readwriter_darwin.go @@ -41,12 +41,13 @@ const ( type Source struct { handle packetHandle mode packetLinkMode + ipv6 bool } // Assert that Source conforms to the packet.ReadWriter interface. var _ packet.ReadWriter = (*Source)(nil) -func NewPacketSource(iface string, vpnMode bool) (*Source, error) { +func NewPacketSource(iface string, vpnMode bool, ipv6 bool) (*Source, error) { handle, err := pcap.OpenLive(iface, defaultSnapLen, false, pcap.BlockForever) if err != nil { return nil, err @@ -55,19 +56,19 @@ func NewPacketSource(iface string, vpnMode bool) (*Source, error) { handle.Close() return nil, err } - return newSource(handle, vpnMode) + return newSource(handle, vpnMode, ipv6) } -func newSource(handle packetHandle, vpnMode bool) (*Source, error) { - mode, err := newPacketLinkMode(handle.LinkType(), vpnMode) +func newSource(handle packetHandle, vpnMode bool, ipv6 bool) (*Source, error) { + mode, err := newPacketLinkMode(handle.LinkType(), vpnMode, ipv6) if err != nil { handle.Close() return nil, err } - return &Source{handle: handle, mode: mode}, nil + return &Source{handle: handle, mode: mode, ipv6: ipv6}, nil } -func newPacketLinkMode(linkType layers.LinkType, vpnMode bool) (packetLinkMode, error) { +func newPacketLinkMode(linkType layers.LinkType, vpnMode bool, ipv6 bool) (packetLinkMode, error) { if !vpnMode { if linkType == layers.LinkTypeEthernet { return packetLinkEthernet, nil @@ -76,8 +77,16 @@ func newPacketLinkMode(linkType layers.LinkType, vpnMode bool) (packetLinkMode, } switch linkType { - case layers.LinkTypeRaw, layers.LinkTypeIPv4: + case layers.LinkTypeRaw: return packetLinkRaw, nil + case layers.LinkTypeIPv4: + if !ipv6 { + return packetLinkRaw, nil + } + case layers.LinkTypeIPv6: + if ipv6 { + return packetLinkRaw, nil + } case layers.LinkTypeNull: return packetLinkNull, nil case layers.LinkTypeLoop: @@ -85,6 +94,7 @@ func newPacketLinkMode(linkType layers.LinkType, vpnMode bool) (packetLinkMode, default: return 0, fmt.Errorf("%w: %s", ErrUnsupportedLinkType, linkType) } + return 0, fmt.Errorf("%w: %s", ErrUnsupportedLinkType, linkType) } func (s *Source) SetBPFFilter(bpfFilter string, _ int) error { @@ -128,17 +138,21 @@ func (s *Source) WritePacketData(pkt []byte) error { func (s *Source) encodePacket(pkt []byte) []byte { switch s.mode { case packetLinkNull: - return appendLoopbackHeader(pkt, binary.LittleEndian) + return appendLoopbackHeaderForFamily(pkt, binary.LittleEndian, s.ipv6) case packetLinkLoop: - return appendLoopbackHeader(pkt, binary.BigEndian) + return appendLoopbackHeaderForFamily(pkt, binary.BigEndian, s.ipv6) default: return pkt } } -func appendLoopbackHeader(pkt []byte, byteOrder binary.ByteOrder) []byte { +func appendLoopbackHeaderForFamily(pkt []byte, byteOrder binary.ByteOrder, ipv6 bool) []byte { result := make([]byte, loopbackLen+len(pkt)) - byteOrder.PutUint32(result[:loopbackLen], uint32(layers.ProtocolFamilyIPv4)) + family := layers.ProtocolFamilyIPv4 + if ipv6 { + family = layers.ProtocolFamilyIPv6Darwin + } + byteOrder.PutUint32(result[:loopbackLen], uint32(family)) copy(result[loopbackLen:], pkt) return result } diff --git a/pkg/packet/afpacket/readwriter_darwin_test.go b/pkg/packet/afpacket/readwriter_darwin_test.go index 6167f0d..8eba003 100644 --- a/pkg/packet/afpacket/readwriter_darwin_test.go +++ b/pkg/packet/afpacket/readwriter_darwin_test.go @@ -56,7 +56,7 @@ func TestDarwinNewSourceRejectsUnsupportedLinkType(t *testing.T) { t.Parallel() handle := &fakePacketHandle{linkType: layers.LinkTypePPP} - _, err := newSource(handle, false) + _, err := newSource(handle, false, false) require.ErrorIs(t, err, ErrUnsupportedLinkType) require.True(t, handle.closed) @@ -65,7 +65,7 @@ func TestDarwinNewSourceRejectsUnsupportedLinkType(t *testing.T) { func TestDarwinSourceSetBPFFilter(t *testing.T) { t.Parallel() handle := &fakePacketHandle{linkType: layers.LinkTypeEthernet} - source, err := newSource(handle, false) + source, err := newSource(handle, false, false) require.NoError(t, err) err = source.SetBPFFilter("tcp", 1518) @@ -80,6 +80,7 @@ func TestDarwinSourceReadPacketData(t *testing.T) { name string linkType layers.LinkType vpnMode bool + ipv6 bool input []byte expected []byte }{ @@ -103,6 +104,14 @@ func TestDarwinSourceReadPacketData(t *testing.T) { input: []byte{0x45, 0x00}, expected: []byte{0x45, 0x00}, }, + { + name: "IPv6", + linkType: layers.LinkTypeIPv6, + vpnMode: true, + ipv6: true, + input: []byte{0x60, 0x00}, + expected: []byte{0x60, 0x00}, + }, { name: "Null", linkType: layers.LinkTypeNull, @@ -132,7 +141,7 @@ func TestDarwinSourceReadPacketData(t *testing.T) { }, }}, } - source, err := newSource(handle, tt.vpnMode) + source, err := newSource(handle, tt.vpnMode, tt.ipv6) require.NoError(t, err) data, ci, err := source.ReadPacketData() @@ -153,7 +162,7 @@ func TestDarwinSourceReadShortLoopbackPacketReturnsError(t *testing.T) { data: []byte{0x02, 0x00, 0x00}, }}, } - source, err := newSource(handle, true) + source, err := newSource(handle, true, false) require.NoError(t, err) _, _, err = source.ReadPacketData() @@ -167,6 +176,7 @@ func TestDarwinSourceWritePacketData(t *testing.T) { name string linkType layers.LinkType vpnMode bool + ipv6 bool expected []byte }{ { @@ -192,13 +202,20 @@ func TestDarwinSourceWritePacketData(t *testing.T) { vpnMode: true, expected: []byte{0x00, 0x00, 0x00, 0x02, 0x45, 0x00}, }, + { + name: "NullIPv6", + linkType: layers.LinkTypeNull, + vpnMode: true, + ipv6: true, + expected: []byte{0x1e, 0x00, 0x00, 0x00, 0x45, 0x00}, + }, } for _, vtt := range tests { tt := vtt t.Run(tt.name, func(t *testing.T) { t.Parallel() handle := &fakePacketHandle{linkType: tt.linkType} - source, err := newSource(handle, tt.vpnMode) + source, err := newSource(handle, tt.vpnMode, tt.ipv6) require.NoError(t, err) err = source.WritePacketData([]byte{0x45, 0x00}) diff --git a/pkg/packet/afpacket/readwriter_other.go b/pkg/packet/afpacket/readwriter_other.go index f4227fa..f6e326b 100644 --- a/pkg/packet/afpacket/readwriter_other.go +++ b/pkg/packet/afpacket/readwriter_other.go @@ -16,7 +16,7 @@ type Source struct{} // Assert that AfPacketSource conforms to the packet.ReadWriter interface var _ packet.ReadWriter = (*Source)(nil) -func NewPacketSource(_ string, _ bool) (*Source, error) { +func NewPacketSource(_ string, _, _ bool) (*Source, error) { return nil, ErrOS } diff --git a/pkg/scan/arp/arp.go b/pkg/scan/arp/arp.go index 4afc2c4..b9a305b 100644 --- a/pkg/scan/arp/arp.go +++ b/pkg/scan/arp/arp.go @@ -1,15 +1,13 @@ -//go:generate go tool easyjson -output_filename result_easyjson.go arp.go - package arp import ( - "fmt" "net" "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/google/gopacket/macs" "github.com/v-byte-cpu/sx/pkg/scan" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" ) type ScanMethod struct { @@ -26,21 +24,6 @@ type ScanMethod struct { // Assert that arp.ScanMethod conforms to the scan.Method interface var _ scan.PacketMethod = (*ScanMethod)(nil) -//easyjson:json -type ScanResult struct { - IP string `json:"ip"` - MAC string `json:"mac"` - Vendor string `json:"vendor"` -} - -func (r *ScanResult) String() string { - return fmt.Sprintf("%-20s %-20s %s", r.IP, r.MAC, r.Vendor) -} - -func (r *ScanResult) ID() string { - return r.IP -} - func NewScanMethod(psrc scan.PacketSource, results scan.ResultChan) *ScanMethod { sm := &ScanMethod{ PacketSource: psrc, @@ -67,7 +50,7 @@ func (s *ScanMethod) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo) err copy(s.rcvMacPrefix[:], s.rcvARP.SourceHwAddress[:3]) hwVendor := macs.ValidMACPrefixMap[s.rcvMacPrefix] - s.results.Put(&ScanResult{ + s.results.Put(&neighbor.ScanResult{ IP: net.IP(s.rcvARP.SourceProtAddress).String(), MAC: net.HardwareAddr(s.rcvARP.SourceHwAddress).String(), Vendor: hwVendor, @@ -95,9 +78,9 @@ func (*PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) erro ProtAddressSize: uint8(4), Operation: layers.ARPRequest, SourceHwAddress: r.SrcMAC, - SourceProtAddress: r.SrcIP, + SourceProtAddress: net.IP(r.SrcIP.AsSlice()), DstHwAddress: net.HardwareAddr{0x00, 0x00, 0x00, 0x00, 0x00, 0x00}, - DstProtAddress: r.DstIP.To4(), + DstProtAddress: net.IP(r.DstIP.AsSlice()), } var opt gopacket.SerializeOptions diff --git a/pkg/scan/arp/arp_test.go b/pkg/scan/arp/arp_test.go index 5ff6c76..3cdfcd1 100644 --- a/pkg/scan/arp/arp_test.go +++ b/pkg/scan/arp/arp_test.go @@ -10,6 +10,7 @@ import ( "github.com/google/gopacket/layers" "github.com/stretchr/testify/assert" "github.com/v-byte-cpu/sx/pkg/scan" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" ) func TestProcessPacketData(t *testing.T) { @@ -60,7 +61,7 @@ func TestProcessPacketData(t *testing.T) { assert.Fail(t, "results chan is empty") return } - arpResult := result.(*ScanResult) + arpResult := result.(*neighbor.ScanResult) assert.Equal(t, net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}.String(), arpResult.MAC) assert.Equal(t, net.IPv4(192, 168, 0, 3).To4().String(), arpResult.IP) diff --git a/pkg/scan/arp/bpf.go b/pkg/scan/arp/bpf.go index 3fe9d55..8ba75e2 100644 --- a/pkg/scan/arp/bpf.go +++ b/pkg/scan/arp/bpf.go @@ -12,8 +12,8 @@ import ( const MaxPacketLength = 64 func BPFFilter(r *scan.Range) (filter string, maxPacketLength int) { - if r.DstSubnet == nil { + if !r.DstPrefix.IsValid() { return "arp", MaxPacketLength } - return fmt.Sprintf("arp src net %s", r.DstSubnet.String()), MaxPacketLength + return fmt.Sprintf("arp src net %s", r.DstPrefix.String()), MaxPacketLength } diff --git a/pkg/scan/arp/bpf_test.go b/pkg/scan/arp/bpf_test.go index 286060e..97ad4fe 100644 --- a/pkg/scan/arp/bpf_test.go +++ b/pkg/scan/arp/bpf_test.go @@ -1,7 +1,7 @@ package arp import ( - "net" + "net/netip" "testing" "github.com/stretchr/testify/assert" @@ -22,13 +22,8 @@ func TestBPFFilter(t *testing.T) { scanRange: &scan.Range{}, }, { - name: "OneSubnet", - scanRange: &scan.Range{ - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0), - Mask: net.CIDRMask(24, 32), - }, - }, + name: "OneSubnet", + scanRange: &scan.Range{DstPrefix: netip.MustParsePrefix("192.168.0.0/24")}, expectedFilter: "arp src net 192.168.0.0/24", }, } diff --git a/pkg/scan/arp/cache.go b/pkg/scan/arp/cache.go deleted file mode 100644 index c3a5c89..0000000 --- a/pkg/scan/arp/cache.go +++ /dev/null @@ -1,96 +0,0 @@ -package arp - -import ( - "bufio" - "context" - "errors" - "fmt" - "io" - "net" - "sync" - - "github.com/v-byte-cpu/sx/pkg/scan" -) - -type Cache struct { - cache map[string]net.HardwareAddr - mu sync.RWMutex -} - -func NewCache() *Cache { - return &Cache{cache: make(map[string]net.HardwareAddr)} -} - -func (c *Cache) Put(ip net.IP, mac net.HardwareAddr) { - c.mu.Lock() - defer c.mu.Unlock() - c.cache[ip.String()] = mac -} - -func (c *Cache) Get(ip net.IP) net.HardwareAddr { - c.mu.RLock() - defer c.mu.RUnlock() - return c.cache[ip.String()] -} - -func (c *Cache) Delete(ip net.IP) { - c.mu.Lock() - defer c.mu.Unlock() - delete(c.cache, ip.String()) -} - -func FillCache(cache *Cache, r io.Reader) error { - scanner := bufio.NewScanner(r) - for scanner.Scan() { - var entry ScanResult - if err := entry.UnmarshalJSON(scanner.Bytes()); err != nil { - return err - } - ip := net.ParseIP(entry.IP) - if ip == nil { - return errors.New("invalid IP") - } - mac, err := net.ParseMAC(entry.MAC) - if err != nil { - return err - } - cache.Put(ip, mac) - } - return scanner.Err() -} - -type cacheReqGenerator struct { - reqgen scan.RequestGenerator - getMAC func(net.IP) net.HardwareAddr -} - -func NewCacheRequestGenerator(reqgen scan.RequestGenerator, gatewayMAC net.HardwareAddr, cache *Cache) scan.RequestGenerator { - result := &cacheReqGenerator{reqgen: reqgen} - result.getMAC = func(ip net.IP) net.HardwareAddr { - if mac := cache.Get(ip); mac != nil { - return mac - } - return gatewayMAC - } - return result -} - -func (g *cacheReqGenerator) GenerateRequests(ctx context.Context, r *scan.Range) (<-chan *scan.Request, error) { - requests, err := g.reqgen.GenerateRequests(ctx, r) - if err != nil { - return nil, err - } - result := make(chan *scan.Request, cap(requests)) - go func() { - defer close(result) - for request := range requests { - if mac := g.getMAC(request.DstIP); mac != nil { - request.DstMAC = mac - } else { - request.Err = fmt.Errorf("no destination MAC address for %s", request.DstIP) - } - result <- request - } - }() - return result, nil -} diff --git a/pkg/scan/arp/cache_test.go b/pkg/scan/arp/cache_test.go deleted file mode 100644 index 4fa203d..0000000 --- a/pkg/scan/arp/cache_test.go +++ /dev/null @@ -1,301 +0,0 @@ -//go:generate go tool mockgen -package arp -destination=mock_request_test.go github.com/v-byte-cpu/sx/pkg/scan RequestGenerator - -package arp - -import ( - "context" - "errors" - "fmt" - "net" - "strings" - "testing" - "time" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "github.com/v-byte-cpu/sx/pkg/scan" - "go.uber.org/mock/gomock" -) - -func TestCachePut(t *testing.T) { - t.Parallel() - cache := NewCache() - cache.Put(net.IPv4(192, 168, 0, 2).To4(), net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}) - mac := cache.Get(net.IPv4(192, 168, 0, 2).To4()) - require.Equal(t, net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, mac) -} - -func TestCacheDelete(t *testing.T) { - t.Parallel() - cache := NewCache() - cache.Put(net.IPv4(192, 168, 0, 2).To4(), net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}) - cache.Delete(net.IPv4(192, 168, 0, 2).To4()) - require.Nil(t, cache.Get(net.IPv4(192, 168, 0, 2).To4())) -} - -func TestFillCache(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - input string - expected []*ipMacPair - err bool - }{ - { - name: "oneIP", - input: `{"ip":"192.168.0.2","mac":"01:02:03:04:05:06"}`, - expected: []*ipMacPair{ - { - ip: net.IPv4(192, 168, 0, 2).To4(), - mac: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "twoIP", - input: strings.Join([]string{ - `{"ip":"192.168.0.2","mac":"01:02:03:04:05:06"}`, - `{"ip":"192.168.0.3","mac":"11:12:13:14:15:16"}`, - }, "\n"), - expected: []*ipMacPair{ - { - ip: net.IPv4(192, 168, 0, 2).To4(), - mac: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - { - ip: net.IPv4(192, 168, 0, 3).To4(), - mac: net.HardwareAddr{0x11, 0x12, 0x13, 0x14, 0x15, 0x16}, - }, - }, - }, - { - name: "invalidJson", - input: `{"ip":"192`, - err: true, - }, - { - name: "invalidIP", - input: `{"ip":"192.1680","mac":"01:02:03:04:05:06"}`, - err: true, - }, - { - name: "invalidMAC", - input: `{"ip":"192.168.0.2","mac":"01:02:03"}`, - err: true, - }, - } - - for _, tt := range tests { - tt := tt - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - cache := NewCache() - err := FillCache(cache, strings.NewReader(tt.input)) - if tt.err { - require.Error(t, err) - return - } - require.NoError(t, err) - - for _, pair := range tt.expected { - mac := cache.Get(pair.ip) - require.Equal(t, pair.mac, mac) - } - }) - } -} - -type ipMacPair struct { - ip net.IP - mac net.HardwareAddr -} - -func TestCacheRequestGenerator(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - gatewayMAC net.HardwareAddr - ipMacPairs []*ipMacPair - requests []*scan.Request - expectedRequests []*scan.Request - }{ - { - name: "oneRequest", - ipMacPairs: []*ipMacPair{ - { - ip: net.IPv4(192, 168, 0, 2).To4(), - mac: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - requests: []*scan.Request{ - {DstIP: net.IPv4(192, 168, 0, 2).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(192, 168, 0, 2).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "twoRequests", - ipMacPairs: []*ipMacPair{ - { - ip: net.IPv4(192, 168, 0, 2).To4(), - mac: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - { - ip: net.IPv4(192, 168, 0, 3).To4(), - mac: net.HardwareAddr{0x11, 0x12, 0x13, 0x14, 0x15, 0x16}, - }, - }, - requests: []*scan.Request{ - {DstIP: net.IPv4(192, 168, 0, 3).To4()}, - {DstIP: net.IPv4(192, 168, 0, 2).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(192, 168, 0, 3).To4(), - DstMAC: net.HardwareAddr{0x11, 0x12, 0x13, 0x14, 0x15, 0x16}, - }, - { - DstIP: net.IPv4(192, 168, 0, 2).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "oneRequestWithGatewayMAC", - gatewayMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - requests: []*scan.Request{ - {DstIP: net.IPv4(10, 168, 0, 2).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(10, 168, 0, 2).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "twoRequestsWithGatewayMAC", - gatewayMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - requests: []*scan.Request{ - {DstIP: net.IPv4(10, 168, 0, 2).To4()}, - {DstIP: net.IPv4(10, 168, 0, 3).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(10, 168, 0, 2).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - { - DstIP: net.IPv4(10, 168, 0, 3).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "twoRequestsWithCacheAndGatewayMAC", - gatewayMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - ipMacPairs: []*ipMacPair{ - { - ip: net.IPv4(10, 168, 0, 2).To4(), - mac: net.HardwareAddr{0x2, 0x3, 0x4, 0x5, 0x6, 0x7}, - }, - }, - requests: []*scan.Request{ - {DstIP: net.IPv4(10, 168, 0, 2).To4()}, - {DstIP: net.IPv4(10, 168, 0, 3).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(10, 168, 0, 2).To4(), - DstMAC: net.HardwareAddr{0x2, 0x3, 0x4, 0x5, 0x6, 0x7}, - }, - { - DstIP: net.IPv4(10, 168, 0, 3).To4(), - DstMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - }, - }, - }, - { - name: "oneRequestWithCacheMiss", - ipMacPairs: []*ipMacPair{}, - requests: []*scan.Request{ - {DstIP: net.IPv4(10, 168, 0, 2).To4()}, - }, - expectedRequests: []*scan.Request{ - { - DstIP: net.IPv4(10, 168, 0, 2).To4(), - Err: fmt.Errorf("no destination MAC address for %s", net.IPv4(10, 168, 0, 2)), - }, - }, - }, - } - - for _, vtt := range tests { - tt := vtt - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - done := make(chan interface{}) - - go func() { - defer close(done) - - // prefill ARP cache - cache := NewCache() - for _, ipMac := range tt.ipMacPairs { - cache.Put(ipMac.ip, ipMac.mac) - } - // prefil input requests - requestsCh := make(chan *scan.Request, len(tt.requests)) - for _, request := range tt.requests { - requestsCh <- request - } - close(requestsCh) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - - ctx := context.Background() - scanRange := &scan.Range{} - reqgen.EXPECT().GenerateRequests(ctx, scanRange).Return(requestsCh, nil) - - cachegen := NewCacheRequestGenerator(reqgen, tt.gatewayMAC, cache) - results, err := cachegen.GenerateRequests(ctx, scanRange) - if !assert.NoError(t, err) { - return - } - - for _, expectedResult := range tt.expectedRequests { - result := <-results - assert.Equal(t, expectedResult, result) - } - - _, ok := <-results - assert.False(t, ok, "results chan is not empty") - - }() - select { - case <-done: - case <-time.After(3 * time.Second): - t.Fatal("test timeout") - } - }) - } -} - -func TestCacheRequestGeneratorReturnsError(t *testing.T) { - t.Parallel() - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - reqgen.EXPECT().GenerateRequests(gomock.Any(), gomock.Any()). - Return(nil, errors.New("request error")) - - cachegen := NewCacheRequestGenerator(reqgen, nil, NewCache()) - _, err := cachegen.GenerateRequests(context.Background(), &scan.Range{}) - require.Error(t, err) -} diff --git a/pkg/scan/arp/result_easyjson.go b/pkg/scan/arp/result_easyjson.go deleted file mode 100644 index da00fa0..0000000 --- a/pkg/scan/arp/result_easyjson.go +++ /dev/null @@ -1,106 +0,0 @@ -// Code generated by easyjson for marshaling/unmarshaling. DO NOT EDIT. - -package arp - -import ( - json "encoding/json" - easyjson "github.com/mailru/easyjson" - jlexer "github.com/mailru/easyjson/jlexer" - jwriter "github.com/mailru/easyjson/jwriter" -) - -// suppress unused package warning -var ( - _ *json.RawMessage - _ *jlexer.Lexer - _ *jwriter.Writer - _ easyjson.Marshaler -) - -func easyjsonD3b49167DecodeGithubComVByteCpuSxPkgScanArp(in *jlexer.Lexer, out *ScanResult) { - isTopLevel := in.IsStart() - if in.IsNull() { - if isTopLevel { - in.Consumed() - } - in.Skip() - return - } - in.Delim('{') - for !in.IsDelim('}') { - key := in.UnsafeFieldName(false) - in.WantColon() - switch key { - case "ip": - if in.IsNull() { - in.Skip() - } else { - out.IP = string(in.String()) - } - case "mac": - if in.IsNull() { - in.Skip() - } else { - out.MAC = string(in.String()) - } - case "vendor": - if in.IsNull() { - in.Skip() - } else { - out.Vendor = string(in.String()) - } - default: - in.SkipRecursive() - } - in.WantComma() - } - in.Delim('}') - if isTopLevel { - in.Consumed() - } -} -func easyjsonD3b49167EncodeGithubComVByteCpuSxPkgScanArp(out *jwriter.Writer, in ScanResult) { - out.RawByte('{') - first := true - _ = first - { - const prefix string = ",\"ip\":" - out.RawString(prefix[1:]) - out.String(string(in.IP)) - } - { - const prefix string = ",\"mac\":" - out.RawString(prefix) - out.String(string(in.MAC)) - } - { - const prefix string = ",\"vendor\":" - out.RawString(prefix) - out.String(string(in.Vendor)) - } - out.RawByte('}') -} - -// MarshalJSON supports json.Marshaler interface -func (v ScanResult) MarshalJSON() ([]byte, error) { - w := jwriter.Writer{} - easyjsonD3b49167EncodeGithubComVByteCpuSxPkgScanArp(&w, v) - return w.Buffer.BuildBytes(), w.Error -} - -// MarshalEasyJSON supports easyjson.Marshaler interface -func (v ScanResult) MarshalEasyJSON(w *jwriter.Writer) { - easyjsonD3b49167EncodeGithubComVByteCpuSxPkgScanArp(w, v) -} - -// UnmarshalJSON supports json.Unmarshaler interface -func (v *ScanResult) UnmarshalJSON(data []byte) error { - r := jlexer.Lexer{Data: data} - easyjsonD3b49167DecodeGithubComVByteCpuSxPkgScanArp(&r, v) - return r.Error() -} - -// UnmarshalEasyJSON supports easyjson.Unmarshaler interface -func (v *ScanResult) UnmarshalEasyJSON(l *jlexer.Lexer) { - easyjsonD3b49167DecodeGithubComVByteCpuSxPkgScanArp(l, v) -} diff --git a/pkg/scan/engine.go b/pkg/scan/engine.go index 4998452..2331f35 100644 --- a/pkg/scan/engine.go +++ b/pkg/scan/engine.go @@ -5,6 +5,7 @@ package scan import ( "context" "net" + "net/netip" "sync" "time" @@ -18,8 +19,9 @@ type PortRange struct { type Range struct { Interface *net.Interface - DstSubnet *net.IPNet - SrcIP net.IP + DstPrefix netip.Prefix + DstZone string + SrcIP netip.Addr SrcMAC net.HardwareAddr Ports []*PortRange } diff --git a/pkg/scan/engine_test.go b/pkg/scan/engine_test.go index 41667d2..a7a7f89 100644 --- a/pkg/scan/engine_test.go +++ b/pkg/scan/engine_test.go @@ -6,6 +6,7 @@ import ( "context" "errors" "net" + "net/netip" "sort" "testing" "time" @@ -110,10 +111,7 @@ func TestPacketEngineStartCollectsAllErrors(t *testing.T) { e := NewPacketEngine(ps, s, r) _, out := e.Start(context.Background(), &Range{ - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 1), - Mask: net.CIDRMask(32, 32), - }, + DstPrefix: netip.MustParsePrefix("192.168.0.1/32"), Ports: []*PortRange{ { StartPort: 888, @@ -136,7 +134,7 @@ func TestPacketSourceReturnsError(t *testing.T) { pktgen := NewMockPacketGenerator(ctrl) scanRange := &Range{ - SrcIP: net.IPv4(192, 168, 0, 1), + SrcIP: ip4(192, 168, 0, 1), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, Ports: []*PortRange{ { @@ -162,7 +160,7 @@ func TestPacketSourceReturnsData(t *testing.T) { pktgen := NewMockPacketGenerator(ctrl) scanRange := &Range{ - SrcIP: net.IPv4(192, 168, 0, 1), + SrcIP: ip4(192, 168, 0, 1), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, Ports: []*PortRange{ { @@ -194,7 +192,7 @@ func TestRateLimitScanner(t *testing.T) { ctrl := gomock.NewController(t) scanner := NewMockScanner(ctrl) - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} + req1 := &Request{DstIP: ip4(192, 168, 0, 1), DstPort: 22} expectedResult := &mockScanResult{"id1"} scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1). Return(expectedResult, nil).AnyTimes() @@ -264,7 +262,7 @@ func TestScanEngineWithScannerError(t *testing.T) { ctx := context.Background() requests := make(chan *Request, 1) - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} + req1 := &Request{DstIP: ip4(192, 168, 0, 1), DstPort: 22} requests <- req1 close(requests) reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). @@ -287,8 +285,8 @@ func TestScanEngineWithResults(t *testing.T) { defer cancel() requests := make(chan *Request, 2) - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} - req2 := &Request{DstIP: net.IPv4(192, 168, 0, 2), DstPort: 22} + req1 := &Request{DstIP: ip4(192, 168, 0, 1), DstPort: 22} + req2 := &Request{DstIP: ip4(192, 168, 0, 2), DstPort: 22} requests <- req1 requests <- req2 close(requests) diff --git a/pkg/scan/generator_test.go b/pkg/scan/generator_test.go index a14dbaf..53d81b9 100644 --- a/pkg/scan/generator_test.go +++ b/pkg/scan/generator_test.go @@ -3,7 +3,6 @@ package scan import ( "context" "errors" - "net" "runtime" "testing" "time" @@ -46,14 +45,14 @@ func TestGeneratorPacketsWithOnePair(t *testing.T) { port := uint16(888) in := make(chan *Request, 1) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}) g := NewPacketGenerator(f) @@ -71,14 +70,14 @@ func TestMultiGeneratorPacketsWithOnePair(t *testing.T) { port := uint16(888) in := make(chan *Request, 1) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}) g := NewPacketMultiGenerator(f, runtime.NumCPU()) @@ -96,18 +95,18 @@ func TestGeneratorPacketsWithTwoPairs(t *testing.T) { port := uint16(888) in := make(chan *Request, 2) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port + 1} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port + 1} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port + 1}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port + 1}) g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) @@ -127,18 +126,18 @@ func TestMultiGeneratorPacketsWithTwoPairs(t *testing.T) { port := uint16(888) in := make(chan *Request, 2) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port + 1} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port + 1} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port + 1}) + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port + 1}) g := NewPacketMultiGenerator(f, runtime.NumCPU()) @@ -179,14 +178,14 @@ func TestGeneratorPacketsReturnsFillError(t *testing.T) { port := uint16(888) in := make(chan *Request, 1) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}). + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}). Return(errors.New("failed request")) g := NewPacketGenerator(f) @@ -224,14 +223,14 @@ func TestMultiGeneratorPacketsReturnsFillError(t *testing.T) { port := uint16(888) in := make(chan *Request, 1) - in <- &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port} + in <- &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port} close(in) ctrl := gomock.NewController(t) f := NewMockPacketFiller(ctrl) f.EXPECT(). Fill(gomock.Not(gomock.Nil()), - &Request{DstIP: net.IPv4(192, 168, 0, 1).To4(), DstPort: port}). + &Request{DstIP: ip4(192, 168, 0, 1), DstPort: port}). Return(errors.New("failed request")) g := NewPacketMultiGenerator(f, runtime.NumCPU()) diff --git a/pkg/scan/icmp/bpf.go b/pkg/scan/icmp/bpf.go index 664627a..3022b5a 100644 --- a/pkg/scan/icmp/bpf.go +++ b/pkg/scan/icmp/bpf.go @@ -12,11 +12,19 @@ const MaxPacketLength = 1518 func BPFFilter(r *scan.Range) (filter string, maxPacketLength int) { var sb strings.Builder + if r.SrcIP.Is6() || (r.DstPrefix.IsValid() && r.DstPrefix.Addr().Is6()) { + sb.WriteString("icmp6 and icmp6[0]!=128") + if r.DstPrefix.IsValid() { + sb.WriteString(" and ip6 src net ") + sb.WriteString(r.DstPrefix.String()) + } + return sb.String(), MaxPacketLength + } // filter ECHO requests sb.WriteString("icmp and icmp[0]!=8") - if r.DstSubnet != nil { + if r.DstPrefix.IsValid() { sb.WriteString(" and ip src net ") - sb.WriteString(r.DstSubnet.String()) + sb.WriteString(r.DstPrefix.String()) } return sb.String(), MaxPacketLength } diff --git a/pkg/scan/icmp/bpf_test.go b/pkg/scan/icmp/bpf_test.go index ecffa9b..8f4bfa4 100644 --- a/pkg/scan/icmp/bpf_test.go +++ b/pkg/scan/icmp/bpf_test.go @@ -1,7 +1,7 @@ package icmp import ( - "net" + "net/netip" "testing" "github.com/stretchr/testify/assert" @@ -22,13 +22,8 @@ func TestBPFFilter(t *testing.T) { scanRange: &scan.Range{}, }, { - name: "OneSubnet", - scanRange: &scan.Range{ - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0), - Mask: net.CIDRMask(24, 32), - }, - }, + name: "OneSubnet", + scanRange: &scan.Range{DstPrefix: netip.MustParsePrefix("192.168.0.0/24")}, expectedFilter: "icmp and icmp[0]!=8 and ip src net 192.168.0.0/24", }, } diff --git a/pkg/scan/icmp/icmp.go b/pkg/scan/icmp/icmp.go index 406c482..c475cba 100644 --- a/pkg/scan/icmp/icmp.go +++ b/pkg/scan/icmp/icmp.go @@ -5,6 +5,8 @@ package icmp import ( "fmt" rand "math/rand/v2" + "net" + "net/netip" "github.com/google/gopacket" "github.com/google/gopacket/layers" @@ -23,12 +25,19 @@ type Response struct { type ScanResult struct { ScanType string `json:"scan"` IP string `json:"ip"` - TTL uint8 `json:"ttl"` + TTL uint8 `json:"ttl,omitempty"` + HopLimit uint8 `json:"hop_limit,omitempty"` ICMP *Response `json:"icmp"` } func (r *ScanResult) String() string { - return fmt.Sprintf("%-20s %-5d %-5d %-5d", r.IP, r.ICMP.Type, r.ICMP.Code, r.TTL) + ttl := r.TTL + width := 20 + if r.HopLimit != 0 { + ttl = r.HopLimit + width = 40 + } + return fmt.Sprintf("%-*s %-5d %-5d %-5d", width, r.IP, r.ICMP.Type, r.ICMP.Code, ttl) } func (r *ScanResult) ID() string { @@ -44,8 +53,16 @@ type ScanMethod struct { // Assert that icmp.ScanMethod conforms to the scan.PacketMethod interface var _ scan.PacketMethod = (*ScanMethod)(nil) -func NewScanMethod(psrc scan.PacketSource, results scan.ResultChan, vpnMode bool) *ScanMethod { - pp := NewPacketProcessor(ScanType, results, vpnMode) +func NewScanMethod(psrc scan.PacketSource, results scan.ResultChan, vpnMode bool, ipv6 ...bool) *ScanMethod { + pp := NewPacketProcessor(ScanType, results, vpnMode, ipv6...) + return newScanMethod(psrc, pp) +} + +func NewScanMethodForFamily(psrc scan.PacketSource, results scan.ResultChan, vpnMode, ipv6 bool, zone string) *ScanMethod { + return newScanMethod(psrc, newPacketProcessor(ScanType, results, vpnMode, ipv6, zone)) +} + +func newScanMethod(psrc scan.PacketSource, pp *PacketProcessor) *ScanMethod { return &ScanMethod{ PacketSource: psrc, Processor: pp, @@ -61,17 +78,42 @@ type PacketProcessor struct { rcvDecoded []gopacket.LayerType rcvEth layers.Ethernet rcvIP layers.IPv4 + rcvIPv6 layers.IPv6 rcvICMP layers.ICMPv4 + rcvICMPv6 layers.ICMPv6 + ipv6 bool + zone string +} + +func NewPacketProcessor(scanType string, results scan.ResultChan, vpnMode bool, ipv6 ...bool) *PacketProcessor { + isIPv6 := false + if len(ipv6) > 0 { + isIPv6 = ipv6[0] + } + return newPacketProcessor(scanType, results, vpnMode, isIPv6, "") } -func NewPacketProcessor(scanType string, results scan.ResultChan, vpnMode bool) *PacketProcessor { - p := &PacketProcessor{scanType: scanType, results: results} +func NewPacketProcessorForFamily(scanType string, results scan.ResultChan, vpnMode, ipv6 bool, zone string) *PacketProcessor { + return newPacketProcessor(scanType, results, vpnMode, ipv6, zone) +} + +func newPacketProcessor(scanType string, results scan.ResultChan, vpnMode, ipv6 bool, zone string) *PacketProcessor { + p := &PacketProcessor{scanType: scanType, results: results, ipv6: ipv6, zone: zone} layerType := layers.LayerTypeEthernet if vpnMode { - layerType = layers.LayerTypeIPv4 + if p.ipv6 { + layerType = layers.LayerTypeIPv6 + } else { + layerType = layers.LayerTypeIPv4 + } + } + var parser *gopacket.DecodingLayerParser + if p.ipv6 { + parser = gopacket.NewDecodingLayerParser(layerType, &p.rcvEth, &p.rcvIPv6, &p.rcvICMPv6) + } else { + parser = gopacket.NewDecodingLayerParser(layerType, &p.rcvEth, &p.rcvIP, &p.rcvICMP) } - parser := gopacket.NewDecodingLayerParser(layerType, &p.rcvEth, &p.rcvIP, &p.rcvICMP) parser.IgnoreUnsupported = true p.parser = parser return p @@ -85,9 +127,25 @@ func (p *PacketProcessor) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo if err = p.parser.DecodeLayers(data, &p.rcvDecoded); err != nil { return } - if !validPacket(p.rcvDecoded) { + if !p.validPacket() { return } + if p.ipv6 { + address, _ := netip.AddrFromSlice(p.rcvIPv6.SrcIP) + if address.IsLinkLocalUnicast() && p.zone != "" { + address = address.WithZone(p.zone) + } + p.results.Put(&ScanResult{ + ScanType: p.scanType, + IP: address.String(), + HopLimit: p.rcvIPv6.HopLimit, + ICMP: &Response{ + Type: p.rcvICMPv6.TypeCode.Type(), + Code: p.rcvICMPv6.TypeCode.Code(), + }, + }) + return nil + } p.results.Put(&ScanResult{ ScanType: p.scanType, @@ -101,6 +159,13 @@ func (p *PacketProcessor) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo return } +func (p *PacketProcessor) validPacket() bool { + if p.ipv6 { + return len(p.rcvDecoded) == 3 || (len(p.rcvDecoded) == 2 && p.rcvDecoded[0] == layers.LayerTypeIPv6) + } + return validPacket(p.rcvDecoded) +} + func validPacket(decoded []gopacket.LayerType) bool { return len(decoded) == 3 || (len(decoded) == 2 && decoded[0] == layers.LayerTypeIPv4) } @@ -114,6 +179,12 @@ type PacketFiller struct { code uint8 payload []byte vpnMode bool + + hopLimit uint8 + nextHeader layers.IPProtocol + payloadLength uint16 + icmpv6Type uint8 + icmpv6Code uint8 } // Assert that icmp.PacketFiller conforms to the scan.PacketFiller interface @@ -171,17 +242,40 @@ func WithVPNmode(vpnMode bool) PacketFillerOption { } } +func WithHopLimit(hopLimit uint8) PacketFillerOption { + return func(f *PacketFiller) { f.hopLimit = hopLimit } +} + +func WithNextHeader(nextHeader uint8) PacketFillerOption { + return func(f *PacketFiller) { f.nextHeader = layers.IPProtocol(nextHeader) } +} + +func WithPayloadLength(payloadLength uint16) PacketFillerOption { + return func(f *PacketFiller) { f.payloadLength = payloadLength } +} + +func WithICMPv6Type(typ uint8) PacketFillerOption { + return func(f *PacketFiller) { f.icmpv6Type = typ } +} + +func WithICMPv6Code(code uint8) PacketFillerOption { + return func(f *PacketFiller) { f.icmpv6Code = code } +} + func NewPacketFiller(opts ...PacketFillerOption) *PacketFiller { payload := make([]byte, 48) fillRandomPayload(payload) f := &PacketFiller{ // typical TTL value for Linux - ttl: 64, - proto: layers.IPProtocolICMPv4, - flags: layers.IPv4DontFragment, - typ: layers.ICMPv4TypeEchoRequest, - code: 0, - payload: payload, + ttl: 64, + proto: layers.IPProtocolICMPv4, + flags: layers.IPv4DontFragment, + typ: layers.ICMPv4TypeEchoRequest, + code: 0, + payload: payload, + hopLimit: 64, + nextHeader: layers.IPProtocolICMPv6, + icmpv6Type: layers.ICMPv6TypeEchoRequest, } for _, o := range opts { o(f) @@ -196,6 +290,9 @@ func fillRandomPayload(payload []byte) { } func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (err error) { + if r.DstIP.Is6() { + return f.fillIPv6(packet, r) + } ip := &layers.IPv4{ Version: 4, @@ -209,8 +306,8 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e TTL: f.ttl, Length: f.length, Protocol: f.proto, - SrcIP: r.SrcIP, - DstIP: r.DstIP, + SrcIP: r.SrcIP.AsSlice(), + DstIP: r.DstIP.AsSlice(), } icmp := &layers.ICMPv4{ @@ -234,3 +331,26 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e } return gopacket.SerializeLayers(packet, opt, eth, ip, icmp, gopacket.Payload(f.payload)) } + +func (f *PacketFiller) fillIPv6(packet gopacket.SerializeBuffer, r *scan.Request) error { + ipv6 := &layers.IPv6{ + Version: 6, + Length: f.payloadLength, + NextHeader: f.nextHeader, + HopLimit: f.hopLimit, + SrcIP: net.IP(r.SrcIP.WithZone("").AsSlice()), + DstIP: net.IP(r.DstIP.WithZone("").AsSlice()), + } + icmpv6 := &layers.ICMPv6{TypeCode: layers.CreateICMPv6TypeCode(f.icmpv6Type, f.icmpv6Code)} + if err := icmpv6.SetNetworkLayerForChecksum(ipv6); err != nil { + return err + } + options := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: f.payloadLength == 0} + layersToSerialize := []gopacket.SerializableLayer{ipv6, icmpv6, gopacket.Payload(f.payload)} + if !f.vpnMode { + layersToSerialize = append([]gopacket.SerializableLayer{&layers.Ethernet{ + SrcMAC: r.SrcMAC, DstMAC: r.DstMAC, EthernetType: layers.EthernetTypeIPv6, + }}, layersToSerialize...) + } + return gopacket.SerializeLayers(packet, options, layersToSerialize...) +} diff --git a/pkg/scan/icmp/icmp_test.go b/pkg/scan/icmp/icmp_test.go index e4f323d..73fd888 100644 --- a/pkg/scan/icmp/icmp_test.go +++ b/pkg/scan/icmp/icmp_test.go @@ -3,6 +3,7 @@ package icmp import ( "context" "net" + "net/netip" "testing" "time" @@ -20,8 +21,8 @@ func TestPacketFillerEthernet(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -62,8 +63,8 @@ func TestPacketFillerIPv4(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1), WithVPNmode(true)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -101,8 +102,8 @@ func TestPacketFillerPayload(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -129,8 +130,8 @@ func TestPacketFillerTTL(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -152,8 +153,8 @@ func TestPacketFillerIPTotalLength(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1), WithPayload([]byte("abc"))) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -175,8 +176,8 @@ func TestPacketFillerIPProtocol(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) @@ -197,8 +198,8 @@ func TestPacketFillerIPFlags(t *testing.T) { WithType(layers.ICMPv4TypeTimestampRequest), WithCode(1)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, }) diff --git a/pkg/scan/icmp/ipv6_test.go b/pkg/scan/icmp/ipv6_test.go new file mode 100644 index 0000000..4613d49 --- /dev/null +++ b/pkg/scan/icmp/ipv6_test.go @@ -0,0 +1,56 @@ +package icmp + +import ( + "context" + "encoding/json" + "net" + "net/netip" + "testing" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" +) + +func TestPacketFillerIPv6(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + err := NewPacketFiller(WithHopLimit(37), WithICMPv6Type(layers.ICMPv6TypeEchoRequest)).Fill(packet, &scan.Request{ + SrcIP: netip.MustParseAddr("2001:db8::1"), DstIP: netip.MustParseAddr("2001:db8::2"), + SrcMAC: net.HardwareAddr{2, 0, 0, 0, 0, 1}, DstMAC: net.HardwareAddr{2, 0, 0, 0, 0, 2}, + }) + require.NoError(t, err) + + decoded := gopacket.NewPacket(packet.Bytes(), layers.LayerTypeEthernet, gopacket.Default) + ipv6 := decoded.Layer(layers.LayerTypeIPv6).(*layers.IPv6) + require.Equal(t, uint8(37), ipv6.HopLimit) + require.Equal(t, layers.IPProtocolICMPv6, ipv6.NextHeader) + require.Equal(t, net.ParseIP("2001:db8::2"), ipv6.DstIP) + icmpv6 := decoded.Layer(layers.LayerTypeICMPv6).(*layers.ICMPv6) + require.Equal(t, uint8(layers.ICMPv6TypeEchoRequest), icmpv6.TypeCode.Type()) +} + +func TestPacketProcessorIPv6ReportsHopLimit(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + eth := &layers.Ethernet{SrcMAC: net.HardwareAddr{2, 0, 0, 0, 0, 2}, DstMAC: net.HardwareAddr{2, 0, 0, 0, 0, 1}, EthernetType: layers.EthernetTypeIPv6} + ipv6 := &layers.IPv6{Version: 6, HopLimit: 51, NextHeader: layers.IPProtocolICMPv6, SrcIP: net.ParseIP("2001:db8::2"), DstIP: net.ParseIP("2001:db8::1")} + icmpv6 := &layers.ICMPv6{TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeEchoReply, 0)} + require.NoError(t, icmpv6.SetNetworkLayerForChecksum(ipv6)) + require.NoError(t, gopacket.SerializeLayers(packet, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, eth, ipv6, icmpv6, gopacket.Payload([]byte{0, 1, 0, 1}))) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + processor := NewPacketProcessor(ScanType, scan.NewResultChan(ctx, 1), false, true) + require.NoError(t, processor.ProcessPacketData(packet.Bytes(), &gopacket.CaptureInfo{})) + result := (<-processor.Results()).(*ScanResult) + require.Equal(t, "2001:db8::2", result.IP) + require.Equal(t, uint8(51), result.HopLimit) + require.Zero(t, result.TTL) + encoded, err := json.Marshal(result) + require.NoError(t, err) + require.JSONEq(t, `{"scan":"icmp","ip":"2001:db8::2","hop_limit":51,"icmp":{"type":129,"code":0}}`, string(encoded)) +} diff --git a/pkg/scan/icmp/result_easyjson.go b/pkg/scan/icmp/result_easyjson.go index 38a1374..2c6217c 100644 --- a/pkg/scan/icmp/result_easyjson.go +++ b/pkg/scan/icmp/result_easyjson.go @@ -49,6 +49,12 @@ func easyjsonD3b49167DecodeGithubComVByteCpuSxPkgScanIcmp(in *jlexer.Lexer, out } else { out.TTL = uint8(in.Uint8()) } + case "hop_limit": + if in.IsNull() { + in.Skip() + } else { + out.HopLimit = uint8(in.Uint8()) + } case "icmp": if in.IsNull() { in.Skip() @@ -83,11 +89,16 @@ func easyjsonD3b49167EncodeGithubComVByteCpuSxPkgScanIcmp(out *jwriter.Writer, i out.RawString(prefix) out.String(string(in.IP)) } - { + if in.TTL != 0 { const prefix string = ",\"ttl\":" out.RawString(prefix) out.Uint8(uint8(in.TTL)) } + if in.HopLimit != 0 { + const prefix string = ",\"hop_limit\":" + out.RawString(prefix) + out.Uint8(uint8(in.HopLimit)) + } { const prefix string = ",\"icmp\":" out.RawString(prefix) diff --git a/pkg/scan/ip_test.go b/pkg/scan/ip_test.go new file mode 100644 index 0000000..7e1a6d3 --- /dev/null +++ b/pkg/scan/ip_test.go @@ -0,0 +1,7 @@ +package scan + +import "net/netip" + +func ip4(a, b, c, d byte) netip.Addr { + return netip.AddrFrom4([4]byte{a, b, c, d}) +} diff --git a/pkg/scan/mock_request_test.go b/pkg/scan/mock_request_test.go index 182b9c5..12cdfd9 100644 --- a/pkg/scan/mock_request_test.go +++ b/pkg/scan/mock_request_test.go @@ -11,7 +11,7 @@ package scan import ( context "context" - net "net" + netip "net/netip" reflect "reflect" gomock "go.uber.org/mock/gomock" @@ -81,10 +81,10 @@ func (m *MockIPGenerator) EXPECT() *MockIPGeneratorMockRecorder { } // IPs mocks base method. -func (m *MockIPGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { +func (m *MockIPGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[netip.Addr], error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "IPs", ctx, r) - ret0, _ := ret[0].(<-chan GeneratorResult[net.IP]) + ret0, _ := ret[0].(<-chan GeneratorResult[netip.Addr]) ret1, _ := ret[1].(error) return ret0, ret1 } @@ -159,7 +159,7 @@ func (m *MockIPContainer) EXPECT() *MockIPContainerMockRecorder { } // Contains mocks base method. -func (m *MockIPContainer) Contains(ip net.IP) (bool, error) { +func (m *MockIPContainer) Contains(ip netip.Addr) (bool, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "Contains", ip) ret0, _ := ret[0].(bool) diff --git a/pkg/scan/ndp/bpf.go b/pkg/scan/ndp/bpf.go new file mode 100644 index 0000000..6f14cf7 --- /dev/null +++ b/pkg/scan/ndp/bpf.go @@ -0,0 +1,17 @@ +package ndp + +import ( + "fmt" + + "github.com/v-byte-cpu/sx/pkg/scan" +) + +const MaxPacketLength = 1518 + +func BPFFilter(r *scan.Range) (string, int) { + filter := "icmp6 and icmp6[0] == 136" + if r.DstPrefix.IsValid() { + filter = fmt.Sprintf("%s and src net %s", filter, r.DstPrefix) + } + return filter, MaxPacketLength +} diff --git a/pkg/scan/ndp/bpf_test.go b/pkg/scan/ndp/bpf_test.go new file mode 100644 index 0000000..703fda3 --- /dev/null +++ b/pkg/scan/ndp/bpf_test.go @@ -0,0 +1,17 @@ +package ndp + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" +) + +func TestBPFFilter(t *testing.T) { + t.Parallel() + + filter, maxPacketLength := BPFFilter(&scan.Range{DstPrefix: netip.MustParsePrefix("2001:db8::/120")}) + require.Equal(t, "icmp6 and icmp6[0] == 136 and src net 2001:db8::/120", filter) + require.Equal(t, 1518, maxPacketLength) +} diff --git a/pkg/scan/ndp/ndp.go b/pkg/scan/ndp/ndp.go new file mode 100644 index 0000000..e8515b4 --- /dev/null +++ b/pkg/scan/ndp/ndp.go @@ -0,0 +1,134 @@ +package ndp + +import ( + "errors" + "net" + "net/netip" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/macs" + "github.com/v-byte-cpu/sx/pkg/scan" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" +) + +const ScanType = "ndp" + +var errIPv6Required = errors.New("NDP requires IPv6 source and destination addresses") + +type ScanMethod struct { + scan.PacketSource + parser *gopacket.DecodingLayerParser + results scan.ResultChan + + decoded []gopacket.LayerType + eth layers.Ethernet + ipv6 layers.IPv6 + icmpv6 layers.ICMPv6 + na layers.ICMPv6NeighborAdvertisement + zone string +} + +var _ scan.PacketMethod = (*ScanMethod)(nil) + +func NewScanMethod(source scan.PacketSource, results scan.ResultChan, zone ...string) *ScanMethod { + method := &ScanMethod{PacketSource: source, results: results} + if len(zone) > 0 { + method.zone = zone[0] + } + method.parser = gopacket.NewDecodingLayerParser( + layers.LayerTypeEthernet, + &method.eth, + &method.ipv6, + &method.icmpv6, + &method.na, + ) + method.parser.IgnoreUnsupported = true + return method +} + +func (m *ScanMethod) Results() <-chan scan.Result { + return m.results.Chan() +} + +func (m *ScanMethod) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo) error { + if err := m.parser.DecodeLayers(data, &m.decoded); err != nil { + return err + } + if len(m.decoded) != 4 || m.ipv6.HopLimit != 255 || + m.icmpv6.TypeCode.Type() != layers.ICMPv6TypeNeighborAdvertisement || m.icmpv6.TypeCode.Code() != 0 { + return nil + } + + mac := m.eth.SrcMAC + for _, option := range m.na.Options { + if option.Type == layers.ICMPv6OptTargetAddress && len(option.Data) >= 6 { + mac = net.HardwareAddr(option.Data[:6]) + break + } + } + var prefix [3]byte + copy(prefix[:], mac) + address, _ := netip.AddrFromSlice(m.na.TargetAddress) + if address.IsLinkLocalUnicast() && m.zone != "" { + address = address.WithZone(m.zone) + } + m.results.Put(&neighbor.ScanResult{ + IP: address.String(), + MAC: mac.String(), + Vendor: macs.ValidMACPrefixMap[prefix], + }) + return nil +} + +type PacketFiller struct{} + +func NewPacketFiller() *PacketFiller { + return &PacketFiller{} +} + +func (*PacketFiller) Fill(packet gopacket.SerializeBuffer, request *scan.Request) error { + if !request.SrcIP.Is6() || !request.DstIP.Is6() { + return errIPv6Required + } + target := request.DstIP.WithZone("").As16() + multicastIP := [16]byte{0xff, 0x02} + multicastIP[11] = 0x01 + multicastIP[12] = 0xff + copy(multicastIP[13:], target[13:]) + multicastMAC := net.HardwareAddr{0x33, 0x33, 0xff, target[13], target[14], target[15]} + + eth := &layers.Ethernet{ + SrcMAC: request.SrcMAC, + DstMAC: multicastMAC, + EthernetType: layers.EthernetTypeIPv6, + } + ipv6 := &layers.IPv6{ + Version: 6, + HopLimit: 255, + NextHeader: layers.IPProtocolICMPv6, + SrcIP: net.IP(request.SrcIP.WithZone("").AsSlice()), + DstIP: net.IP(multicastIP[:]), + } + icmpv6 := &layers.ICMPv6{ + TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeNeighborSolicitation, 0), + } + if err := icmpv6.SetNetworkLayerForChecksum(ipv6); err != nil { + return err + } + ns := &layers.ICMPv6NeighborSolicitation{ + TargetAddress: net.IP(target[:]), + Options: layers.ICMPv6Options{{ + Type: layers.ICMPv6OptSourceAddress, + Data: request.SrcMAC, + }}, + } + return gopacket.SerializeLayers( + packet, + gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, + eth, + ipv6, + icmpv6, + ns, + ) +} diff --git a/pkg/scan/ndp/ndp_test.go b/pkg/scan/ndp/ndp_test.go new file mode 100644 index 0000000..70602f3 --- /dev/null +++ b/pkg/scan/ndp/ndp_test.go @@ -0,0 +1,85 @@ +package ndp + +import ( + "context" + "net" + "net/netip" + "testing" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" +) + +func TestPacketFillerBuildsNeighborSolicitation(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + filler := NewPacketFiller() + err := filler.Fill(packet, &scan.Request{ + SrcIP: netip.MustParseAddr("fe80::2"), + DstIP: netip.MustParseAddr("fe80::1"), + SrcMAC: net.HardwareAddr{0x02, 0, 0, 0, 0, 2}, + }) + require.NoError(t, err) + + decoded := gopacket.NewPacket(packet.Bytes(), layers.LayerTypeEthernet, gopacket.Default) + eth := decoded.Layer(layers.LayerTypeEthernet).(*layers.Ethernet) + require.Equal(t, net.HardwareAddr{0x33, 0x33, 0xff, 0, 0, 1}, eth.DstMAC) + require.Equal(t, layers.EthernetTypeIPv6, eth.EthernetType) + + ipv6 := decoded.Layer(layers.LayerTypeIPv6).(*layers.IPv6) + require.Equal(t, uint8(255), ipv6.HopLimit) + require.Equal(t, layers.IPProtocolICMPv6, ipv6.NextHeader) + require.Equal(t, net.ParseIP("ff02::1:ff00:1"), ipv6.DstIP) + + icmp := decoded.Layer(layers.LayerTypeICMPv6).(*layers.ICMPv6) + require.Equal(t, uint8(layers.ICMPv6TypeNeighborSolicitation), icmp.TypeCode.Type()) + + ns := decoded.Layer(layers.LayerTypeICMPv6NeighborSolicitation).(*layers.ICMPv6NeighborSolicitation) + require.Equal(t, net.ParseIP("fe80::1"), ns.TargetAddress) + require.Len(t, ns.Options, 1) + require.Equal(t, layers.ICMPv6OptSourceAddress, ns.Options[0].Type) + require.Equal(t, []byte{0x02, 0, 0, 0, 0, 2}, ns.Options[0].Data) +} + +func TestProcessPacketDataEmitsNeighbor(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + eth := &layers.Ethernet{ + SrcMAC: net.HardwareAddr{0x02, 0, 0, 0, 0, 1}, + DstMAC: net.HardwareAddr{0x02, 0, 0, 0, 0, 2}, + EthernetType: layers.EthernetTypeIPv6, + } + ipv6 := &layers.IPv6{ + Version: 6, + HopLimit: 255, + NextHeader: layers.IPProtocolICMPv6, + SrcIP: net.ParseIP("fe80::1"), + DstIP: net.ParseIP("fe80::2"), + } + icmp := &layers.ICMPv6{TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeNeighborAdvertisement, 0)} + require.NoError(t, icmp.SetNetworkLayerForChecksum(ipv6)) + na := &layers.ICMPv6NeighborAdvertisement{ + Flags: 0x60, + TargetAddress: net.ParseIP("fe80::1"), + Options: layers.ICMPv6Options{{ + Type: layers.ICMPv6OptTargetAddress, + Data: []byte{0x02, 0, 0, 0, 0, 1}, + }}, + } + require.NoError(t, gopacket.SerializeLayers(packet, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, eth, ipv6, icmp, na)) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + results := scan.NewResultChan(ctx, 1) + method := NewScanMethod(nil, results, "en0") + require.NoError(t, method.ProcessPacketData(packet.Bytes(), &gopacket.CaptureInfo{})) + + result := (<-method.Results()).(*neighbor.ScanResult) + require.Equal(t, "fe80::1%en0", result.IP) + require.Equal(t, "02:00:00:00:00:01", result.MAC) +} diff --git a/pkg/scan/neighbor/cache.go b/pkg/scan/neighbor/cache.go new file mode 100644 index 0000000..7a75bc9 --- /dev/null +++ b/pkg/scan/neighbor/cache.go @@ -0,0 +1,124 @@ +package neighbor + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/netip" + "sync" + + "github.com/v-byte-cpu/sx/pkg/scan" +) + +type Cache struct { + cache map[netip.Addr]net.HardwareAddr + mu sync.RWMutex +} + +func NewCache() *Cache { + return &Cache{cache: make(map[netip.Addr]net.HardwareAddr)} +} + +func (c *Cache) Put(ip netip.Addr, mac net.HardwareAddr) { + c.mu.Lock() + defer c.mu.Unlock() + c.cache[cacheKey(ip)] = mac +} + +func (c *Cache) Get(ip netip.Addr) net.HardwareAddr { + c.mu.RLock() + defer c.mu.RUnlock() + key := cacheKey(ip) + if mac := c.cache[key]; mac != nil { + return mac + } + if key.Zone() != "" { + return c.cache[key.WithZone("")] + } + return nil +} + +func (c *Cache) Delete(ip netip.Addr) { + c.mu.Lock() + defer c.mu.Unlock() + delete(c.cache, cacheKey(ip)) +} + +func cacheKey(ip netip.Addr) netip.Addr { + return ip.Unmap() +} + +func FillCache(cache *Cache, input io.Reader) error { + scanner := bufio.NewScanner(input) + line := 0 + for scanner.Scan() { + line++ + var entry ScanResult + if err := json.Unmarshal(scanner.Bytes(), &entry); err != nil { + return fmt.Errorf("neighbor cache: line %d: %w", line, err) + } + ip, err := netip.ParseAddr(entry.IP) + if err != nil { + return fmt.Errorf("neighbor cache: line %d: %w", line, errors.New("invalid IP")) + } + mac, err := net.ParseMAC(entry.MAC) + if err != nil { + return fmt.Errorf("neighbor cache: line %d: %w", line, err) + } + cache.Put(ip, mac) + } + return scanner.Err() +} + +type cacheRequestGenerator struct { + requestGenerator scan.RequestGenerator + gatewayMAC net.HardwareAddr + cache *Cache +} + +func NewCacheRequestGenerator( + requestGenerator scan.RequestGenerator, + gatewayMAC net.HardwareAddr, + cache *Cache, +) scan.RequestGenerator { + return &cacheRequestGenerator{ + requestGenerator: requestGenerator, + gatewayMAC: gatewayMAC, + cache: cache, + } +} + +func (g *cacheRequestGenerator) GenerateRequests( + ctx context.Context, + r *scan.Range, +) (<-chan *scan.Request, error) { + requests, err := g.requestGenerator.GenerateRequests(ctx, r) + if err != nil { + return nil, err + } + result := make(chan *scan.Request, cap(requests)) + go func() { + defer close(result) + for request := range requests { + mac := g.cache.Get(request.DstIP) + if mac == nil { + mac = g.gatewayMAC + } + if mac == nil { + request.Err = fmt.Errorf("no destination MAC address for %s", request.DstIP) + } else { + request.DstMAC = mac + } + select { + case <-ctx.Done(): + return + case result <- request: + } + } + }() + return result, nil +} diff --git a/pkg/scan/neighbor/cache_test.go b/pkg/scan/neighbor/cache_test.go new file mode 100644 index 0000000..b80bd91 --- /dev/null +++ b/pkg/scan/neighbor/cache_test.go @@ -0,0 +1,72 @@ +//go:generate go tool mockgen -package neighbor -destination=mock_request_test.go github.com/v-byte-cpu/sx/pkg/scan RequestGenerator + +package neighbor + +import ( + "context" + "net" + "net/netip" + "strings" + "testing" + + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" + "go.uber.org/mock/gomock" +) + +func TestCacheStoresIPv4AndIPv6Neighbors(t *testing.T) { + t.Parallel() + + cache := NewCache() + ipv4MAC := net.HardwareAddr{0x00, 0x11, 0x22, 0x33, 0x44, 0x55} + ipv6MAC := net.HardwareAddr{0x00, 0xaa, 0xbb, 0xcc, 0xdd, 0xee} + cache.Put(netip.MustParseAddr("192.0.2.1"), ipv4MAC) + cache.Put(netip.MustParseAddr("fe80::1%en0"), ipv6MAC) + otherMAC := net.HardwareAddr{0x00, 1, 1, 1, 1, 1} + cache.Put(netip.MustParseAddr("fe80::1%en1"), otherMAC) + fallbackMAC := net.HardwareAddr{0x00, 2, 2, 2, 2, 2} + cache.Put(netip.MustParseAddr("fe80::2"), fallbackMAC) + + require.Equal(t, ipv4MAC, cache.Get(netip.MustParseAddr("192.0.2.1"))) + require.Equal(t, ipv6MAC, cache.Get(netip.MustParseAddr("fe80::1%en0"))) + require.Equal(t, otherMAC, cache.Get(netip.MustParseAddr("fe80::1%en1"))) + require.Equal(t, fallbackMAC, cache.Get(netip.MustParseAddr("fe80::2%en0"))) +} + +func TestCacheRequestGeneratorUsesNeighborThenGateway(t *testing.T) { + t.Parallel() + + controller := gomock.NewController(t) + upstream := NewMockRequestGenerator(controller) + scanRange := &scan.Range{} + requests := make(chan *scan.Request, 2) + requests <- &scan.Request{DstIP: netip.MustParseAddr("2001:db8::1")} + requests <- &scan.Request{DstIP: netip.MustParseAddr("2001:db8::2")} + close(requests) + upstream.EXPECT().GenerateRequests(gomock.Any(), scanRange).Return((<-chan *scan.Request)(requests), nil) + + directMAC := net.HardwareAddr{0x00, 1, 2, 3, 4, 5} + gatewayMAC := net.HardwareAddr{0x00, 6, 7, 8, 9, 10} + cache := NewCache() + cache.Put(netip.MustParseAddr("2001:db8::1"), directMAC) + generated, err := NewCacheRequestGenerator(upstream, gatewayMAC, cache).GenerateRequests(context.Background(), scanRange) + require.NoError(t, err) + + first := <-generated + second := <-generated + require.Equal(t, directMAC, net.HardwareAddr(first.DstMAC)) + require.Equal(t, gatewayMAC, net.HardwareAddr(second.DstMAC)) + require.NoError(t, first.Err) + require.NoError(t, second.Err) +} + +func TestFillCacheReportsInvalidAddressLine(t *testing.T) { + t.Parallel() + + err := FillCache(NewCache(), strings.NewReader(strings.Join([]string{ + `{"ip":"192.0.2.1","mac":"00:11:22:33:44:55"}`, + `{"ip":"not-an-ip","mac":"00:11:22:33:44:55"}`, + }, "\n"))) + + require.EqualError(t, err, "neighbor cache: line 2: invalid IP") +} diff --git a/pkg/scan/arp/mock_request_test.go b/pkg/scan/neighbor/mock_request_test.go similarity index 90% rename from pkg/scan/arp/mock_request_test.go rename to pkg/scan/neighbor/mock_request_test.go index 9913ca1..4ed7f4f 100644 --- a/pkg/scan/arp/mock_request_test.go +++ b/pkg/scan/neighbor/mock_request_test.go @@ -3,11 +3,11 @@ // // Generated by this command: // -// mockgen -package arp -destination=mock_request_test.go github.com/v-byte-cpu/sx/pkg/scan RequestGenerator +// mockgen -package neighbor -destination=mock_request_test.go github.com/v-byte-cpu/sx/pkg/scan RequestGenerator // -// Package arp is a generated GoMock package. -package arp +// Package neighbor is a generated GoMock package. +package neighbor import ( context "context" diff --git a/pkg/scan/neighbor/result.go b/pkg/scan/neighbor/result.go new file mode 100644 index 0000000..9dc80cb --- /dev/null +++ b/pkg/scan/neighbor/result.go @@ -0,0 +1,30 @@ +package neighbor + +import ( + "encoding/json" + "fmt" + "strings" +) + +type ScanResult struct { + IP string `json:"ip"` + MAC string `json:"mac"` + Vendor string `json:"vendor"` +} + +func (r *ScanResult) String() string { + width := 20 + if strings.ContainsRune(r.IP, ':') { + width = 40 + } + return fmt.Sprintf("%-*s %-20s %s", width, r.IP, r.MAC, r.Vendor) +} + +func (r *ScanResult) ID() string { + return r.IP +} + +func (r *ScanResult) MarshalJSON() ([]byte, error) { + type result ScanResult + return json.Marshal((*result)(r)) +} diff --git a/pkg/scan/request.go b/pkg/scan/request.go index 72a464c..5027e93 100644 --- a/pkg/scan/request.go +++ b/pkg/scan/request.go @@ -9,7 +9,7 @@ import ( "errors" "io" "math/big" - "net" + "net/netip" "time" ) @@ -23,8 +23,8 @@ var ( type Request struct { Meta map[string]interface{} - SrcIP net.IP - DstIP net.IP + SrcIP netip.Addr + DstIP netip.Addr SrcMAC []byte DstMAC []byte DstPort uint16 @@ -97,7 +97,7 @@ func validatePorts(ports []*PortRange) error { // IPGenerator produces IP addresses for a scan range. type IPGenerator interface { // IPs generates the IP addresses described by r until completion or context cancellation. - IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) + IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[netip.Addr], error) } func NewIPGenerator() IPGenerator { @@ -106,31 +106,41 @@ func NewIPGenerator() IPGenerator { type ipGenerator struct{} -func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { - if r.DstSubnet == nil { +func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[netip.Addr], error) { + if !r.DstPrefix.IsValid() { return nil, ErrSubnet } - ipnet := r.DstSubnet - ones, bits := ipnet.Mask.Size() - it, err := newRangeIterator(1 << (bits - ones)) + prefix := r.DstPrefix.Masked() + bits := prefix.Addr().BitLen() + hostBits := bits - prefix.Bits() + if hostBits > 32 { + return nil, errRangeSize + } + it, err := newRangeIterator(1 << hostBits) if err != nil { return nil, err } - baseIP := big.NewInt(0).SetBytes(ipnet.IP.Mask(ipnet.Mask)) + baseIP := big.NewInt(0).SetBytes(prefix.Addr().AsSlice()) baseIP.Sub(baseIP, big.NewInt(1)) - out := make(chan GeneratorResult[net.IP], 100) + out := make(chan GeneratorResult[netip.Addr], 100) go func() { defer close(out) for { i := it.Int() baseIP.Add(baseIP, i) - // TODO IPv6 - ipaddr := baseIP.FillBytes(make([]byte, 4)) + addr, ok := netip.AddrFromSlice(baseIP.FillBytes(make([]byte, bits/8))) baseIP.Sub(baseIP, i) + if !ok { + sendContext(ctx, out, GeneratorResult[netip.Addr]{Err: ErrIP}) + return + } + if r.DstZone != "" { + addr = addr.WithZone(r.DstZone) + } - if !sendContext(ctx, out, GeneratorResult[net.IP]{Value: ipaddr}) { + if !sendContext(ctx, out, GeneratorResult[netip.Addr]{Value: addr}) { return } @@ -242,29 +252,12 @@ func (rg *fileIPPortGenerator) GenerateRequests(ctx context.Context, r *Range) ( defer close(out) defer input.Close() scanner := bufio.NewScanner(input) - var entry IPPort for scanner.Scan() { - entry.IP = "" - entry.Port = 0 - if err := entry.UnmarshalJSON(scanner.Bytes()); err != nil { - sendContext(ctx, out, &Request{Err: ErrJSON}) + request, fatal := parseFileIPPortRequest(scanner.Bytes(), r) + if !sendContext(ctx, out, request) { return } - ip := net.ParseIP(entry.IP) - if ip == nil { - if !sendContext(ctx, out, &Request{Err: ErrIP}) { - return - } - continue - } - if !isValidPort(entry.Port) { - if !sendContext(ctx, out, &Request{Err: ErrPort}) { - return - } - continue - } - if !sendContext(ctx, out, &Request{ - SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: ip, DstPort: uint16(entry.Port)}) { + if fatal { return } } @@ -275,6 +268,27 @@ func (rg *fileIPPortGenerator) GenerateRequests(ctx context.Context, r *Range) ( return out, nil } +func parseFileIPPortRequest(data []byte, r *Range) (*Request, bool) { + var entry IPPort + if err := entry.UnmarshalJSON(data); err != nil { + return &Request{Err: ErrJSON}, true + } + ip, err := netip.ParseAddr(entry.IP) + if err != nil { + return &Request{Err: ErrIP}, false + } + if !isValidPort(entry.Port) { + return &Request{Err: ErrPort}, false + } + ip = ip.Unmap() + if r.SrcIP.IsValid() && r.SrcIP.Is4() != ip.Is4() { + return &Request{Err: ErrIP}, false + } + return &Request{ + SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: ip, DstPort: uint16(entry.Port), + }, false +} + func isValidPort(port int) bool { return port > 0 && port <= 0xFFFF } @@ -287,12 +301,12 @@ func NewFileIPGenerator(openFile OpenFileFunc) IPGenerator { return &fileIPGenerator{openFile} } -func (g *fileIPGenerator) IPs(ctx context.Context, _ *Range) (<-chan GeneratorResult[net.IP], error) { +func (g *fileIPGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[netip.Addr], error) { input, err := g.openFile() if err != nil { return nil, err } - out := make(chan GeneratorResult[net.IP]) + out := make(chan GeneratorResult[netip.Addr]) go func() { defer close(out) defer input.Close() @@ -301,20 +315,25 @@ func (g *fileIPGenerator) IPs(ctx context.Context, _ *Range) (<-chan GeneratorRe for scanner.Scan() { entry = IPPort{} if err := entry.UnmarshalJSON(scanner.Bytes()); err != nil { - sendContext(ctx, out, GeneratorResult[net.IP]{Err: ErrJSON}) + sendContext(ctx, out, GeneratorResult[netip.Addr]{Err: ErrJSON}) + return + } + ip, err := netip.ParseAddr(entry.IP) + if err != nil { + sendContext(ctx, out, GeneratorResult[netip.Addr]{Err: ErrIP}) return } - ip := net.ParseIP(entry.IP) - if ip == nil { - sendContext(ctx, out, GeneratorResult[net.IP]{Err: ErrIP}) + ip = ip.Unmap() + if r.SrcIP.IsValid() && r.SrcIP.Is4() != ip.Is4() { + sendContext(ctx, out, GeneratorResult[netip.Addr]{Err: ErrIP}) return } - if !sendContext(ctx, out, GeneratorResult[net.IP]{Value: ip}) { + if !sendContext(ctx, out, GeneratorResult[netip.Addr]{Value: ip}) { return } } if err = scanner.Err(); err != nil { - sendContext(ctx, out, GeneratorResult[net.IP]{Err: err}) + sendContext(ctx, out, GeneratorResult[netip.Addr]{Err: err}) } }() return out, nil @@ -370,7 +389,8 @@ func readRequest(ctx context.Context, requests <-chan *Request) (request *Reques } type IPContainer interface { - Contains(ip net.IP) (bool, error) + // Contains reports whether ip belongs to the container. + Contains(ip netip.Addr) (bool, error) } type filterIPRequestGenerator struct { diff --git a/pkg/scan/request_test.go b/pkg/scan/request_test.go index 6fffe93..e7e5c84 100644 --- a/pkg/scan/request_test.go +++ b/pkg/scan/request_test.go @@ -1,12 +1,12 @@ package scan import ( - "bytes" "context" "errors" "io" "math/big" "net" + "net/netip" "sort" "strings" "testing" @@ -18,12 +18,9 @@ import ( func newScanRange(opts ...scanRangeOption) *Range { sr := &Range{ - SrcIP: net.IPv4(192, 168, 0, 3), - SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0), - Mask: net.CIDRMask(24, 32), - }, + SrcIP: ip4(192, 168, 0, 3), + SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, + DstPrefix: netip.PrefixFrom(ip4(192, 168, 0, 0), 24), Ports: []*PortRange{ { StartPort: 22, @@ -45,15 +42,15 @@ func withPorts(ports []*PortRange) scanRangeOption { } } -func withSubnet(subnet *net.IPNet) scanRangeOption { +func withPrefix(subnet netip.Prefix) scanRangeOption { return func(sr *Range) { - sr.DstSubnet = subnet + sr.DstPrefix = subnet } } func newScanRequest(opts ...scanRequestOption) *Request { r := &Request{ - SrcIP: net.IPv4(192, 168, 0, 3), + SrcIP: ip4(192, 168, 0, 3), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, } for _, o := range opts { @@ -64,7 +61,7 @@ func newScanRequest(opts ...scanRequestOption) *Request { type scanRequestOption func(sr *Request) -func withDstIP(dstIP net.IP) scanRequestOption { +func withDstIP(dstIP netip.Addr) scanRequestOption { return func(sr *Request) { sr.DstIP = dstIP } @@ -96,8 +93,8 @@ func generatedPort(value uint16) GeneratorResult[uint16] { return GeneratorResult[uint16]{Value: value} } -func generatedIP(value net.IP) GeneratorResult[net.IP] { - return GeneratorResult[net.IP]{Value: value} +func generatedIP(value netip.Addr) GeneratorResult[netip.Addr] { + return GeneratorResult[netip.Addr]{Value: value} } func generationError[T any](err error) GeneratorResult[T] { @@ -279,43 +276,43 @@ func TestIPGenerator(t *testing.T) { tests := []struct { name string scanRange *Range - expected []GeneratorResult[net.IP] + expected []GeneratorResult[netip.Addr] err bool }{ { name: "NilSubnet", - scanRange: newScanRange(withSubnet(nil)), + scanRange: newScanRange(withPrefix(netip.Prefix{})), err: true, }, { name: "OneIP", scanRange: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1), Mask: net.CIDRMask(32, 32)}), + withPrefix(netip.PrefixFrom(ip4(192, 168, 0, 1), 32)), ), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1).To4()), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), }, }, { name: "TwoIPs", scanRange: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(1, 0, 0, 1), Mask: net.CIDRMask(31, 32)}), + withPrefix(netip.PrefixFrom(ip4(1, 0, 0, 1), 31)), ), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(1, 0, 0, 0).To4()), - generatedIP(net.IPv4(1, 0, 0, 1).To4()), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(1, 0, 0, 0)), + generatedIP(ip4(1, 0, 0, 1)), }, }, { name: "FourIPs", scanRange: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 1), Mask: net.CIDRMask(30, 32)}), + withPrefix(netip.PrefixFrom(ip4(10, 0, 0, 1), 30)), ), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(10, 0, 0, 0).To4()), - generatedIP(net.IPv4(10, 0, 0, 1).To4()), - generatedIP(net.IPv4(10, 0, 0, 2).To4()), - generatedIP(net.IPv4(10, 0, 0, 3).To4()), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(10, 0, 0, 0)), + generatedIP(ip4(10, 0, 0, 1)), + generatedIP(ip4(10, 0, 0, 2)), + generatedIP(ip4(10, 0, 0, 3)), }, }, } @@ -334,7 +331,7 @@ func TestIPGenerator(t *testing.T) { require.NoError(t, err) result := collectChannel(ips) sort.Slice(result, func(i, j int) bool { - return bytes.Compare(result[i].Value, result[j].Value) < 1 + return result[i].Value.Less(result[j].Value) }) require.Equal(t, tt.expected, result) }) @@ -354,68 +351,68 @@ func TestIPPortGenerator(t *testing.T) { tests := []struct { name string - ips []GeneratorResult[net.IP] + ips []GeneratorResult[netip.Addr] ports []GeneratorResult[uint16] expected []*Request }{ { name: "OneIpOnePort", - ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ips: []GeneratorResult[netip.Addr]{generatedIP(ip4(192, 168, 0, 1))}, ports: []GeneratorResult[uint16]{generatedPort(888)}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(888)), }, }, { name: "OneIpTwoPorts", - ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ips: []GeneratorResult[netip.Addr]{generatedIP(ip4(192, 168, 0, 1))}, ports: []GeneratorResult[uint16]{generatedPort(888), generatedPort(889)}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(889)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(889)), }, }, { name: "ThreeIpsOnePort", - ips: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generatedIP(net.IPv4(192, 168, 0, 2)), - generatedIP(net.IPv4(192, 168, 0, 3)), + ips: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generatedIP(ip4(192, 168, 0, 2)), + generatedIP(ip4(192, 168, 0, 3)), }, ports: []GeneratorResult[uint16]{generatedPort(888)}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2)), withDstPort(888)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 3)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 2)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 3)), withDstPort(888)), }, }, { name: "TwoIpsTwoPorts", - ips: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generatedIP(net.IPv4(192, 168, 0, 2)), + ips: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generatedIP(ip4(192, 168, 0, 2)), }, ports: []GeneratorResult[uint16]{generatedPort(888), generatedPort(889)}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2)), withDstPort(888)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(889)), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2)), withDstPort(889)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 2)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(889)), + newScanRequest(withDstIP(ip4(192, 168, 0, 2)), withDstPort(889)), }, }, { name: "IPError", - ips: []GeneratorResult[net.IP]{ - generationError[net.IP](errors.New("ip error")), + ips: []GeneratorResult[netip.Addr]{ + generationError[netip.Addr](errors.New("ip error")), }, ports: []GeneratorResult[uint16]{generatedPort(888)}, expected: []*Request{ - newScanRequest(withDstIP(nil), withDstPort(888), withError(errors.New("ip error"))), + newScanRequest(withDstIP(netip.Addr{}), withDstPort(888), withError(errors.New("ip error"))), }, }, { name: "PortError", - ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ips: []GeneratorResult[netip.Addr]{generatedIP(ip4(192, 168, 0, 1))}, ports: []GeneratorResult[uint16]{ generationError[uint16](errors.New("port error")), }, @@ -425,14 +422,14 @@ func TestIPPortGenerator(t *testing.T) { }, { name: "ValidPortAfterPortError", - ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ips: []GeneratorResult[netip.Addr]{generatedIP(ip4(192, 168, 0, 1))}, ports: []GeneratorResult[uint16]{ generationError[uint16](errors.New("port error")), generatedPort(888), }, expected: []*Request{ {Err: errors.New("port error")}, - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), + newScanRequest(withDstIP(ip4(192, 168, 0, 1)), withDstPort(888)), }, }, } @@ -448,8 +445,8 @@ func TestIPPortGenerator(t *testing.T) { ctx := context.Background() scanRange := newScanRange() ipgen.EXPECT().IPs(ctx, scanRange). - DoAndReturn(func(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { - ips := make(chan GeneratorResult[net.IP], len(tt.ips)) + DoAndReturn(func(ctx context.Context, r *Range) (<-chan GeneratorResult[netip.Addr], error) { + ips := make(chan GeneratorResult[netip.Addr], len(tt.ips)) for _, ip := range tt.ips { ips <- ip } @@ -526,38 +523,63 @@ func TestIPRequestGenerator(t *testing.T) { }{ { name: "NilSubnet", - input: newScanRange(withSubnet(nil)), + input: newScanRange(withPrefix(netip.Prefix{})), err: true, }, { name: "OneIP", input: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(32, 32)}), + withPrefix(netip.PrefixFrom(ip4(192, 168, 0, 1), 32)), + ), + expected: []*Request{ + newScanRequest(withDstIP(ip4(192, 168, 0, 1))), + }, + }, + { + name: "OneIPv6", + input: newScanRange( + withPrefix(netip.MustParsePrefix("2001:db8::1/128")), ), expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), + newScanRequest(withDstIP(netip.MustParseAddr("2001:db8::1"))), }, }, + { + name: "ScopedIPv6", + input: func() *Range { + r := newScanRange(withPrefix(netip.MustParsePrefix("fe80::1/128"))) + r.DstZone = "en0" + return r + }(), + expected: []*Request{ + newScanRequest(withDstIP(netip.MustParseAddr("fe80::1%en0"))), + }, + }, + { + name: "IPv6RangeTooLarge", + input: newScanRange(withPrefix(netip.MustParsePrefix("2001:db8::/95"))), + err: true, + }, { name: "TwoIPs", input: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(31, 32)}), + withPrefix(netip.PrefixFrom(ip4(192, 168, 0, 1), 31)), ), expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 0).To4())), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), + newScanRequest(withDstIP(ip4(192, 168, 0, 0))), + newScanRequest(withDstIP(ip4(192, 168, 0, 1))), }, }, { name: "FourIPs", input: newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(30, 32)}), + withPrefix(netip.PrefixFrom(ip4(192, 168, 0, 1), 30)), ), expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 0).To4())), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2).To4())), - newScanRequest(withDstIP(net.IPv4(192, 168, 0, 3).To4())), + newScanRequest(withDstIP(ip4(192, 168, 0, 0))), + newScanRequest(withDstIP(ip4(192, 168, 0, 1))), + newScanRequest(withDstIP(ip4(192, 168, 0, 2))), + newScanRequest(withDstIP(ip4(192, 168, 0, 3))), }, }, } @@ -576,9 +598,7 @@ func TestIPRequestGenerator(t *testing.T) { require.NoError(t, err) result := collectChannel(pairs) sort.Slice(result, func(i, j int) bool { - return bytes.Compare( - result[i].DstIP, - result[j].DstIP) < 1 + return result[i].DstIP.Less(result[j].DstIP) }) require.Equal(t, tt.expected, result) }) @@ -608,14 +628,14 @@ func TestFileIPPortGenerator(t *testing.T) { name: "OneIPPort", input: `{"ip":"192.168.0.1","port":888}`, expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, }, }, { name: "OneIPPortWithUnknownField", input: `{"ip":"192.168.0.1","port":888,"abc":"field"}`, expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, }, }, { @@ -625,8 +645,8 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.2","port":222}`, }, "\n"), expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - {DstIP: net.IPv4(192, 168, 0, 2), DstPort: 222}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 2), DstPort: 222}, }, }, { @@ -643,7 +663,7 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192`, }, "\n"), expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, {Err: ErrJSON}, }, }, @@ -655,7 +675,7 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.3","port":888}`, }, "\n"), expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, {Err: ErrJSON}, }, }, @@ -680,7 +700,7 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.3"}`, }, "\n"), expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, {Err: ErrPort}, }, }, @@ -691,7 +711,7 @@ func TestFileIPPortGenerator(t *testing.T) { `{"port":888}`, }, "\n"), expected: []*Request{ - {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: ip4(192, 168, 0, 1), DstPort: 888}, {Err: ErrIP}, }, }, @@ -699,14 +719,14 @@ func TestFileIPPortGenerator(t *testing.T) { name: "OneIPPortWithSrcIPandSrcMAC", input: `{"ip":"192.168.0.1","port":888}`, scanRange: &Range{ - SrcIP: net.IPv4(192, 168, 0, 3), + SrcIP: ip4(192, 168, 0, 3), SrcMAC: net.HardwareAddr{0x01, 0x02, 0x03, 0x04, 0x05, 0x06}, }, expected: []*Request{ { - SrcIP: net.IPv4(192, 168, 0, 3), + SrcIP: ip4(192, 168, 0, 3), SrcMAC: net.HardwareAddr{0x01, 0x02, 0x03, 0x04, 0x05, 0x06}, - DstIP: net.IPv4(192, 168, 0, 1), + DstIP: ip4(192, 168, 0, 1), DstPort: 888, }, }, @@ -747,20 +767,20 @@ func TestFileIPGenerator(t *testing.T) { tests := []struct { name string input string - expected []GeneratorResult[net.IP] + expected []GeneratorResult[netip.Addr] }{ { name: "OneIP", input: `{"ip":"192.168.0.1"}`, - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), }, }, { name: "OneIPWithUnknownField", input: `{"ip":"192.168.0.1","abc":"field"}`, - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), }, }, { @@ -769,16 +789,16 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192.168.0.1"}`, `{"ip":"192.168.0.2"}`, }, "\n"), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generatedIP(net.IPv4(192, 168, 0, 2)), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generatedIP(ip4(192, 168, 0, 2)), }, }, { name: "InvalidJSON", input: `{"ip":"192`, - expected: []GeneratorResult[net.IP]{ - generationError[net.IP](ErrJSON), + expected: []GeneratorResult[netip.Addr]{ + generationError[netip.Addr](ErrJSON), }, }, { @@ -787,9 +807,9 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"ip":"192`, }, "\n"), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generationError[net.IP](ErrJSON), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generationError[netip.Addr](ErrJSON), }, }, { @@ -799,16 +819,16 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192`, `{"ip":"192.168.0.3","port":888}`, }, "\n"), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generationError[net.IP](ErrJSON), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generationError[netip.Addr](ErrJSON), }, }, { name: "InvalidIP", input: `{"ip":"192.168.0.1111"}`, - expected: []GeneratorResult[net.IP]{ - generationError[net.IP](ErrIP), + expected: []GeneratorResult[netip.Addr]{ + generationError[netip.Addr](ErrIP), }, }, { @@ -817,9 +837,9 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192.168.0.1"}`, `{}`, }, "\n"), - expected: []GeneratorResult[net.IP]{ - generatedIP(net.IPv4(192, 168, 0, 1)), - generationError[net.IP](ErrIP), + expected: []GeneratorResult[netip.Addr]{ + generatedIP(ip4(192, 168, 0, 1)), + generationError[netip.Addr](ErrIP), }, }, } @@ -930,48 +950,48 @@ func TestFilterIPRequestGenerator(t *testing.T) { { name: "EmptyFilter", input: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), }, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), }, }, { name: "OneIPFilter", input: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), }, filtered: []bool{true, false}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), }, }, { name: "OneIPFilterMiddle", input: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), + newScanRequest(withDstIP(ip4(10, 0, 3, 3))), }, filtered: []bool{false, true, false}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 3, 3))), }, }, { name: "TwoIPFilter", input: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), - newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), + newScanRequest(withDstIP(ip4(10, 0, 1, 1))), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), + newScanRequest(withDstIP(ip4(10, 0, 3, 3))), }, filtered: []bool{true, false, true}, expected: []*Request{ - newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), + newScanRequest(withDstIP(ip4(10, 0, 2, 2))), }, }, } @@ -990,7 +1010,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { } close(input) r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + withPrefix(netip.PrefixFrom(ip4(10, 0, 0, 0), 8)), ) delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). Return(input, nil) @@ -1022,7 +1042,7 @@ func TestFilterIPRequestGeneratorWithGeneratorError(t *testing.T) { delegate := NewMockRequestGenerator(ctrl) r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + withPrefix(netip.PrefixFrom(ip4(10, 0, 0, 0), 8)), ) delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). Return(nil, errors.New("generate error")) @@ -1041,10 +1061,10 @@ func TestFilterIPRequestGeneratorWithIPContainerError(t *testing.T) { delegate := NewMockRequestGenerator(ctrl) r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + withPrefix(netip.PrefixFrom(ip4(10, 0, 0, 0), 8)), ) input := make(chan *Request, 1) - input <- newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())) + input <- newScanRequest(withDstIP(ip4(10, 0, 1, 1))) close(input) delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). Return(input, nil) @@ -1059,6 +1079,6 @@ func TestFilterIPRequestGeneratorWithIPContainerError(t *testing.T) { result := collectChannel(requests) require.Equal(t, []*Request{ newScanRequest( - withDstIP(net.IPv4(10, 0, 1, 1).To4()), + withDstIP(ip4(10, 0, 1, 1)), withError(errors.New("ip container error")))}, result) } diff --git a/pkg/scan/socks5/ipv6_test.go b/pkg/scan/socks5/ipv6_test.go new file mode 100644 index 0000000..c51fe65 --- /dev/null +++ b/pkg/scan/socks5/ipv6_test.go @@ -0,0 +1,19 @@ +package socks5 + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" +) + +func TestDestinationIPv6(t *testing.T) { + t.Parallel() + require.Equal(t, "[2001:db8::1]:1080", destination(&scan.Request{ + DstIP: netip.MustParseAddr("2001:db8::1"), DstPort: 1080, + })) + require.Equal(t, "[fe80::1%en0]:1080", destination(&scan.Request{ + DstIP: netip.MustParseAddr("fe80::1%en0"), DstPort: 1080, + })) +} diff --git a/pkg/scan/socks5/socks5.go b/pkg/scan/socks5/socks5.go index 170bdb1..abbb4df 100644 --- a/pkg/scan/socks5/socks5.go +++ b/pkg/scan/socks5/socks5.go @@ -5,6 +5,8 @@ import ( "encoding/json" "fmt" "net" + "strconv" + "strings" "time" "github.com/v-byte-cpu/sx/pkg/scan" @@ -27,11 +29,15 @@ type ScanResult struct { } func (r *ScanResult) String() string { - return fmt.Sprintf("%-20s %-5d", r.IP, r.Port) + width := 20 + if strings.ContainsRune(r.IP, ':') { + width = 40 + } + return fmt.Sprintf("%-*s %-5d", width, r.IP, r.Port) } func (r *ScanResult) ID() string { - return fmt.Sprintf("%s:%d", r.IP, r.Port) + return net.JoinHostPort(r.IP, strconv.Itoa(int(r.Port))) } func (r *ScanResult) MarshalJSON() ([]byte, error) { @@ -78,7 +84,7 @@ func NewScanner(opts ...ScannerOption) *Scanner { func (s *Scanner) Scan(ctx context.Context, r *scan.Request) (result scan.Result, err error) { var conn net.Conn - if conn, err = s.dialer.DialContext(ctx, "tcp", fmt.Sprintf("%s:%d", r.DstIP, r.DstPort)); err != nil { + if conn, err = s.dialer.DialContext(ctx, "tcp", destination(r)); err != nil { return } defer conn.Close() @@ -126,6 +132,14 @@ func (s *Scanner) Scan(ctx context.Context, r *scan.Request) (result scan.Result return } +func destination(r *scan.Request) string { + return (&net.TCPAddr{ + IP: net.IP(r.DstIP.WithZone("").AsSlice()), + Port: int(r.DstPort), + Zone: r.DstIP.Zone(), + }).String() +} + type socksConn struct { conn net.Conn timeout time.Duration diff --git a/pkg/scan/tcp/bpf.go b/pkg/scan/tcp/bpf.go index 93a6653..7047037 100644 --- a/pkg/scan/tcp/bpf.go +++ b/pkg/scan/tcp/bpf.go @@ -14,9 +14,13 @@ const MaxPacketLength = 1518 func BPFFilter(r *scan.Range) (filter string, maxPacketLength int) { var sb strings.Builder sb.WriteString("tcp") - if r.DstSubnet != nil { - sb.WriteString(" and ip src net ") - sb.WriteString(r.DstSubnet.String()) + if r.DstPrefix.IsValid() { + if r.DstPrefix.Addr().Is6() { + sb.WriteString(" and ip6 src net ") + } else { + sb.WriteString(" and ip src net ") + } + sb.WriteString(r.DstPrefix.String()) } if len(r.Ports) > 0 { sb.WriteString(" and (") @@ -32,5 +36,8 @@ func BPFFilter(r *scan.Range) (filter string, maxPacketLength int) { func SYNACKBPFFilter(r *scan.Range) (filter string, maxPacketLength int) { filter, maxPacketLength = BPFFilter(r) + if r.SrcIP.Is6() || (r.DstPrefix.IsValid() && r.DstPrefix.Addr().Is6()) { + return filter, maxPacketLength + } return filter + " and tcp[13] == 18", maxPacketLength } diff --git a/pkg/scan/tcp/bpf_test.go b/pkg/scan/tcp/bpf_test.go index 8732fef..7a13148 100644 --- a/pkg/scan/tcp/bpf_test.go +++ b/pkg/scan/tcp/bpf_test.go @@ -1,7 +1,7 @@ package tcp import ( - "net" + "net/netip" "testing" "github.com/stretchr/testify/assert" @@ -22,13 +22,8 @@ func TestBPFFilter(t *testing.T) { scanRange: &scan.Range{}, }, { - name: "OneSubnet", - scanRange: &scan.Range{ - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0), - Mask: net.CIDRMask(24, 32), - }, - }, + name: "OneSubnet", + scanRange: &scan.Range{DstPrefix: netip.MustParsePrefix("192.168.0.0/24")}, expectedFilter: "tcp and ip src net 192.168.0.0/24", }, { @@ -98,13 +93,8 @@ func TestSYNACKBPFFilter(t *testing.T) { scanRange: &scan.Range{}, }, { - name: "OneSubnet", - scanRange: &scan.Range{ - DstSubnet: &net.IPNet{ - IP: net.IPv4(192, 168, 0, 0), - Mask: net.CIDRMask(24, 32), - }, - }, + name: "OneSubnet", + scanRange: &scan.Range{DstPrefix: netip.MustParsePrefix("192.168.0.0/24")}, expectedFilter: "tcp and ip src net 192.168.0.0/24 and tcp[13] == 18", }, { diff --git a/pkg/scan/tcp/ipv6_test.go b/pkg/scan/tcp/ipv6_test.go new file mode 100644 index 0000000..394ba80 --- /dev/null +++ b/pkg/scan/tcp/ipv6_test.go @@ -0,0 +1,70 @@ +package tcp + +import ( + "context" + "net" + "net/netip" + "testing" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" +) + +func TestPacketFillerIPv6(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + err := NewPacketFiller(WithSYN(), WithHopLimit(33)).Fill(packet, &scan.Request{ + SrcIP: netip.MustParseAddr("2001:db8::1"), DstIP: netip.MustParseAddr("2001:db8::2"), DstPort: 443, + SrcMAC: net.HardwareAddr{2, 0, 0, 0, 0, 1}, DstMAC: net.HardwareAddr{2, 0, 0, 0, 0, 2}, + }) + require.NoError(t, err) + + decoded := gopacket.NewPacket(packet.Bytes(), layers.LayerTypeEthernet, gopacket.Default) + ipv6 := decoded.Layer(layers.LayerTypeIPv6).(*layers.IPv6) + require.Equal(t, uint8(33), ipv6.HopLimit) + require.Equal(t, layers.IPProtocolTCP, ipv6.NextHeader) + tcp := decoded.Layer(layers.LayerTypeTCP).(*layers.TCP) + require.True(t, tcp.SYN) + require.Equal(t, layers.TCPPort(443), tcp.DstPort) +} + +func TestBPFFilterIPv6(t *testing.T) { + t.Parallel() + filter, _ := BPFFilter(&scan.Range{DstPrefix: netip.MustParsePrefix("2001:db8::/120")}) + require.Equal(t, "tcp and ip6 src net 2001:db8::/120", filter) +} + +func TestSYNACKBPFFilterIPv6(t *testing.T) { + t.Parallel() + filter, _ := SYNACKBPFFilter(&scan.Range{DstPrefix: netip.MustParsePrefix("2001:db8::/120")}) + require.Equal(t, "tcp and ip6 src net 2001:db8::/120", filter) +} + +func TestSYNACKBPFFilterIPv6Source(t *testing.T) { + t.Parallel() + filter, _ := SYNACKBPFFilter(&scan.Range{SrcIP: netip.MustParseAddr("2001:db8::1")}) + require.Equal(t, "tcp", filter) +} + +func TestScanMethodProcessesScopedIPv6Response(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + eth := &layers.Ethernet{SrcMAC: net.HardwareAddr{2, 0, 0, 0, 0, 2}, DstMAC: net.HardwareAddr{2, 0, 0, 0, 0, 1}, EthernetType: layers.EthernetTypeIPv6} + ipv6 := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolTCP, SrcIP: net.ParseIP("fe80::2"), DstIP: net.ParseIP("fe80::1")} + tcpLayer := &layers.TCP{SrcPort: 443, DstPort: 40000, SYN: true, ACK: true} + require.NoError(t, tcpLayer.SetNetworkLayerForChecksum(ipv6)) + require.NoError(t, gopacket.SerializeLayers(packet, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, eth, ipv6, tcpLayer)) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + results := scan.NewResultChan(ctx, 1) + method := NewScanMethod(SYNScanType, nil, results, WithScanIPv6(true), WithScanZone("en0")) + require.NoError(t, method.ProcessPacketData(packet.Bytes(), &gopacket.CaptureInfo{})) + result := (<-method.Results()).(*ScanResult) + require.Equal(t, "fe80::2%en0", result.IP) + require.Equal(t, uint16(443), result.Port) +} diff --git a/pkg/scan/tcp/tcp.go b/pkg/scan/tcp/tcp.go index 89e69c9..2d2a249 100644 --- a/pkg/scan/tcp/tcp.go +++ b/pkg/scan/tcp/tcp.go @@ -5,6 +5,9 @@ package tcp import ( "fmt" "math/rand" + "net" + "net/netip" + "strconv" "strings" "github.com/google/gopacket" @@ -29,11 +32,15 @@ type ScanResult struct { } func (r *ScanResult) String() string { - return fmt.Sprintf("%-20s %-5d %s", r.IP, r.Port, r.Flags) + width := 20 + if strings.ContainsRune(r.IP, ':') { + width = 40 + } + return fmt.Sprintf("%-*s %-5d %s", width, r.IP, r.Port, r.Flags) } func (r *ScanResult) ID() string { - return fmt.Sprintf("%s:%d", r.IP, r.Port) + return net.JoinHostPort(r.IP, strconv.Itoa(int(r.Port))) } type PacketFilterFunc func(pkt *layers.TCP) bool @@ -93,7 +100,10 @@ type ScanMethod struct { rcvDecoded []gopacket.LayerType rcvEth layers.Ethernet rcvIP layers.IPv4 + rcvIPv6 layers.IPv6 rcvTCP layers.TCP + ipv6 bool + zone string } // Assert that tcp.ScanMethod conforms to the scan.PacketMethod interface @@ -119,6 +129,14 @@ func WithScanVPNmode(vpnMode bool) ScanMethodOption { } } +func WithScanIPv6(ipv6 bool) ScanMethodOption { + return func(s *ScanMethod) { s.ipv6 = ipv6 } +} + +func WithScanZone(zone string) ScanMethodOption { + return func(s *ScanMethod) { s.zone = zone } +} + func NewScanMethod(scanType string, psrc scan.PacketSource, results scan.ResultChan, opts ...ScanMethodOption) *ScanMethod { sm := &ScanMethod{ @@ -135,9 +153,18 @@ func NewScanMethod(scanType string, psrc scan.PacketSource, layerType := layers.LayerTypeEthernet if sm.vpnMode { - layerType = layers.LayerTypeIPv4 + if sm.ipv6 { + layerType = layers.LayerTypeIPv6 + } else { + layerType = layers.LayerTypeIPv4 + } + } + var parser *gopacket.DecodingLayerParser + if sm.ipv6 { + parser = gopacket.NewDecodingLayerParser(layerType, &sm.rcvEth, &sm.rcvIPv6, &sm.rcvTCP) + } else { + parser = gopacket.NewDecodingLayerParser(layerType, &sm.rcvEth, &sm.rcvIP, &sm.rcvTCP) } - parser := gopacket.NewDecodingLayerParser(layerType, &sm.rcvEth, &sm.rcvIP, &sm.rcvTCP) parser.IgnoreUnsupported = true sm.parser = parser return sm @@ -151,14 +178,22 @@ func (s *ScanMethod) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo) (er if err = s.parser.DecodeLayers(data, &s.rcvDecoded); err != nil { return } - if !validPacket(s.rcvDecoded) { + if !s.validPacket() { return } if s.pktFilter(&s.rcvTCP) { + srcIP := s.rcvIP.SrcIP + if s.ipv6 { + srcIP = s.rcvIPv6.SrcIP + } + address, _ := netip.AddrFromSlice(srcIP) + if address.IsLinkLocalUnicast() && s.zone != "" { + address = address.WithZone(s.zone) + } s.results.Put(&ScanResult{ ScanType: s.scanType, - IP: s.rcvIP.SrcIP.String(), + IP: address.String(), Port: uint16(s.rcvTCP.SrcPort), Flags: s.pktFlags(&s.rcvTCP), }) @@ -166,6 +201,13 @@ func (s *ScanMethod) ProcessPacketData(data []byte, _ *gopacket.CaptureInfo) (er return } +func (s *ScanMethod) validPacket() bool { + if s.ipv6 { + return len(s.rcvDecoded) == 3 || (len(s.rcvDecoded) == 2 && s.rcvDecoded[0] == layers.LayerTypeIPv6) + } + return validPacket(s.rcvDecoded) +} + func validPacket(decoded []gopacket.LayerType) bool { return len(decoded) == 3 || (len(decoded) == 2 && decoded[0] == layers.LayerTypeIPv4) } @@ -181,7 +223,10 @@ type PacketFiller struct { CWR bool NS bool - vpnMode bool + vpnMode bool + hopLimit uint8 + nextHeader layers.IPProtocol + payloadLength uint16 } // Assert that tcp.PacketFiller conforms to the scan.PacketFiller interface @@ -249,8 +294,20 @@ func WithFillerVPNmode(vpnMode bool) PacketFillerOption { } } +func WithHopLimit(hopLimit uint8) PacketFillerOption { + return func(f *PacketFiller) { f.hopLimit = hopLimit } +} + +func WithNextHeader(nextHeader uint8) PacketFillerOption { + return func(f *PacketFiller) { f.nextHeader = layers.IPProtocol(nextHeader) } +} + +func WithPayloadLength(payloadLength uint16) PacketFillerOption { + return func(f *PacketFiller) { f.payloadLength = payloadLength } +} + func NewPacketFiller(opts ...PacketFillerOption) *PacketFiller { - f := &PacketFiller{} + f := &PacketFiller{hopLimit: 64, nextHeader: layers.IPProtocolTCP} for _, o := range opts { o(f) } @@ -258,6 +315,9 @@ func NewPacketFiller(opts ...PacketFillerOption) *PacketFiller { } func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (err error) { + if r.DstIP.Is6() { + return f.fillIPv6(packet, r) + } ip := &layers.IPv4{ Version: 4, @@ -268,8 +328,8 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e Flags: layers.IPv4DontFragment, TTL: 64, Protocol: layers.IPProtocolTCP, - SrcIP: r.SrcIP, - DstIP: r.DstIP, + SrcIP: r.SrcIP.AsSlice(), + DstIP: r.DstIP.AsSlice(), } tcp := &layers.TCP{ // emulate Linux default ephemeral ports range: 32768 60999 @@ -319,3 +379,40 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e } return gopacket.SerializeLayers(packet, opt, eth, ip, tcp) } + +func (f *PacketFiller) fillIPv6(packet gopacket.SerializeBuffer, r *scan.Request) error { + ipv6 := &layers.IPv6{ + Version: 6, + Length: f.payloadLength, + NextHeader: f.nextHeader, + HopLimit: f.hopLimit, + SrcIP: net.IP(r.SrcIP.WithZone("").AsSlice()), + DstIP: net.IP(r.DstIP.WithZone("").AsSlice()), + } + tcp := f.newTCPLayer(r) + if err := tcp.SetNetworkLayerForChecksum(ipv6); err != nil { + return err + } + options := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: f.payloadLength == 0} + layersToSerialize := []gopacket.SerializableLayer{ipv6, tcp} + if !f.vpnMode { + layersToSerialize = append([]gopacket.SerializableLayer{&layers.Ethernet{ + SrcMAC: r.SrcMAC, DstMAC: r.DstMAC, EthernetType: layers.EthernetTypeIPv6, + }}, layersToSerialize...) + } + return gopacket.SerializeLayers(packet, options, layersToSerialize...) +} + +func (f *PacketFiller) newTCPLayer(r *scan.Request) *layers.TCP { + return &layers.TCP{ + SrcPort: layers.TCPPort(32768 + rand.Intn(61000-32768)), + DstPort: layers.TCPPort(r.DstPort), + Seq: rand.Uint32(), SYN: f.SYN, ACK: f.ACK, FIN: f.FIN, RST: f.RST, + PSH: f.PSH, URG: f.URG, ECE: f.ECE, CWR: f.CWR, NS: f.NS, Window: 64240, + Options: []layers.TCPOption{ + {OptionType: layers.TCPOptionKindMSS, OptionLength: 4, OptionData: []byte{0x05, 0xb4}}, + {OptionType: layers.TCPOptionKindSACKPermitted, OptionLength: 2}, + {OptionType: layers.TCPOptionKindWindowScale, OptionLength: 3, OptionData: []byte{7}}, + }, + } +} diff --git a/pkg/scan/tcp/tcp_test.go b/pkg/scan/tcp/tcp_test.go index 0a9a097..497a5ac 100644 --- a/pkg/scan/tcp/tcp_test.go +++ b/pkg/scan/tcp/tcp_test.go @@ -3,6 +3,7 @@ package tcp import ( "context" "net" + "net/netip" "runtime" "testing" "time" @@ -12,7 +13,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/v-byte-cpu/sx/pkg/scan" - "github.com/v-byte-cpu/sx/pkg/scan/arp" + "github.com/v-byte-cpu/sx/pkg/scan/neighbor" ) func TestPacketFillerEthernet(t *testing.T) { @@ -85,8 +86,8 @@ func TestPacketFillerEthernet(t *testing.T) { packet := gopacket.NewSerializeBuffer() err := tt.filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -197,8 +198,8 @@ func TestPacketFillerIPv4(t *testing.T) { packet := gopacket.NewSerializeBuffer() err := tt.filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -482,12 +483,12 @@ func TestAllFlags(t *testing.T) { type mockIPGeneratorFunc func( ctx context.Context, r *scan.Range, -) (<-chan scan.GeneratorResult[net.IP], error) +) (<-chan scan.GeneratorResult[netip.Addr], error) func (f mockIPGeneratorFunc) IPs( ctx context.Context, r *scan.Range, -) (<-chan scan.GeneratorResult[net.IP], error) { +) (<-chan scan.GeneratorResult[netip.Addr], error) { return f(ctx, r) } @@ -506,28 +507,28 @@ func BenchmarkTCPScanEngine(b *testing.B) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() - dstIP := net.IPv4(192, 168, 0, 3).To4() + dstIP := netip.MustParseAddr("192.168.0.3") ipgen := mockIPGeneratorFunc(func( ctx context.Context, _ *scan.Range, - ) (<-chan scan.GeneratorResult[net.IP], error) { - out := make(chan scan.GeneratorResult[net.IP], 100) + ) (<-chan scan.GeneratorResult[netip.Addr], error) { + out := make(chan scan.GeneratorResult[netip.Addr], 100) go func() { defer close(out) for i := 0; i < b.N; i++ { select { case <-ctx.Done(): return - case out <- scan.GeneratorResult[net.IP]{Value: dstIP}: + case out <- scan.GeneratorResult[netip.Addr]{Value: dstIP}: } } }() return out, nil }) - reqgen := arp.NewCacheRequestGenerator( + reqgen := neighbor.NewCacheRequestGenerator( scan.NewIPPortGenerator(ipgen, scan.NewPortGenerator()), net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, - arp.NewCache()) + neighbor.NewCache()) pktgen := scan.NewPacketMultiGenerator(NewPacketFiller(), runtime.NumCPU()) psrc := scan.NewPacketSource(reqgen, pktgen) results := scan.NewResultChan(ctx, 1000) @@ -535,7 +536,7 @@ func BenchmarkTCPScanEngine(b *testing.B) { engine := scan.SetupPacketEngine(&nullPacketReadWriter{}, sm) done, _ := engine.Start(ctx, &scan.Range{ - SrcIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, Ports: []*scan.PortRange{ { diff --git a/pkg/scan/udp/ipv6_test.go b/pkg/scan/udp/ipv6_test.go new file mode 100644 index 0000000..de60806 --- /dev/null +++ b/pkg/scan/udp/ipv6_test.go @@ -0,0 +1,31 @@ +package udp + +import ( + "net" + "net/netip" + "testing" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/stretchr/testify/require" + "github.com/v-byte-cpu/sx/pkg/scan" +) + +func TestPacketFillerIPv6(t *testing.T) { + t.Parallel() + + packet := gopacket.NewSerializeBuffer() + err := NewPacketFiller(WithHopLimit(42), WithPayload([]byte("dns"))).Fill(packet, &scan.Request{ + SrcIP: netip.MustParseAddr("2001:db8::1"), DstIP: netip.MustParseAddr("2001:db8::2"), DstPort: 53, + SrcMAC: net.HardwareAddr{2, 0, 0, 0, 0, 1}, DstMAC: net.HardwareAddr{2, 0, 0, 0, 0, 2}, + }) + require.NoError(t, err) + + decoded := gopacket.NewPacket(packet.Bytes(), layers.LayerTypeEthernet, gopacket.Default) + ipv6 := decoded.Layer(layers.LayerTypeIPv6).(*layers.IPv6) + require.Equal(t, uint8(42), ipv6.HopLimit) + require.Equal(t, layers.IPProtocolUDP, ipv6.NextHeader) + udp := decoded.Layer(layers.LayerTypeUDP).(*layers.UDP) + require.Equal(t, layers.UDPPort(53), udp.DstPort) + require.Equal(t, []byte("dns"), udp.Payload) +} diff --git a/pkg/scan/udp/udp.go b/pkg/scan/udp/udp.go index 6a0a1a4..099a560 100644 --- a/pkg/scan/udp/udp.go +++ b/pkg/scan/udp/udp.go @@ -2,6 +2,7 @@ package udp import ( "math/rand" + "net" "github.com/google/gopacket" "github.com/google/gopacket/layers" @@ -25,8 +26,16 @@ type ScanMethod struct { // Assert that udp.ScanMethod conforms to the scan.PacketMethod interface var _ scan.PacketMethod = (*ScanMethod)(nil) -func NewScanMethod(psrc scan.PacketSource, results scan.ResultChan, vpnMode bool) *ScanMethod { - pp := icmp.NewPacketProcessor(ScanType, results, vpnMode) +func NewScanMethod(psrc scan.PacketSource, results scan.ResultChan, vpnMode bool, ipv6 ...bool) *ScanMethod { + pp := icmp.NewPacketProcessor(ScanType, results, vpnMode, ipv6...) + return newScanMethod(psrc, pp) +} + +func NewScanMethodForFamily(psrc scan.PacketSource, results scan.ResultChan, vpnMode, ipv6 bool, zone string) *ScanMethod { + return newScanMethod(psrc, icmp.NewPacketProcessorForFamily(ScanType, results, vpnMode, ipv6, zone)) +} + +func newScanMethod(psrc scan.PacketSource, pp *icmp.PacketProcessor) *ScanMethod { return &ScanMethod{ PacketSource: psrc, Processor: pp, @@ -41,6 +50,10 @@ type PacketFiller struct { flags layers.IPv4Flag payload []byte vpnMode bool + + hopLimit uint8 + nextHeader layers.IPProtocol + payloadLength uint16 } // Assert that udp.PacketFiller conforms to the scan.PacketFiller interface @@ -86,12 +99,26 @@ func WithVPNmode(vpnMode bool) PacketFillerOption { } } +func WithHopLimit(hopLimit uint8) PacketFillerOption { + return func(f *PacketFiller) { f.hopLimit = hopLimit } +} + +func WithNextHeader(nextHeader uint8) PacketFillerOption { + return func(f *PacketFiller) { f.nextHeader = layers.IPProtocol(nextHeader) } +} + +func WithPayloadLength(payloadLength uint16) PacketFillerOption { + return func(f *PacketFiller) { f.payloadLength = payloadLength } +} + func NewPacketFiller(opts ...PacketFillerOption) *PacketFiller { f := &PacketFiller{ // typical TTL value for Linux - ttl: 64, - proto: layers.IPProtocolUDP, - flags: layers.IPv4DontFragment, + ttl: 64, + proto: layers.IPProtocolUDP, + flags: layers.IPv4DontFragment, + hopLimit: 64, + nextHeader: layers.IPProtocolUDP, } for _, o := range opts { o(f) @@ -100,6 +127,9 @@ func NewPacketFiller(opts ...PacketFillerOption) *PacketFiller { } func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (err error) { + if r.DstIP.Is6() { + return f.fillIPv6(packet, r) + } ip := &layers.IPv4{ Version: 4, @@ -113,8 +143,8 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e TTL: f.ttl, Length: f.length, Protocol: f.proto, - SrcIP: r.SrcIP, - DstIP: r.DstIP, + SrcIP: r.SrcIP.AsSlice(), + DstIP: r.DstIP.AsSlice(), } udp := &layers.UDP{ @@ -142,3 +172,29 @@ func (f *PacketFiller) Fill(packet gopacket.SerializeBuffer, r *scan.Request) (e } return gopacket.SerializeLayers(packet, opt, eth, ip, udp, gopacket.Payload(f.payload)) } + +func (f *PacketFiller) fillIPv6(packet gopacket.SerializeBuffer, r *scan.Request) error { + ipv6 := &layers.IPv6{ + Version: 6, + Length: f.payloadLength, + NextHeader: f.nextHeader, + HopLimit: f.hopLimit, + SrcIP: net.IP(r.SrcIP.WithZone("").AsSlice()), + DstIP: net.IP(r.DstIP.WithZone("").AsSlice()), + } + udp := &layers.UDP{ + SrcPort: layers.UDPPort(32768 + rand.Intn(61000-32768)), + DstPort: layers.UDPPort(r.DstPort), + } + if err := udp.SetNetworkLayerForChecksum(ipv6); err != nil { + return err + } + options := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: f.payloadLength == 0} + layersToSerialize := []gopacket.SerializableLayer{ipv6, udp, gopacket.Payload(f.payload)} + if !f.vpnMode { + layersToSerialize = append([]gopacket.SerializableLayer{&layers.Ethernet{ + SrcMAC: r.SrcMAC, DstMAC: r.DstMAC, EthernetType: layers.EthernetTypeIPv6, + }}, layersToSerialize...) + } + return gopacket.SerializeLayers(packet, options, layersToSerialize...) +} diff --git a/pkg/scan/udp/udp_test.go b/pkg/scan/udp/udp_test.go index 0c371e3..35c6577 100644 --- a/pkg/scan/udp/udp_test.go +++ b/pkg/scan/udp/udp_test.go @@ -3,6 +3,7 @@ package udp import ( "context" "net" + "net/netip" "testing" "time" @@ -20,8 +21,8 @@ func TestPacketFillerEthernet(t *testing.T) { filler := NewPacketFiller() packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -63,8 +64,8 @@ func TestPacketFillerIPv4(t *testing.T) { filler := NewPacketFiller(WithVPNmode(true)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -103,8 +104,8 @@ func TestPacketFillerPayload(t *testing.T) { filler := NewPacketFiller(WithPayload([]byte("abc"))) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -139,8 +140,8 @@ func TestPacketFillerTTL(t *testing.T) { filler := NewPacketFiller(WithTTL(37)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -170,8 +171,8 @@ func TestPacketFillerIPTotalLength(t *testing.T) { filler := NewPacketFiller(WithIPTotalLength(57)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -201,8 +202,8 @@ func TestPacketFillerIPProtocol(t *testing.T) { filler := NewPacketFiller(WithIPProtocol(37)) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567, @@ -223,8 +224,8 @@ func TestPacketFillerIPFlags(t *testing.T) { filler := NewPacketFiller(WithIPFlags(uint8(layers.IPv4DontFragment | layers.IPv4MoreFragments))) packet := gopacket.NewSerializeBuffer() err := filler.Fill(packet, &scan.Request{ - SrcIP: net.IPv4(192, 168, 0, 3).To4(), - DstIP: net.IPv4(192, 168, 0, 2).To4(), + SrcIP: netip.MustParseAddr("192.168.0.3"), + DstIP: netip.MustParseAddr("192.168.0.2"), SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, DstMAC: net.HardwareAddr{0x10, 0x11, 0x12, 0x13, 0x14, 0x15}, DstPort: 4567,