diff --git a/arp/handler.go b/arp/handler.go index 897e5d1..cce5e52 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -51,7 +51,7 @@ func (h *Handler) Reset(cfg HandlerConfig) error { return errors.New("invalid Handler query or pending config") } *h = Handler{ - connID: h.connID, + connID: h.connID + 1, ourHWAddr: h.ourHWAddr[:0], ourProtoAddr: h.ourProtoAddr[:0], htype: cfg.HardwareType, diff --git a/internet/definitions.go b/internet/definitions.go index 8246f97..010a896 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -118,7 +118,7 @@ func (h *handlers) tryHandleError(node *node, err error) (discardedGracefully bo func (h *handlers) nodeByProto(proto uint16) *node { for i := range h.nodes { node := &h.nodes[i] - if node.proto == proto { + if node.proto == proto && !node.IsInvalid() { return node } } @@ -128,7 +128,7 @@ func (h *handlers) nodeByProto(proto uint16) *node { func (h *handlers) nodeByPort(port uint16) *node { for i := range h.nodes { node := &h.nodes[i] - if node.port == port { + if node.port == port && !node.IsInvalid() { return node } } @@ -138,7 +138,7 @@ func (h *handlers) nodeByPort(port uint16) *node { func (h *handlers) nodeByPortProto(port uint16, protocol uint16) *node { for i := range h.nodes { node := &h.nodes[i] - if node.port == port && node.proto == protocol { + if node.port == port && node.proto == protocol && !node.IsInvalid() { return node } } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 5533736..84430c9 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -33,6 +33,7 @@ type StackAsync struct { dhcpUDP internet.StackUDPPort dhcp dhcpv4.Client dhcpResults DHCPResults + subnet netip.Prefix // Local subnet for ARP resolution. dnsUDP internet.StackUDPPort dns dns.Client @@ -112,6 +113,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error { if err != nil { return err } + // err = s.resetARP() if err != nil { return err @@ -135,10 +137,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error { } // Now setup stacks. - err = s.link.Register(&s.arp) // ARP. - if err != nil { - return err - } + // ARP registered in resetARP. err = s.link.Register(&s.ip) // IPv4 | IPv6 if err != nil { return err @@ -201,6 +200,10 @@ func (s *StackAsync) Prand32() uint32 { func (s *StackAsync) SetIPAddr(addr netip.Addr) error { s.mu.Lock() defer s.mu.Unlock() + return s.setIPAddr(addr) +} + +func (s *StackAsync) setIPAddr(addr netip.Addr) error { err := s.ip.SetAddr(addr) if err != nil { return err @@ -245,7 +248,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP s.mu.Lock() defer s.mu.Unlock() var mac []byte - if s.dhcpResults.Subnet.Contains(addrp.Addr()) { + if s.subnet.Contains(addrp.Addr()) { mac = make([]byte, 6) ip := addrp.Addr().As4() // StartQuery starts an ARP query for addresses in this network. @@ -454,11 +457,15 @@ func (s *StackAsync) ReadStatistics(stats *Statistics) { // AssimilateDHCPResults sets the stack's following parameters: // - IPv4 address. // - DNS server. +// - Subnet (for ARP resolution of local addresses). func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error { stack.mu.Lock() defer stack.mu.Unlock() + if results.Subnet.IsValid() { + stack.subnet = results.Subnet + } if results.AssignedAddr.IsValid() { - err := stack.SetIPAddr(results.AssignedAddr) + err := stack.setIPAddr(results.AssignedAddr) if err != nil { return err } diff --git a/x/xnet/xnet_arp_test.go b/x/xnet/xnet_arp_test.go new file mode 100644 index 0000000..859135e --- /dev/null +++ b/x/xnet/xnet_arp_test.go @@ -0,0 +1,49 @@ +package xnet + +import ( + "bytes" + "net/netip" + "testing" +) + +func TestARPLocal(t *testing.T) { + const mtu = 1500 + const seed = 1 + s1, s2, c1, c2 := newTCPStacks(t, seed, mtu) + routerHw := [6]byte{1, 2, 3, 4, 5, 6} + // Most common case: we have a router in between computers. + s1.SetGateway6(routerHw) + s2.SetGateway6(routerHw) + addr1 := netip.AddrPortFrom(s1.Addr(), 1024) // dialer, client. + addr2 := netip.AddrPortFrom(s2.Addr(), 80) // listener, server. + err := s1.AssimilateDHCPResults(&DHCPResults{ + Router: netip.AddrFrom4([4]byte{10, 0, 0, 255}), + BroadcastAddr: netip.AddrFrom4([4]byte{255, 255, 255, 255}), + AssignedAddr: s1.Addr(), + Subnet: netip.PrefixFrom(s2.Addr(), 24), // Subnet containing s2 will force an ARP on s1. + TRenewal: 1000, + TRebind: 1000, + TLease: 1000, + }) + if err != nil { + t.Fatal(err) + } + hw2 := s2.HardwareAddress() + err = s1.DialTCP(c1, addr1.Port(), addr2) // addr2 MAC address is unknown and must be resolved by stack. + if err != nil { + t.Fatal(err) + } + err = s2.ListenTCP(c2, addr2.Port()) + if err != nil { + t.Fatal(err) + } + tst := testerFrom(t, mtu) + _ = tst + tst.ARPExchangeOnly(s1, s2) + hwaddr, err := s1.arp.QueryResult(addr2.Addr().AsSlice()) + if err != nil { + t.Fatal(err) + } else if !bytes.Equal(hwaddr[:], hw2[:]) { + t.Errorf("expected hardware address %x, got %x", hw2, hwaddr) + } +} diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index d066224..3e19ddc 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -2,11 +2,13 @@ package xnet import ( "bytes" + "encoding/binary" "errors" "math/rand" "net/netip" "testing" + "github.com/soypat/lneto/arp" "github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/internet/pcap" "github.com/soypat/lneto/tcp" @@ -24,9 +26,7 @@ func TestStackAsyncTCP_multipacket(t *testing.T) { const svPort = 8080 const maxPktLen = 30 client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) - tst := tester{ - t: t, buf: make([]byte, MTU), - } + tst := testerFrom(t, MTU) rng := rand.New(rand.NewSource(seed)) client2, sv2, clconn2, svconn2 := newTCPStacks(t, seed, MTU) _, _, _, _ = client2, sv2, clconn2, svconn2 @@ -58,10 +58,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) { const MTU = 1500 const svPort = 80 client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) - - tst := tester{ - t: t, buf: make([]byte, MTU), - } + tst := testerFrom(t, MTU) tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337) sendData := []byte("hello") @@ -80,7 +77,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) { func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1, c2 *tcp.Conn) { s1, s2 = new(StackAsync), new(StackAsync) c1, c2 = new(tcp.Conn), new(tcp.Conn) - byte1 := byte(randSeed) / 4 + byte1 := byte(randSeed)/4 - 1 err := s1.Reset(StackConfig{ Hostname: "Stack1", RandSeed: randSeed, @@ -127,6 +124,13 @@ func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1 return s1, s2, c1, c2 } +func testerFrom(t *testing.T, mtu int) *tester { + return &tester{ + t: t, + buf: make([]byte, mtu), + } +} + type tester struct { t *testing.T cap pcap.PacketBreakdown @@ -334,6 +338,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs t.Error("expected no data sent and got data") return } + defer setzero(buf[:n]) tst.buf = tst.buf[:n] tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) @@ -374,7 +379,121 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs if err != nil { t.Fatal(err) } +} + +func (tst *tester) ARPExchangeOnly(querying, target *StackAsync) { + t := tst.t + t.Helper() + buf := tst.buf[:cap(tst.buf)] + + // === PHASE 1: ARP Request from querying stack === + n, err := querying.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal(err) + } else if n == 0 { + t.Error("zero bits sent by ARP querying stack") + return + } + + tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) + if err != nil { + t.Fatal(err) + } + tst.buf = tst.buf[:n] + + qHw := querying.HardwareAddress() + tgtHw := target.HardwareAddress() + broadcast := ethernet.BroadcastAddr() + qIP := querying.Addr() + tgtIP := target.Addr() + + // Validate Ethernet layer (request is broadcast) + if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) { + t.Errorf("request: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) + } + if !bytes.Equal(broadcast[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) { + t.Errorf("request: expected broadcast ethernet dst addr, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) + } + + // Validate ARP request fields + // ARP fields: FieldClassSrc with 6 octets = HW addr, 4 octets = proto addr + // occurrence 0 = sender, occurrence 1 = target + if tst.getARPOperation() != arp.OpRequest { + t.Errorf("request: expected ARP OpRequest, got %d", tst.getARPOperation()) + } + if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) { + t.Errorf("request: mismatched ARP sender HW") + } + if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) { + t.Errorf("request: mismatched ARP sender proto") + } + if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) { + t.Errorf("request: mismatched ARP target proto") + } + + // Deliver request to target + err = target.Demux(buf[:n], 0) + if err != nil { + t.Fatal("target demux request:", err) + } setzero(buf[:n]) + + // === PHASE 2: ARP Reply from target stack === + buf = tst.buf[:cap(tst.buf)] + n, err = target.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal(err) + } else if n == 0 { + t.Error("zero bits sent by ARP target stack (no reply)") + return + } + + tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) + if err != nil { + t.Fatal(err) + } + tst.buf = tst.buf[:n] + + // Validate Ethernet layer (reply is unicast to querying) + if !bytes.Equal(tgtHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) { + t.Errorf("reply: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) + } + if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) { + t.Errorf("reply: expected unicast to querying, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) + } + + // Validate ARP reply fields + if tst.getARPOperation() != arp.OpReply { + t.Errorf("reply: expected ARP OpReply, got %d", tst.getARPOperation()) + } + if !bytes.Equal(tgtHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) { + t.Errorf("reply: mismatched ARP sender HW (should be target's MAC)") + } + if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) { + t.Errorf("reply: mismatched ARP sender proto (should be target's IP)") + } + if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 1)) { + t.Errorf("reply: mismatched ARP target HW (should be querying's MAC)") + } + if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) { + t.Errorf("reply: mismatched ARP target proto (should be querying's IP)") + } + + // Deliver reply to querying stack + err = querying.Demux(buf[:n], 0) + if err != nil { + t.Fatal("querying demux reply:", err) + } + setzero(buf[:n]) + + // === PHASE 3: Verify querying stack learned target's MAC === + resolvedHw, err := querying.ResultResolveHardwareAddress6(tgtIP) + if err != nil { + t.Fatalf("ARP query result failed: %v", err) + } + if resolvedHw != tgtHw { + t.Errorf("ARP resolved wrong MAC: got %x, want %x", resolvedHw, tgtHw) + } } func (tst *tester) getTCPFrame() tcp.Frame { @@ -457,3 +576,34 @@ func setzero[T ~[]E, E any](s T) { s[i] = zero } } + +// getFieldByClassLen finds a field by protocol, class, and octet length. +// occurrence specifies which match to return (0 = first, 1 = second, etc.) +// This is needed for ARP where sender and target fields share the same class. +func (tst *tester) getFieldByClassLen(proto any, class pcap.FieldClass, octetLen, occurrence int) []byte { + tst.t.Helper() + frm := getProtoFrame(tst.frmbuf, proto) + if frm == nil { + tst.t.Fatalf("no frame for proto %v found", proto) + } + count := 0 + for _, field := range frm.Fields { + if field.Class == class && field.BitLength == octetLen*8 { + if count == occurrence { + bitoff := frm.PacketBitOffset + field.FrameBitOffset + return tst.buf[bitoff/8 : bitoff/8+field.BitLength/8] + } + count++ + } + } + tst.t.Fatalf("field (proto=%v, class=%v, octets=%d, occurrence=%d) not found", proto, class, octetLen, occurrence) + return nil +} + +func (tst *tester) getARPOperation() arp.Operation { + tst.t.Helper() + // ARP has 3 FieldClassType fields: Hardware type (0), Protocol type (1), Opcode (2) + // All are 2 bytes, so we need occurrence=2 to get Opcode. + data := tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassType, 2, 2) + return arp.Operation(binary.BigEndian.Uint16(data)) +}