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
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).
+8 -8
View File
@@ -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 {
+84 -15
View File
@@ -3,9 +3,11 @@ package arp
import (
"bytes"
"errors"
"log/slog"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
)
type Handler struct {
@@ -33,6 +35,14 @@ func (h *Handler) Protocol() uint64 { return uint64(ethernet.TypeARP) }
func (h *Handler) ConnectionID() *uint64 { return &h.connID }
func (h *Handler) UpdateProtoAddr(protoAddr []byte) error {
if len(protoAddr) != len(h.ourProtoAddr) {
return errors.New("mismatch ARP proto size")
}
copy(h.ourProtoAddr, protoAddr)
return nil
}
func (h *Handler) Reset(cfg HandlerConfig) error {
if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 ||
len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 {
@@ -63,9 +73,22 @@ func (h *Handler) Reset(cfg HandlerConfig) error {
type queryResult struct {
protoaddr []byte
hwaddr []byte
dstHw []byte
querysent bool
}
func (qr *queryResult) destroy() {
*qr = queryResult{protoaddr: qr.protoaddr[:0], hwaddr: qr.hwaddr[:0]}
}
func (qr *queryResult) response() []byte {
if len(qr.hwaddr) == 0 {
return nil
}
return qr.hwaddr[:]
}
func (qr *queryResult) isInvalid() bool { return len(qr.protoaddr) == 0 }
// AbortPending drops pending queries and incoming requests.
func (h *Handler) AbortPending() {
h.pendingResponse = h.pendingResponse[:0]
@@ -81,31 +104,69 @@ func (h *Handler) QueryResult(protoAddr []byte) (hwAddr []byte, err error) {
if bytes.Equal(protoAddr, h.queries[i].protoaddr) {
if !h.queries[i].querysent {
return nil, errors.New("query not yet sent")
} else if len(h.queries[i].hwaddr) == 0 {
}
mac := h.queries[i].response()
if mac == nil {
return nil, errors.New("no response yet")
}
return h.queries[i].hwaddr, nil
return mac, nil
}
}
return nil, errors.New("query not exist or dropped")
}
func (h *Handler) StartQuery(proto []byte) error {
func (h *Handler) DiscardQuery(protoAddr []byte) error {
for i := range h.queries {
q := &h.queries[i]
if bytes.Equal(protoAddr, q.protoaddr) {
q.destroy()
return nil
}
}
return errors.New("query not found")
}
func (h *Handler) compactQueries() {
validOff := 0
for i := 0; i < len(h.queries); i++ {
if h.queries[i].isInvalid() {
h.queries[validOff] = h.queries[i]
validOff++
}
}
h.queries = h.queries[:validOff]
}
// StartQuery queues a query to perform over ARP for the protocol address `proto`.
// The user can additionally specify an dstHWAddr to write query result to on completion.
// If dstHWAddr is nil then query still occurs but no external buffer is written on query completion.
// dstHWAddr must be zeroed out (invalid MAC).
func (h *Handler) StartQuery(dstHWAddr, proto []byte) error {
if len(h.queries) == cap(h.queries) {
h.compactQueries()
if len(h.queries) == cap(h.queries) {
return errors.New("too many ongoing queries")
}
}
if len(proto) != len(h.ourProtoAddr) {
return errors.New("bad protocol address length")
} else if len(h.queries) == cap(h.queries) {
return errors.New("too many ongoing queries")
} else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) {
return errors.New("mismatch hardware size")
} else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) {
return errors.New("write-to buffer must be zeroed out")
}
h.queries = h.queries[:len(h.queries)+1]
q := &h.queries[len(h.queries)-1]
q.hwaddr = q.hwaddr[:0]
q.querysent = false
q.protoaddr = append(q.protoaddr[:0], proto...)
*q = queryResult{
protoaddr: append(q.protoaddr[:0], proto...),
hwaddr: q.hwaddr[:0],
dstHw: dstHWAddr,
}
return nil
}
func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) {
b := eth[frameOffset:]
func (h *Handler) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
b := carrierData[offsetToFrame:]
n := h.expectSize()
if len(b) < n {
return 0, errShortARP
@@ -120,7 +181,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) {
copy(hwsender, h.ourHWAddr)
n := copy(b, afrm.Clip().RawData())
tgt, _ := afrm.Target()
trySetEthernetDst(eth[:frameOffset], tgt)
trySetEthernetDst(carrierData[:offsetToFrame], tgt)
return n, nil
}
for i := range h.queries {
@@ -139,7 +200,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) {
hwTarget[j] = 0
}
broadcast := ethernet.BroadcastAddr()
trySetEthernetDst(eth[:frameOffset], broadcast[:])
trySetEthernetDst(carrierData[:offsetToFrame], broadcast[:])
return n, nil
}
}
@@ -181,8 +242,16 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
case OpReply:
hwaddr, protoaddr := afrm.Sender()
for i := range h.queries {
if len(h.queries[i].hwaddr) == 0 && bytes.Equal(h.queries[i].protoaddr, protoaddr) {
h.queries[i].hwaddr = append(h.queries[i].hwaddr[:0], hwaddr...)
q := &h.queries[i]
mac := q.response()
if mac == nil && bytes.Equal(q.protoaddr, protoaddr) {
q.hwaddr = append(q.hwaddr, hwaddr...)
if q.dstHw != nil {
if !internal.IsZeroed(q.dstHw...) {
slog.Error("race-condition:ARP-reused-buffer")
}
copy(q.dstHw, hwaddr) // External write to user buffer.
}
return nil
}
}
@@ -194,7 +263,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
}
func trySetEthernetDst(ethFrame []byte, dst []byte) {
if len(ethFrame) > 14 {
if len(ethFrame) >= 14 {
copy(ethFrame[:6], dst)
}
}
+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) 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
}
+8 -8
View File
@@ -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 {
+4 -4
View File
@@ -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
}
+2 -2
View File
@@ -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
+5 -5
View File
@@ -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) {
+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 {
lg.Error("handle", slog.String("err", err.Error()))
} else if nw > 0 {
@@ -230,7 +230,7 @@ func (stack *Stack) Recv(b []byte) error {
}
func (stack *Stack) Send(b []byte) (int, error) {
return stack.ethernet.Encapsulate(b, 0)
return stack.ethernet.Encapsulate(b, -1, 0)
}
func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, error) {
@@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err
if err != nil {
return nil, err
}
err = stack.tcpports.Register(&listener)
err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting.
if err != nil {
return nil, err
}
@@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error
if err != nil {
return nil, err
}
err = stack.tcpports.Register(conn)
err = stack.tcpports.Register(conn) // Passive MAC with no listening.
if err != nil {
return nil, err
}
+4 -4
View File
@@ -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
}
+1 -1
View File
@@ -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 {
+11
View File
@@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) {
copy(dstaddr, dst)
return nil
}
// IsZeroed returns true if all arguments are set to their zero value.
func IsZeroed[T comparable](a ...T) bool {
var z T
for i := range a {
if a[i] != z {
return false
}
}
return true
}
+159 -86
View File
@@ -2,25 +2,30 @@ package internet
import (
"errors"
"log/slog"
"math"
"net"
"slices"
)
// StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols,
// from Ethernet to IP to TCP, practically any protocol can be expressed as a StackNode and function completely.
type StackNode interface {
// Encapsulate writes the stack node's frame into carrierData[frameOffset:]
// Encapsulate writes the stack node's frame into carrierData[offsetToFrame:]
// along with any other frame or payload the stack node encapsulates.
// The returned integer is amount of bytes written such that carrierData[frameOffset:frameOffset+n]
// contains written data. Data inside carrierData[:frameOffset] usually contains data necessary for
// The returned integer is amount of bytes written such that carrierData[offsetToFrame:offsetToFrame+n]
// contains written data. Data inside carrierData[:offsetToFrame] usually contains data necessary for
// a StackNode to correctly emit valid frame data: such is the case for TCP packets which require IP
// frame data for checksum calculation. Thus StackNodes must provide fields in their own frame
// required by sub-stacknodes for correct encapsulation; in the case of IPv4/6 this means including fields
// used in pseudo-header checksum like local IP (see [ipv4.CRCWriteUDPPseudo]).
//
// offsetToIP is the offset to the IP frame, if present, else its value should be -1.
// The relation offsetToIP<=offsetToFrame should always hold.
//
// When [net.ErrClosed] is returned the StackNode should be discarded and any written data passed up normally.
// Errors returned by Encapsulate are "extraordinary" and should not be returned unless the StackNode is receiving invalid carrierData/frameOffset.
Encapsulate(carrierData []byte, frameOffset int) (int, error)
Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error)
// Demux reads from the argument buffer where frameOffset is the offset of this StackNode's frame first byte.
// The stack node then dispatches(demuxes) the encapsulated frames to its corresponding sub-stack-node(s).
Demux(carrierData []byte, frameOffset int) error
@@ -36,9 +41,156 @@ type node struct {
currConnID uint64
connID *uint64
demux func([]byte, int) error
encapsulate func([]byte, int) (int, error)
encapsulate func([]byte, int, int) (int, error)
proto uint16
port uint16
// remoteAddr will be set on active(outbound) port connections
// that require an ARP to set the remoteAddr beforehand.
remoteAddr []byte
}
type handlers struct {
context string
logger
nodes []node
}
func (h *handlers) reset(context string, maxNodes int) {
h.nodes = slices.Grow(h.nodes[:0], maxNodes)
h.context = context
}
func (h *handlers) registerByProto(n node) error {
err := h.prepAdd()
if err != nil {
return err
}
if h.nodeByProto(n.proto) != nil {
return errProtoRegistered
}
h.nodes = append(h.nodes, n)
return nil
}
func (h *handlers) registerByPortProto(n node) error {
err := h.prepAdd()
if err != nil {
return err
}
if h.nodeByPortProto(n.port, n.proto) != nil {
return errProtoRegistered
}
h.nodes = append(h.nodes, n)
return nil
}
func (h *handlers) prepAdd() error {
if h.full() {
h.compact()
if h.full() {
return errNodesFull
}
}
return nil
}
func (h *handlers) full() bool { return cap(h.nodes) == len(h.nodes) }
func (h *handlers) compact() {
nilOff := 0
for i := 0; i < len(h.nodes); i++ {
if !h.nodes[i].IsInvalid() {
h.nodes[nilOff] = h.nodes[i]
nilOff++
}
}
h.nodes = h.nodes[:nilOff]
}
func (h *handlers) tryHandleError(node *node, err error) (discardedGracefully bool) {
if err != nil && (err == net.ErrClosed || node.IsInvalid()) {
node.destroy()
discardedGracefully = true
}
return discardedGracefully
}
func (h *handlers) nodeByProto(proto uint16) *node {
for i := range h.nodes {
node := &h.nodes[i]
if node.proto == proto && !node.IsInvalid() {
return node
}
}
return nil
}
func (h *handlers) nodeByPort(port uint16) *node {
for i := range h.nodes {
node := &h.nodes[i]
if node.port == port && !node.IsInvalid() {
return node
}
}
return nil
}
func (h *handlers) nodeByPortProto(port uint16, protocol uint16) *node {
for i := range h.nodes {
node := &h.nodes[i]
if node.port == port && node.proto == protocol && !node.IsInvalid() {
return node
}
}
return nil
}
func (h *handlers) demuxByProto(buf []byte, offset int, proto uint16) (*node, error) {
node := h.nodeByProto(proto)
if node == nil {
return nil, nil
}
err := node.demux(buf, offset)
if h.tryHandleError(node, err) {
err = nil
}
return node, err
}
func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, error) {
node := h.nodeByPort(port)
if node == nil {
return nil, nil
}
err := node.demux(buf, offset)
if h.tryHandleError(node, err) {
err = nil
node = nil // Node is destroyed in tryHandleError and invalidated.
}
return node, err
}
// encapsulateAny finds a node suitable to write and encapsulates the package.
// If no data is sent it returns the last error encountered.
func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) {
for i := range h.nodes {
node := &h.nodes[i]
if node.IsInvalid() {
continue
}
n, err = node.encapsulate(buf, offsetIP, offsetThisFrame)
if h.tryHandleError(node, err) {
err = nil // CLOSE error handled gracefully by deleting node.
node = nil // Node is destroyed in tryHandleError and invalidated.
}
if n > 0 {
return node, n, err
} else if err != nil {
// Make sure not to hang on one handler that keeps returning an error.
h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error()))
}
}
return nil, 0, err // Return last written error.
}
var (
@@ -50,32 +202,6 @@ var (
_ = net.ErrClosed
)
func registerNode(nodesPtr *[]node, h node) error {
if cap(*nodesPtr)-len(*nodesPtr) <= 0 {
*nodesPtr = nodesCompact(*nodesPtr)
}
if cap(*nodesPtr)-len(*nodesPtr) <= 0 {
return errNodesFull
}
*nodesPtr = append(*nodesPtr, h)
return nil
}
func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) {
if err != nil {
if nodeIdx >= len(*nodesPtr) {
panic("unreachable")
}
nodes := *nodesPtr
if checkNodeErr(&nodes[nodeIdx], err) {
// *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1)
(*nodesPtr)[nodeIdx] = node{} // 'Delete' node without modifying slice length.
discarded = true
}
}
return discarded
}
func (node *node) IsInvalid() bool {
return node.demux == nil || node.encapsulate == nil || (node.connID != nil && node.currConnID != *node.connID)
}
@@ -84,7 +210,7 @@ func checkNodeErr(node *node, err error) (discard bool) {
return node.IsInvalid() || (err != nil && err == net.ErrClosed)
}
func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node {
func nodeFromStackNode(s StackNode, port uint16, protocol uint64, remoteAddr []byte) node {
if protocol > math.MaxUint16 {
panic(">16bit protocol number unsupported")
}
@@ -100,64 +226,11 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node {
encapsulate: s.Encapsulate,
proto: uint16(protocol),
port: port,
remoteAddr: remoteAddr, // SHARED MEMORY- used to signal.
}
}
func getNode(nodes []node, port uint16, protocol uint16) (node *node) {
for i := range nodes {
node := &nodes[i]
if node.port == port && node.proto == protocol {
return node
}
}
return nil
}
func getEncapsulateNode(nodes *[]node, carrierData []byte, frameOffset int) (nodeIdx int, written int, err error) {
destroyed := false
for i := range *nodes {
node := &(*nodes)[i]
if node.IsInvalid() {
destroyed = true
node.destroy()
continue
}
written, err = node.encapsulate(carrierData, frameOffset)
if written > 0 {
return i, written, err
} else if err != nil {
}
}
if destroyed {
*nodes = nodesCompact(*nodes)
}
return -1, 0, nil
}
// destroy removes all references to underlying StackNode. Allows garbage collection of node if possible.
func (n *node) destroy() {
*n = node{}
}
func getNodeByProto(nodes []node, protocol uint16) int {
for i := range nodes {
node := &nodes[i]
if node.proto == protocol {
return i
}
}
return -1
}
func nodesCompact(nodes []node) []node {
nilOff := 0
for i := 0; i < len(nodes); i++ {
if !nodes[i].IsInvalid() {
nodes[nilOff] = nodes[i]
nilOff++
}
}
return nodes[:nilOff]
}
+2 -2
View File
@@ -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)
}
+25 -43
View File
@@ -6,7 +6,6 @@ import (
"log/slog"
"math"
"net"
"slices"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
@@ -14,11 +13,10 @@ import (
type StackEthernet struct {
connID uint64
handlers []node
logger
mac [6]byte
gwmac [6]byte
mtu uint16
handlers handlers
mac [6]byte
gwmac [6]byte
mtu uint16
}
func (ls *StackEthernet) SetGateway6(gw [6]byte) {
@@ -43,11 +41,10 @@ func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error {
} else if maxNodes <= 0 {
return errZeroMaxNodesArg
}
ls.handlers = slices.Grow(ls.handlers[:0], maxNodes)
ls.handlers.reset("StackEthernet", maxNodes)
*ls = StackEthernet{
connID: ls.connID + 1,
handlers: ls.handlers,
logger: ls.logger,
mac: mac,
gwmac: gateway,
mtu: uint16(mtu),
@@ -68,18 +65,7 @@ func (ls *StackEthernet) Register(h StackNode) error {
if proto > math.MaxUint16 || proto <= 1500 {
return errInvalidProto
}
eproto := uint16(proto)
for i := range ls.handlers {
hgot := &ls.handlers[i]
if hgot.proto == eproto {
return errProtoRegistered
}
}
return registerNode(&ls.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
proto: eproto,
})
return ls.handlers.registerByProto(nodeFromStackNode(h, 0, proto, nil))
}
func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) {
@@ -98,21 +84,17 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error)
if vld.HasError() {
return vld.ErrPop()
}
for i := range ls.handlers {
h := &ls.handlers[i]
if h.proto == uint16(etype) {
return h.demux(efrm.Payload(), 0)
}
if h, err := ls.handlers.demuxByProto(efrm.Payload(), 0, uint16(etype)); h != nil {
return err
}
DROP:
ls.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String()))
ls.handlers.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String()))
return lneto.ErrPacketDrop
}
func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int, err error) {
func (ls *StackEthernet) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
mtu := ls.mtu
dst := carrierData[frameOffset:]
dst := carrierData[offsetToFrame:]
if len(dst) < int(mtu) {
return 0, io.ErrShortBuffer
}
@@ -121,19 +103,19 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int
return 0, err
}
*efrm.DestinationHardwareAddr() = ls.gwmac
for i := range ls.handlers {
h := &ls.handlers[i]
n, err = h.encapsulate(dst[:mtu], 14)
if err != nil {
ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error()))
continue
}
if n > 0 {
// Found packet
*efrm.SourceHardwareAddr() = ls.mac
efrm.SetEtherType(ethernet.Type(h.proto))
return n + 14, nil
}
var h *node
// Children (IP/ARP) start at offset 14 (after ethernet header).
// For IP: offsetToIP=14, offsetToFrame=14
// For ARP: offsetToIP=-1, offsetToFrame=14 (but ARP ignores offsetToIP)
// Clip carrierData to MTU to prevent writes beyond MTU limit.
mtuLimit := offsetToFrame + int(mtu)
h, n, err = ls.handlers.encapsulateAny(carrierData[:mtuLimit], offsetToFrame+14, offsetToFrame+14)
if n == 0 {
return n, err
}
return 0, err
// Found packet
*efrm.SourceHardwareAddr() = ls.mac
efrm.SetEtherType(ethernet.Type(h.proto))
n += 14
return n, err
}
+54 -79
View File
@@ -5,7 +5,6 @@ import (
"io"
"log/slog"
"net/netip"
"slices"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
@@ -19,13 +18,11 @@ import (
var _ StackNode = (*StackIP)(nil)
type StackIP struct {
connID uint64
ipID uint16
ip [4]byte
validator lneto.Validator
handlers []node
pendingICMP [][]byte
logger
connID uint64
ipID uint16
ip [4]byte
validator lneto.Validator
handlers handlers
}
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
@@ -36,14 +33,12 @@ func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
if err != nil {
return err
}
sb.handlers = slices.Grow(sb.handlers[:0], maxNodes)
sb.handlers.reset("StackIP", maxNodes)
*sb = StackIP{
connID: sb.connID + 1,
validator: sb.validator,
handlers: sb.handlers,
logger: sb.logger,
ip: sb.ip,
pendingICMP: make([][]byte, maxNodes*4),
connID: sb.connID + 1,
validator: sb.validator,
handlers: sb.handlers,
ip: sb.ip,
}
return nil
}
@@ -73,11 +68,11 @@ func (sb *StackIP) Addr() netip.Addr {
}
func (sb *StackIP) SetLogger(logger *slog.Logger) {
sb.logger.log = logger
sb.handlers.log = logger
}
func (sb *StackIP) Demux(carrierData []byte, offset int) error {
sb.info("StackIP.Demux:start")
sb.handlers.info("StackIP.Demux:start")
frame := carrierData[offset:] // we don't care about carrier data in IP.
ifrm, err := ipv4.NewFrame(frame)
if err != nil {
@@ -96,7 +91,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
gotCRC := ifrm.CRC()
wantCRC := ifrm.CalculateHeaderCRC()
if gotCRC != wantCRC {
sb.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC)))
sb.handlers.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC)))
return errors.New("IPv4 CRC mismatch")
}
off := ifrm.HeaderLength()
@@ -105,10 +100,11 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
if proto == lneto.IPProtoICMP {
return sb.recvicmp(ifrm.RawData(), ifrm.HeaderLength())
}
nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if nodeIdx < 0 {
node := sb.handlers.nodeByProto(uint16(proto))
// nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if node == nil {
// Drop packet.
sb.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String()))
sb.handlers.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String()))
return nil
}
// Incoming CRC Validation of common IP Protocols.
@@ -135,17 +131,17 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
return errors.New("UDP CRC mismatch")
}
}
sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = sb.handlers[nodeIdx].demux(frame[:totalLen], off)
if handleNodeError(&sb.handlers, nodeIdx, err) {
sb.info("ipclose", slog.String("proto", proto.String()))
sb.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = node.demux(frame[:totalLen], off)
if sb.handlers.tryHandleError(node, err) {
sb.handlers.info("ipclose", slog.String("proto", proto.String()))
err = nil
}
return err
}
func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) {
frame := carrierData[frameOffset:]
func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
frame := carrierData[offsetToFrame:]
if len(frame) < 256 {
return 0, io.ErrShortBuffer
}
@@ -162,46 +158,37 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
ifrm.SetTTL(64)
*ifrm.SourceAddr() = sb.ip
sb.ipID = id
for i := range sb.handlers {
h := &sb.handlers[i]
proto := lneto.IPProto(h.proto)
n, err := h.encapsulate(frame[:], headerlen)
if err != nil {
if handleNodeError(&sb.handlers, i, err) {
println("IP NODE REMOVED", proto.String(), h.port)
h.destroy()
}
sb.error("StackIP:encapsulate", slog.String("proto", proto.String()), slog.String("err", err.Error()))
continue
} else if n == 0 {
continue
}
totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto)
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
// Calculate CRC for our newly generated packet.
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, _ := tcp.NewFrame(ifrm.Payload())
tfrm.CRCWrite(&crc)
tfrm.SetCRC(crc.Sum16())
case lneto.IPProtoUDP:
ifrm.CRCWriteUDPPseudo(&crc)
ufrm, _ := udp.NewFrame(ifrm.Payload())
ufrm.SetLength(uint16(n))
ufrm.CRCWriteIPv4(&crc)
ufrm.SetCRC(crc.Sum16())
if n != int(ufrm.Length()) {
sb.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length())))
return 0, errors.New("invalid UDP length")
}
}
return totalLen, nil
// Children (TCP/UDP) start at offset headerlen (20 bytes after IP header start).
// offsetToIP is 0 relative to this slice (frame), children's frame starts at headerlen.
node, n, err := sb.handlers.encapsulateAny(carrierData, offsetToFrame, offsetToFrame+headerlen)
if n == 0 {
return n, err
}
return 0, nil
proto := lneto.IPProto(node.proto)
totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto)
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
// Calculate CRC for our newly generated packet.
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, _ := tcp.NewFrame(ifrm.Payload())
tfrm.CRCWrite(&crc)
tfrm.SetCRC(crc.Sum16())
case lneto.IPProtoUDP:
ifrm.CRCWriteUDPPseudo(&crc)
ufrm, _ := udp.NewFrame(ifrm.Payload())
ufrm.SetLength(uint16(n))
ufrm.CRCWriteIPv4(&crc)
ufrm.SetCRC(crc.Sum16())
if n != int(ufrm.Length()) {
sb.handlers.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length())))
return 0, errors.New("invalid UDP length")
}
}
return totalLen, err
}
func (sb *StackIP) Register(h StackNode) error {
@@ -209,19 +196,7 @@ func (sb *StackIP) Register(h StackNode) error {
if proto > 255 {
return errInvalidProto
}
connID := h.ConnectionID()
var currConnID uint64
if connID != nil {
currConnID = *connID
}
return registerNode(&sb.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
proto: uint16(proto),
port: h.LocalPort(),
currConnID: currConnID,
connID: connID,
})
return sb.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil))
}
func (sb *StackIP) recvicmp(carrierData []byte, offset int) error {
+84 -54
View File
@@ -2,18 +2,23 @@ package internet
import (
"encoding/binary"
"errors"
"io"
"log/slog"
"math"
"slices"
"strconv"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
)
type StackPorts struct {
connID uint64
handlers []node
handlers handlers
dstPortOff uint16
protocol uint16
// stores last node to demux/encapsulate.
}
func (ps *StackPorts) ResetUDP(maxNodes int) error {
@@ -30,7 +35,7 @@ func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int)
} else if maxNodes <= 0 {
return errZeroMaxNodesArg
}
ps.handlers = slices.Grow(ps.handlers[:0], maxNodes)
ps.handlers.reset("StackPorts(proto="+strconv.Itoa(int(protocol))+")", maxNodes)
*ps = StackPorts{
connID: ps.connID + 1,
handlers: ps.handlers,
@@ -46,24 +51,11 @@ func (ps *StackPorts) Protocol() uint64 { return uint64(ps.protocol) }
func (ps *StackPorts) ConnectionID() *uint64 { return &ps.connID }
func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) {
if int(ps.dstPortOff)+offset+2 > len(b) {
func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) {
return 0, io.ErrShortBuffer
}
var i int
for i = 0; i < len(ps.handlers); i++ {
if ps.handlers[i].IsInvalid() {
continue
}
n, err = ps.handlers[i].encapsulate(b, offset)
if err != nil || n > 0 {
if ps.handleResult(i, n, err) {
err = nil // Handler discarded. Keep looking for other handlers.
continue
}
break
}
}
_, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame)
return n, err
}
@@ -72,52 +64,90 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) {
return io.ErrShortBuffer
}
port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:])
var i int
for i = 0; i < len(ps.handlers); i++ {
if port != ps.handlers[i].port {
continue
}
err = ps.handlers[i].demux(b, offset)
if err != nil {
if ps.handleResult(i, 0, err) {
err = nil // Handler discarded. Keep looking for other maybe available handlers.
continue
}
break
}
}
ps.handleResult(i, 0, err)
_, err = ps.handlers.demuxByPort(b, offset, port)
return err
}
// Register registers a port StackNode on StackPorts.
// If dstMAC is set to non-nil, length six buffer then
func (ps *StackPorts) Register(h StackNode) error {
port := h.LocalPort()
proto := h.Protocol()
if port <= 0 {
return errZeroPort
} else if proto != uint64(ps.protocol) {
return errInvalidProto
}
var cid uint64
cidPtr := h.ConnectionID()
if cidPtr != nil {
cid = *cidPtr
}
return registerNode(&ps.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
port: port,
currConnID: cid,
connID: cidPtr,
proto: uint16(proto),
})
return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil))
}
func (ps *StackPorts) handleResult(handlerIdx, n int, err error) (discarded bool) {
if handleNodeError(&ps.handlers, handlerIdx, err) {
discarded = true
println("DISCARD", handlerIdx, "witherr", err.Error())
}
return discarded
// StackPortsMACFiltered is a StackPorts implementation but that avoids calling encapsulate on nodes
// with a non-nil MAC address registered via Register method that is set to all zero values.
// If the address is set to nil no filtering occurs. MAC Address is set automatically on the ethernet frame by StackPortsMACFiltered when non-nil.
type StackPortsMACFiltered struct {
sp StackPorts
}
func (mfsp *StackPortsMACFiltered) Register(h StackNode, addr []byte) error {
port := h.LocalPort()
proto := h.Protocol()
if port <= 0 {
return errZeroPort
} else if proto != uint64(mfsp.sp.protocol) {
return errInvalidProto
} else if addr != nil && len(addr) != 6 {
return errors.New("invalid MAC")
}
return mfsp.sp.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, addr))
}
func (ps *StackPortsMACFiltered) ResetUDP(maxNodes int) error {
return ps.sp.ResetUDP(maxNodes)
}
func (ps *StackPortsMACFiltered) ResetTCP(maxNodes int) error {
return ps.sp.ResetTCP(maxNodes)
}
func (ps *StackPortsMACFiltered) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) error {
return ps.sp.Reset(protocol, dstPortOffset, maxNodes)
}
func (ps *StackPortsMACFiltered) LocalPort() uint16 { return 0 }
func (ps *StackPortsMACFiltered) Protocol() uint64 { return uint64(ps.sp.protocol) }
func (ps *StackPortsMACFiltered) ConnectionID() *uint64 { return &ps.sp.connID }
func (ps *StackPortsMACFiltered) Demux(b []byte, offset int) (err error) {
// No MAC Filtering on ingress. TODO?
return ps.sp.Demux(b, offset)
}
func (ps *StackPortsMACFiltered) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
if int(ps.sp.dstPortOff)+offsetToFrame+2 > len(carrierData) {
return 0, io.ErrShortBuffer
}
h := &ps.sp.handlers
for i := range h.nodes {
node := &h.nodes[i]
if node.IsInvalid() || (len(node.remoteAddr) > 0 && internal.IsZeroed(node.remoteAddr...)) {
continue
}
n, err = node.encapsulate(carrierData, offsetToIP, offsetToFrame)
if h.tryHandleError(node, err) {
err = nil // CLOSE error handled gracefully by deleting node.
}
if n > 0 {
if len(node.remoteAddr) == 6 && offsetToIP >= 14 {
efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:])
*efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr)
}
return n, err
} else if err != nil {
// Make sure not to hang on one handler that keeps returning an error.
h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error()))
}
}
return 0, err // Return last written error.
}
+7 -6
View File
@@ -17,7 +17,7 @@ type StackUDPPort struct {
}
func (sudp *StackUDPPort) SetStackNode(node StackNode, raddr []byte, rmport uint16) {
sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol())
sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol(), raddr)
sudp.rmport = rmport
sudp.raddr = append(sudp.raddr[:0], raddr...)
}
@@ -61,24 +61,25 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error {
return err
}
func (sudp *StackUDPPort) Encapsulate(carrierData []byte, frameOffset int) (int, error) {
func (sudp *StackUDPPort) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
if sudp.h.IsInvalid() {
sudp.h.destroy()
return 0, net.ErrClosed
}
ufrm, err := udp.NewFrame(carrierData[frameOffset:])
ufrm, err := udp.NewFrame(carrierData[offsetToFrame:])
if err != nil {
return 0, err
}
ufrm.SetSourcePort(sudp.h.port)
ufrm.SetDestinationPort(sudp.rmport)
if len(sudp.raddr) > 0 && frameOffset >= 20 {
err = internal.SetIPAddrs(carrierData, 0, nil, sudp.raddr)
if len(sudp.raddr) > 0 && offsetToIP >= 0 {
err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, nil, sudp.raddr)
if err != nil {
return 0, err
}
}
n, err := sudp.h.encapsulate(carrierData, frameOffset+8)
// Child payload starts 8 bytes after UDP header start.
n, err := sudp.h.encapsulate(carrierData, offsetToIP, offsetToFrame+8)
if n == 0 {
if err != nil {
slog.Error("stackudp:encapsulate", slog.String("err", err.Error()))
+1 -1
View File
@@ -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 {
+2 -2
View File
@@ -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
+22 -8
View File
@@ -246,7 +246,7 @@ func (conn *Conn) checkPipeOpen() error {
if conn.abortErr != nil {
return conn.abortErr
}
state := conn.State()
state := conn.h.State()
if state.IsClosed() {
return net.ErrClosed
}
@@ -278,23 +278,27 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) {
return nil
}
func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) {
func (conn *Conn) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
conn.mu.Lock()
defer conn.mu.Unlock()
if len(conn.remoteAddr) == 0 {
return 0, errNoRemoteAddr
}
raddr, _, _, _, err := internal.GetIPAddr(buf[:off])
if offsetToIP < 0 {
return 0, errNoRemoteAddr // No IP layer present.
}
ipFrame := carrierData[offsetToIP:offsetToFrame]
raddr, _, _, _, err := internal.GetIPAddr(ipFrame)
if err != nil {
return 0, err
} else if len(raddr) != len(conn.remoteAddr) {
return 0, errMismatchedIPVersion
}
n, err = conn.h.Send(buf[off:])
n, err = conn.h.Send(carrierData[offsetToFrame:])
if err != nil {
return 0, err
}
err = internal.SetIPAddrs(buf[:off], conn.ipID, nil, conn.remoteAddr)
err = internal.SetIPAddrs(ipFrame, conn.ipID, nil, conn.remoteAddr)
if err != nil {
return 0, err
}
@@ -328,11 +332,11 @@ func (conn *Conn) reset(h Handler) {
func (conn *Conn) SetDeadline(t time.Time) error {
conn.mu.Lock()
defer conn.mu.Unlock()
err := conn.SetReadDeadline(t)
err := conn.setReadDeadline(t)
if err != nil {
return err
}
return conn.SetWriteDeadline(t)
return conn.setWriteDeadline(t)
}
// SetReadDeadline sets the deadline for future Read calls
@@ -340,7 +344,11 @@ func (conn *Conn) SetDeadline(t time.Time) error {
func (conn *Conn) SetReadDeadline(t time.Time) error {
conn.mu.Lock()
defer conn.mu.Unlock()
conn.trace("TCPConn.SetReadDeadline:start")
return conn.setReadDeadline(t)
}
func (conn *Conn) setReadDeadline(t time.Time) error {
conn.trace("TCPConn.setReadDeadline:start")
err := conn.checkPipeOpen()
if err == nil {
conn.rdead = t
@@ -354,6 +362,12 @@ func (conn *Conn) SetReadDeadline(t time.Time) error {
// some of the data was successfully written.
// A zero value for t means Write will not time out.
func (conn *Conn) SetWriteDeadline(t time.Time) error {
conn.mu.Lock()
defer conn.mu.Unlock()
return conn.setWriteDeadline(t)
}
func (conn *Conn) setWriteDeadline(t time.Time) error {
conn.trace("TCPConn.SetWriteDeadline:start")
err := conn.checkPipeOpen()
if err == nil {
+42 -18
View File
@@ -28,11 +28,12 @@ type StackAsync struct {
ip internet.StackIP
arp arp.Handler
udps internet.StackPorts
tcps internet.StackPorts
tcps internet.StackPortsMACFiltered
dhcpUDP internet.StackUDPPort
dhcp dhcpv4.Client
dhcpResults DHCPResults
subnet netip.Prefix // Local subnet for ARP resolution.
dnsUDP internet.StackUDPPort
dns dns.Client
@@ -73,11 +74,11 @@ func (s *StackAsync) Demux(carrierData []byte, etherOff int) error {
return s.link.Demux(carrierData, etherOff)
}
func (s *StackAsync) Encapsulate(carrierData []byte, etherOff int) (int, error) {
func (s *StackAsync) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
s.mu.Lock()
defer s.mu.Unlock()
n, err := s.link.Encapsulate(carrierData, etherOff)
n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame)
s.totalsent += uint64(n)
return n, err
}
@@ -112,6 +113,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error {
if err != nil {
return err
}
//
err = s.resetARP()
if err != nil {
return err
@@ -135,10 +137,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error {
}
// Now setup stacks.
err = s.link.Register(&s.arp) // ARP.
if err != nil {
return err
}
// ARP registered in resetARP.
err = s.link.Register(&s.ip) // IPv4 | IPv6
if err != nil {
return err
@@ -169,7 +168,7 @@ func (s *StackAsync) resetARP() error {
if addr.Is6() {
proto = ethernet.TypeIPv6
}
return s.arp.Reset(arp.HandlerConfig{
err := s.arp.Reset(arp.HandlerConfig{
HardwareAddr: mac[:],
ProtocolAddr: addr.AsSlice(),
MaxQueries: 3,
@@ -177,6 +176,14 @@ func (s *StackAsync) resetARP() error {
HardwareType: 1,
ProtocolType: proto,
})
if err != nil {
return err
}
err = s.link.Register(&s.arp)
if err != nil {
return err
}
return nil
}
// Prand32 generates a pseudo random 32-bit unsigned integer from the internal state and advances the seed.
@@ -193,11 +200,17 @@ func (s *StackAsync) Prand32() uint32 {
func (s *StackAsync) SetIPAddr(addr netip.Addr) error {
s.mu.Lock()
defer s.mu.Unlock()
return s.setIPAddr(addr)
}
func (s *StackAsync) setIPAddr(addr netip.Addr) error {
err := s.ip.SetAddr(addr)
if err != nil {
return err
}
return s.resetARP()
ip := addr.As4()
err = s.arp.UpdateProtoAddr(ip[:])
return err
}
func (s *StackAsync) Addr() netip.Addr {
@@ -234,11 +247,23 @@ func (s *StackAsync) Gateway6() [6]byte {
func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort) (err error) {
s.mu.Lock()
defer s.mu.Unlock()
var mac []byte
if s.subnet.Contains(addrp.Addr()) {
mac = make([]byte, 6)
ip := addrp.Addr().As4()
// StartQuery starts an ARP query for addresses in this network.
// On finishing query MAC is set and thus the StackPort will allow encapsulating
// data on that connection.
err = s.arp.StartQuery(mac, ip[:])
if err != nil {
return err
}
}
err = conn.OpenActive(localPort, addrp, tcp.Value(s.Prand32()))
if err != nil {
return err
}
err = s.tcps.Register(conn)
err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network.
if err != nil {
conn.Abort()
return err
@@ -253,7 +278,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) {
if err != nil {
return err
}
err = s.tcps.Register(conn)
err = s.tcps.Register(conn, nil)
if err != nil {
conn.Abort()
return err
@@ -376,7 +401,7 @@ func (s *StackAsync) StartResolveHardwareAddress6(ip netip.Addr) error {
return errors.New("unsupported or invalid IP address")
}
addr := ip.As4()
return s.arp.StartQuery(addr[:])
return s.arp.StartQuery(nil, addr[:])
}
// ResultResolveHardwareAddress6
@@ -432,16 +457,15 @@ func (s *StackAsync) ReadStatistics(stats *Statistics) {
// AssimilateDHCPResults sets the stack's following parameters:
// - IPv4 address.
// - DNS server.
// - Subnet (for ARP resolution of local addresses).
func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error {
stack.mu.Lock()
defer stack.mu.Unlock()
if results.Subnet.IsValid() {
stack.subnet = results.Subnet
}
if results.AssignedAddr.IsValid() {
err := stack.ip.SetAddr(results.AssignedAddr)
if err != nil {
return err
}
// Reset ARP handler with new IP address so it can respond to ARP requests.
err = stack.resetARP()
err := stack.setIPAddr(results.AssignedAddr)
if err != nil {
return err
}
+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 (
"bytes"
"encoding/binary"
"errors"
"math/rand"
"net/netip"
"testing"
"github.com/soypat/lneto/arp"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internet/pcap"
"github.com/soypat/lneto/tcp"
@@ -24,9 +26,7 @@ func TestStackAsyncTCP_multipacket(t *testing.T) {
const svPort = 8080
const maxPktLen = 30
client, sv, clconn, svconn := newTCPStacks(t, seed, MTU)
tst := tester{
t: t, buf: make([]byte, MTU),
}
tst := testerFrom(t, MTU)
rng := rand.New(rand.NewSource(seed))
client2, sv2, clconn2, svconn2 := newTCPStacks(t, seed, MTU)
_, _, _, _ = client2, sv2, clconn2, svconn2
@@ -58,10 +58,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) {
const MTU = 1500
const svPort = 80
client, sv, clconn, svconn := newTCPStacks(t, seed, MTU)
tst := tester{
t: t, buf: make([]byte, MTU),
}
tst := testerFrom(t, MTU)
tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337)
sendData := []byte("hello")
@@ -80,7 +77,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) {
func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1, c2 *tcp.Conn) {
s1, s2 = new(StackAsync), new(StackAsync)
c1, c2 = new(tcp.Conn), new(tcp.Conn)
byte1 := byte(randSeed) / 4
byte1 := byte(randSeed)/4 - 1
err := s1.Reset(StackConfig{
Hostname: "Stack1",
RandSeed: randSeed,
@@ -127,6 +124,13 @@ func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1
return s1, s2, c1, c2
}
func testerFrom(t *testing.T, mtu int) *tester {
return &tester{
t: t,
buf: make([]byte, mtu),
}
}
type tester struct {
t *testing.T
cap pcap.PacketBreakdown
@@ -322,7 +326,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs
default:
panic("OOB")
}
n, err := src.Encapsulate(buf[:], 0)
n, err := src.Encapsulate(buf[:], -1, 0)
if err != nil {
t.Fatal(err)
} else if n == 0 {
@@ -334,6 +338,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs
t.Error("expected no data sent and got data")
return
}
defer setzero(buf[:n])
tst.buf = tst.buf[:n]
tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0)
@@ -374,7 +379,121 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs
if err != nil {
t.Fatal(err)
}
}
func (tst *tester) ARPExchangeOnly(querying, target *StackAsync) {
t := tst.t
t.Helper()
buf := tst.buf[:cap(tst.buf)]
// === PHASE 1: ARP Request from querying stack ===
n, err := querying.Encapsulate(buf[:], -1, 0)
if err != nil {
t.Fatal(err)
} else if n == 0 {
t.Error("zero bits sent by ARP querying stack")
return
}
tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0)
if err != nil {
t.Fatal(err)
}
tst.buf = tst.buf[:n]
qHw := querying.HardwareAddress()
tgtHw := target.HardwareAddress()
broadcast := ethernet.BroadcastAddr()
qIP := querying.Addr()
tgtIP := target.Addr()
// Validate Ethernet layer (request is broadcast)
if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) {
t.Errorf("request: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc))
}
if !bytes.Equal(broadcast[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) {
t.Errorf("request: expected broadcast ethernet dst addr, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst))
}
// Validate ARP request fields
// ARP fields: FieldClassSrc with 6 octets = HW addr, 4 octets = proto addr
// occurrence 0 = sender, occurrence 1 = target
if tst.getARPOperation() != arp.OpRequest {
t.Errorf("request: expected ARP OpRequest, got %d", tst.getARPOperation())
}
if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) {
t.Errorf("request: mismatched ARP sender HW")
}
if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) {
t.Errorf("request: mismatched ARP sender proto")
}
if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) {
t.Errorf("request: mismatched ARP target proto")
}
// Deliver request to target
err = target.Demux(buf[:n], 0)
if err != nil {
t.Fatal("target demux request:", err)
}
setzero(buf[:n])
// === PHASE 2: ARP Reply from target stack ===
buf = tst.buf[:cap(tst.buf)]
n, err = target.Encapsulate(buf[:], -1, 0)
if err != nil {
t.Fatal(err)
} else if n == 0 {
t.Error("zero bits sent by ARP target stack (no reply)")
return
}
tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0)
if err != nil {
t.Fatal(err)
}
tst.buf = tst.buf[:n]
// Validate Ethernet layer (reply is unicast to querying)
if !bytes.Equal(tgtHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) {
t.Errorf("reply: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc))
}
if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) {
t.Errorf("reply: expected unicast to querying, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst))
}
// Validate ARP reply fields
if tst.getARPOperation() != arp.OpReply {
t.Errorf("reply: expected ARP OpReply, got %d", tst.getARPOperation())
}
if !bytes.Equal(tgtHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) {
t.Errorf("reply: mismatched ARP sender HW (should be target's MAC)")
}
if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) {
t.Errorf("reply: mismatched ARP sender proto (should be target's IP)")
}
if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 1)) {
t.Errorf("reply: mismatched ARP target HW (should be querying's MAC)")
}
if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) {
t.Errorf("reply: mismatched ARP target proto (should be querying's IP)")
}
// Deliver reply to querying stack
err = querying.Demux(buf[:n], 0)
if err != nil {
t.Fatal("querying demux reply:", err)
}
setzero(buf[:n])
// === PHASE 3: Verify querying stack learned target's MAC ===
resolvedHw, err := querying.ResultResolveHardwareAddress6(tgtIP)
if err != nil {
t.Fatalf("ARP query result failed: %v", err)
}
if resolvedHw != tgtHw {
t.Errorf("ARP resolved wrong MAC: got %x, want %x", resolvedHw, tgtHw)
}
}
func (tst *tester) getTCPFrame() tcp.Frame {
@@ -457,3 +576,34 @@ func setzero[T ~[]E, E any](s T) {
s[i] = zero
}
}
// getFieldByClassLen finds a field by protocol, class, and octet length.
// occurrence specifies which match to return (0 = first, 1 = second, etc.)
// This is needed for ARP where sender and target fields share the same class.
func (tst *tester) getFieldByClassLen(proto any, class pcap.FieldClass, octetLen, occurrence int) []byte {
tst.t.Helper()
frm := getProtoFrame(tst.frmbuf, proto)
if frm == nil {
tst.t.Fatalf("no frame for proto %v found", proto)
}
count := 0
for _, field := range frm.Fields {
if field.Class == class && field.BitLength == octetLen*8 {
if count == occurrence {
bitoff := frm.PacketBitOffset + field.FrameBitOffset
return tst.buf[bitoff/8 : bitoff/8+field.BitLength/8]
}
count++
}
}
tst.t.Fatalf("field (proto=%v, class=%v, octets=%d, occurrence=%d) not found", proto, class, octetLen, occurrence)
return nil
}
func (tst *tester) getARPOperation() arp.Operation {
tst.t.Helper()
// ARP has 3 FieldClassType fields: Hardware type (0), Protocol type (1), Opcode (2)
// All are 2 bytes, so we need occurrence=2 to get Opcode.
data := tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassType, 2, 2)
return arp.Operation(binary.BigEndian.Uint16(data))
}