diff --git a/gosnmp.go b/gosnmp.go index 90781fd..88c912f 100644 --- a/gosnmp.go +++ b/gosnmp.go @@ -266,105 +266,6 @@ const ( // Public Functions (main interface) // -// Connect creates and opens a socket. Because UDP is a connectionless -// protocol, you won't know if the remote host is responding until you send -// packets. Neither will you know if the host is regularly disappearing and reappearing. -// -// For historical reasons (ie this is part of the public API), the method won't -// be renamed to Dial(). -func (x *GoSNMP) Connect() error { - return x.connect("") -} - -// ConnectIPv4 forces an IPv4-only connection -func (x *GoSNMP) ConnectIPv4() error { - return x.connect("4") -} - -// ConnectIPv6 forces an IPv6-only connection -func (x *GoSNMP) ConnectIPv6() error { - return x.connect("6") -} - -// Close closes the underlaying connection. -// -// This method is safe to call multiple times and from concurrent goroutines. -// Only the first call will close the connection; subsequent calls are no-ops. -func (x *GoSNMP) Close() error { - x.mu.Lock() - defer x.mu.Unlock() - - if x.Conn == nil { - return nil - } - - err := x.Conn.Close() - x.Conn = nil - return err -} - -// connect to address addr on the given network -// -// https://golang.org/pkg/net/#Dial gives acceptable network values as: -// -// "tcp", "tcp4" (IPv4-only), "tcp6" (IPv6-only), "udp", "udp4" (IPv4-only),"udp6" (IPv6-only), "ip", -// "ip4" (IPv4-only), "ip6" (IPv6-only), "unix", "unixgram" and "unixpacket" -func (x *GoSNMP) connect(networkSuffix string) error { - err := x.validateParameters() - if err != nil { - return err - } - - x.Transport += networkSuffix - if err = x.netConnect(); err != nil { - return fmt.Errorf("error establishing connection to host: %w", err) - } - - if err = x.seedIDs(); err != nil { - return err - } - - x.rxBuf = new([rxBufSize]byte) - - return nil -} - -// Performs the real socket opening network operation. This can be used to do a -// reconnect (needed for TCP) -func (x *GoSNMP) netConnect() error { - var err error - var localAddr net.Addr - addr := net.JoinHostPort(x.Target, strconv.Itoa(int(x.Port))) - - switch x.Transport { - case "udp", "udp4", "udp6": - if localAddr, err = net.ResolveUDPAddr(x.Transport, x.LocalAddr); err != nil { - return err - } - if addr4 := localAddr.(*net.UDPAddr).IP.To4(); addr4 != nil { - x.Transport = "udp4" - } - if x.UseUnconnectedUDPSocket { - x.uaddr, err = net.ResolveUDPAddr(x.Transport, addr) - if err != nil { - return err - } - x.Conn, err = net.ListenUDP(x.Transport, localAddr.(*net.UDPAddr)) - return err - } - case "tcp", "tcp4", "tcp6": - if localAddr, err = net.ResolveTCPAddr(x.Transport, x.LocalAddr); err != nil { - return err - } - if addr4 := localAddr.(*net.TCPAddr).IP.To4(); addr4 != nil { - x.Transport = "tcp4" - } - } - dialer := net.Dialer{Timeout: x.Timeout, LocalAddr: localAddr, Control: x.Control} - x.Conn, err = dialer.DialContext(x.Context, x.Transport, addr) - return err -} - func (x *GoSNMP) validateParameters() error { if x.Transport == "" { x.Transport = udp diff --git a/marshal.go b/marshal.go index 5ecdc53..b1feb5e 100644 --- a/marshal.go +++ b/marshal.go @@ -10,7 +10,6 @@ import ( "errors" "fmt" "io" - "net" "runtime" "strings" "time" @@ -243,13 +242,7 @@ sendRetry: if x.Logger.enabled() { x.Logger.Printf("SENDING PACKET: %s", packetOut.SafeString()) } - // If using UDP and unconnected socket, send packet directly to stored address. - if uconn, ok := x.Conn.(net.PacketConn); ok && x.uaddr != nil { - _, err = uconn.WriteTo(outBuf, x.uaddr) - } else { - _, err = x.Conn.Write(outBuf) - } - if err != nil { + if err = x.write(outBuf); err != nil { continue } if x.OnSent != nil { @@ -481,30 +474,3 @@ func (x *GoSNMP) send(packetOut *SnmpPacket) (result *SnmpPacket, err error) { func (packet *SnmpPacket) MarshalMsg() ([]byte, error) { return packet.marshalMsg() } - -// receive response from network and read into a byte array -func (x *GoSNMP) receive() ([]byte, error) { - var n int - var err error - // If we are using UDP and unconnected socket, read the packet and - // disregard the source address. - if uconn, ok := x.Conn.(net.PacketConn); ok { - n, _, err = uconn.ReadFrom(x.rxBuf[:]) - } else { - n, err = x.Conn.Read(x.rxBuf[:]) - } - if err == io.EOF { - return nil, err - } else if err != nil { - return nil, fmt.Errorf("error reading from socket: %w", err) - } - - if n == rxBufSize { - // This should never happen unless we're using something like a unix domain socket. - return nil, fmt.Errorf("response buffer too small") - } - - resp := make([]byte, n) - copy(resp, x.rxBuf[:n]) - return resp, nil -} diff --git a/transport.go b/transport.go new file mode 100644 index 0000000..d3e28f4 --- /dev/null +++ b/transport.go @@ -0,0 +1,149 @@ +// Copyright 2012 The GoSNMP Authors. All rights reserved. Use of this +// source code is governed by a BSD-style license that can be found in the +// LICENSE file. + +package gosnmp + +import ( + "fmt" + "io" + "net" + "strconv" +) + +// Connect creates and opens a socket. Because UDP is a connectionless +// protocol, you won't know if the remote host is responding until you send +// packets. Neither will you know if the host is regularly disappearing and reappearing. +// +// For historical reasons (ie this is part of the public API), the method won't +// be renamed to Dial(). +func (x *GoSNMP) Connect() error { + return x.connect("") +} + +// ConnectIPv4 forces an IPv4-only connection +func (x *GoSNMP) ConnectIPv4() error { + return x.connect("4") +} + +// ConnectIPv6 forces an IPv6-only connection +func (x *GoSNMP) ConnectIPv6() error { + return x.connect("6") +} + +// Close closes the underlaying connection. +// +// This method is safe to call multiple times and from concurrent goroutines. +// Only the first call will close the connection; subsequent calls are no-ops. +func (x *GoSNMP) Close() error { + x.mu.Lock() + defer x.mu.Unlock() + + if x.Conn == nil { + return nil + } + + err := x.Conn.Close() + x.Conn = nil + return err +} + +// connect to address addr on the given network +// +// https://golang.org/pkg/net/#Dial gives acceptable network values as: +// +// "tcp", "tcp4" (IPv4-only), "tcp6" (IPv6-only), "udp", "udp4" (IPv4-only),"udp6" (IPv6-only), "ip", +// "ip4" (IPv4-only), "ip6" (IPv6-only), "unix", "unixgram" and "unixpacket" +func (x *GoSNMP) connect(networkSuffix string) error { + err := x.validateParameters() + if err != nil { + return err + } + + x.Transport += networkSuffix + if err = x.netConnect(); err != nil { + return fmt.Errorf("error establishing connection to host: %w", err) + } + + if err = x.seedIDs(); err != nil { + return err + } + + x.rxBuf = new([rxBufSize]byte) + + return nil +} + +// Performs the real socket opening network operation. This can be used to do a +// reconnect (needed for TCP) +func (x *GoSNMP) netConnect() error { + var err error + var localAddr net.Addr + addr := net.JoinHostPort(x.Target, strconv.Itoa(int(x.Port))) + + switch x.Transport { + case "udp", "udp4", "udp6": + if localAddr, err = net.ResolveUDPAddr(x.Transport, x.LocalAddr); err != nil { + return err + } + if addr4 := localAddr.(*net.UDPAddr).IP.To4(); addr4 != nil { + x.Transport = "udp4" + } + if x.UseUnconnectedUDPSocket { + x.uaddr, err = net.ResolveUDPAddr(x.Transport, addr) + if err != nil { + return err + } + x.Conn, err = net.ListenUDP(x.Transport, localAddr.(*net.UDPAddr)) + return err + } + case "tcp", "tcp4", "tcp6": + if localAddr, err = net.ResolveTCPAddr(x.Transport, x.LocalAddr); err != nil { + return err + } + if addr4 := localAddr.(*net.TCPAddr).IP.To4(); addr4 != nil { + x.Transport = "tcp4" + } + } + dialer := net.Dialer{Timeout: x.Timeout, LocalAddr: localAddr, Control: x.Control} + x.Conn, err = dialer.DialContext(x.Context, x.Transport, addr) + return err +} + +// write sends b to the agent: to its address when the client uses an +// unconnected UDP socket, on the connection otherwise. +func (x *GoSNMP) write(b []byte) error { + if uconn, ok := x.Conn.(net.PacketConn); ok && x.uaddr != nil { + _, err := uconn.WriteTo(b, x.uaddr) + return err + } + _, err := x.Conn.Write(b) + return err +} + +// receive response from network and read into a byte array +func (x *GoSNMP) receive() ([]byte, error) { + var n int + var err error + // If we are using UDP and unconnected socket, read the packet and + // disregard the source address. + if uconn, ok := x.Conn.(net.PacketConn); ok { + n, _, err = uconn.ReadFrom(x.rxBuf[:]) + } else { + n, err = x.Conn.Read(x.rxBuf[:]) + } + if err == io.EOF { + return nil, err + } else if err != nil { + return nil, fmt.Errorf("error reading from socket: %w", err) + } + + if n == rxBufSize { + // This should never happen unless we're using something like a unix domain socket. + return nil, fmt.Errorf("response buffer too small") + } + + resp := make([]byte, n) + copy(resp, x.rxBuf[:n]) + return resp, nil +} diff --git a/transport_test.go b/transport_test.go new file mode 100644 index 0000000..c3bb55e --- /dev/null +++ b/transport_test.go @@ -0,0 +1,172 @@ +// Copyright 2026 Netdata Inc. All rights reserved. Use of this +// source code is governed by a BSD-style license that can be found in the +// LICENSE file. + +package gosnmp + +import ( + "context" + "errors" + "net" + "runtime" + "strconv" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// connectResult is what TestConnect observes of a Connect. +type connectResult struct { + err string // "" for no error + transport string + connected bool // Conn is set + local string // the IP of Conn's local address + controls []string +} + +// TestConnect pins how Connect, ConnectIPv4 and ConnectIPv6 validate the +// client, choose the network and dial, known bugs included, on loopback +// sockets. +func TestConnect(t *testing.T) { + udpPort := unusedUDPPort(t) + tcp, err := net.Listen("tcp4", "127.0.0.1:0") + require.NoError(t, err) + defer tcp.Close() + tcpPort := tcp.Addr().(*net.TCPAddr).Port + + recordControl := func(x *GoSNMP, got *connectResult, sleep time.Duration) { + x.Control = func(network, _ string, _ syscall.RawConn) error { + got.controls = append(got.controls, network) + time.Sleep(sleep) + return nil + } + } + canceled, cancel := context.WithCancel(context.Background()) + cancel() + + tests := map[string]struct { + setup func(x *GoSNMP, got *connectResult) + connect func(x *GoSNMP) error + unix bool // the dial checks its context after a UDP connect only on Unix + want connectResult + }{ + "udp": { + want: connectResult{transport: "udp", connected: true, local: "127.0.0.1"}, + }, + "udp with an IPv4 local address": { + setup: func(x *GoSNMP, _ *connectResult) { x.LocalAddr = "127.0.0.1:0" }, + want: connectResult{transport: "udp4", connected: true, local: "127.0.0.1"}, + }, + "tcp with an IPv4 local address": { + setup: func(x *GoSNMP, _ *connectResult) { + x.Transport, x.LocalAddr, x.Port = "tcp", "127.0.0.1:0", uint16(tcpPort) //nolint:gosec // a loopback port + }, + want: connectResult{transport: "tcp4", connected: true, local: "127.0.0.1"}, + }, + "ConnectIPv4": { + connect: (*GoSNMP).ConnectIPv4, + want: connectResult{transport: "udp4", connected: true, local: "127.0.0.1"}, + }, + "ConnectIPv4 twice": { + connect: func(x *GoSNMP) error { + if err := x.ConnectIPv4(); err != nil { + return err + } + _ = x.Conn.Close() + return x.ConnectIPv4() + }, + want: connectResult{ + err: "error establishing connection to host: dial udp44: unknown network udp44", + transport: "udp44", connected: false, + }, + }, + "ConnectIPv6 to an IPv4 target": { + connect: (*GoSNMP).ConnectIPv6, + want: connectResult{ + err: "error establishing connection to host: dial udp6: address 127.0.0.1: no suitable address found", + transport: "udp6", + }, + }, + "invalid parameters": { + setup: func(x *GoSNMP, _ *connectResult) { x.MaxOids = -1 }, + want: connectResult{err: "field MaxOids cannot be less than 0", transport: "udp"}, + }, + "Control": { + setup: func(x *GoSNMP, got *connectResult) { recordControl(x, got, 0) }, + want: connectResult{transport: "udp", connected: true, local: "127.0.0.1", controls: []string{"udp4"}}, + }, + "dial timeout": { + unix: true, + setup: func(x *GoSNMP, got *connectResult) { + x.Timeout = 20 * time.Millisecond + recordControl(x, got, 200*time.Millisecond) + }, + want: connectResult{ + err: "error establishing connection to host: dial udp :0->127.0.0.1:" + strconv.Itoa(int(udpPort)) + ": i/o timeout", + transport: "udp", controls: []string{"udp4"}, + }, + }, + "canceled context": { + setup: func(x *GoSNMP, _ *connectResult) { x.Context = canceled }, + want: connectResult{ + err: "error establishing connection to host: dial udp :0->127.0.0.1:" + strconv.Itoa(int(udpPort)) + ": operation was canceled", + transport: "udp", + }, + }, + } + for name, tc := range tests { + t.Run(name, func(t *testing.T) { + if tc.unix && runtime.GOOS == "windows" { + t.Skip("Windows dials UDP without checking the context after connect") + } + x := &GoSNMP{Target: "127.0.0.1", Port: udpPort, Version: Version2c, Community: "public", Timeout: time.Second} + var got connectResult + if tc.setup != nil { + tc.setup(x, &got) + } + connect := tc.connect + if connect == nil { + connect = (*GoSNMP).Connect + } + if err := connect(x); err != nil { + got.err = err.Error() + } + got.transport = x.Transport + if x.Conn != nil { + got.connected = true + host, _, err := net.SplitHostPort(x.Conn.LocalAddr().String()) + require.NoError(t, err) + got.local = host + _ = x.Conn.Close() + } + assert.Equal(t, tc.want, got) + }) + } +} + +// closeErrConn is a connection whose Close fails. +type closeErrConn struct { + net.Conn + closes int +} + +var errCloseConn = errors.New("codec close failure") + +func (c *closeErrConn) Close() error { + c.closes++ + return errCloseConn +} + +// TestCloseOnce pins that Close closes the connection once, returns its error +// and leaves no connection, and that a client without one closes cleanly. +func TestCloseOnce(t *testing.T) { + c := &closeErrConn{} + x := &GoSNMP{Conn: c} + assert.ErrorIs(t, x.Close(), errCloseConn) + assert.Nil(t, x.Conn) + assert.NoError(t, x.Close()) + assert.Equal(t, 1, c.closes) +}