mirror of
https://github.com/soypat/lneto.git
synced 2026-08-07 16:33:40 +00:00
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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]
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
}
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user