Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions pkg/js/libs/smb/smb.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package smb
import (
"context"
"fmt"
"net"
"time"

"github.com/praetorian-inc/fingerprintx/pkg/plugins"
Expand Down Expand Up @@ -49,19 +50,24 @@ 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
}
// try to get SMBv2/v3 info
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
}
Expand Down
22 changes: 22 additions & 0 deletions pkg/js/libs/smb/smb_private.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
}
87 changes: 87 additions & 0 deletions pkg/js/libs/smb/smb_test.go
Original file line number Diff line number Diff line change
@@ -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
}
Loading