Skip to content
Open
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
4 changes: 4 additions & 0 deletions packages/gateway-v2/discovery_handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -180,3 +180,7 @@ func writeRPCJSON(w http.ResponseWriter, status int, payload any) {
func writeRPCError(w http.ResponseWriter, status int, message string) {
writeRPCJSON(w, status, sshExecErrorResponse{Error: sshExecErrorBody{Message: message}})
}

func writeRPCErrorWithKind(w http.ResponseWriter, status int, message string, kind string) {
writeRPCJSON(w, status, sshExecErrorResponse{Error: sshExecErrorBody{Message: message, Kind: kind}})
}
33 changes: 27 additions & 6 deletions packages/gateway-v2/ssh_handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ type sshExecErrorResponse struct {

type sshExecErrorBody struct {
Message string `json:"message"`
// Set only by the test-connection handler; absent on every other RPC and on older gateways.
Kind string `json:"kind,omitempty"`
}

func parseSSHExecPrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
Expand All @@ -53,16 +55,27 @@ func parseSSHExecPrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
return ssh.ParsePrivateKey([]byte(privateKey))
}

func buildSSHExecAuth(env sshExecEnvelope) ([]ssh.AuthMethod, error) {
// The ssh package runs these callbacks only once the server offered the method, so onAttempt firing is what
// proves a credential was actually sent.
func buildSSHExecAuth(env sshExecEnvelope, onAttempt func()) ([]ssh.AuthMethod, error) {
if onAttempt == nil {
onAttempt = func() {}
}
switch env.AuthMethod {
case "password":
return []ssh.AuthMethod{ssh.Password(env.Password)}, nil
return []ssh.AuthMethod{ssh.PasswordCallback(func() (string, error) {
onAttempt()
return env.Password, nil
})}, nil
case "public-key":
signer, err := parseSSHExecPrivateKey(env.PrivateKey, env.Passphrase)
if err != nil {
return nil, fmt.Errorf("failed to parse private key: %w", err)
}
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
return []ssh.AuthMethod{ssh.PublicKeysCallback(func() ([]ssh.Signer, error) {
onAttempt()
return []ssh.Signer{signer}, nil
})}, nil
case "certificate":
signer, err := parseSSHExecPrivateKey(env.PrivateKey, env.Passphrase)
if err != nil {
Expand All @@ -80,14 +93,18 @@ func buildSSHExecAuth(env sshExecEnvelope) ([]ssh.AuthMethod, error) {
if err != nil {
return nil, fmt.Errorf("failed to create certificate signer: %w", err)
}
return []ssh.AuthMethod{ssh.PublicKeys(certSigner)}, nil
return []ssh.AuthMethod{ssh.PublicKeysCallback(func() ([]ssh.Signer, error) {
onAttempt()
return []ssh.Signer{certSigner}, nil
})}, nil
default:
return nil, fmt.Errorf("invalid auth method: %s", env.AuthMethod)
}
}

func doSSHExec(targetHost string, targetPort int, env sshExecEnvelope) (sshExecResult, error) {
authMethods, err := buildSSHExecAuth(env)
credentialOffered := false
authMethods, err := buildSSHExecAuth(env, func() { credentialOffered = true })
if err != nil {
return sshExecResult{}, err
}
Expand All @@ -104,7 +121,11 @@ func doSSHExec(targetHost string, targetPort int, env sshExecEnvelope) (sshExecR
Timeout: timeout,
})
if err != nil {
return sshExecResult{}, fmt.Errorf("failed to dial target SSH server: %w", err)
err = fmt.Errorf("failed to dial target SSH server: %w", err)
if credentialOffered {
return sshExecResult{}, authFailure(err)
}
return sshExecResult{}, connectFailure(err)
}
defer client.Close()

Expand Down
72 changes: 72 additions & 0 deletions packages/gateway-v2/test_connection_failure_kind.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
package gatewayv2

import (
"errors"
"io"
"net"
"os"
)

// A refused credential stops the heartbeat schedule; an unreachable target keeps retrying. Probes dial and then
// authenticate, so each tags the phase it failed in rather than the classification being read back out of the
// driver's error text, which would need new codes for every account type.
type testConnFailureKind string

const (
failureKindAuth testConnFailureKind = "auth"
failureKindTransport testConnFailureKind = "transport"
failureKindUnknown testConnFailureKind = "unknown"
)

type probeError struct {
kind testConnFailureKind
err error
}

func (e *probeError) Error() string { return e.err.Error() }
func (e *probeError) Unwrap() error { return e.err }

func connectFailure(err error) error {
if err == nil {
return nil
}
return &probeError{kind: failureKindTransport, err: err}
}

// A network error at this point is the connection dying mid-exchange, not the credential being refused.
func authFailure(err error) error {
if err == nil {
return nil
}
if isNetworkError(err) {
return &probeError{kind: failureKindTransport, err: err}
}
return &probeError{kind: failureKindAuth, err: err}
}

func isNetworkError(err error) bool {
var netErr net.Error
if errors.As(err, &netErr) {
return true
}
var opErr *net.OpError
if errors.As(err, &opErr) {
return true
}
var dnsErr *net.DNSError
if errors.As(err, &dnsErr) {
return true
}
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return true
}
return errors.Is(err, os.ErrDeadlineExceeded)
}

func classifyTestConnFailure(err error) testConnFailureKind {
var probeErr *probeError
if errors.As(err, &probeErr) {
return probeErr.kind
}
return failureKindUnknown
}
71 changes: 71 additions & 0 deletions packages/gateway-v2/test_connection_failure_kind_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
package gatewayv2

import (
"errors"
"fmt"
"net"
"testing"
)

func TestClassifyTestConnFailure(t *testing.T) {
cases := []struct {
name string
err error
want testConnFailureKind
}{
{
name: "dial failure",
err: connectFailure(&net.OpError{Op: "dial", Err: errors.New("connection refused")}),
want: failureKindTransport,
},
{
name: "refused credential",
err: authFailure(errors.New("password authentication failed for user \"pam\"")),
want: failureKindAuth,
},
{
name: "connection dropped mid-authentication",
err: authFailure(&net.OpError{Op: "read", Err: errors.New("connection reset by peer")}),
want: failureKindTransport,
},
{
name: "wrapped by a caller",
err: fmt.Errorf("test connection: %w", authFailure(errors.New("login failed"))),
want: failureKindAuth,
},
{
name: "untagged",
err: errors.New("unsupported SQL dialect"),
want: failureKindUnknown,
},
}

for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := classifyTestConnFailure(tc.err); got != tc.want {
t.Fatalf("classifyTestConnFailure() = %q, want %q", got, tc.want)
}
})
}
}

func TestSSHPhaseDependsOnCredentialBeingOffered(t *testing.T) {
offered := false
methods, err := buildSSHExecAuth(sshExecEnvelope{AuthMethod: "password", Password: "pw"}, func() { offered = true })
if err != nil {
t.Fatalf("buildSSHExecAuth: %v", err)
}
if len(methods) != 1 {
t.Fatalf("expected one auth method, got %d", len(methods))
}
if offered {
t.Fatal("building the auth method must not count as offering a credential")
}
}

func TestProbeErrorPreservesMessage(t *testing.T) {
const message = "redis authentication failed: WRONGPASS"
if got := authFailure(errors.New(message)).Error(); got != message {
t.Fatalf("Error() = %q, want %q", got, message)
}
}
Loading
Loading