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..aca916b 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -3,6 +3,7 @@ package arp import ( "bytes" "errors" + "log/slog" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -63,9 +64,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 +95,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 && !allZeros(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 +172,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 +191,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 +233,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 !allZeros(q.dstHw) { + slog.Error("race-condition:ARP-reused-buffer") + } + copy(q.dstHw, hwaddr) // External write to user buffer. + } return nil } } @@ -198,3 +258,12 @@ func trySetEthernetDst(ethFrame []byte, dst []byte) { copy(ethFrame[:6], dst) } } + +func allZeros(b []byte) bool { + for i := range b { + if b[i] != 0 { + return false + } + } + return true +} 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..f81c4c2 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) { 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/internet/definitions.go b/internet/definitions.go index 683d34b..5b47ee0 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -11,18 +11,21 @@ import ( // 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 @@ -38,10 +41,12 @@ 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 []byte + // remoteAddr will be set on active(outbound) port connections + // that require an ARP to set the remoteAddr beforehand. + remoteAddr []byte } type handlers struct { @@ -166,13 +171,13 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro // 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, offset int) (_ *node, n int, err error) { +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, offset) + n, err = node.encapsulate(buf, offsetIP, offsetThisFrame) if h.tryHandleError(node, err) { err = nil // CLOSE error handled gracefully by deleting node. } 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 46d6794..fd1edcc 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -92,9 +92,9 @@ DROP: 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 } @@ -104,7 +104,12 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int } *efrm.DestinationHardwareAddr() = ls.gwmac var h *node - h, n, err = ls.handlers.encapsulateAny(dst[:mtu], 14) + // 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 } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index e7748e9..a4ec5ce 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -140,8 +140,8 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { 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 } @@ -158,7 +158,9 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) ifrm.SetTTL(64) *ifrm.SourceAddr() = sb.ip sb.ipID = id - node, n, err := sb.handlers.encapsulateAny(frame, headerlen) + // 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 } diff --git a/internet/stack-ports.go b/internet/stack-ports.go index 87eece3..dd5e304 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -46,11 +46,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 } - _, n, err = ps.handlers.encapsulateAny(b, offset) + _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) return n, err } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index c63ee1b..e64615f 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -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..3dff2db 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -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 } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 6e4e09b..7946153 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -73,11 +73,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 } @@ -376,7 +376,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 diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index 086bf99..d066224 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -322,7 +322,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 {