diff --git a/README.md b/README.md index 68e38fb..77cc069 100644 --- a/README.md +++ b/README.md @@ -57,11 +57,11 @@ The following interface is implemented by networking stack nodes and the stack t ```go type StackNode interface { // Encapsulate receives a buffer the receiver must fill with data. - // The receiver's start byte is at carrierData[frameOffset]. - Encapsulate(carrierData []byte, frameOffset int) (int, error) + // The receiver's start byte is at carrierData[offsetToFrame]. + Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) // Demux receives a buffer the receiver must decode and pass on to corresponding child StackNode(s). - // The receiver's start byte is at carrierData[frameOffset]. - Demux(carrierData []byte, frameOffset int) error + // The receiver's start byte is at carrierData[offsetToFrame]. + Demux(carrierData []byte, offsetToFrame int) error // LocalPort returns the port of the node if applicable or zero. Used for UDP/TCP nodes. LocalPort() uint16 // Protocol returns the protocol of this node if applicable or zero. Usually either a ethernet.Type (EtherType) or lneto.IPProto (IP Protocol number). diff --git a/arp/arp_test.go b/arp/arp_test.go index 6b10a33..59bb875 100644 --- a/arp/arp_test.go +++ b/arp/arp_test.go @@ -34,13 +34,13 @@ func TestHandler(t *testing.T) { t.Fatal(err) } var buf, discard [64]byte - n, err := c1.Encapsulate(buf[:], 0) + n, err := c1.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal("error on should be nop send:", err) } else if n > 0 { t.Fatal("should not send if no query") } - n, err = c2.Encapsulate(buf[:], 0) + n, err = c2.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal("error on should be nop send:", err) } else if n > 0 { @@ -50,11 +50,11 @@ func TestHandler(t *testing.T) { // Perform ARP exchange. expectHWAddr := c2.ourHWAddr queryAddr := c2.ourProtoAddr - err = c1.StartQuery(queryAddr) + err = c1.StartQuery(nil, queryAddr) if err != nil { t.Fatal(err) } - n, err = c1.Encapsulate(buf[:], 0) // Send Request. + n, err = c1.Encapsulate(buf[:], -1, 0) // Send Request. if err != nil { t.Fatal(err) } else if n == 0 { @@ -66,14 +66,14 @@ func TestHandler(t *testing.T) { t.Fatal(err) } - n, err = c2.Encapsulate(buf[:], 0) // Send response. + n, err = c2.Encapsulate(buf[:], -1, 0) // Send response. if err != nil { t.Fatal(err) } else if n == 0 { t.Fatal("got no response to request") } validateARP(t, buf[:]) - n, err = c2.Encapsulate(discard[:], 0) // Double tap check, should send nothing. + n, err = c2.Encapsulate(discard[:], -1, 0) // Double tap check, should send nothing. if err != nil { t.Fatal("double tap send error:", err) } else if n > 0 { @@ -90,13 +90,13 @@ func TestHandler(t *testing.T) { } else if !bytes.Equal(hwaddr, expectHWAddr) { log.Fatalf("expected to get hwaddr %x!=%x", hwaddr, expectHWAddr) } - n, err = c1.Encapsulate(buf[:], 0) + n, err = c1.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n > 0 { t.Fatal("expected no data") } - n, err = c2.Encapsulate(buf[:], 0) + n, err = c2.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n > 0 { diff --git a/arp/handler.go b/arp/handler.go index e761ad1..cce5e52 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -3,9 +3,11 @@ package arp import ( "bytes" "errors" + "log/slog" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type Handler struct { @@ -33,6 +35,14 @@ func (h *Handler) Protocol() uint64 { return uint64(ethernet.TypeARP) } func (h *Handler) ConnectionID() *uint64 { return &h.connID } +func (h *Handler) UpdateProtoAddr(protoAddr []byte) error { + if len(protoAddr) != len(h.ourProtoAddr) { + return errors.New("mismatch ARP proto size") + } + copy(h.ourProtoAddr, protoAddr) + return nil +} + func (h *Handler) Reset(cfg HandlerConfig) error { if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 || len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 { @@ -63,9 +73,22 @@ func (h *Handler) Reset(cfg HandlerConfig) error { type queryResult struct { protoaddr []byte hwaddr []byte + dstHw []byte querysent bool } +func (qr *queryResult) destroy() { + *qr = queryResult{protoaddr: qr.protoaddr[:0], hwaddr: qr.hwaddr[:0]} +} + +func (qr *queryResult) response() []byte { + if len(qr.hwaddr) == 0 { + return nil + } + return qr.hwaddr[:] +} +func (qr *queryResult) isInvalid() bool { return len(qr.protoaddr) == 0 } + // AbortPending drops pending queries and incoming requests. func (h *Handler) AbortPending() { h.pendingResponse = h.pendingResponse[:0] @@ -81,31 +104,69 @@ func (h *Handler) QueryResult(protoAddr []byte) (hwAddr []byte, err error) { if bytes.Equal(protoAddr, h.queries[i].protoaddr) { if !h.queries[i].querysent { return nil, errors.New("query not yet sent") - } else if len(h.queries[i].hwaddr) == 0 { + } + mac := h.queries[i].response() + if mac == nil { return nil, errors.New("no response yet") } - return h.queries[i].hwaddr, nil + return mac, nil } } return nil, errors.New("query not exist or dropped") } -func (h *Handler) StartQuery(proto []byte) error { +func (h *Handler) DiscardQuery(protoAddr []byte) error { + for i := range h.queries { + q := &h.queries[i] + if bytes.Equal(protoAddr, q.protoaddr) { + q.destroy() + return nil + } + } + return errors.New("query not found") +} + +func (h *Handler) compactQueries() { + validOff := 0 + for i := 0; i < len(h.queries); i++ { + if h.queries[i].isInvalid() { + h.queries[validOff] = h.queries[i] + validOff++ + } + } + h.queries = h.queries[:validOff] +} + +// StartQuery queues a query to perform over ARP for the protocol address `proto`. +// The user can additionally specify an dstHWAddr to write query result to on completion. +// If dstHWAddr is nil then query still occurs but no external buffer is written on query completion. +// dstHWAddr must be zeroed out (invalid MAC). +func (h *Handler) StartQuery(dstHWAddr, proto []byte) error { + if len(h.queries) == cap(h.queries) { + h.compactQueries() + if len(h.queries) == cap(h.queries) { + return errors.New("too many ongoing queries") + } + } if len(proto) != len(h.ourProtoAddr) { return errors.New("bad protocol address length") - } else if len(h.queries) == cap(h.queries) { - return errors.New("too many ongoing queries") + } else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) { + return errors.New("mismatch hardware size") + } else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) { + return errors.New("write-to buffer must be zeroed out") } h.queries = h.queries[:len(h.queries)+1] q := &h.queries[len(h.queries)-1] - q.hwaddr = q.hwaddr[:0] - q.querysent = false - q.protoaddr = append(q.protoaddr[:0], proto...) + *q = queryResult{ + protoaddr: append(q.protoaddr[:0], proto...), + hwaddr: q.hwaddr[:0], + dstHw: dstHWAddr, + } return nil } -func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { - b := eth[frameOffset:] +func (h *Handler) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + b := carrierData[offsetToFrame:] n := h.expectSize() if len(b) < n { return 0, errShortARP @@ -120,7 +181,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { copy(hwsender, h.ourHWAddr) n := copy(b, afrm.Clip().RawData()) tgt, _ := afrm.Target() - trySetEthernetDst(eth[:frameOffset], tgt) + trySetEthernetDst(carrierData[:offsetToFrame], tgt) return n, nil } for i := range h.queries { @@ -139,7 +200,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { hwTarget[j] = 0 } broadcast := ethernet.BroadcastAddr() - trySetEthernetDst(eth[:frameOffset], broadcast[:]) + trySetEthernetDst(carrierData[:offsetToFrame], broadcast[:]) return n, nil } } @@ -181,8 +242,16 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { case OpReply: hwaddr, protoaddr := afrm.Sender() for i := range h.queries { - if len(h.queries[i].hwaddr) == 0 && bytes.Equal(h.queries[i].protoaddr, protoaddr) { - h.queries[i].hwaddr = append(h.queries[i].hwaddr[:0], hwaddr...) + q := &h.queries[i] + mac := q.response() + if mac == nil && bytes.Equal(q.protoaddr, protoaddr) { + q.hwaddr = append(q.hwaddr, hwaddr...) + if q.dstHw != nil { + if !internal.IsZeroed(q.dstHw...) { + slog.Error("race-condition:ARP-reused-buffer") + } + copy(q.dstHw, hwaddr) // External write to user buffer. + } return nil } } @@ -194,7 +263,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { } func trySetEthernetDst(ethFrame []byte, dst []byte) { - if len(ethFrame) > 14 { + if len(ethFrame) >= 14 { copy(ethFrame[:6], dst) } } diff --git a/dhcpv4/client.go b/dhcpv4/client.go index eb6ab5f..2c73312 100644 --- a/dhcpv4/client.go +++ b/dhcpv4/client.go @@ -104,11 +104,11 @@ func (c *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } func (c *Client) LocalPort() uint16 { return DefaultClientPort } func (c *Client) ConnectionID() *uint64 { return &c.connID } -func (c *Client) setIP(b []byte, frameOffset int) { - if frameOffset < 28 { - return // Not an IP/UDP frame. +func (c *Client) setIP(carrierFrame []byte, offsetToIP int) { + if offsetToIP < 0 { + return // No IP layer present. } - ifrm, _ := ipv4.NewFrame(b) + ifrm, _ := ipv4.NewFrame(carrierFrame[offsetToIP:]) ifrm.SetID((uint16(c.currentXID) ^ uint16(c.currentXID>>16)) + uint16(c.state)) if c.state > StateInit { // Match server ToS since some routers drop DHCP requests if no ToS set apparently? @@ -124,7 +124,7 @@ func (c *Client) setIP(b []byte, frameOffset int) { } } -func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.isClosed() { return 0, net.ErrClosed } else if c.state == StateSelecting && !c.offer.valid { @@ -134,7 +134,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) } else if c.state == StateRequesting { return 0, nil // Currently awaiting ACK. } - dst := carrierFrame[frameOffset:] + dst := carrierData[offsetToFrame:] frm, err := NewFrame(dst) if err != nil { return 0, err @@ -194,7 +194,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) opts[numOpts] = byte(OptEnd) numOpts++ c.setHeader(frm) - c.setIP(carrierFrame, frameOffset) + c.setIP(carrierData, offsetToIP) c.state = nextState return OptionsOffset + numOpts, nil } diff --git a/dhcpv4/dhcp_test.go b/dhcpv4/dhcp_test.go index 66f0990..fe69f7a 100644 --- a/dhcpv4/dhcp_test.go +++ b/dhcpv4/dhcp_test.go @@ -28,7 +28,7 @@ func TestClientServer(t *testing.T) { // CLIENT DISCOVER. assertClState(StateInit) var buf [1024]byte - n, err := cl.Encapsulate(buf[:], 0) + n, err := cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -40,7 +40,7 @@ func TestClientServer(t *testing.T) { t.Fatal(err) } // SERVER REPLY OFFER - n, err = sv.Encapsulate(buf[:], 0) + n, err = sv.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -53,7 +53,7 @@ func TestClientServer(t *testing.T) { assertClState(StateSelecting) // CLIENT SEND OUT REQUEST. - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -66,7 +66,7 @@ func TestClientServer(t *testing.T) { } // SERVER REPLIES WITH ACK. - n, err = sv.Encapsulate(buf[:], 0) + n, err = sv.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -99,13 +99,13 @@ func TestExample(t *testing.T) { }) buf := make([]byte, 2048) buf2 := make([]byte, len(buf)) - n, err := cl.Encapsulate(buf, 0) + n, err := cl.Encapsulate(buf, -1, 0) if err != nil { t.Fatal(err) } else if n <= 0 { t.Fatal("no data sent out by client after starting request") } - n, err = cl.Encapsulate(buf2, 0) + n, err = cl.Encapsulate(buf2, -1, 0) if err != nil { t.Error("client encaps double tap after discover:", err) } @@ -141,13 +141,13 @@ func TestExample(t *testing.T) { t.Fatal(err) } - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n <= 0 { t.Fatal("no data written from client in response to offer") } - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Error("encapsulate double tap after request:", err) } else if n > 0 { diff --git a/dhcpv4/server.go b/dhcpv4/server.go index df2a143..88bceb0 100644 --- a/dhcpv4/server.go +++ b/dhcpv4/server.go @@ -159,9 +159,9 @@ func (sv *Server) Demux(carrierData []byte, frameOffset int) error { return nil } -func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) { - carrierIsIP := frameOffset >= 28 - dfrm, err := NewFrame(carrierData[frameOffset:]) +func (sv *Server) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + carrierIsIP := offsetToIP >= 0 + dfrm, err := NewFrame(carrierData[offsetToFrame:]) optBuf := dfrm.OptionsPayload()[:] if err != nil { return 0, err @@ -220,7 +220,7 @@ func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) copy(dfrm.CHAddrAs6()[:], client.hwaddr[:]) dfrm.SetMagicCookie(MagicCookie) if carrierIsIP { - err = internal.SetIPAddrs(carrierData, 0, sv.siaddr[:], client.addr[:]) + err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, sv.siaddr[:], client.addr[:]) if err != nil { return 0, err } diff --git a/dns/client.go b/dns/client.go index c143112..802e143 100644 --- a/dns/client.go +++ b/dns/client.go @@ -43,7 +43,7 @@ func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error { return nil } -func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.isClosed() { return 0, net.ErrClosed } else if c.state != dnsSendQuery { @@ -51,7 +51,7 @@ func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { } msg := &c.msg - frame := carrierData[frameOffset:] + frame := carrierData[offsetToFrame:] msglen := msg.Len() if msglen > uint16(len(frame)) { return 0, errCalcLen diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 3b00d6e..8a423d1 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -199,7 +199,7 @@ func run() (err error) { prevState = state clear(buf) - nwrite, err := stack.Encapsulate(buf[:], 0) + nwrite, err := stack.Encapsulate(buf[:], -1, 0) if err != nil { fmt.Println("ERR:ENCAPSULATE", err) } else if nwrite > 0 { @@ -267,10 +267,10 @@ func (s *Stack) Demux(b []byte, _ int) (err error) { return s.link.Demux(b, 0) } -func (s *Stack) Encapsulate(b []byte, _ int) (int, error) { - n, err := s.link.Encapsulate(b, 0) +func (s *Stack) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame) if n > 0 { - iframes, errpcap := s.shark.CaptureEthernet(s.aux[:0], b[:n], 0) + iframes, errpcap := s.shark.CaptureEthernet(s.aux[:0], carrierData[:n], 0) if errpcap != nil { fmt.Println("OU", iframes, errpcap.Error()) } else { @@ -426,7 +426,7 @@ func (s *Stack) StartResolveHardwareAddress6(ip netip.Addr) error { return errors.New("unsupported or invalid IP address") } addr := ip.As4() - return s.arp.StartQuery(addr[:]) + return s.arp.StartQuery(nil, addr[:]) } func (s *Stack) ResultResolveHardwareAddress6(ip netip.Addr) (hw [6]byte, err error) { diff --git a/examples/stack/main.go b/examples/stack/main.go index 8d84603..9370c3e 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -107,7 +107,7 @@ func main() { } } - nw, err := stack.ethernet.Encapsulate(buf[:], 0) + nw, err := stack.ethernet.Encapsulate(buf[:], -1, 0) if err != nil { lg.Error("handle", slog.String("err", err.Error())) } else if nw > 0 { @@ -230,7 +230,7 @@ func (stack *Stack) Recv(b []byte) error { } func (stack *Stack) Send(b []byte) (int, error) { - return stack.ethernet.Encapsulate(b, 0) + return stack.ethernet.Encapsulate(b, -1, 0) } func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, error) { @@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err if err != nil { return nil, err } - err = stack.tcpports.Register(&listener) + err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting. if err != nil { return nil, err } @@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error if err != nil { return nil, err } - err = stack.tcpports.Register(conn) + err = stack.tcpports.Register(conn) // Passive MAC with no listening. if err != nil { return nil, err } diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go index 8e124bb..39ace54 100644 --- a/examples/stackbasic/main.go +++ b/examples/stackbasic/main.go @@ -232,7 +232,7 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s type handler struct { raddr []byte recv func([]byte, int) error - handle func([]byte, int) (int, error) + handle func([]byte, int, int) (int, error) proto ethernet.Type lport uint16 } @@ -295,7 +295,7 @@ func (ls *LinkStack) HandleEth(dst []byte) (n int, err error) { copy(efrm.DestinationHardwareAddr()[:], ls.gwmac[:]) // default set the gateway. for i := range ls.handlers { h := &ls.handlers[i] - n, err = h.handle(dst[:mtu], 14) + n, err = h.handle(dst[:mtu], 14, 14) if err != nil { ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) continue @@ -322,8 +322,8 @@ func (as *ARPStack) Recv(EtherFrame []byte, arpOff int) error { return as.handler.Demux(EtherFrame, arpOff) } -func (as *ARPStack) Handle(EtherFrame []byte, arpOff int) (int, error) { - n, err := as.handler.Encapsulate(EtherFrame, arpOff) +func (as *ARPStack) Handle(EtherFrame []byte, offsetToIP, arpOff int) (int, error) { + n, err := as.handler.Encapsulate(EtherFrame, offsetToIP, arpOff) if err != nil || n == 0 { return 0, err } diff --git a/examples/xnet/main.go b/examples/xnet/main.go index 1408ab2..2061de0 100644 --- a/examples/xnet/main.go +++ b/examples/xnet/main.go @@ -110,7 +110,7 @@ func run() (err error) { var frames []pcap.Frame for { clear(buf) - nwrite, err := stack.Encapsulate(buf[:], 0) + nwrite, err := stack.Encapsulate(buf[:], -1, 0) if err != nil { fmt.Println("ERR:ENCAPSULATE", err) } else if nwrite > 0 { diff --git a/internal/ip.go b/internal/ip.go index ac86109..77a9263 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { copy(dstaddr, dst) return nil } + +// IsZeroed returns true if all arguments are set to their zero value. +func IsZeroed[T comparable](a ...T) bool { + var z T + for i := range a { + if a[i] != z { + return false + } + } + return true +} diff --git a/internet/definitions.go b/internet/definitions.go index 55a4fb9..010a896 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -2,25 +2,30 @@ package internet import ( "errors" + "log/slog" "math" "net" + "slices" ) // StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols, // from Ethernet to IP to TCP, practically any protocol can be expressed as a StackNode and function completely. type StackNode interface { - // Encapsulate writes the stack node's frame into carrierData[frameOffset:] + // Encapsulate writes the stack node's frame into carrierData[offsetToFrame:] // along with any other frame or payload the stack node encapsulates. - // The returned integer is amount of bytes written such that carrierData[frameOffset:frameOffset+n] - // contains written data. Data inside carrierData[:frameOffset] usually contains data necessary for + // The returned integer is amount of bytes written such that carrierData[offsetToFrame:offsetToFrame+n] + // contains written data. Data inside carrierData[:offsetToFrame] usually contains data necessary for // a StackNode to correctly emit valid frame data: such is the case for TCP packets which require IP // frame data for checksum calculation. Thus StackNodes must provide fields in their own frame // required by sub-stacknodes for correct encapsulation; in the case of IPv4/6 this means including fields // used in pseudo-header checksum like local IP (see [ipv4.CRCWriteUDPPseudo]). // + // offsetToIP is the offset to the IP frame, if present, else its value should be -1. + // The relation offsetToIP<=offsetToFrame should always hold. + // // When [net.ErrClosed] is returned the StackNode should be discarded and any written data passed up normally. // Errors returned by Encapsulate are "extraordinary" and should not be returned unless the StackNode is receiving invalid carrierData/frameOffset. - Encapsulate(carrierData []byte, frameOffset int) (int, error) + Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) // Demux reads from the argument buffer where frameOffset is the offset of this StackNode's frame first byte. // The stack node then dispatches(demuxes) the encapsulated frames to its corresponding sub-stack-node(s). Demux(carrierData []byte, frameOffset int) error @@ -36,9 +41,156 @@ type node struct { currConnID uint64 connID *uint64 demux func([]byte, int) error - encapsulate func([]byte, int) (int, error) + encapsulate func([]byte, int, int) (int, error) proto uint16 port uint16 + // remoteAddr will be set on active(outbound) port connections + // that require an ARP to set the remoteAddr beforehand. + remoteAddr []byte +} + +type handlers struct { + context string + logger + nodes []node +} + +func (h *handlers) reset(context string, maxNodes int) { + h.nodes = slices.Grow(h.nodes[:0], maxNodes) + h.context = context +} + +func (h *handlers) registerByProto(n node) error { + err := h.prepAdd() + if err != nil { + return err + } + if h.nodeByProto(n.proto) != nil { + return errProtoRegistered + } + h.nodes = append(h.nodes, n) + return nil +} + +func (h *handlers) registerByPortProto(n node) error { + err := h.prepAdd() + if err != nil { + return err + } + if h.nodeByPortProto(n.port, n.proto) != nil { + return errProtoRegistered + } + h.nodes = append(h.nodes, n) + return nil +} + +func (h *handlers) prepAdd() error { + if h.full() { + h.compact() + if h.full() { + return errNodesFull + } + } + return nil +} + +func (h *handlers) full() bool { return cap(h.nodes) == len(h.nodes) } + +func (h *handlers) compact() { + nilOff := 0 + for i := 0; i < len(h.nodes); i++ { + if !h.nodes[i].IsInvalid() { + h.nodes[nilOff] = h.nodes[i] + nilOff++ + } + } + h.nodes = h.nodes[:nilOff] +} + +func (h *handlers) tryHandleError(node *node, err error) (discardedGracefully bool) { + if err != nil && (err == net.ErrClosed || node.IsInvalid()) { + node.destroy() + discardedGracefully = true + } + return discardedGracefully +} + +func (h *handlers) nodeByProto(proto uint16) *node { + for i := range h.nodes { + node := &h.nodes[i] + if node.proto == proto && !node.IsInvalid() { + return node + } + } + return nil +} + +func (h *handlers) nodeByPort(port uint16) *node { + for i := range h.nodes { + node := &h.nodes[i] + if node.port == port && !node.IsInvalid() { + return node + } + } + return nil +} + +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 && !node.IsInvalid() { + return node + } + } + return nil +} + +func (h *handlers) demuxByProto(buf []byte, offset int, proto uint16) (*node, error) { + node := h.nodeByProto(proto) + if node == nil { + return nil, nil + } + err := node.demux(buf, offset) + if h.tryHandleError(node, err) { + err = nil + } + return node, err +} + +func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, error) { + node := h.nodeByPort(port) + if node == nil { + return nil, nil + } + err := node.demux(buf, offset) + if h.tryHandleError(node, err) { + err = nil + node = nil // Node is destroyed in tryHandleError and invalidated. + } + return node, err +} + +// encapsulateAny finds a node suitable to write and encapsulates the package. +// If no data is sent it returns the last error encountered. +func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) { + for i := range h.nodes { + node := &h.nodes[i] + if node.IsInvalid() { + continue + } + n, err = node.encapsulate(buf, offsetIP, offsetThisFrame) + if h.tryHandleError(node, err) { + err = nil // CLOSE error handled gracefully by deleting node. + node = nil // Node is destroyed in tryHandleError and invalidated. + } + if n > 0 { + return node, n, err + } else if err != nil { + // Make sure not to hang on one handler that keeps returning an error. + h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error())) + } + } + return nil, 0, err // Return last written error. } var ( @@ -50,32 +202,6 @@ var ( _ = net.ErrClosed ) -func registerNode(nodesPtr *[]node, h node) error { - if cap(*nodesPtr)-len(*nodesPtr) <= 0 { - *nodesPtr = nodesCompact(*nodesPtr) - } - if cap(*nodesPtr)-len(*nodesPtr) <= 0 { - return errNodesFull - } - *nodesPtr = append(*nodesPtr, h) - return nil -} - -func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) { - if err != nil { - if nodeIdx >= len(*nodesPtr) { - panic("unreachable") - } - nodes := *nodesPtr - if checkNodeErr(&nodes[nodeIdx], err) { - // *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1) - (*nodesPtr)[nodeIdx] = node{} // 'Delete' node without modifying slice length. - discarded = true - } - } - return discarded -} - func (node *node) IsInvalid() bool { return node.demux == nil || node.encapsulate == nil || (node.connID != nil && node.currConnID != *node.connID) } @@ -84,7 +210,7 @@ func checkNodeErr(node *node, err error) (discard bool) { return node.IsInvalid() || (err != nil && err == net.ErrClosed) } -func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node { +func nodeFromStackNode(s StackNode, port uint16, protocol uint64, remoteAddr []byte) node { if protocol > math.MaxUint16 { panic(">16bit protocol number unsupported") } @@ -100,64 +226,11 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node { encapsulate: s.Encapsulate, proto: uint16(protocol), port: port, + remoteAddr: remoteAddr, // SHARED MEMORY- used to signal. } } -func getNode(nodes []node, port uint16, protocol uint16) (node *node) { - for i := range nodes { - node := &nodes[i] - if node.port == port && node.proto == protocol { - return node - } - } - return nil -} - -func getEncapsulateNode(nodes *[]node, carrierData []byte, frameOffset int) (nodeIdx int, written int, err error) { - destroyed := false - for i := range *nodes { - node := &(*nodes)[i] - if node.IsInvalid() { - destroyed = true - node.destroy() - continue - } - written, err = node.encapsulate(carrierData, frameOffset) - if written > 0 { - return i, written, err - } else if err != nil { - - } - } - if destroyed { - *nodes = nodesCompact(*nodes) - } - return -1, 0, nil -} - // destroy removes all references to underlying StackNode. Allows garbage collection of node if possible. func (n *node) destroy() { *n = node{} } - -func getNodeByProto(nodes []node, protocol uint16) int { - for i := range nodes { - node := &nodes[i] - if node.proto == protocol { - return i - } - } - - return -1 -} - -func nodesCompact(nodes []node) []node { - nilOff := 0 - for i := 0; i < len(nodes); i++ { - if !nodes[i].IsInvalid() { - nodes[nilOff] = nodes[i] - nilOff++ - } - } - return nodes[:nilOff] -} diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index 8cac79e..9b1aa14 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -94,7 +94,7 @@ func (listener *NodeTCPListener) TryAccept() (*tcp.Conn, error) { } // Encapsulate implements [StackNode]. -func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset int) (int, error) { +func (listener *NodeTCPListener) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if listener.isClosed() { return 0, net.ErrClosed } @@ -102,7 +102,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset if conn == nil { continue } - n, err := conn.Encapsulate(carrierData, tcpFrameOffset) + n, err := conn.Encapsulate(carrierData, offsetToIP, offsetToFrame) if err != nil { err = listener.maintainConn(listener.accepted, i, err) } diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index c2ebaad..fd1edcc 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -6,7 +6,6 @@ import ( "log/slog" "math" "net" - "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -14,11 +13,10 @@ import ( type StackEthernet struct { connID uint64 - handlers []node - logger - mac [6]byte - gwmac [6]byte - mtu uint16 + handlers handlers + mac [6]byte + gwmac [6]byte + mtu uint16 } func (ls *StackEthernet) SetGateway6(gw [6]byte) { @@ -43,11 +41,10 @@ func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error { } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ls.handlers = slices.Grow(ls.handlers[:0], maxNodes) + ls.handlers.reset("StackEthernet", maxNodes) *ls = StackEthernet{ connID: ls.connID + 1, handlers: ls.handlers, - logger: ls.logger, mac: mac, gwmac: gateway, mtu: uint16(mtu), @@ -68,18 +65,7 @@ func (ls *StackEthernet) Register(h StackNode) error { if proto > math.MaxUint16 || proto <= 1500 { return errInvalidProto } - eproto := uint16(proto) - for i := range ls.handlers { - hgot := &ls.handlers[i] - if hgot.proto == eproto { - return errProtoRegistered - } - } - return registerNode(&ls.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - proto: eproto, - }) + return ls.handlers.registerByProto(nodeFromStackNode(h, 0, proto, nil)) } func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) { @@ -98,21 +84,17 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) if vld.HasError() { return vld.ErrPop() } - - for i := range ls.handlers { - h := &ls.handlers[i] - if h.proto == uint16(etype) { - return h.demux(efrm.Payload(), 0) - } + if h, err := ls.handlers.demuxByProto(efrm.Payload(), 0, uint16(etype)); h != nil { + return err } DROP: - ls.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String())) + ls.handlers.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String())) return lneto.ErrPacketDrop } -func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int, err error) { +func (ls *StackEthernet) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { mtu := ls.mtu - dst := carrierData[frameOffset:] + dst := carrierData[offsetToFrame:] if len(dst) < int(mtu) { return 0, io.ErrShortBuffer } @@ -121,19 +103,19 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int return 0, err } *efrm.DestinationHardwareAddr() = ls.gwmac - for i := range ls.handlers { - h := &ls.handlers[i] - n, err = h.encapsulate(dst[:mtu], 14) - if err != nil { - ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) - continue - } - if n > 0 { - // Found packet - *efrm.SourceHardwareAddr() = ls.mac - efrm.SetEtherType(ethernet.Type(h.proto)) - return n + 14, nil - } + var h *node + // Children (IP/ARP) start at offset 14 (after ethernet header). + // For IP: offsetToIP=14, offsetToFrame=14 + // For ARP: offsetToIP=-1, offsetToFrame=14 (but ARP ignores offsetToIP) + // Clip carrierData to MTU to prevent writes beyond MTU limit. + mtuLimit := offsetToFrame + int(mtu) + h, n, err = ls.handlers.encapsulateAny(carrierData[:mtuLimit], offsetToFrame+14, offsetToFrame+14) + if n == 0 { + return n, err } - return 0, err + // Found packet + *efrm.SourceHardwareAddr() = ls.mac + efrm.SetEtherType(ethernet.Type(h.proto)) + n += 14 + return n, err } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 5ac4549..a4ec5ce 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -5,7 +5,6 @@ import ( "io" "log/slog" "net/netip" - "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -19,13 +18,11 @@ import ( var _ StackNode = (*StackIP)(nil) type StackIP struct { - connID uint64 - ipID uint16 - ip [4]byte - validator lneto.Validator - handlers []node - pendingICMP [][]byte - logger + connID uint64 + ipID uint16 + ip [4]byte + validator lneto.Validator + handlers handlers } func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { @@ -36,14 +33,12 @@ func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { if err != nil { return err } - sb.handlers = slices.Grow(sb.handlers[:0], maxNodes) + sb.handlers.reset("StackIP", maxNodes) *sb = StackIP{ - connID: sb.connID + 1, - validator: sb.validator, - handlers: sb.handlers, - logger: sb.logger, - ip: sb.ip, - pendingICMP: make([][]byte, maxNodes*4), + connID: sb.connID + 1, + validator: sb.validator, + handlers: sb.handlers, + ip: sb.ip, } return nil } @@ -73,11 +68,11 @@ func (sb *StackIP) Addr() netip.Addr { } func (sb *StackIP) SetLogger(logger *slog.Logger) { - sb.logger.log = logger + sb.handlers.log = logger } func (sb *StackIP) Demux(carrierData []byte, offset int) error { - sb.info("StackIP.Demux:start") + sb.handlers.info("StackIP.Demux:start") frame := carrierData[offset:] // we don't care about carrier data in IP. ifrm, err := ipv4.NewFrame(frame) if err != nil { @@ -96,7 +91,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { gotCRC := ifrm.CRC() wantCRC := ifrm.CalculateHeaderCRC() if gotCRC != wantCRC { - sb.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC))) + sb.handlers.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC))) return errors.New("IPv4 CRC mismatch") } off := ifrm.HeaderLength() @@ -105,10 +100,11 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { if proto == lneto.IPProtoICMP { return sb.recvicmp(ifrm.RawData(), ifrm.HeaderLength()) } - nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) - if nodeIdx < 0 { + node := sb.handlers.nodeByProto(uint16(proto)) + // nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) + if node == nil { // Drop packet. - sb.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String())) + sb.handlers.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String())) return nil } // Incoming CRC Validation of common IP Protocols. @@ -135,17 +131,17 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { return errors.New("UDP CRC mismatch") } } - sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) - err = sb.handlers[nodeIdx].demux(frame[:totalLen], off) - if handleNodeError(&sb.handlers, nodeIdx, err) { - sb.info("ipclose", slog.String("proto", proto.String())) + sb.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) + err = node.demux(frame[:totalLen], off) + if sb.handlers.tryHandleError(node, err) { + sb.handlers.info("ipclose", slog.String("proto", proto.String())) err = nil } return err } -func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) { - frame := carrierData[frameOffset:] +func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + frame := carrierData[offsetToFrame:] if len(frame) < 256 { return 0, io.ErrShortBuffer } @@ -162,46 +158,37 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) ifrm.SetTTL(64) *ifrm.SourceAddr() = sb.ip sb.ipID = id - for i := range sb.handlers { - h := &sb.handlers[i] - proto := lneto.IPProto(h.proto) - n, err := h.encapsulate(frame[:], headerlen) - if err != nil { - if handleNodeError(&sb.handlers, i, err) { - println("IP NODE REMOVED", proto.String(), h.port) - h.destroy() - } - sb.error("StackIP:encapsulate", slog.String("proto", proto.String()), slog.String("err", err.Error())) - continue - } else if n == 0 { - continue - } - totalLen := n + headerlen - ifrm.SetTotalLength(uint16(totalLen)) - ifrm.SetProtocol(proto) - ifrm.SetCRC(ifrm.CalculateHeaderCRC()) - // Calculate CRC for our newly generated packet. - var crc lneto.CRC791 - switch proto { - case lneto.IPProtoTCP: - ifrm.CRCWriteTCPPseudo(&crc) - tfrm, _ := tcp.NewFrame(ifrm.Payload()) - tfrm.CRCWrite(&crc) - tfrm.SetCRC(crc.Sum16()) - case lneto.IPProtoUDP: - ifrm.CRCWriteUDPPseudo(&crc) - ufrm, _ := udp.NewFrame(ifrm.Payload()) - ufrm.SetLength(uint16(n)) - ufrm.CRCWriteIPv4(&crc) - ufrm.SetCRC(crc.Sum16()) - if n != int(ufrm.Length()) { - sb.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length()))) - return 0, errors.New("invalid UDP length") - } - } - return totalLen, nil + // Children (TCP/UDP) start at offset headerlen (20 bytes after IP header start). + // offsetToIP is 0 relative to this slice (frame), children's frame starts at headerlen. + node, n, err := sb.handlers.encapsulateAny(carrierData, offsetToFrame, offsetToFrame+headerlen) + if n == 0 { + return n, err } - return 0, nil + proto := lneto.IPProto(node.proto) + totalLen := n + headerlen + ifrm.SetTotalLength(uint16(totalLen)) + ifrm.SetProtocol(proto) + ifrm.SetCRC(ifrm.CalculateHeaderCRC()) + // Calculate CRC for our newly generated packet. + var crc lneto.CRC791 + switch proto { + case lneto.IPProtoTCP: + ifrm.CRCWriteTCPPseudo(&crc) + tfrm, _ := tcp.NewFrame(ifrm.Payload()) + tfrm.CRCWrite(&crc) + tfrm.SetCRC(crc.Sum16()) + case lneto.IPProtoUDP: + ifrm.CRCWriteUDPPseudo(&crc) + ufrm, _ := udp.NewFrame(ifrm.Payload()) + ufrm.SetLength(uint16(n)) + ufrm.CRCWriteIPv4(&crc) + ufrm.SetCRC(crc.Sum16()) + if n != int(ufrm.Length()) { + sb.handlers.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length()))) + return 0, errors.New("invalid UDP length") + } + } + return totalLen, err } func (sb *StackIP) Register(h StackNode) error { @@ -209,19 +196,7 @@ func (sb *StackIP) Register(h StackNode) error { if proto > 255 { return errInvalidProto } - connID := h.ConnectionID() - var currConnID uint64 - if connID != nil { - currConnID = *connID - } - return registerNode(&sb.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - proto: uint16(proto), - port: h.LocalPort(), - currConnID: currConnID, - connID: connID, - }) + return sb.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil)) } func (sb *StackIP) recvicmp(carrierData []byte, offset int) error { diff --git a/internet/stack-ports.go b/internet/stack-ports.go index bd51284..413cc09 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -2,18 +2,23 @@ package internet import ( "encoding/binary" + "errors" "io" + "log/slog" "math" - "slices" + "strconv" "github.com/soypat/lneto" + "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type StackPorts struct { connID uint64 - handlers []node + handlers handlers dstPortOff uint16 protocol uint16 + // stores last node to demux/encapsulate. } func (ps *StackPorts) ResetUDP(maxNodes int) error { @@ -30,7 +35,7 @@ func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ps.handlers = slices.Grow(ps.handlers[:0], maxNodes) + ps.handlers.reset("StackPorts(proto="+strconv.Itoa(int(protocol))+")", maxNodes) *ps = StackPorts{ connID: ps.connID + 1, handlers: ps.handlers, @@ -46,24 +51,11 @@ func (ps *StackPorts) Protocol() uint64 { return uint64(ps.protocol) } func (ps *StackPorts) ConnectionID() *uint64 { return &ps.connID } -func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) { - if int(ps.dstPortOff)+offset+2 > len(b) { +func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { + if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) { return 0, io.ErrShortBuffer } - var i int - for i = 0; i < len(ps.handlers); i++ { - if ps.handlers[i].IsInvalid() { - continue - } - n, err = ps.handlers[i].encapsulate(b, offset) - if err != nil || n > 0 { - if ps.handleResult(i, n, err) { - err = nil // Handler discarded. Keep looking for other handlers. - continue - } - break - } - } + _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) return n, err } @@ -72,52 +64,90 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { return io.ErrShortBuffer } port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) - var i int - for i = 0; i < len(ps.handlers); i++ { - if port != ps.handlers[i].port { - continue - } - err = ps.handlers[i].demux(b, offset) - if err != nil { - if ps.handleResult(i, 0, err) { - err = nil // Handler discarded. Keep looking for other maybe available handlers. - continue - } - break - } - } - ps.handleResult(i, 0, err) + _, err = ps.handlers.demuxByPort(b, offset, port) return err } +// Register registers a port StackNode on StackPorts. +// If dstMAC is set to non-nil, length six buffer then func (ps *StackPorts) Register(h StackNode) error { port := h.LocalPort() proto := h.Protocol() - if port <= 0 { return errZeroPort } else if proto != uint64(ps.protocol) { return errInvalidProto } - var cid uint64 - cidPtr := h.ConnectionID() - if cidPtr != nil { - cid = *cidPtr - } - return registerNode(&ps.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - port: port, - currConnID: cid, - connID: cidPtr, - proto: uint16(proto), - }) + return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil)) } -func (ps *StackPorts) handleResult(handlerIdx, n int, err error) (discarded bool) { - if handleNodeError(&ps.handlers, handlerIdx, err) { - discarded = true - println("DISCARD", handlerIdx, "witherr", err.Error()) - } - return discarded +// StackPortsMACFiltered is a StackPorts implementation but that avoids calling encapsulate on nodes +// with a non-nil MAC address registered via Register method that is set to all zero values. +// If the address is set to nil no filtering occurs. MAC Address is set automatically on the ethernet frame by StackPortsMACFiltered when non-nil. +type StackPortsMACFiltered struct { + sp StackPorts +} + +func (mfsp *StackPortsMACFiltered) Register(h StackNode, addr []byte) error { + port := h.LocalPort() + proto := h.Protocol() + if port <= 0 { + return errZeroPort + } else if proto != uint64(mfsp.sp.protocol) { + return errInvalidProto + } else if addr != nil && len(addr) != 6 { + return errors.New("invalid MAC") + } + return mfsp.sp.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, addr)) +} + +func (ps *StackPortsMACFiltered) ResetUDP(maxNodes int) error { + return ps.sp.ResetUDP(maxNodes) +} + +func (ps *StackPortsMACFiltered) ResetTCP(maxNodes int) error { + return ps.sp.ResetTCP(maxNodes) +} + +func (ps *StackPortsMACFiltered) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) error { + return ps.sp.Reset(protocol, dstPortOffset, maxNodes) +} + +func (ps *StackPortsMACFiltered) LocalPort() uint16 { return 0 } + +func (ps *StackPortsMACFiltered) Protocol() uint64 { return uint64(ps.sp.protocol) } + +func (ps *StackPortsMACFiltered) ConnectionID() *uint64 { return &ps.sp.connID } + +func (ps *StackPortsMACFiltered) Demux(b []byte, offset int) (err error) { + // No MAC Filtering on ingress. TODO? + return ps.sp.Demux(b, offset) +} + +func (ps *StackPortsMACFiltered) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { + if int(ps.sp.dstPortOff)+offsetToFrame+2 > len(carrierData) { + return 0, io.ErrShortBuffer + } + h := &ps.sp.handlers + for i := range h.nodes { + node := &h.nodes[i] + if node.IsInvalid() || (len(node.remoteAddr) > 0 && internal.IsZeroed(node.remoteAddr...)) { + continue + } + n, err = node.encapsulate(carrierData, offsetToIP, offsetToFrame) + if h.tryHandleError(node, err) { + err = nil // CLOSE error handled gracefully by deleting node. + } + if n > 0 { + if len(node.remoteAddr) == 6 && offsetToIP >= 14 { + efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) + *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) + } + return n, err + } else if err != nil { + // Make sure not to hang on one handler that keeps returning an error. + h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error())) + } + } + return 0, err // Return last written error. } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index aaa4ae8..e64615f 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -17,7 +17,7 @@ type StackUDPPort struct { } func (sudp *StackUDPPort) SetStackNode(node StackNode, raddr []byte, rmport uint16) { - sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol()) + sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol(), raddr) sudp.rmport = rmport sudp.raddr = append(sudp.raddr[:0], raddr...) } @@ -61,24 +61,25 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error { return err } -func (sudp *StackUDPPort) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (sudp *StackUDPPort) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if sudp.h.IsInvalid() { sudp.h.destroy() return 0, net.ErrClosed } - ufrm, err := udp.NewFrame(carrierData[frameOffset:]) + ufrm, err := udp.NewFrame(carrierData[offsetToFrame:]) if err != nil { return 0, err } ufrm.SetSourcePort(sudp.h.port) ufrm.SetDestinationPort(sudp.rmport) - if len(sudp.raddr) > 0 && frameOffset >= 20 { - err = internal.SetIPAddrs(carrierData, 0, nil, sudp.raddr) + if len(sudp.raddr) > 0 && offsetToIP >= 0 { + err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, nil, sudp.raddr) if err != nil { return 0, err } } - n, err := sudp.h.encapsulate(carrierData, frameOffset+8) + // Child payload starts 8 bytes after UDP header start. + n, err := sudp.h.encapsulate(carrierData, offsetToIP, offsetToFrame+8) if n == 0 { if err != nil { slog.Error("stackudp:encapsulate", slog.String("err", err.Error())) diff --git a/internet/stackbasic_test.go b/internet/stackbasic_test.go index 39a704a..2e93702 100644 --- a/internet/stackbasic_test.go +++ b/internet/stackbasic_test.go @@ -44,7 +44,7 @@ func TestBasicStack2(t *testing.T) { func expectExchange(t *testing.T, from, to *StackIP, buf []byte) { t.Helper() - n, err := from.Encapsulate(buf, 0) + n, err := from.Encapsulate(buf, -1, 0) if err != nil { t.Error("expectExchange:encapsulate:", err) } else if n == 0 { diff --git a/ntp/client.go b/ntp/client.go index 8c65018..2d6c387 100644 --- a/ntp/client.go +++ b/ntp/client.go @@ -50,11 +50,11 @@ func (c *Client) ConnectionID() *uint64 { return &c.connID } -func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.IsDone() { return 0, nil } - payload := carrierData[frameOffset:] + payload := carrierData[offsetToFrame:] frm, err := NewFrame(payload) if err != nil { return 0, err diff --git a/tcp/conn.go b/tcp/conn.go index 62c7d7a..c7577d6 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -246,7 +246,7 @@ func (conn *Conn) checkPipeOpen() error { if conn.abortErr != nil { return conn.abortErr } - state := conn.State() + state := conn.h.State() if state.IsClosed() { return net.ErrClosed } @@ -278,23 +278,27 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) { return nil } -func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) { +func (conn *Conn) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { conn.mu.Lock() defer conn.mu.Unlock() if len(conn.remoteAddr) == 0 { return 0, errNoRemoteAddr } - raddr, _, _, _, err := internal.GetIPAddr(buf[:off]) + if offsetToIP < 0 { + return 0, errNoRemoteAddr // No IP layer present. + } + ipFrame := carrierData[offsetToIP:offsetToFrame] + raddr, _, _, _, err := internal.GetIPAddr(ipFrame) if err != nil { return 0, err } else if len(raddr) != len(conn.remoteAddr) { return 0, errMismatchedIPVersion } - n, err = conn.h.Send(buf[off:]) + n, err = conn.h.Send(carrierData[offsetToFrame:]) if err != nil { return 0, err } - err = internal.SetIPAddrs(buf[:off], conn.ipID, nil, conn.remoteAddr) + err = internal.SetIPAddrs(ipFrame, conn.ipID, nil, conn.remoteAddr) if err != nil { return 0, err } @@ -328,11 +332,11 @@ func (conn *Conn) reset(h Handler) { func (conn *Conn) SetDeadline(t time.Time) error { conn.mu.Lock() defer conn.mu.Unlock() - err := conn.SetReadDeadline(t) + err := conn.setReadDeadline(t) if err != nil { return err } - return conn.SetWriteDeadline(t) + return conn.setWriteDeadline(t) } // SetReadDeadline sets the deadline for future Read calls @@ -340,7 +344,11 @@ func (conn *Conn) SetDeadline(t time.Time) error { func (conn *Conn) SetReadDeadline(t time.Time) error { conn.mu.Lock() defer conn.mu.Unlock() - conn.trace("TCPConn.SetReadDeadline:start") + return conn.setReadDeadline(t) +} + +func (conn *Conn) setReadDeadline(t time.Time) error { + conn.trace("TCPConn.setReadDeadline:start") err := conn.checkPipeOpen() if err == nil { conn.rdead = t @@ -354,6 +362,12 @@ func (conn *Conn) SetReadDeadline(t time.Time) error { // some of the data was successfully written. // A zero value for t means Write will not time out. func (conn *Conn) SetWriteDeadline(t time.Time) error { + conn.mu.Lock() + defer conn.mu.Unlock() + return conn.setWriteDeadline(t) +} + +func (conn *Conn) setWriteDeadline(t time.Time) error { conn.trace("TCPConn.SetWriteDeadline:start") err := conn.checkPipeOpen() if err == nil { diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 6e4e09b..84430c9 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -28,11 +28,12 @@ type StackAsync struct { ip internet.StackIP arp arp.Handler udps internet.StackPorts - tcps internet.StackPorts + tcps internet.StackPortsMACFiltered dhcpUDP internet.StackUDPPort dhcp dhcpv4.Client dhcpResults DHCPResults + subnet netip.Prefix // Local subnet for ARP resolution. dnsUDP internet.StackUDPPort dns dns.Client @@ -73,11 +74,11 @@ func (s *StackAsync) Demux(carrierData []byte, etherOff int) error { return s.link.Demux(carrierData, etherOff) } -func (s *StackAsync) Encapsulate(carrierData []byte, etherOff int) (int, error) { +func (s *StackAsync) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { s.mu.Lock() defer s.mu.Unlock() - n, err := s.link.Encapsulate(carrierData, etherOff) + n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame) s.totalsent += uint64(n) return n, err } @@ -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 @@ -169,7 +168,7 @@ func (s *StackAsync) resetARP() error { if addr.Is6() { proto = ethernet.TypeIPv6 } - return s.arp.Reset(arp.HandlerConfig{ + err := s.arp.Reset(arp.HandlerConfig{ HardwareAddr: mac[:], ProtocolAddr: addr.AsSlice(), MaxQueries: 3, @@ -177,6 +176,14 @@ func (s *StackAsync) resetARP() error { HardwareType: 1, ProtocolType: proto, }) + if err != nil { + return err + } + err = s.link.Register(&s.arp) + if err != nil { + return err + } + return nil } // Prand32 generates a pseudo random 32-bit unsigned integer from the internal state and advances the seed. @@ -193,11 +200,17 @@ 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 } - return s.resetARP() + ip := addr.As4() + err = s.arp.UpdateProtoAddr(ip[:]) + return err } func (s *StackAsync) Addr() netip.Addr { @@ -234,11 +247,23 @@ func (s *StackAsync) Gateway6() [6]byte { func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort) (err error) { s.mu.Lock() defer s.mu.Unlock() + var mac []byte + if s.subnet.Contains(addrp.Addr()) { + mac = make([]byte, 6) + ip := addrp.Addr().As4() + // StartQuery starts an ARP query for addresses in this network. + // On finishing query MAC is set and thus the StackPort will allow encapsulating + // data on that connection. + err = s.arp.StartQuery(mac, ip[:]) + if err != nil { + return err + } + } err = conn.OpenActive(localPort, addrp, tcp.Value(s.Prand32())) if err != nil { return err } - err = s.tcps.Register(conn) + err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network. if err != nil { conn.Abort() return err @@ -253,7 +278,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { if err != nil { return err } - err = s.tcps.Register(conn) + err = s.tcps.Register(conn, nil) if err != nil { conn.Abort() return err @@ -376,7 +401,7 @@ func (s *StackAsync) StartResolveHardwareAddress6(ip netip.Addr) error { return errors.New("unsupported or invalid IP address") } addr := ip.As4() - return s.arp.StartQuery(addr[:]) + return s.arp.StartQuery(nil, addr[:]) } // ResultResolveHardwareAddress6 @@ -432,16 +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.ip.SetAddr(results.AssignedAddr) - if err != nil { - return err - } - // Reset ARP handler with new IP address so it can respond to ARP requests. - err = stack.resetARP() + 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 086bf99..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 @@ -322,7 +326,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs default: panic("OOB") } - n, err := src.Encapsulate(buf[:], 0) + n, err := src.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -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)) +}