diff --git a/api/client/credentials_test.go b/api/client/credentials_test.go index d30a9c68ae6..6c832ac6b3c 100644 --- a/api/client/credentials_test.go +++ b/api/client/credentials_test.go @@ -190,7 +190,7 @@ func getExpectedTLSConfig(t *testing.T) *tls.Config { } func getExpectedSSHConfig(t *testing.T) *ssh.ClientConfig { - config, err := sshutils.SSHClientConfig(sshCert, keyPEM, [][]byte{sshCACert}) + config, err := sshutils.ProxyClientSSHConfig(sshCert, keyPEM, [][]byte{sshCACert}) require.NoError(t, err) return config diff --git a/api/client/identityfile.go b/api/client/identityfile.go index 3bcbb5d1a77..4e895aedb31 100644 --- a/api/client/identityfile.go +++ b/api/client/identityfile.go @@ -86,10 +86,11 @@ func (i *IdentityFile) TLSConfig() (*tls.Config, error) { // SSHClientConfig returns the identity file's associated SSHClientConfig. func (i *IdentityFile) SSHClientConfig() (*ssh.ClientConfig, error) { - ssh, err := sshutils.SSHClientConfig(i.Certs.SSH, i.PrivateKey, i.CACerts.SSH) + ssh, err := sshutils.ProxyClientSSHConfig(i.Certs.SSH, i.PrivateKey, i.CACerts.SSH) if err != nil { return nil, trace.Wrap(err) } + return ssh, nil } diff --git a/api/client/profile.go b/api/client/profile.go index 52faaa9826a..70678f701f5 100644 --- a/api/client/profile.go +++ b/api/client/profile.go @@ -131,7 +131,7 @@ func (p *Profile) SSHClientConfig() (*ssh.ClientConfig, error) { return nil, trace.Wrap(err) } - ssh, err := sshutils.SSHClientConfig(cert, key, [][]byte{caCerts}) + ssh, err := sshutils.ProxyClientSSHConfig(cert, key, [][]byte{caCerts}) if err != nil { return nil, trace.Wrap(err) } diff --git a/api/utils/sshutils/ssh.go b/api/utils/sshutils/ssh.go index 185bd22d8b5..7d7628db2a1 100644 --- a/api/utils/sshutils/ssh.go +++ b/api/utils/sshutils/ssh.go @@ -47,9 +47,12 @@ func ParseCertificate(buf []byte) (*ssh.Certificate, error) { return cert, nil } -// SSHClientConfig returns an ssh.ClientConfig with SSH credentials from this +// ProxyClientSSHConfig returns an ssh.ClientConfig with SSH credentials from this // Key and HostKeyCallback matching SSH CAs in the Key. -func SSHClientConfig(sshCert, privKey []byte, caCerts [][]byte) (*ssh.ClientConfig, error) { +// +// The config is set up to authenticate to proxy with the first available principal. +// +func ProxyClientSSHConfig(sshCert, privKey []byte, caCerts [][]byte) (*ssh.ClientConfig, error) { cert, err := ParseCertificate(sshCert) if err != nil { return nil, trace.Wrap(err, "failed to extract username from SSH certificate") @@ -57,16 +60,22 @@ func SSHClientConfig(sshCert, privKey []byte, caCerts [][]byte) (*ssh.ClientConf authMethod, err := AsAuthMethod(cert, privKey) if err != nil { - return nil, trace.Wrap(err, "failed to convert identity file to auth method") + return nil, trace.Wrap(err, "failed to convert key pair to auth method") } hostKeyCallback, err := HostKeyCallback(caCerts) if err != nil { - return nil, trace.Wrap(err, "failed to convert identity file to HostKeyCallback") + return nil, trace.Wrap(err, "failed to convert certificate authorities to HostKeyCallback") + } + + // The KeyId is not always a valid principal, so we use the first valid principal instead. + user := cert.KeyId + if len(cert.ValidPrincipals) > 0 { + user = cert.ValidPrincipals[0] } return &ssh.ClientConfig{ - User: cert.KeyId, + User: user, Auth: []ssh.AuthMethod{authMethod}, HostKeyCallback: hostKeyCallback, Timeout: defaults.DefaultDialTimeout, diff --git a/lib/client/interfaces.go b/lib/client/interfaces.go index 552d34335fa..6700d8d3dc4 100644 --- a/lib/client/interfaces.go +++ b/lib/client/interfaces.go @@ -218,10 +218,23 @@ func (k *Key) clientTLSConfig(cipherSuites []uint16, tlsCertRaw []byte) (*tls.Co return tlsConfig, nil } -// ClientSSHConfig returns an ssh.ClientConfig with SSH credentials from this +// ProxyClientSSHConfig returns an ssh.ClientConfig with SSH credentials from this // Key and HostKeyCallback matching SSH CAs in the Key. -func (k *Key) ClientSSHConfig() (*ssh.ClientConfig, error) { - return sshutils.SSHClientConfig(k.Cert, k.Priv, k.SSHCAs()) +// +// The config is set up to authenticate to proxy with the first available principal +// and ( if keyStore != nil ) trust local SSH CAs without asking for public keys. +// +func (k *Key) ProxyClientSSHConfig(keyStore LocalKeyStore) (*ssh.ClientConfig, error) { + sshConfig, err := sshutils.ProxyClientSSHConfig(k.Cert, k.Priv, k.SSHCAs()) + if err != nil { + return nil, trace.Wrap(err) + } + + if keyStore != nil { + sshConfig.HostKeyCallback = NewKeyStoreCertChecker(keyStore) + } + + return sshConfig, nil } // CertUsername returns the name of the Teleport user encoded in the SSH certificate. @@ -403,27 +416,6 @@ func (k *Key) HostKeyCallback() (ssh.HostKeyCallback, error) { return sshutils.HostKeyCallback(k.SSHCAs()) } -// ProxyClientSSHConfig returns an ssh.ClientConfig with SSH credentials from this -// Key and HostKeyCallback matching SSH CAs in the Key. -// -// The config is set up to authenticate to proxy with the first -// available principal and trust local SSH CAs without asking -// for public keys. -// -func ProxyClientSSHConfig(k *Key, keyStore LocalKeyStore) (*ssh.ClientConfig, error) { - sshConfig, err := k.ClientSSHConfig() - if err != nil { - return nil, trace.Wrap(err) - } - principals, err := k.CertPrincipals() - if err != nil { - return nil, trace.Wrap(err) - } - sshConfig.User = principals[0] - sshConfig.HostKeyCallback = NewKeyStoreCertChecker(keyStore) - return sshConfig, nil -} - // RootClusterName extracts the root cluster name from the issuer // of the Teleport TLS certificate. func (k *Key) RootClusterName() (string, error) { diff --git a/lib/client/keystore_test.go b/lib/client/keystore_test.go index d5c2a0e57b1..98c913faa05 100644 --- a/lib/client/keystore_test.go +++ b/lib/client/keystore_test.go @@ -227,7 +227,7 @@ func TestProxySSHConfig(t *testing.T) { err = s.store.AddKnownHostKeys("127.0.0.1", []ssh.PublicKey{caPub}) require.NoError(t, err) - clientConfig, err := ProxyClientSSHConfig(key, s.store) + clientConfig, err := key.ProxyClientSSHConfig(s.store) require.NoError(t, err) called := atomic.NewInt32(0) diff --git a/tool/tctl/common/tctl.go b/tool/tctl/common/tctl.go index 9fad049d67d..9c45fca26f5 100644 --- a/tool/tctl/common/tctl.go +++ b/tool/tctl/common/tctl.go @@ -384,7 +384,7 @@ func applyConfig(ccf *GlobalCLIFlags, cfg *service.Config) (*AuthServiceClientCo if err != nil { return nil, trace.Wrap(err) } - authConfig.SSH, err = key.ClientSSHConfig() + authConfig.SSH, err = key.ProxyClientSSHConfig(nil) if err != nil { return nil, trace.Wrap(err) } @@ -462,7 +462,7 @@ func loadConfigFromProfile(ccf *GlobalCLIFlags, cfg *service.Config) (*AuthServi return nil, trace.Wrap(err) } authConfig.TLS.InsecureSkipVerify = ccf.Insecure - authConfig.SSH, err = client.ProxyClientSSHConfig(key, keyStore) + authConfig.SSH, err = key.ProxyClientSSHConfig(keyStore) if err != nil { return nil, trace.Wrap(err) }