Merge pull request #12 from soypat/local-arp-tcp

Fix Local ARP
This commit is contained in:
Pat Whittingslow
2025-12-31 12:03:41 -03:00
committed by GitHub
24 changed files with 748 additions and 370 deletions
+4 -4
View File
@@ -57,11 +57,11 @@ The following interface is implemented by networking stack nodes and the stack t
```go ```go
type StackNode interface { type StackNode interface {
// Encapsulate receives a buffer the receiver must fill with data. // Encapsulate receives a buffer the receiver must fill with data.
// The receiver's start byte is at carrierData[frameOffset]. // The receiver's start byte is at carrierData[offsetToFrame].
Encapsulate(carrierData []byte, frameOffset int) (int, error) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error)
// Demux receives a buffer the receiver must decode and pass on to corresponding child StackNode(s). // 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]. // The receiver's start byte is at carrierData[offsetToFrame].
Demux(carrierData []byte, frameOffset int) error Demux(carrierData []byte, offsetToFrame int) error
// LocalPort returns the port of the node if applicable or zero. Used for UDP/TCP nodes. // LocalPort returns the port of the node if applicable or zero. Used for UDP/TCP nodes.
LocalPort() uint16 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). // Protocol returns the protocol of this node if applicable or zero. Usually either a ethernet.Type (EtherType) or lneto.IPProto (IP Protocol number).
+8 -8
View File
@@ -34,13 +34,13 @@ func TestHandler(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
var buf, discard [64]byte var buf, discard [64]byte
n, err := c1.Encapsulate(buf[:], 0) n, err := c1.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal("error on should be nop send:", err) t.Fatal("error on should be nop send:", err)
} else if n > 0 { } else if n > 0 {
t.Fatal("should not send if no query") t.Fatal("should not send if no query")
} }
n, err = c2.Encapsulate(buf[:], 0) n, err = c2.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal("error on should be nop send:", err) t.Fatal("error on should be nop send:", err)
} else if n > 0 { } else if n > 0 {
@@ -50,11 +50,11 @@ func TestHandler(t *testing.T) {
// Perform ARP exchange. // Perform ARP exchange.
expectHWAddr := c2.ourHWAddr expectHWAddr := c2.ourHWAddr
queryAddr := c2.ourProtoAddr queryAddr := c2.ourProtoAddr
err = c1.StartQuery(queryAddr) err = c1.StartQuery(nil, queryAddr)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
n, err = c1.Encapsulate(buf[:], 0) // Send Request. n, err = c1.Encapsulate(buf[:], -1, 0) // Send Request.
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
@@ -66,14 +66,14 @@ func TestHandler(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
n, err = c2.Encapsulate(buf[:], 0) // Send response. n, err = c2.Encapsulate(buf[:], -1, 0) // Send response.
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
t.Fatal("got no response to request") t.Fatal("got no response to request")
} }
validateARP(t, buf[:]) 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 { if err != nil {
t.Fatal("double tap send error:", err) t.Fatal("double tap send error:", err)
} else if n > 0 { } else if n > 0 {
@@ -90,13 +90,13 @@ func TestHandler(t *testing.T) {
} else if !bytes.Equal(hwaddr, expectHWAddr) { } else if !bytes.Equal(hwaddr, expectHWAddr) {
log.Fatalf("expected to get hwaddr %x!=%x", 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 { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n > 0 { } else if n > 0 {
t.Fatal("expected no data") t.Fatal("expected no data")
} }
n, err = c2.Encapsulate(buf[:], 0) n, err = c2.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n > 0 { } else if n > 0 {
+84 -15
View File
@@ -3,9 +3,11 @@ package arp
import ( import (
"bytes" "bytes"
"errors" "errors"
"log/slog"
"github.com/soypat/lneto" "github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
) )
type Handler struct { 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) 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 { func (h *Handler) Reset(cfg HandlerConfig) error {
if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 || if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 ||
len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 { len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 {
@@ -63,9 +73,22 @@ func (h *Handler) Reset(cfg HandlerConfig) error {
type queryResult struct { type queryResult struct {
protoaddr []byte protoaddr []byte
hwaddr []byte hwaddr []byte
dstHw []byte
querysent bool 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. // AbortPending drops pending queries and incoming requests.
func (h *Handler) AbortPending() { func (h *Handler) AbortPending() {
h.pendingResponse = h.pendingResponse[:0] 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 bytes.Equal(protoAddr, h.queries[i].protoaddr) {
if !h.queries[i].querysent { if !h.queries[i].querysent {
return nil, errors.New("query not yet sent") 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 nil, errors.New("no response yet")
} }
return h.queries[i].hwaddr, nil return mac, nil
} }
} }
return nil, errors.New("query not exist or dropped") 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) { if len(proto) != len(h.ourProtoAddr) {
return errors.New("bad protocol address length") return errors.New("bad protocol address length")
} else if len(h.queries) == cap(h.queries) { } else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) {
return errors.New("too many ongoing queries") 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] h.queries = h.queries[:len(h.queries)+1]
q := &h.queries[len(h.queries)-1] q := &h.queries[len(h.queries)-1]
q.hwaddr = q.hwaddr[:0] *q = queryResult{
q.querysent = false protoaddr: append(q.protoaddr[:0], proto...),
q.protoaddr = append(q.protoaddr[:0], proto...) hwaddr: q.hwaddr[:0],
dstHw: dstHWAddr,
}
return nil return nil
} }
func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { func (h *Handler) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
b := eth[frameOffset:] b := carrierData[offsetToFrame:]
n := h.expectSize() n := h.expectSize()
if len(b) < n { if len(b) < n {
return 0, errShortARP return 0, errShortARP
@@ -120,7 +181,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) {
copy(hwsender, h.ourHWAddr) copy(hwsender, h.ourHWAddr)
n := copy(b, afrm.Clip().RawData()) n := copy(b, afrm.Clip().RawData())
tgt, _ := afrm.Target() tgt, _ := afrm.Target()
trySetEthernetDst(eth[:frameOffset], tgt) trySetEthernetDst(carrierData[:offsetToFrame], tgt)
return n, nil return n, nil
} }
for i := range h.queries { for i := range h.queries {
@@ -139,7 +200,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) {
hwTarget[j] = 0 hwTarget[j] = 0
} }
broadcast := ethernet.BroadcastAddr() broadcast := ethernet.BroadcastAddr()
trySetEthernetDst(eth[:frameOffset], broadcast[:]) trySetEthernetDst(carrierData[:offsetToFrame], broadcast[:])
return n, nil return n, nil
} }
} }
@@ -181,8 +242,16 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
case OpReply: case OpReply:
hwaddr, protoaddr := afrm.Sender() hwaddr, protoaddr := afrm.Sender()
for i := range h.queries { for i := range h.queries {
if len(h.queries[i].hwaddr) == 0 && bytes.Equal(h.queries[i].protoaddr, protoaddr) { q := &h.queries[i]
h.queries[i].hwaddr = append(h.queries[i].hwaddr[:0], hwaddr...) 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 return nil
} }
} }
@@ -194,7 +263,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
} }
func trySetEthernetDst(ethFrame []byte, dst []byte) { func trySetEthernetDst(ethFrame []byte, dst []byte) {
if len(ethFrame) > 14 { if len(ethFrame) >= 14 {
copy(ethFrame[:6], dst) copy(ethFrame[:6], dst)
} }
} }
+7 -7
View File
@@ -104,11 +104,11 @@ func (c *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) }
func (c *Client) LocalPort() uint16 { return DefaultClientPort } func (c *Client) LocalPort() uint16 { return DefaultClientPort }
func (c *Client) ConnectionID() *uint64 { return &c.connID } func (c *Client) ConnectionID() *uint64 { return &c.connID }
func (c *Client) setIP(b []byte, frameOffset int) { func (c *Client) setIP(carrierFrame []byte, offsetToIP int) {
if frameOffset < 28 { if offsetToIP < 0 {
return // Not an IP/UDP frame. 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)) ifrm.SetID((uint16(c.currentXID) ^ uint16(c.currentXID>>16)) + uint16(c.state))
if c.state > StateInit { if c.state > StateInit {
// Match server ToS since some routers drop DHCP requests if no ToS set apparently? // 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() { if c.isClosed() {
return 0, net.ErrClosed return 0, net.ErrClosed
} else if c.state == StateSelecting && !c.offer.valid { } 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 { } else if c.state == StateRequesting {
return 0, nil // Currently awaiting ACK. return 0, nil // Currently awaiting ACK.
} }
dst := carrierFrame[frameOffset:] dst := carrierData[offsetToFrame:]
frm, err := NewFrame(dst) frm, err := NewFrame(dst)
if err != nil { if err != nil {
return 0, err return 0, err
@@ -194,7 +194,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
opts[numOpts] = byte(OptEnd) opts[numOpts] = byte(OptEnd)
numOpts++ numOpts++
c.setHeader(frm) c.setHeader(frm)
c.setIP(carrierFrame, frameOffset) c.setIP(carrierData, offsetToIP)
c.state = nextState c.state = nextState
return OptionsOffset + numOpts, nil return OptionsOffset + numOpts, nil
} }
+8 -8
View File
@@ -28,7 +28,7 @@ func TestClientServer(t *testing.T) {
// CLIENT DISCOVER. // CLIENT DISCOVER.
assertClState(StateInit) assertClState(StateInit)
var buf [1024]byte var buf [1024]byte
n, err := cl.Encapsulate(buf[:], 0) n, err := cl.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
@@ -40,7 +40,7 @@ func TestClientServer(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
// SERVER REPLY OFFER // SERVER REPLY OFFER
n, err = sv.Encapsulate(buf[:], 0) n, err = sv.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
@@ -53,7 +53,7 @@ func TestClientServer(t *testing.T) {
assertClState(StateSelecting) assertClState(StateSelecting)
// CLIENT SEND OUT REQUEST. // CLIENT SEND OUT REQUEST.
n, err = cl.Encapsulate(buf[:], 0) n, err = cl.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
@@ -66,7 +66,7 @@ func TestClientServer(t *testing.T) {
} }
// SERVER REPLIES WITH ACK. // SERVER REPLIES WITH ACK.
n, err = sv.Encapsulate(buf[:], 0) n, err = sv.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } else if n == 0 {
@@ -99,13 +99,13 @@ func TestExample(t *testing.T) {
}) })
buf := make([]byte, 2048) buf := make([]byte, 2048)
buf2 := make([]byte, len(buf)) buf2 := make([]byte, len(buf))
n, err := cl.Encapsulate(buf, 0) n, err := cl.Encapsulate(buf, -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n <= 0 { } else if n <= 0 {
t.Fatal("no data sent out by client after starting request") 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 { if err != nil {
t.Error("client encaps double tap after discover:", err) t.Error("client encaps double tap after discover:", err)
} }
@@ -141,13 +141,13 @@ func TestExample(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
n, err = cl.Encapsulate(buf[:], 0) n, err = cl.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n <= 0 { } else if n <= 0 {
t.Fatal("no data written from client in response to offer") 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 { if err != nil {
t.Error("encapsulate double tap after request:", err) t.Error("encapsulate double tap after request:", err)
} else if n > 0 { } else if n > 0 {
+4 -4
View File
@@ -159,9 +159,9 @@ func (sv *Server) Demux(carrierData []byte, frameOffset int) error {
return nil return nil
} }
func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) { func (sv *Server) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
carrierIsIP := frameOffset >= 28 carrierIsIP := offsetToIP >= 0
dfrm, err := NewFrame(carrierData[frameOffset:]) dfrm, err := NewFrame(carrierData[offsetToFrame:])
optBuf := dfrm.OptionsPayload()[:] optBuf := dfrm.OptionsPayload()[:]
if err != nil { if err != nil {
return 0, err return 0, err
@@ -220,7 +220,7 @@ func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error)
copy(dfrm.CHAddrAs6()[:], client.hwaddr[:]) copy(dfrm.CHAddrAs6()[:], client.hwaddr[:])
dfrm.SetMagicCookie(MagicCookie) dfrm.SetMagicCookie(MagicCookie)
if carrierIsIP { if carrierIsIP {
err = internal.SetIPAddrs(carrierData, 0, sv.siaddr[:], client.addr[:]) err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, sv.siaddr[:], client.addr[:])
if err != nil { if err != nil {
return 0, err return 0, err
} }
+2 -2
View File
@@ -43,7 +43,7 @@ func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error {
return nil 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() { if c.isClosed() {
return 0, net.ErrClosed return 0, net.ErrClosed
} else if c.state != dnsSendQuery { } else if c.state != dnsSendQuery {
@@ -51,7 +51,7 @@ func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) {
} }
msg := &c.msg msg := &c.msg
frame := carrierData[frameOffset:] frame := carrierData[offsetToFrame:]
msglen := msg.Len() msglen := msg.Len()
if msglen > uint16(len(frame)) { if msglen > uint16(len(frame)) {
return 0, errCalcLen return 0, errCalcLen
+5 -5
View File
@@ -199,7 +199,7 @@ func run() (err error) {
prevState = state prevState = state
clear(buf) clear(buf)
nwrite, err := stack.Encapsulate(buf[:], 0) nwrite, err := stack.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
fmt.Println("ERR:ENCAPSULATE", err) fmt.Println("ERR:ENCAPSULATE", err)
} else if nwrite > 0 { } else if nwrite > 0 {
@@ -267,10 +267,10 @@ func (s *Stack) Demux(b []byte, _ int) (err error) {
return s.link.Demux(b, 0) return s.link.Demux(b, 0)
} }
func (s *Stack) Encapsulate(b []byte, _ int) (int, error) { func (s *Stack) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
n, err := s.link.Encapsulate(b, 0) n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame)
if n > 0 { 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 { if errpcap != nil {
fmt.Println("OU", iframes, errpcap.Error()) fmt.Println("OU", iframes, errpcap.Error())
} else { } else {
@@ -426,7 +426,7 @@ func (s *Stack) StartResolveHardwareAddress6(ip netip.Addr) error {
return errors.New("unsupported or invalid IP address") return errors.New("unsupported or invalid IP address")
} }
addr := ip.As4() 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) { func (s *Stack) ResultResolveHardwareAddress6(ip netip.Addr) (hw [6]byte, err error) {
+4 -4
View File
@@ -107,7 +107,7 @@ func main() {
} }
} }
nw, err := stack.ethernet.Encapsulate(buf[:], 0) nw, err := stack.ethernet.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
lg.Error("handle", slog.String("err", err.Error())) lg.Error("handle", slog.String("err", err.Error()))
} else if nw > 0 { } else if nw > 0 {
@@ -230,7 +230,7 @@ func (stack *Stack) Recv(b []byte) error {
} }
func (stack *Stack) Send(b []byte) (int, 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) { 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 { if err != nil {
return nil, err return nil, err
} }
err = stack.tcpports.Register(&listener) err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting.
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error
if err != nil { if err != nil {
return nil, err return nil, err
} }
err = stack.tcpports.Register(conn) err = stack.tcpports.Register(conn) // Passive MAC with no listening.
if err != nil { if err != nil {
return nil, err return nil, err
} }
+4 -4
View File
@@ -232,7 +232,7 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s
type handler struct { type handler struct {
raddr []byte raddr []byte
recv func([]byte, int) error recv func([]byte, int) error
handle func([]byte, int) (int, error) handle func([]byte, int, int) (int, error)
proto ethernet.Type proto ethernet.Type
lport uint16 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. copy(efrm.DestinationHardwareAddr()[:], ls.gwmac[:]) // default set the gateway.
for i := range ls.handlers { for i := range ls.handlers {
h := &ls.handlers[i] h := &ls.handlers[i]
n, err = h.handle(dst[:mtu], 14) n, err = h.handle(dst[:mtu], 14, 14)
if err != nil { if err != nil {
ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error()))
continue continue
@@ -322,8 +322,8 @@ func (as *ARPStack) Recv(EtherFrame []byte, arpOff int) error {
return as.handler.Demux(EtherFrame, arpOff) return as.handler.Demux(EtherFrame, arpOff)
} }
func (as *ARPStack) Handle(EtherFrame []byte, arpOff int) (int, error) { func (as *ARPStack) Handle(EtherFrame []byte, offsetToIP, arpOff int) (int, error) {
n, err := as.handler.Encapsulate(EtherFrame, arpOff) n, err := as.handler.Encapsulate(EtherFrame, offsetToIP, arpOff)
if err != nil || n == 0 { if err != nil || n == 0 {
return 0, err return 0, err
} }
+1 -1
View File
@@ -110,7 +110,7 @@ func run() (err error) {
var frames []pcap.Frame var frames []pcap.Frame
for { for {
clear(buf) clear(buf)
nwrite, err := stack.Encapsulate(buf[:], 0) nwrite, err := stack.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
fmt.Println("ERR:ENCAPSULATE", err) fmt.Println("ERR:ENCAPSULATE", err)
} else if nwrite > 0 { } else if nwrite > 0 {
+11
View File
@@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) {
copy(dstaddr, dst) copy(dstaddr, dst)
return nil 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
}
+159 -86
View File
@@ -2,25 +2,30 @@ package internet
import ( import (
"errors" "errors"
"log/slog"
"math" "math"
"net" "net"
"slices"
) )
// StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols, // 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. // from Ethernet to IP to TCP, practically any protocol can be expressed as a StackNode and function completely.
type StackNode interface { 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. // 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] // The returned integer is amount of bytes written such that carrierData[offsetToFrame:offsetToFrame+n]
// contains written data. Data inside carrierData[:frameOffset] usually contains data necessary for // 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 // 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 // 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 // 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]). // 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. // 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. // 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. // 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). // The stack node then dispatches(demuxes) the encapsulated frames to its corresponding sub-stack-node(s).
Demux(carrierData []byte, frameOffset int) error Demux(carrierData []byte, frameOffset int) error
@@ -36,9 +41,156 @@ type node struct {
currConnID uint64 currConnID uint64
connID *uint64 connID *uint64
demux func([]byte, int) error demux func([]byte, int) error
encapsulate func([]byte, int) (int, error) encapsulate func([]byte, int, int) (int, error)
proto uint16 proto uint16
port 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 ( var (
@@ -50,32 +202,6 @@ var (
_ = net.ErrClosed _ = 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 { func (node *node) IsInvalid() bool {
return node.demux == nil || node.encapsulate == nil || (node.connID != nil && node.currConnID != *node.connID) 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) 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 { if protocol > math.MaxUint16 {
panic(">16bit protocol number unsupported") panic(">16bit protocol number unsupported")
} }
@@ -100,64 +226,11 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node {
encapsulate: s.Encapsulate, encapsulate: s.Encapsulate,
proto: uint16(protocol), proto: uint16(protocol),
port: port, 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. // destroy removes all references to underlying StackNode. Allows garbage collection of node if possible.
func (n *node) destroy() { func (n *node) destroy() {
*n = node{} *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]
}
+2 -2
View File
@@ -94,7 +94,7 @@ func (listener *NodeTCPListener) TryAccept() (*tcp.Conn, error) {
} }
// Encapsulate implements [StackNode]. // 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() { if listener.isClosed() {
return 0, net.ErrClosed return 0, net.ErrClosed
} }
@@ -102,7 +102,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset
if conn == nil { if conn == nil {
continue continue
} }
n, err := conn.Encapsulate(carrierData, tcpFrameOffset) n, err := conn.Encapsulate(carrierData, offsetToIP, offsetToFrame)
if err != nil { if err != nil {
err = listener.maintainConn(listener.accepted, i, err) err = listener.maintainConn(listener.accepted, i, err)
} }
+19 -37
View File
@@ -6,7 +6,6 @@ import (
"log/slog" "log/slog"
"math" "math"
"net" "net"
"slices"
"github.com/soypat/lneto" "github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/ethernet"
@@ -14,8 +13,7 @@ import (
type StackEthernet struct { type StackEthernet struct {
connID uint64 connID uint64
handlers []node handlers handlers
logger
mac [6]byte mac [6]byte
gwmac [6]byte gwmac [6]byte
mtu uint16 mtu uint16
@@ -43,11 +41,10 @@ func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error {
} else if maxNodes <= 0 { } else if maxNodes <= 0 {
return errZeroMaxNodesArg return errZeroMaxNodesArg
} }
ls.handlers = slices.Grow(ls.handlers[:0], maxNodes) ls.handlers.reset("StackEthernet", maxNodes)
*ls = StackEthernet{ *ls = StackEthernet{
connID: ls.connID + 1, connID: ls.connID + 1,
handlers: ls.handlers, handlers: ls.handlers,
logger: ls.logger,
mac: mac, mac: mac,
gwmac: gateway, gwmac: gateway,
mtu: uint16(mtu), mtu: uint16(mtu),
@@ -68,18 +65,7 @@ func (ls *StackEthernet) Register(h StackNode) error {
if proto > math.MaxUint16 || proto <= 1500 { if proto > math.MaxUint16 || proto <= 1500 {
return errInvalidProto return errInvalidProto
} }
eproto := uint16(proto) return ls.handlers.registerByProto(nodeFromStackNode(h, 0, proto, nil))
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,
})
} }
func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) { 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() { if vld.HasError() {
return vld.ErrPop() return vld.ErrPop()
} }
if h, err := ls.handlers.demuxByProto(efrm.Payload(), 0, uint16(etype)); h != nil {
for i := range ls.handlers { return err
h := &ls.handlers[i]
if h.proto == uint16(etype) {
return h.demux(efrm.Payload(), 0)
}
} }
DROP: 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 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 mtu := ls.mtu
dst := carrierData[frameOffset:] dst := carrierData[offsetToFrame:]
if len(dst) < int(mtu) { if len(dst) < int(mtu) {
return 0, io.ErrShortBuffer return 0, io.ErrShortBuffer
} }
@@ -121,19 +103,19 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int
return 0, err return 0, err
} }
*efrm.DestinationHardwareAddr() = ls.gwmac *efrm.DestinationHardwareAddr() = ls.gwmac
for i := range ls.handlers { var h *node
h := &ls.handlers[i] // Children (IP/ARP) start at offset 14 (after ethernet header).
n, err = h.encapsulate(dst[:mtu], 14) // For IP: offsetToIP=14, offsetToFrame=14
if err != nil { // For ARP: offsetToIP=-1, offsetToFrame=14 (but ARP ignores offsetToIP)
ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) // Clip carrierData to MTU to prevent writes beyond MTU limit.
continue mtuLimit := offsetToFrame + int(mtu)
h, n, err = ls.handlers.encapsulateAny(carrierData[:mtuLimit], offsetToFrame+14, offsetToFrame+14)
if n == 0 {
return n, err
} }
if n > 0 {
// Found packet // Found packet
*efrm.SourceHardwareAddr() = ls.mac *efrm.SourceHardwareAddr() = ls.mac
efrm.SetEtherType(ethernet.Type(h.proto)) efrm.SetEtherType(ethernet.Type(h.proto))
return n + 14, nil n += 14
} return n, err
}
return 0, err
} }
+24 -49
View File
@@ -5,7 +5,6 @@ import (
"io" "io"
"log/slog" "log/slog"
"net/netip" "net/netip"
"slices"
"github.com/soypat/lneto" "github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/ethernet"
@@ -23,9 +22,7 @@ type StackIP struct {
ipID uint16 ipID uint16
ip [4]byte ip [4]byte
validator lneto.Validator validator lneto.Validator
handlers []node handlers handlers
pendingICMP [][]byte
logger
} }
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { 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 { if err != nil {
return err return err
} }
sb.handlers = slices.Grow(sb.handlers[:0], maxNodes) sb.handlers.reset("StackIP", maxNodes)
*sb = StackIP{ *sb = StackIP{
connID: sb.connID + 1, connID: sb.connID + 1,
validator: sb.validator, validator: sb.validator,
handlers: sb.handlers, handlers: sb.handlers,
logger: sb.logger,
ip: sb.ip, ip: sb.ip,
pendingICMP: make([][]byte, maxNodes*4),
} }
return nil return nil
} }
@@ -73,11 +68,11 @@ func (sb *StackIP) Addr() netip.Addr {
} }
func (sb *StackIP) SetLogger(logger *slog.Logger) { func (sb *StackIP) SetLogger(logger *slog.Logger) {
sb.logger.log = logger sb.handlers.log = logger
} }
func (sb *StackIP) Demux(carrierData []byte, offset int) error { 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. frame := carrierData[offset:] // we don't care about carrier data in IP.
ifrm, err := ipv4.NewFrame(frame) ifrm, err := ipv4.NewFrame(frame)
if err != nil { if err != nil {
@@ -96,7 +91,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
gotCRC := ifrm.CRC() gotCRC := ifrm.CRC()
wantCRC := ifrm.CalculateHeaderCRC() wantCRC := ifrm.CalculateHeaderCRC()
if gotCRC != wantCRC { 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") return errors.New("IPv4 CRC mismatch")
} }
off := ifrm.HeaderLength() off := ifrm.HeaderLength()
@@ -105,10 +100,11 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
if proto == lneto.IPProtoICMP { if proto == lneto.IPProtoICMP {
return sb.recvicmp(ifrm.RawData(), ifrm.HeaderLength()) return sb.recvicmp(ifrm.RawData(), ifrm.HeaderLength())
} }
nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) node := sb.handlers.nodeByProto(uint16(proto))
if nodeIdx < 0 { // nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if node == nil {
// Drop packet. // 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 return nil
} }
// Incoming CRC Validation of common IP Protocols. // 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") return errors.New("UDP CRC mismatch")
} }
} }
sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) sb.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = sb.handlers[nodeIdx].demux(frame[:totalLen], off) err = node.demux(frame[:totalLen], off)
if handleNodeError(&sb.handlers, nodeIdx, err) { if sb.handlers.tryHandleError(node, err) {
sb.info("ipclose", slog.String("proto", proto.String())) sb.handlers.info("ipclose", slog.String("proto", proto.String()))
err = nil err = nil
} }
return err return err
} }
func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) { func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
frame := carrierData[frameOffset:] frame := carrierData[offsetToFrame:]
if len(frame) < 256 { if len(frame) < 256 {
return 0, io.ErrShortBuffer return 0, io.ErrShortBuffer
} }
@@ -162,20 +158,13 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
ifrm.SetTTL(64) ifrm.SetTTL(64)
*ifrm.SourceAddr() = sb.ip *ifrm.SourceAddr() = sb.ip
sb.ipID = id sb.ipID = id
for i := range sb.handlers { // Children (TCP/UDP) start at offset headerlen (20 bytes after IP header start).
h := &sb.handlers[i] // offsetToIP is 0 relative to this slice (frame), children's frame starts at headerlen.
proto := lneto.IPProto(h.proto) node, n, err := sb.handlers.encapsulateAny(carrierData, offsetToFrame, offsetToFrame+headerlen)
n, err := h.encapsulate(frame[:], headerlen) if n == 0 {
if err != nil { return n, err
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
} }
proto := lneto.IPProto(node.proto)
totalLen := n + headerlen totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen)) ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto) ifrm.SetProtocol(proto)
@@ -195,13 +184,11 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
ufrm.CRCWriteIPv4(&crc) ufrm.CRCWriteIPv4(&crc)
ufrm.SetCRC(crc.Sum16()) ufrm.SetCRC(crc.Sum16())
if n != int(ufrm.Length()) { if n != int(ufrm.Length()) {
sb.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", 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 0, errors.New("invalid UDP length")
} }
} }
return totalLen, nil return totalLen, err
}
return 0, nil
} }
func (sb *StackIP) Register(h StackNode) error { func (sb *StackIP) Register(h StackNode) error {
@@ -209,19 +196,7 @@ func (sb *StackIP) Register(h StackNode) error {
if proto > 255 { if proto > 255 {
return errInvalidProto return errInvalidProto
} }
connID := h.ConnectionID() return sb.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil))
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,
})
} }
func (sb *StackIP) recvicmp(carrierData []byte, offset int) error { func (sb *StackIP) recvicmp(carrierData []byte, offset int) error {
+84 -54
View File
@@ -2,18 +2,23 @@ package internet
import ( import (
"encoding/binary" "encoding/binary"
"errors"
"io" "io"
"log/slog"
"math" "math"
"slices" "strconv"
"github.com/soypat/lneto" "github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
) )
type StackPorts struct { type StackPorts struct {
connID uint64 connID uint64
handlers []node handlers handlers
dstPortOff uint16 dstPortOff uint16
protocol uint16 protocol uint16
// stores last node to demux/encapsulate.
} }
func (ps *StackPorts) ResetUDP(maxNodes int) error { 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 { } else if maxNodes <= 0 {
return errZeroMaxNodesArg return errZeroMaxNodesArg
} }
ps.handlers = slices.Grow(ps.handlers[:0], maxNodes) ps.handlers.reset("StackPorts(proto="+strconv.Itoa(int(protocol))+")", maxNodes)
*ps = StackPorts{ *ps = StackPorts{
connID: ps.connID + 1, connID: ps.connID + 1,
handlers: ps.handlers, 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) ConnectionID() *uint64 { return &ps.connID }
func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) { func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
if int(ps.dstPortOff)+offset+2 > len(b) { if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) {
return 0, io.ErrShortBuffer return 0, io.ErrShortBuffer
} }
var i int _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame)
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
}
}
return n, err return n, err
} }
@@ -72,52 +64,90 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) {
return io.ErrShortBuffer return io.ErrShortBuffer
} }
port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:])
var i int _, err = ps.handlers.demuxByPort(b, offset, port)
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)
return err 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 { func (ps *StackPorts) Register(h StackNode) error {
port := h.LocalPort() port := h.LocalPort()
proto := h.Protocol() proto := h.Protocol()
if port <= 0 { if port <= 0 {
return errZeroPort return errZeroPort
} else if proto != uint64(ps.protocol) { } else if proto != uint64(ps.protocol) {
return errInvalidProto return errInvalidProto
} }
var cid uint64 return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil))
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),
})
} }
func (ps *StackPorts) handleResult(handlerIdx, n int, err error) (discarded bool) { // StackPortsMACFiltered is a StackPorts implementation but that avoids calling encapsulate on nodes
if handleNodeError(&ps.handlers, handlerIdx, err) { // with a non-nil MAC address registered via Register method that is set to all zero values.
discarded = true // If the address is set to nil no filtering occurs. MAC Address is set automatically on the ethernet frame by StackPortsMACFiltered when non-nil.
println("DISCARD", handlerIdx, "witherr", err.Error()) type StackPortsMACFiltered struct {
} sp StackPorts
return discarded }
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.
} }
+7 -6
View File
@@ -17,7 +17,7 @@ type StackUDPPort struct {
} }
func (sudp *StackUDPPort) SetStackNode(node StackNode, raddr []byte, rmport uint16) { 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.rmport = rmport
sudp.raddr = append(sudp.raddr[:0], raddr...) sudp.raddr = append(sudp.raddr[:0], raddr...)
} }
@@ -61,24 +61,25 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error {
return err 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() { if sudp.h.IsInvalid() {
sudp.h.destroy() sudp.h.destroy()
return 0, net.ErrClosed return 0, net.ErrClosed
} }
ufrm, err := udp.NewFrame(carrierData[frameOffset:]) ufrm, err := udp.NewFrame(carrierData[offsetToFrame:])
if err != nil { if err != nil {
return 0, err return 0, err
} }
ufrm.SetSourcePort(sudp.h.port) ufrm.SetSourcePort(sudp.h.port)
ufrm.SetDestinationPort(sudp.rmport) ufrm.SetDestinationPort(sudp.rmport)
if len(sudp.raddr) > 0 && frameOffset >= 20 { if len(sudp.raddr) > 0 && offsetToIP >= 0 {
err = internal.SetIPAddrs(carrierData, 0, nil, sudp.raddr) err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, nil, sudp.raddr)
if err != nil { if err != nil {
return 0, err 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 n == 0 {
if err != nil { if err != nil {
slog.Error("stackudp:encapsulate", slog.String("err", err.Error())) slog.Error("stackudp:encapsulate", slog.String("err", err.Error()))
+1 -1
View File
@@ -44,7 +44,7 @@ func TestBasicStack2(t *testing.T) {
func expectExchange(t *testing.T, from, to *StackIP, buf []byte) { func expectExchange(t *testing.T, from, to *StackIP, buf []byte) {
t.Helper() t.Helper()
n, err := from.Encapsulate(buf, 0) n, err := from.Encapsulate(buf, -1, 0)
if err != nil { if err != nil {
t.Error("expectExchange:encapsulate:", err) t.Error("expectExchange:encapsulate:", err)
} else if n == 0 { } else if n == 0 {
+2 -2
View File
@@ -50,11 +50,11 @@ func (c *Client) ConnectionID() *uint64 {
return &c.connID 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() { if c.IsDone() {
return 0, nil return 0, nil
} }
payload := carrierData[frameOffset:] payload := carrierData[offsetToFrame:]
frm, err := NewFrame(payload) frm, err := NewFrame(payload)
if err != nil { if err != nil {
return 0, err return 0, err
+22 -8
View File
@@ -246,7 +246,7 @@ func (conn *Conn) checkPipeOpen() error {
if conn.abortErr != nil { if conn.abortErr != nil {
return conn.abortErr return conn.abortErr
} }
state := conn.State() state := conn.h.State()
if state.IsClosed() { if state.IsClosed() {
return net.ErrClosed return net.ErrClosed
} }
@@ -278,23 +278,27 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) {
return nil 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() conn.mu.Lock()
defer conn.mu.Unlock() defer conn.mu.Unlock()
if len(conn.remoteAddr) == 0 { if len(conn.remoteAddr) == 0 {
return 0, errNoRemoteAddr 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 { if err != nil {
return 0, err return 0, err
} else if len(raddr) != len(conn.remoteAddr) { } else if len(raddr) != len(conn.remoteAddr) {
return 0, errMismatchedIPVersion return 0, errMismatchedIPVersion
} }
n, err = conn.h.Send(buf[off:]) n, err = conn.h.Send(carrierData[offsetToFrame:])
if err != nil { if err != nil {
return 0, err 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 { if err != nil {
return 0, err return 0, err
} }
@@ -328,11 +332,11 @@ func (conn *Conn) reset(h Handler) {
func (conn *Conn) SetDeadline(t time.Time) error { func (conn *Conn) SetDeadline(t time.Time) error {
conn.mu.Lock() conn.mu.Lock()
defer conn.mu.Unlock() defer conn.mu.Unlock()
err := conn.SetReadDeadline(t) err := conn.setReadDeadline(t)
if err != nil { if err != nil {
return err return err
} }
return conn.SetWriteDeadline(t) return conn.setWriteDeadline(t)
} }
// SetReadDeadline sets the deadline for future Read calls // 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 { func (conn *Conn) SetReadDeadline(t time.Time) error {
conn.mu.Lock() conn.mu.Lock()
defer conn.mu.Unlock() 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() err := conn.checkPipeOpen()
if err == nil { if err == nil {
conn.rdead = t conn.rdead = t
@@ -354,6 +362,12 @@ func (conn *Conn) SetReadDeadline(t time.Time) error {
// some of the data was successfully written. // some of the data was successfully written.
// A zero value for t means Write will not time out. // A zero value for t means Write will not time out.
func (conn *Conn) SetWriteDeadline(t time.Time) error { 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") conn.trace("TCPConn.SetWriteDeadline:start")
err := conn.checkPipeOpen() err := conn.checkPipeOpen()
if err == nil { if err == nil {
+42 -18
View File
@@ -28,11 +28,12 @@ type StackAsync struct {
ip internet.StackIP ip internet.StackIP
arp arp.Handler arp arp.Handler
udps internet.StackPorts udps internet.StackPorts
tcps internet.StackPorts tcps internet.StackPortsMACFiltered
dhcpUDP internet.StackUDPPort dhcpUDP internet.StackUDPPort
dhcp dhcpv4.Client dhcp dhcpv4.Client
dhcpResults DHCPResults dhcpResults DHCPResults
subnet netip.Prefix // Local subnet for ARP resolution.
dnsUDP internet.StackUDPPort dnsUDP internet.StackUDPPort
dns dns.Client dns dns.Client
@@ -73,11 +74,11 @@ func (s *StackAsync) Demux(carrierData []byte, etherOff int) error {
return s.link.Demux(carrierData, etherOff) 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() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
n, err := s.link.Encapsulate(carrierData, etherOff) n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame)
s.totalsent += uint64(n) s.totalsent += uint64(n)
return n, err return n, err
} }
@@ -112,6 +113,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error {
if err != nil { if err != nil {
return err return err
} }
//
err = s.resetARP() err = s.resetARP()
if err != nil { if err != nil {
return err return err
@@ -135,10 +137,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error {
} }
// Now setup stacks. // Now setup stacks.
err = s.link.Register(&s.arp) // ARP. // ARP registered in resetARP.
if err != nil {
return err
}
err = s.link.Register(&s.ip) // IPv4 | IPv6 err = s.link.Register(&s.ip) // IPv4 | IPv6
if err != nil { if err != nil {
return err return err
@@ -169,7 +168,7 @@ func (s *StackAsync) resetARP() error {
if addr.Is6() { if addr.Is6() {
proto = ethernet.TypeIPv6 proto = ethernet.TypeIPv6
} }
return s.arp.Reset(arp.HandlerConfig{ err := s.arp.Reset(arp.HandlerConfig{
HardwareAddr: mac[:], HardwareAddr: mac[:],
ProtocolAddr: addr.AsSlice(), ProtocolAddr: addr.AsSlice(),
MaxQueries: 3, MaxQueries: 3,
@@ -177,6 +176,14 @@ func (s *StackAsync) resetARP() error {
HardwareType: 1, HardwareType: 1,
ProtocolType: proto, 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. // 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 { func (s *StackAsync) SetIPAddr(addr netip.Addr) error {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
return s.setIPAddr(addr)
}
func (s *StackAsync) setIPAddr(addr netip.Addr) error {
err := s.ip.SetAddr(addr) err := s.ip.SetAddr(addr)
if err != nil { if err != nil {
return err return err
} }
return s.resetARP() ip := addr.As4()
err = s.arp.UpdateProtoAddr(ip[:])
return err
} }
func (s *StackAsync) Addr() netip.Addr { 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) { func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort) (err error) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() 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())) err = conn.OpenActive(localPort, addrp, tcp.Value(s.Prand32()))
if err != nil { if err != nil {
return err 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 { if err != nil {
conn.Abort() conn.Abort()
return err return err
@@ -253,7 +278,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) {
if err != nil { if err != nil {
return err return err
} }
err = s.tcps.Register(conn) err = s.tcps.Register(conn, nil)
if err != nil { if err != nil {
conn.Abort() conn.Abort()
return err return err
@@ -376,7 +401,7 @@ func (s *StackAsync) StartResolveHardwareAddress6(ip netip.Addr) error {
return errors.New("unsupported or invalid IP address") return errors.New("unsupported or invalid IP address")
} }
addr := ip.As4() addr := ip.As4()
return s.arp.StartQuery(addr[:]) return s.arp.StartQuery(nil, addr[:])
} }
// ResultResolveHardwareAddress6 // ResultResolveHardwareAddress6
@@ -432,16 +457,15 @@ func (s *StackAsync) ReadStatistics(stats *Statistics) {
// AssimilateDHCPResults sets the stack's following parameters: // AssimilateDHCPResults sets the stack's following parameters:
// - IPv4 address. // - IPv4 address.
// - DNS server. // - DNS server.
// - Subnet (for ARP resolution of local addresses).
func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error { func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error {
stack.mu.Lock() stack.mu.Lock()
defer stack.mu.Unlock() defer stack.mu.Unlock()
if results.AssignedAddr.IsValid() { if results.Subnet.IsValid() {
err := stack.ip.SetAddr(results.AssignedAddr) stack.subnet = results.Subnet
if err != nil {
return err
} }
// Reset ARP handler with new IP address so it can respond to ARP requests. if results.AssignedAddr.IsValid() {
err = stack.resetARP() err := stack.setIPAddr(results.AssignedAddr)
if err != nil { if err != nil {
return err return err
} }
+49
View File
@@ -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)
}
}
+159 -9
View File
@@ -2,11 +2,13 @@ package xnet
import ( import (
"bytes" "bytes"
"encoding/binary"
"errors" "errors"
"math/rand" "math/rand"
"net/netip" "net/netip"
"testing" "testing"
"github.com/soypat/lneto/arp"
"github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internet/pcap" "github.com/soypat/lneto/internet/pcap"
"github.com/soypat/lneto/tcp" "github.com/soypat/lneto/tcp"
@@ -24,9 +26,7 @@ func TestStackAsyncTCP_multipacket(t *testing.T) {
const svPort = 8080 const svPort = 8080
const maxPktLen = 30 const maxPktLen = 30
client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) client, sv, clconn, svconn := newTCPStacks(t, seed, MTU)
tst := tester{ tst := testerFrom(t, MTU)
t: t, buf: make([]byte, MTU),
}
rng := rand.New(rand.NewSource(seed)) rng := rand.New(rand.NewSource(seed))
client2, sv2, clconn2, svconn2 := newTCPStacks(t, seed, MTU) client2, sv2, clconn2, svconn2 := newTCPStacks(t, seed, MTU)
_, _, _, _ = client2, sv2, clconn2, svconn2 _, _, _, _ = client2, sv2, clconn2, svconn2
@@ -58,10 +58,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) {
const MTU = 1500 const MTU = 1500
const svPort = 80 const svPort = 80
client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) client, sv, clconn, svconn := newTCPStacks(t, seed, MTU)
tst := testerFrom(t, MTU)
tst := tester{
t: t, buf: make([]byte, MTU),
}
tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337) tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337)
sendData := []byte("hello") 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) { func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1, c2 *tcp.Conn) {
s1, s2 = new(StackAsync), new(StackAsync) s1, s2 = new(StackAsync), new(StackAsync)
c1, c2 = new(tcp.Conn), new(tcp.Conn) c1, c2 = new(tcp.Conn), new(tcp.Conn)
byte1 := byte(randSeed) / 4 byte1 := byte(randSeed)/4 - 1
err := s1.Reset(StackConfig{ err := s1.Reset(StackConfig{
Hostname: "Stack1", Hostname: "Stack1",
RandSeed: randSeed, RandSeed: randSeed,
@@ -127,6 +124,13 @@ func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1
return s1, s2, c1, c2 return s1, s2, c1, c2
} }
func testerFrom(t *testing.T, mtu int) *tester {
return &tester{
t: t,
buf: make([]byte, mtu),
}
}
type tester struct { type tester struct {
t *testing.T t *testing.T
cap pcap.PacketBreakdown cap pcap.PacketBreakdown
@@ -322,7 +326,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs
default: default:
panic("OOB") panic("OOB")
} }
n, err := src.Encapsulate(buf[:], 0) n, err := src.Encapsulate(buf[:], -1, 0)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} else if n == 0 { } 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") t.Error("expected no data sent and got data")
return return
} }
defer setzero(buf[:n])
tst.buf = tst.buf[:n] tst.buf = tst.buf[:n]
tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) 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 { if err != nil {
t.Fatal(err) 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]) 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 { func (tst *tester) getTCPFrame() tcp.Frame {
@@ -457,3 +576,34 @@ func setzero[T ~[]E, E any](s T) {
s[i] = zero 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))
}