diff --git a/pkg/js/libs/smb/smb.go b/pkg/js/libs/smb/smb.go index c5f85824eb..803074f2b0 100644 --- a/pkg/js/libs/smb/smb.go +++ b/pkg/js/libs/smb/smb.go @@ -3,6 +3,7 @@ package smb import ( "context" "fmt" + "net" "time" "github.com/praetorian-inc/fingerprintx/pkg/plugins" @@ -49,7 +50,11 @@ func connectSMBInfoMode(ctx context.Context, executionId string, host string, po if dialer == nil { return nil, fmt.Errorf("dialers not initialized for %s", executionId) } - conn, err := dialer.Fastdialer.Dial(ctx, "tcp", fmt.Sprintf("%s:%d", host, port)) + address := net.JoinHostPort(host, fmt.Sprintf("%d", port)) + dialSMBInfo := func(ctx context.Context) (net.Conn, error) { + return dialer.Fastdialer.Dial(ctx, "tcp", address) + } + conn, err := dialSMBInfo(ctx) if err != nil { return nil, err } @@ -57,11 +62,12 @@ func connectSMBInfoMode(ctx context.Context, executionId string, host string, po result, err := getSMBInfo(conn, true, false) _ = conn.Close() // close regardless of error if err == nil { + updateSMBv1Support(ctx, result, dialSMBInfo) return result, nil } // try to negotiate SMBv1 - conn, err = dialer.Fastdialer.Dial(ctx, "tcp", fmt.Sprintf("%s:%d", host, port)) + conn, err = dialSMBInfo(ctx) if err != nil { return nil, err } diff --git a/pkg/js/libs/smb/smb_private.go b/pkg/js/libs/smb/smb_private.go index 955ff0f72a..2aee04e6a4 100644 --- a/pkg/js/libs/smb/smb_private.go +++ b/pkg/js/libs/smb/smb_private.go @@ -12,6 +12,8 @@ import ( zgrabsmb "github.com/zmap/zgrab2/lib/smb/smb" ) +type smbInfoDialFunc func(context.Context) (net.Conn, error) + // ==== private helper functions/methods ==== // collectSMBv2Metadata collects metadata for SMBv2 services. @@ -53,3 +55,23 @@ func getSMBInfo(conn net.Conn, setupSession, v1 bool) (*zgrabsmb.SMBLog, error) } return result, nil } + +func updateSMBv1Support(ctx context.Context, result *zgrabsmb.SMBLog, dial smbInfoDialFunc) { + if result == nil || result.SupportV1 { + return + } + + conn, err := dial(ctx) + if err != nil { + return + } + defer func() { + _ = conn.Close() + }() + + v1Result, err := getSMBInfo(conn, false, true) + if err != nil || v1Result == nil || !v1Result.SupportV1 { + return + } + result.SupportV1 = true +} diff --git a/pkg/js/libs/smb/smb_test.go b/pkg/js/libs/smb/smb_test.go new file mode 100644 index 0000000000..856809c8e9 --- /dev/null +++ b/pkg/js/libs/smb/smb_test.go @@ -0,0 +1,87 @@ +package smb + +import ( + "context" + "encoding/binary" + "errors" + "io" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" + zgrabsmb "github.com/zmap/zgrab2/lib/smb/smb" +) + +func TestUpdateSMBv1SupportPreservesNegotiatedSMB2Version(t *testing.T) { + version := &zgrabsmb.SMBVersions{ + Major: 2, + Minor: 1, + VerString: "SMB 2.1", + } + result := &zgrabsmb.SMBLog{ + Version: version, + } + dial, done := newSMBv1ProbeDialer() + + updateSMBv1Support(context.Background(), result, dial) + + require.True(t, result.SupportV1) + require.Same(t, version, result.Version) + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("SMBv1 probe was not performed") + } +} + +func TestUpdateSMBv1SupportIgnoresProbeErrors(t *testing.T) { + result := &zgrabsmb.SMBLog{} + + updateSMBv1Support(context.Background(), result, func(context.Context) (net.Conn, error) { + return nil, errors.New("dial failed") + }) + + require.False(t, result.SupportV1) +} + +func newSMBv1ProbeDialer() (smbInfoDialFunc, <-chan error) { + done := make(chan error, 1) + return func(context.Context) (net.Conn, error) { + clientConn, serverConn := net.Pipe() + go func() { + defer close(done) + done <- serveSMBv1Probe(serverConn) + }() + return clientConn, nil + }, done +} + +func serveSMBv1Probe(conn net.Conn) error { + defer func() { + _ = conn.Close() + }() + + var requestSize uint32 + if err := binary.Read(conn, binary.BigEndian, &requestSize); err != nil { + return err + } + request := make([]byte, requestSize) + if _, err := io.ReadFull(conn, request); err != nil { + return err + } + if len(request) < 4 { + return io.ErrUnexpectedEOF + } + if string(request[:4]) != zgrabsmb.ProtocolSmb { + return errors.New("expected SMBv1 negotiate request") + } + + response := []byte(zgrabsmb.ProtocolSmb) + if err := binary.Write(conn, binary.BigEndian, uint32(len(response))); err != nil { + return err + } + _, err := conn.Write(response) + return err +}