mirror of
https://github.com/soypat/lneto.git
synced 2026-08-23 07:59:05 +00:00
apply some of @ddirect suggestions
This commit is contained in:
+3
-11
@@ -7,6 +7,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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Handler struct {
|
type Handler struct {
|
||||||
@@ -143,7 +144,7 @@ func (h *Handler) StartQuery(dstHWAddr, proto []byte) error {
|
|||||||
return errors.New("bad protocol address length")
|
return errors.New("bad protocol address length")
|
||||||
} else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) {
|
} else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) {
|
||||||
return errors.New("mismatch hardware size")
|
return errors.New("mismatch hardware size")
|
||||||
} else if dstHWAddr != nil && !allZeros(dstHWAddr) {
|
} else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) {
|
||||||
return errors.New("write-to buffer must be zeroed out")
|
return errors.New("write-to buffer must be zeroed out")
|
||||||
}
|
}
|
||||||
h.queries = h.queries[:len(h.queries)+1]
|
h.queries = h.queries[:len(h.queries)+1]
|
||||||
@@ -238,7 +239,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
|
|||||||
if mac == nil && bytes.Equal(q.protoaddr, protoaddr) {
|
if mac == nil && bytes.Equal(q.protoaddr, protoaddr) {
|
||||||
q.hwaddr = append(q.hwaddr, hwaddr...)
|
q.hwaddr = append(q.hwaddr, hwaddr...)
|
||||||
if q.dstHw != nil {
|
if q.dstHw != nil {
|
||||||
if !allZeros(q.dstHw) {
|
if !internal.IsZeroed(q.dstHw...) {
|
||||||
slog.Error("race-condition:ARP-reused-buffer")
|
slog.Error("race-condition:ARP-reused-buffer")
|
||||||
}
|
}
|
||||||
copy(q.dstHw, hwaddr) // External write to user buffer.
|
copy(q.dstHw, hwaddr) // External write to user buffer.
|
||||||
@@ -258,12 +259,3 @@ func trySetEthernetDst(ethFrame []byte, dst []byte) {
|
|||||||
copy(ethFrame[:6], dst)
|
copy(ethFrame[:6], dst)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func allZeros(b []byte) bool {
|
|
||||||
for i := range b {
|
|
||||||
if b[i] != 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -359,7 +359,7 @@ func (s *Stack) StartLookupIP(host string) error {
|
|||||||
var u internet.StackUDPPort
|
var u internet.StackUDPPort
|
||||||
dns4 := dnsSrvs.As4()
|
dns4 := dnsSrvs.As4()
|
||||||
u.SetStackNode(&s.dns, dns4[:], dns.ServerPort)
|
u.SetStackNode(&s.dns, dns4[:], dns.ServerPort)
|
||||||
err = s.udps.Register(&u, nil)
|
err = s.udps.Register(&u)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -405,7 +405,7 @@ func (s *Stack) BeginDHCPRequest(request [4]byte) error {
|
|||||||
}
|
}
|
||||||
var u internet.StackUDPPort
|
var u internet.StackUDPPort
|
||||||
u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort)
|
u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort)
|
||||||
err = s.udps.Register(&u, nil)
|
err = s.udps.Register(&u)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -417,7 +417,7 @@ func (s *Stack) StartNTP(addr netip.Addr) error {
|
|||||||
var u internet.StackUDPPort
|
var u internet.StackUDPPort
|
||||||
addr4 := addr.As4()
|
addr4 := addr.As4()
|
||||||
u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort)
|
u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort)
|
||||||
err := s.udps.Register(&u, nil)
|
err := s.udps.Register(&u)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
err = stack.tcpports.Register(&listener, nil) // Passive TCP requires no MAC setting.
|
err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting.
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
err = stack.tcpports.Register(conn, nil) // Passive MAC with no listening.
|
err = stack.tcpports.Register(conn) // Passive MAC with no listening.
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) {
|
|||||||
copy(dstaddr, dst)
|
copy(dstaddr, dst)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// IsZeroed returns true if all arguments are set to their zero value.
|
||||||
|
func IsZeroed[T comparable](a ...T) bool {
|
||||||
|
var z T
|
||||||
|
for i := range a {
|
||||||
|
if a[i] != z {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|||||||
+4
-11
@@ -165,6 +165,7 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro
|
|||||||
err := node.demux(buf, offset)
|
err := node.demux(buf, offset)
|
||||||
if h.tryHandleError(node, err) {
|
if h.tryHandleError(node, err) {
|
||||||
err = nil
|
err = nil
|
||||||
|
node = nil // Node is destroyed in tryHandleError and invalidated.
|
||||||
}
|
}
|
||||||
return node, err
|
return node, err
|
||||||
}
|
}
|
||||||
@@ -174,12 +175,13 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro
|
|||||||
func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) {
|
func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) {
|
||||||
for i := range h.nodes {
|
for i := range h.nodes {
|
||||||
node := &h.nodes[i]
|
node := &h.nodes[i]
|
||||||
if node.IsInvalid() || (len(node.remoteAddr) > 0 && isAllZeros(node.remoteAddr)) {
|
if node.IsInvalid() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
n, err = node.encapsulate(buf, offsetIP, offsetThisFrame)
|
n, err = node.encapsulate(buf, offsetIP, offsetThisFrame)
|
||||||
if h.tryHandleError(node, err) {
|
if h.tryHandleError(node, err) {
|
||||||
err = nil // CLOSE error handled gracefully by deleting node.
|
err = nil // CLOSE error handled gracefully by deleting node.
|
||||||
|
node = nil // Node is destroyed in tryHandleError and invalidated.
|
||||||
}
|
}
|
||||||
if n > 0 {
|
if n > 0 {
|
||||||
return node, n, err
|
return node, n, err
|
||||||
@@ -191,15 +193,6 @@ func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_
|
|||||||
return nil, 0, err // Return last written error.
|
return nil, 0, err // Return last written error.
|
||||||
}
|
}
|
||||||
|
|
||||||
func isAllZeros(b []byte) bool {
|
|
||||||
for i := range b {
|
|
||||||
if b[i] != 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
var (
|
||||||
errZeroMaxNodesArg = errors.New("zero max nodes arg")
|
errZeroMaxNodesArg = errors.New("zero max nodes arg")
|
||||||
errZeroPort = errors.New("port must be greater than zero")
|
errZeroPort = errors.New("port must be greater than zero")
|
||||||
|
|||||||
+76
-9
@@ -4,11 +4,13 @@ import (
|
|||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
|
"log/slog"
|
||||||
"math"
|
"math"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
"github.com/soypat/lneto"
|
"github.com/soypat/lneto"
|
||||||
"github.com/soypat/lneto/ethernet"
|
"github.com/soypat/lneto/ethernet"
|
||||||
|
"github.com/soypat/lneto/internal"
|
||||||
)
|
)
|
||||||
|
|
||||||
type StackPorts struct {
|
type StackPorts struct {
|
||||||
@@ -16,6 +18,7 @@ type StackPorts struct {
|
|||||||
handlers handlers
|
handlers handlers
|
||||||
dstPortOff uint16
|
dstPortOff uint16
|
||||||
protocol uint16
|
protocol uint16
|
||||||
|
// stores last node to demux/encapsulate.
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ps *StackPorts) ResetUDP(maxNodes int) error {
|
func (ps *StackPorts) ResetUDP(maxNodes int) error {
|
||||||
@@ -52,12 +55,7 @@ func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame
|
|||||||
if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) {
|
if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) {
|
||||||
return 0, io.ErrShortBuffer
|
return 0, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
var node *node
|
_, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame)
|
||||||
node, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame)
|
|
||||||
if n > 0 && len(node.remoteAddr) == 6 && offsetToIP >= 14 {
|
|
||||||
efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:])
|
|
||||||
*efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr)
|
|
||||||
}
|
|
||||||
return n, err
|
return n, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -72,15 +70,84 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) {
|
|||||||
|
|
||||||
// Register registers a port StackNode on StackPorts.
|
// Register registers a port StackNode on StackPorts.
|
||||||
// If dstMAC is set to non-nil, length six buffer then
|
// If dstMAC is set to non-nil, length six buffer then
|
||||||
func (ps *StackPorts) Register(h StackNode, dstMAC []byte) error {
|
func (ps *StackPorts) Register(h StackNode) error {
|
||||||
port := h.LocalPort()
|
port := h.LocalPort()
|
||||||
proto := h.Protocol()
|
proto := h.Protocol()
|
||||||
if port <= 0 {
|
if port <= 0 {
|
||||||
return errZeroPort
|
return errZeroPort
|
||||||
} else if proto != uint64(ps.protocol) {
|
} else if proto != uint64(ps.protocol) {
|
||||||
return errInvalidProto
|
return errInvalidProto
|
||||||
} else if dstMAC != nil && len(dstMAC) != 6 {
|
}
|
||||||
|
return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 errors.New("invalid MAC")
|
||||||
}
|
}
|
||||||
return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, dstMAC))
|
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.
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -250,7 +250,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network.
|
err = s.tcps.Register(conn) // MAC is set later on by ARP response arriving to our network.
|
||||||
if err != nil {
|
if err != nil {
|
||||||
conn.Abort()
|
conn.Abort()
|
||||||
return err
|
return err
|
||||||
@@ -265,7 +265,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
err = s.tcps.Register(conn, nil)
|
err = s.tcps.Register(conn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
conn.Abort()
|
conn.Abort()
|
||||||
return err
|
return err
|
||||||
@@ -306,7 +306,7 @@ func (s *StackAsync) StartLookupIP(host string) error {
|
|||||||
}
|
}
|
||||||
dns4 := s.dnssv.As4()
|
dns4 := s.dnssv.As4()
|
||||||
s.dnsUDP.SetStackNode(&s.dns, dns4[:], dns.ServerPort)
|
s.dnsUDP.SetStackNode(&s.dns, dns4[:], dns.ServerPort)
|
||||||
err = s.udps.Register(&s.dnsUDP, nil)
|
err = s.udps.Register(&s.dnsUDP)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -353,7 +353,7 @@ func (s *StackAsync) StartDHCPv4Request(request [4]byte) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
s.dhcpUDP.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort)
|
s.dhcpUDP.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort)
|
||||||
err = s.udps.Register(&s.dhcpUDP, nil)
|
err = s.udps.Register(&s.dhcpUDP)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -367,7 +367,7 @@ func (s *StackAsync) StartNTP(addr netip.Addr) error {
|
|||||||
|
|
||||||
addr4 := addr.As4()
|
addr4 := addr.As4()
|
||||||
s.ntpUDP.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort)
|
s.ntpUDP.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort)
|
||||||
err := s.udps.Register(&s.ntpUDP, nil)
|
err := s.udps.Register(&s.ntpUDP)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user