fix: clean up misc TODO comments and improve UDP source filtering (#89)

- http/httpraw: cap header slice growth to pre-allocated capacity, fix
  benchmark to reuse Header across iterations (0 allocs/op)
- internal/ring: replace TODO panic comments with invariant explanations
- tcp/conn: document truncated-frame offset check
- tcp/handler: replace TODO with net.ErrClosed on RST-closed connection
- internet/stack-udpport: filter incoming packets by remote IP address
  when configured via SetStackNode; skip filter for IPv4 multicast
  destinations (class D) so mDNS and similar protocols work correctly

Generated with LLM assistance.

Signed-off-by: Marvin Drees <marvin.drees@9elements.com>
This commit is contained in:
Marvin Drees
2026-05-20 19:29:09 +02:00
committed by GitHub
parent edf302baf8
commit 46d4c06b9f
11 changed files with 76 additions and 14 deletions
+10 -4
View File
@@ -112,13 +112,19 @@ func BenchmarkParseBytes(b *testing.B) {
asRequest = false asRequest = false
) )
req, _ := http.NewRequest(wantMethod, wantURI, strings.NewReader(wantMessage)) req, _ := http.NewRequest(wantMethod, wantURI, strings.NewReader(wantMessage))
var buf bytes.Buffer var rawBuf bytes.Buffer
req.Write(&buf) req.Write(&rawBuf)
data := buf.Bytes() data := rawBuf.Bytes()
// hdr is declared outside the loop so that ParseBytes can reuse the
// backing arrays on Reset (headers slice and data buffer) without
// allocating on every iteration. Declaring it inside the loop causes
// two allocs per iteration: one for the headers slice (make in reset)
// and one for the data buffer (append in readFromBytes).
var hdr Header
b.StartTimer() b.StartTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
var hdr Header
err := hdr.ParseBytes(asRequest, data) err := hdr.ParseBytes(asRequest, data)
if err != nil { if err != nil {
b.Fatal(err) b.Fatal(err)
+7 -1
View File
@@ -120,7 +120,13 @@ func (hb *headerBuf) free() int { return cap(hb.buf) - len(hb.buf) }
func (hb *headerBuf) parseNextHeaders(ss *scannerState) { func (hb *headerBuf) parseNextHeaders(ss *scannerState) {
debuglog("http:nexthdr:loop") debuglog("http:nexthdr:loop")
for kv := hb.next(ss); kv.isValid(); kv = hb.next(ss) { for kv := hb.next(ss); kv.isValid(); kv = hb.next(ss) {
hb.headers = append(hb.headers, kv) // TODO(HEAP): inc=16B slice growth when capacity exceeded if len(hb.headers) == cap(hb.headers) {
// Refuse to grow the headers slice: caller must pre-allocate
// sufficient capacity via reset or use a larger initial size.
ss.err = errOOM
return
}
hb.headers = append(hb.headers, kv)
} }
debuglog("http:nexthdr:done") debuglog("http:nexthdr:done")
} }
+12
View File
@@ -26,6 +26,18 @@ func GetIPAddr(buf []byte) (src, dst []byte, id, ipEndOff uint16, err error) {
return src, dst, id, ipEndOff, err return src, dst, id, ipEndOff, err
} }
// IsMulticastIPAddr reports whether addr is an IPv4 or IPv6 multicast address.
func IsMulticastIPAddr(addr []byte) bool {
switch len(addr) {
case 4:
return addr[0]&0xf0 == 0xe0
case 16:
return addr[0] == 0xff
default:
return false
}
}
func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) {
var dstaddr, srcaddr []byte var dstaddr, srcaddr []byte
version := buf[0] >> 4 version := buf[0] >> 4
+26
View File
@@ -0,0 +1,26 @@
package internal
import "testing"
func TestIsMulticastIPAddr(t *testing.T) {
tests := []struct {
name string
addr []byte
want bool
}{
{"ipv4 multicast", []byte{224, 0, 0, 1}, true},
{"ipv4 multicast upper", []byte{239, 255, 255, 255}, true},
{"ipv4 unicast", []byte{192, 0, 2, 1}, false},
{"ipv6 multicast", []byte{0xff, 0x02, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}, true},
{"ipv6 unicast", []byte{0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}, false},
{"invalid length", []byte{224}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsMulticastIPAddr(tt.addr); got != tt.want {
t.Fatalf("got %t; want %t", got, tt.want)
}
})
}
}
+2 -2
View File
@@ -69,7 +69,7 @@ func (r *Ring) Write(b []byte) (int, error) {
n := copy(r.Buf[r.End:r.Off], b) n := copy(r.Buf[r.End:r.Off], b)
r.End += n r.End += n
if r.End <= 0 { if r.End <= 0 {
panic("zero end after write") // TODO: remove panics after validation. panic("zero end after write") // invariant: End must be >0 after writing into midFree region
} }
return n, nil return n, nil
} else if r.End == 0 { } else if r.End == 0 {
@@ -87,7 +87,7 @@ func (r *Ring) Write(b []byte) (int, error) {
n += n2 n += n2
} }
if r.End <= 0 { if r.End <= 0 {
panic("zero end after write") panic("zero end after write") // invariant: End must be >0 after appending to the tail region
} }
return n, nil return n, nil
} }
+1 -1
View File
@@ -101,7 +101,7 @@ func (si4 *stackip4) demux4(carrierData []byte, offset int) error {
} }
dst := ifrm.DestinationAddr() dst := ifrm.DestinationAddr()
if si4.ip4 != ([4]byte{}) && *dst != si4.ip4 { if si4.ip4 != ([4]byte{}) && *dst != si4.ip4 {
if !si4.acceptMulticast || dst[0]&0xF0 != 0xE0 { if !si4.acceptMulticast || !internal.IsMulticastIPAddr(dst[:]) {
si4.handlers.debug("ip:not-for-us") si4.handlers.debug("ip:not-for-us")
return lneto.ErrPacketDrop // Not meant for us. return lneto.ErrPacketDrop // Not meant for us.
} }
+2 -1
View File
@@ -5,6 +5,7 @@ import (
"github.com/soypat/lneto" "github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv6" "github.com/soypat/lneto/ipv6"
"github.com/soypat/lneto/tcp" "github.com/soypat/lneto/tcp"
"github.com/soypat/lneto/udp" "github.com/soypat/lneto/udp"
@@ -89,7 +90,7 @@ func (si6 *stackip6) demux6(carrierData []byte, offset int) error {
} }
dst := ifrm.DestinationAddr() dst := ifrm.DestinationAddr()
if si6.ip6 != ([16]byte{}) && *dst != si6.ip6 { if si6.ip6 != ([16]byte{}) && *dst != si6.ip6 {
if !si6.acceptMulticast || dst[0] != 0xFF { if !si6.acceptMulticast || !internal.IsMulticastIPAddr(dst[:]) {
si6.handlers.debug("ip6:not-for-us") si6.handlers.debug("ip6:not-for-us")
return lneto.ErrPacketDrop return lneto.ErrPacketDrop
} }
+9 -1
View File
@@ -45,7 +45,15 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error {
if dst != sudp.h.lport { if dst != sudp.h.lport {
return lneto.ErrPacketDrop // Not meant for us. return lneto.ErrPacketDrop // Not meant for us.
} }
// TODO remote ip address handling. if len(sudp.raddr) > 0 && !internal.IsMulticastIPAddr(sudp.raddr) {
srcIP, _, _, _, err := internal.GetIPAddr(carrierData[:frameOffset])
if err != nil {
return err
}
if !internal.BytesEqual(srcIP, sudp.raddr) {
return lneto.ErrPacketDrop
}
}
src := ufrm.SourcePort() src := ufrm.SourcePort()
if sudp.rmport != 0 && src != sudp.rmport { if sudp.rmport != 0 && src != sudp.rmport {
+3 -1
View File
@@ -362,7 +362,9 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) {
conn.mu.Lock() conn.mu.Lock()
defer conn.mu.Unlock() defer conn.mu.Unlock()
if off >= len(buf) { if off >= len(buf) {
return lneto.ErrTruncatedFrame // TODO: this check is bad. // off is the IP header length; if it equals or exceeds the frame
// length, there are zero bytes of TCP payload — drop the frame.
return lneto.ErrTruncatedFrame
} }
raddr, _, id, _, err := internal.GetIPAddr(buf[:off]) raddr, _, id, _, err := internal.GetIPAddr(buf[:off])
if err != nil { if err != nil {
+1 -2
View File
@@ -174,8 +174,7 @@ func (h *Handler) Recv(incomingPacket []byte) error {
err = h.scb.Recv(segIncoming) err = h.scb.Recv(segIncoming)
if err != nil { if err != nil {
if h.scb.State() == StateClosed { if h.scb.State() == StateClosed {
// TODO(soypat): Should return EOF/ErrClosed? err = net.ErrClosed // Connection closed by RST; signal caller to tear down.
err = net.ErrClosed //err // Connection closed by reset.
} }
return err return err
} }
+3 -1
View File
@@ -297,7 +297,9 @@ func newMDNSStack(t *testing.T, hostname string, seed int64,
t.Fatal(hostname, "mdns configure:", err) t.Fatal(hostname, "mdns configure:", err)
} }
err = stack.RegisterUDP4(&client, addr.As4(), mdns.Port) var remote [4]byte
copy(remote[:], mdnsCfg.MulticastAddr)
err = stack.RegisterUDP4(&client, remote, mdns.Port)
if err != nil { if err != nil {
t.Fatal(hostname, "register udp:", err) t.Fatal(hostname, "register udp:", err)
} }