diff --git a/api/breaker/breaker.go b/api/breaker/breaker.go index 43e0653a28b..4130bcc8723 100644 --- a/api/breaker/breaker.go +++ b/api/breaker/breaker.go @@ -21,7 +21,7 @@ import ( "time" "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/trace" "github.com/jonboulle/clockwork" @@ -263,7 +263,7 @@ func (c *Config) CheckAndSetDefaults() error { }) } - c.TrippedPeriod = utils.NewSeventhJitter()(c.TrippedPeriod) + c.TrippedPeriod = retryutils.NewSeventhJitter()(c.TrippedPeriod) return nil } diff --git a/api/go.mod b/api/go.mod index 2b3c850e893..dca671ab349 100644 --- a/api/go.mod +++ b/api/go.mod @@ -3,6 +3,7 @@ module github.com/gravitational/teleport/api go 1.18 require ( + github.com/go-piv/piv-go v1.10.0 github.com/gogo/protobuf v1.3.2 github.com/golang/protobuf v1.5.2 github.com/google/go-cmp v0.5.8 diff --git a/api/go.sum b/api/go.sum index 975243a0c12..8aec050edcc 100644 --- a/api/go.sum +++ b/api/go.sum @@ -76,6 +76,8 @@ github.com/go-logr/logr v1.2.3 h1:2DntVwHkVopvECVRSlL5PSo9eG+cAkDCuckLubN+rq0= github.com/go-logr/logr v1.2.3/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A= github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/go-piv/piv-go v1.10.0 h1:P1Y1VjBI5DnXW0+YkKmTuh5opWnMIrKriUaIOblee9Q= +github.com/go-piv/piv-go v1.10.0/go.mod h1:NZ2zmjVkfFaL/CF8cVQ/pXdXtuj110zEKGdJM6fJZZM= github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= diff --git a/api/identityfile/identityfile.go b/api/identityfile/identityfile.go index 261e7f43d93..2ac4596ca8f 100644 --- a/api/identityfile/identityfile.go +++ b/api/identityfile/identityfile.go @@ -25,6 +25,7 @@ import ( "fmt" "io" "os" + "regexp" "strings" "github.com/gravitational/teleport/api/utils/keypaths" @@ -261,7 +262,7 @@ func decodeIdentityFile(idFile io.Reader) (*IdentityFile, error) { // are copied out of the scanner's buffer. All others are ignored. for scanln() { switch { - case hasPrefix("ssh"): + case isSSHCert(line): ident.Certs.SSH = cloneln() case hasPrefix("@cert-authority"): ident.CACerts.SSH = append(ident.CACerts.SSH, cloneln()) @@ -301,3 +302,15 @@ func decodeIdentityFile(idFile io.Reader) (*IdentityFile, error) { } return &ident, nil } + +// OpenSSH cert types look like "-cert-v@openssh.com". +// Currently, we only use "ssh-rsa-cert-v01@openssh.com" & "ecdsa-sha2-nistp256-cert-v01@openssh.com". +var sshCertTypeRegex = regexp.MustCompile(`^[a-z0-9\-]+-cert-v[0-9]{2}@openssh\.com$`) + +// Check if the given data has an ssh cert type prefix as it's first part. +func isSSHCert(data []byte) bool { + // ssh certs should look like " ", + // so we check if the first element matches a known ssh cert type. + sshCertType := bytes.Split(data, []byte(" "))[0] + return sshCertTypeRegex.Match(sshCertType) +} diff --git a/api/identityfile/identityfile_test.go b/api/identityfile/identityfile_test.go index 108e7dabae7..dbea91944f9 100644 --- a/api/identityfile/identityfile_test.go +++ b/api/identityfile/identityfile_test.go @@ -15,15 +15,15 @@ limitations under the License. */ -package identityfile_test +package identityfile import ( "os" "path/filepath" "testing" - "github.com/gravitational/teleport/api/identityfile" "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" ) // TestIdentityFileBasics verifies basic profile operations such as @@ -31,33 +31,61 @@ import ( func TestIdentityFileBasics(t *testing.T) { t.Parallel() path := filepath.Join(t.TempDir(), "file") - writeIDFile := &identityfile.IdentityFile{ + writeIDFile := &IdentityFile{ PrivateKey: []byte("-----BEGIN RSA PRIVATE KEY-----\nkey\n-----END RSA PRIVATE KEY-----\n"), - Certs: identityfile.Certs{ - SSH: []byte("ssh ssh-cert"), + Certs: Certs{ + SSH: []byte(ssh.CertAlgoRSAv01), TLS: []byte("-----BEGIN CERTIFICATE-----\ntls-cert\n-----END CERTIFICATE-----\n"), }, - CACerts: identityfile.CACerts{ + CACerts: CACerts{ SSH: [][]byte{[]byte("@cert-authority ssh-cacerts")}, TLS: [][]byte{[]byte("-----BEGIN CERTIFICATE-----\ntls-cacerts\n-----END CERTIFICATE-----\n")}, }, } // Write identity file - err := identityfile.Write(writeIDFile, path) + err := Write(writeIDFile, path) require.NoError(t, err) // Read identity file from file - readIDFile, err := identityfile.ReadFile(path) + readIDFile, err := ReadFile(path) require.NoError(t, err) // Read identity file from string s, err := os.ReadFile(path) require.NoError(t, err) - fromStringIDFile, err := identityfile.FromString(string(s)) + fromStringIDFile, err := FromString(string(s)) require.NoError(t, err) // Check that read and write values are equal require.Equal(t, writeIDFile, readIDFile) require.Equal(t, writeIDFile, fromStringIDFile) } + +func TestIsSSHCert(t *testing.T) { + for _, tc := range []struct { + certType string + expectBool bool + }{ + { + certType: "opensesame@openssh.com", + expectBool: false, + }, { + certType: ssh.CertAlgoRSAv01, + expectBool: true, + }, { + certType: ssh.CertAlgoECDSA256v01, + expectBool: true, + }, { + certType: ssh.CertAlgoED25519v01, + expectBool: true, + }, + } { + t.Run(tc.certType, func(t *testing.T) { + certData := append([]byte(tc.certType), []byte(" AAAA...")...) + isSSHCert := isSSHCert(certData) + require.Equal(t, tc.expectBool, isSSHCert) + }) + } + +} diff --git a/api/utils/jitter.go b/api/utils/jitter.go deleted file mode 100644 index f52d34ea481..00000000000 --- a/api/utils/jitter.go +++ /dev/null @@ -1,41 +0,0 @@ -// Copyright 2022 Gravitational, Inc -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package utils - -import ( - "math/rand" - "sync" - "time" -) - -type Jitter func(time.Duration) time.Duration - -// NewSeventhJitter builds a new jitter on the range [6n/7,n). Prefer smaller -// jitters such as this when jittering periodic operations (e.g. cert rotation -// checks) since large jitters result in significantly increased load. -func NewSeventhJitter() Jitter { - var mu sync.Mutex - rng := rand.New(rand.NewSource(time.Now().UnixNano())) - return func(d time.Duration) time.Duration { - // values less than 1 cause rng to panic, and some logic - // relies on treating zero duration as non-blocking case. - if d < 1 { - return 0 - } - mu.Lock() - defer mu.Unlock() - return (6 * d / 7) + time.Duration(rng.Int63n(int64(d))/7) - } -} diff --git a/api/utils/keys/privatekey.go b/api/utils/keys/privatekey.go index 109c329c35a..d5876036857 100644 --- a/api/utils/keys/privatekey.go +++ b/api/utils/keys/privatekey.go @@ -20,6 +20,8 @@ package keys import ( "bytes" "crypto" + "crypto/ecdsa" + "crypto/ed25519" "crypto/rsa" "crypto/tls" "crypto/x509" @@ -27,29 +29,38 @@ import ( "fmt" "os" - "github.com/gravitational/teleport/api/utils/sshutils/ppk" "github.com/gravitational/trace" "golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh/agent" + + "github.com/gravitational/teleport/api/utils/sshutils/ppk" ) const ( - PKCS1PrivateKeyType = "RSA PRIVATE KEY" - PKCS8PrivateKeyType = "PRIVATE KEY" - ECPrivateKeyType = "EC PRIVATE KEY" + PKCS1PrivateKeyType = "RSA PRIVATE KEY" + PKCS8PrivateKeyType = "PRIVATE KEY" + ECPrivateKeyType = "EC PRIVATE KEY" + pivYubiKeyPrivateKeyType = "PIV YUBIKEY PRIVATE KEY" ) +type cryptoPublicKeyI interface { + Equal(x crypto.PublicKey) bool +} + // PrivateKey implements crypto.Signer with additional helper methods. The underlying // private key may be a standard crypto.Signer implemented in the standard library // (aka *rsa.PrivateKey, *ecdsa.PrivateKey, or ed25519.PrivateKey), or it may be a // custom implementation for a non-standard private key, such as a hardware key. type PrivateKey struct { - Signer + crypto.Signer + // sshPub is the public key in ssh.PublicKey form. sshPub ssh.PublicKey + // keyPEM is PEM-encoded private key data which can be parsed with ParsePrivateKey. + keyPEM []byte } // NewPrivateKey returns a new PrivateKey for the given crypto.Signer. -func NewPrivateKey(signer Signer) (*PrivateKey, error) { +func NewPrivateKey(signer crypto.Signer, keyPEM []byte) (*PrivateKey, error) { sshPub, err := ssh.NewPublicKey(signer.Public()) if err != nil { return nil, trace.Wrap(err) @@ -58,6 +69,7 @@ func NewPrivateKey(signer Signer) (*PrivateKey, error) { return &PrivateKey{ Signer: signer, sshPub: sshPub, + keyPEM: keyPEM, }, nil } @@ -71,6 +83,59 @@ func (k *PrivateKey) MarshalSSHPublicKey() []byte { return ssh.MarshalAuthorizedKey(k.sshPub) } +// PrivateKeyPEM returns PEM encoded private key data. This may be data necessary +// to retrieve the key, such as a YubiKey serial number and slot, or it can be a +// PKCS marshaled private key. +// +// The resulting PEM encoded data should only be decoded with ParsePrivateKey to +// prevent errors from parsing non PKCS marshaled keys, such as a PIV key. +func (k *PrivateKey) PrivateKeyPEM() []byte { + return k.keyPEM +} + +// TLSCertificate parses the given TLS certificate(s) paired with the private key +// to rerturn a tls.Certificate, ready to be used in a TLS handshake. +func (k *PrivateKey) TLSCertificate(certPEMBlock []byte) (tls.Certificate, error) { + cert := tls.Certificate{ + PrivateKey: k.Signer, + } + + var skippedBlockTypes []string + for { + var certDERBlock *pem.Block + certDERBlock, certPEMBlock = pem.Decode(certPEMBlock) + if certDERBlock == nil { + break + } + if certDERBlock.Type == "CERTIFICATE" { + cert.Certificate = append(cert.Certificate, certDERBlock.Bytes) + } else { + skippedBlockTypes = append(skippedBlockTypes, certDERBlock.Type) + } + } + + if len(cert.Certificate) == 0 { + if len(skippedBlockTypes) == 0 { + return tls.Certificate{}, trace.BadParameter("tls: failed to find any PEM data in certificate input") + } + return tls.Certificate{}, trace.BadParameter("tls: failed to find \"CERTIFICATE\" PEM block in certificate input after skipping PEM blocks of the following types: %v", skippedBlockTypes) + } + + // Check that the certificate's public key matches this private key. + x509Cert, err := x509.ParseCertificate(cert.Certificate[0]) + if err != nil { + return tls.Certificate{}, trace.Wrap(err) + } + + if keyPub, ok := k.Public().(cryptoPublicKeyI); !ok { + return tls.Certificate{}, trace.BadParameter("private key does not contain a valid public key") + } else if !keyPub.Equal(x509Cert.PublicKey) { + return tls.Certificate{}, trace.BadParameter("private key does not match certificate's public key") + } + + return cert, nil +} + // agentKeyComment is used to generate an agent key comment. type agentKeyComment struct { user string @@ -83,8 +148,18 @@ func (a *agentKeyComment) String() string { // AsAgentKey converts PrivateKey to a agent.AddedKey. If the given PrivateKey is not // supported as an agent key, a trace.NotImplemented error is returned. func (k *PrivateKey) AsAgentKey(sshCert *ssh.Certificate) (agent.AddedKey, error) { - signer, ok := k.Signer.(*StandardSigner) - if !ok { + switch k.Signer.(type) { + case *rsa.PrivateKey, *ecdsa.PrivateKey, ed25519.PrivateKey: + // put a teleport identifier along with the teleport user into the comment field + comment := agentKeyComment{user: sshCert.KeyId} + return agent.AddedKey{ + PrivateKey: k.Signer, + Certificate: sshCert, + Comment: comment.String(), + LifetimeSecs: 0, + ConfirmBeforeUse: false, + }, nil + default: // We return a not implemented error because agent.AddedKey only // supports plain RSA, ECDSA, and ED25519 keys. Non-standard private // keys, like hardware-based private keys, will require custom solutions @@ -93,25 +168,11 @@ func (k *PrivateKey) AsAgentKey(sshCert *ssh.Certificate) (agent.AddedKey, error // caller the ability to handle the error gracefully. return agent.AddedKey{}, trace.NotImplemented("cannot create an agent key using private key signer of type %T", k.Signer) } - - // put a teleport identifier along with the teleport user into the comment field - comment := agentKeyComment{user: sshCert.KeyId} - return agent.AddedKey{ - PrivateKey: signer.Signer, - Certificate: sshCert, - Comment: comment.String(), - LifetimeSecs: 0, - ConfirmBeforeUse: false, - }, nil } // PPKFile returns a PuTTY PPK-formatted keypair func (k *PrivateKey) PPKFile() ([]byte, error) { - signer, ok := k.Signer.(*StandardSigner) - if !ok { - return nil, trace.BadParameter("cannot use private key of type %T as rsa.PrivateKey", k) - } - rsaKey, ok := signer.Signer.(*rsa.PrivateKey) + rsaKey, ok := k.Signer.(*rsa.PrivateKey) if !ok { return nil, trace.BadParameter("cannot use private key of type %T as rsa.PrivateKey", k) } @@ -128,21 +189,10 @@ func (k *PrivateKey) PPKFile() ([]byte, error) { // This is used by some integrations which currently only support raw RSA private keys, // like Kubernetes, MongoDB, and PPK files for windows. func (k *PrivateKey) RSAPrivateKeyPEM() ([]byte, error) { - signer := k.GetBaseSigner() - if _, ok := signer.(*rsa.PrivateKey); !ok { - return nil, trace.BadParameter("cannot get rsa key PEM for private key of type %T", signer) - } - return k.PrivateKeyPEM(), nil -} - -// GetBaseSigner is a helper method to return the actual nested crypto.Signer for this PrivateKey. -func (k *PrivateKey) GetBaseSigner() crypto.Signer { - switch signer := k.Signer.(type) { - case *StandardSigner: - return signer.Signer - default: - return signer + if _, ok := k.Signer.(*rsa.PrivateKey); !ok { + return nil, trace.BadParameter("cannot get rsa key PEM for private key of type %T", k.Signer) } + return k.keyPEM, nil } // LoadPrivateKey returns the PrivateKey for the given key file. @@ -172,13 +222,13 @@ func ParsePrivateKey(keyPEM []byte) (*PrivateKey, error) { if err != nil { return nil, trace.Wrap(err) } - return NewPrivateKey(newStandardSigner(cryptoSigner, keyPEM)) + return NewPrivateKey(cryptoSigner, keyPEM) case ECPrivateKeyType: cryptoSigner, err := x509.ParseECPrivateKey(block.Bytes) if err != nil { return nil, trace.Wrap(err) } - return NewPrivateKey(newStandardSigner(cryptoSigner, keyPEM)) + return NewPrivateKey(cryptoSigner, keyPEM) case PKCS8PrivateKeyType: priv, err := x509.ParsePKCS8PrivateKey(block.Bytes) if err != nil { @@ -188,7 +238,13 @@ func ParsePrivateKey(keyPEM []byte) (*PrivateKey, error) { if !ok { return nil, trace.BadParameter("x509.ParsePKCS8PrivateKey returned an invalid private key of type %T", priv) } - return NewPrivateKey(newStandardSigner(cryptoSigner, keyPEM)) + return NewPrivateKey(cryptoSigner, keyPEM) + case pivYubiKeyPrivateKeyType: + priv, err := parseYubiKeyPrivateKeyData(block.Bytes) + if err != nil { + return nil, trace.Wrap(err) + } + return NewPrivateKey(priv, keyPEM) default: return nil, trace.BadParameter("unexpected private key PEM type %q", block.Type) } diff --git a/api/utils/keys/privatekey_test.go b/api/utils/keys/privatekey_test.go index c21cacc6dd3..17f48c52a17 100644 --- a/api/utils/keys/privatekey_test.go +++ b/api/utils/keys/privatekey_test.go @@ -17,10 +17,12 @@ limitations under the License. package keys import ( + "bytes" "crypto/ecdsa" "crypto/ed25519" "crypto/rsa" "crypto/tls" + "encoding/pem" "testing" "github.com/gravitational/trace" @@ -58,7 +60,7 @@ func TestParsePrivateKey(t *testing.T) { assertKey: func(tt require.TestingT, key interface{}, i2 ...interface{}) { privateKey, ok := key.(*PrivateKey) require.True(t, ok) - require.IsType(t, &rsa.PrivateKey{}, privateKey.GetBaseSigner()) + require.IsType(t, &rsa.PrivateKey{}, privateKey.Signer) }, }, { @@ -68,7 +70,7 @@ func TestParsePrivateKey(t *testing.T) { assertKey: func(tt require.TestingT, key interface{}, i2 ...interface{}) { privateKey, ok := key.(*PrivateKey) require.True(t, ok) - require.IsType(t, &ecdsa.PrivateKey{}, privateKey.GetBaseSigner()) + require.IsType(t, &ecdsa.PrivateKey{}, privateKey.Signer) }, }, { @@ -78,7 +80,7 @@ func TestParsePrivateKey(t *testing.T) { assertKey: func(tt require.TestingT, key interface{}, i2 ...interface{}) { privateKey, ok := key.(*PrivateKey) require.True(t, ok) - require.IsType(t, ed25519.PrivateKey{}, privateKey.GetBaseSigner()) + require.IsType(t, ed25519.PrivateKey{}, privateKey.Signer) }, }, } { @@ -92,13 +94,38 @@ func TestParsePrivateKey(t *testing.T) { // TestX509KeyPair tests that X509KeyPair returns the same value as tls.X509KeyPair. func TestX509KeyPair(t *testing.T) { - expectCert, err := tls.X509KeyPair(rsaCertPEM, rsaKeyPEM) - require.NoError(t, err) + for _, tc := range []struct { + desc string + keyPEM []byte + certPEM []byte + }{ + { + desc: "rsa cert", + keyPEM: rsaKeyPEM, + certPEM: rsaCertPEM, + }, { + desc: "rsa certs", + keyPEM: rsaKeyPEM, + certPEM: func() []byte { + // encode two certs into certPEM. + rsaCertPEMDuplicated := new(bytes.Buffer) + der, _ := pem.Decode(rsaCertPEM) + pem.Encode(rsaCertPEMDuplicated, der) + pem.Encode(rsaCertPEMDuplicated, der) + return rsaCertPEMDuplicated.Bytes() + }(), + }, + } { + t.Run(tc.desc, func(t *testing.T) { + expectCert, err := tls.X509KeyPair(tc.certPEM, tc.keyPEM) + require.NoError(t, err) - tlsCert, err := X509KeyPair(rsaCertPEM, rsaKeyPEM) - require.NoError(t, err) + tlsCert, err := X509KeyPair(tc.certPEM, tc.keyPEM) + require.NoError(t, err) - require.Equal(t, expectCert, tlsCert) + require.Equal(t, expectCert, tlsCert) + }) + } } var ( diff --git a/api/utils/keys/signer.go b/api/utils/keys/signer.go deleted file mode 100644 index 15705a9016a..00000000000 --- a/api/utils/keys/signer.go +++ /dev/null @@ -1,87 +0,0 @@ -/* -Copyright 2022 Gravitational, Inc. - -Licensed under the Apache License, Version 2.0 (the "License"); -you may not use this file except in compliance with the License. -You may obtain a copy of the License at - - http://www.apache.org/licenses/LICENSE-2.0 - -Unless required by applicable law or agreed to in writing, software -distributed under the License is distributed on an "AS IS" BASIS, -WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -See the License for the specific language governing permissions and -limitations under the License. -*/ - -package keys - -import ( - "crypto" - "crypto/tls" - "crypto/x509" - "encoding/pem" - - "github.com/gravitational/trace" -) - -// Signer implements crypto.Signer with additional helper methods. -type Signer interface { - crypto.Signer - - // PrivateKeyPEM returns PEM encoded private key data. This may be data necessary - // to retrieve the key, such as a Yubikey serial number and slot, or it can be a - // PKCS marshaled private key. - // - // The resulting PEM encoded data should only be decoded with ParsePrivateKey to - // prevent errors from parsing non PKCS marshaled keys, such as a PIV key. - PrivateKeyPEM() []byte - - // TLSCertificate parses the given TLS certificate paired with the private key - // to rerturn a tls.Certificate, ready to be used in a TLS handshake. - TLSCertificate(tlsCert []byte) (tls.Certificate, error) -} - -// StandardSigner is a shared Signer implementation for standard crypto.PrivateKey -// implemenations, which are *rsa.PrivateKey, *ecdsa.PrivateKey, and ed25519.PrivateKey. -type StandardSigner struct { - // Signer is an *rsa.PrivateKey, *ecdsa.PrivateKey, or ed25519.PrivateKey. - crypto.Signer - // keyPEM is the PEM-encoded private key. - keyPEM []byte -} - -// NewStandardSigner creates a new StandardSigner from the given *rsa.PrivateKey, *ecdsa.PrivateKey or ed25519.PrivateKey. -func NewStandardSigner(signer crypto.Signer) (*StandardSigner, error) { - keyDER, err := x509.MarshalPKCS8PrivateKey(signer) - if err != nil { - return nil, trace.Wrap(err) - } - - keyPEM := pem.EncodeToMemory(&pem.Block{ - Type: PKCS8PrivateKeyType, - Headers: nil, - Bytes: keyDER, - }) - - return newStandardSigner(signer, keyPEM), nil -} - -func newStandardSigner(signer crypto.Signer, keyPEM []byte) *StandardSigner { - return &StandardSigner{ - Signer: signer, - keyPEM: keyPEM, - } -} - -// PrivateKeyPEM returns the PEM-encoded private key. -func (s *StandardSigner) PrivateKeyPEM() []byte { - return s.keyPEM -} - -// TLSCertificate parses the given TLS certificate paired with the private key -// to return a tls.Certificate, ready to be used in a TLS handshake. -func (s *StandardSigner) TLSCertificate(certRaw []byte) (tls.Certificate, error) { - cert, err := tls.X509KeyPair(certRaw, s.keyPEM) - return cert, trace.Wrap(err) -} diff --git a/api/utils/keys/yubikey.go b/api/utils/keys/yubikey.go new file mode 100644 index 00000000000..430bd905b26 --- /dev/null +++ b/api/utils/keys/yubikey.go @@ -0,0 +1,414 @@ +//go:build libpcsclite +// +build libpcsclite + +/* +Copyright 2022 Gravitational, Inc. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package keys + +import ( + "context" + "crypto" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/json" + "encoding/pem" + "io" + "math/big" + "strings" + "time" + + "github.com/go-piv/piv-go/piv" + "github.com/gravitational/trace" + + "github.com/gravitational/teleport/api" + "github.com/gravitational/teleport/api/utils/retryutils" +) + +const ( + // PIVCardTypeYubiKey is the PIV card type assigned to yubiKeys. + PIVCardTypeYubiKey = "yubikey" +) + +var ( + // We use slot 9a for Teleport Clients which require `private_key_policy: hardware_key`. + pivSlotNoTouch = piv.SlotAuthentication + // We use slot 9c for Teleport Clients which require `private_key_policy: hardware_key_touch`. + pivSlotWithTouch = piv.SlotSignature +) + +// GetOrGenerateYubiKeyPrivateKey connects to a connected yubiKey and gets a private key +// matching the given touch requirement. This private key will either be newly generated +// or previously generated by a Teleport client and reused. +func GetOrGenerateYubiKeyPrivateKey(ctx context.Context, touchRequired bool) (*PrivateKey, error) { + // Use the first yubiKey we find. + y, err := findYubiKey(ctx, 0) + if err != nil { + return nil, trace.Wrap(err) + } + + // Get the correct PIV slot and Touch policy for the given touch requirement: + // - No Touch = 9a + TouchPolicyNever + // - Touch = 9c + TouchPolicyCached + pivSlot := pivSlotNoTouch + touchPolicy := piv.TouchPolicyNever + if touchRequired { + pivSlot = pivSlotWithTouch + touchPolicy = piv.TouchPolicyCached + } + + // First, check if there is already a private key set up by a Teleport Client. + priv, err := y.getPrivateKey(ctx, pivSlot) + if err != nil { + // Generate a new private key on the PIV slot. + if priv, err = y.generatePrivateKey(ctx, pivSlot, touchPolicy); err != nil { + return nil, trace.Wrap(err) + } + } + + keyPEM, err := priv.keyPEM() + if err != nil { + return nil, trace.Wrap(err) + } + + return NewPrivateKey(priv, keyPEM) +} + +// YubiKeyPrivateKey is a YubiKey PIV private key. Cryptographical operations open +// a new temporary connection to the PIV card to perform the operation. +type YubiKeyPrivateKey struct { + // yubiKey is a specific yubiKey PIV module. + *yubiKey + pivSlot piv.Slot + pub crypto.PublicKey + // ctx is used when opening a connection to the PIV module, + // which occurs with a retry loop. + ctx context.Context +} + +// yubiKeyPrivateKeyData is marshalable data used to retrieve a specific yubiKey PIV private key. +type yubiKeyPrivateKeyData struct { + SerialNumber uint32 `json:"serial_number"` + SlotKey uint32 `json:"slot_key"` +} + +func newYubiKeyPrivateKey(ctx context.Context, y *yubiKey, slot piv.Slot, pub crypto.PublicKey) (*YubiKeyPrivateKey, error) { + return &YubiKeyPrivateKey{ + yubiKey: y, + pivSlot: slot, + pub: pub, + ctx: ctx, + }, nil +} + +func parseYubiKeyPrivateKeyData(keyDataBytes []byte) (*YubiKeyPrivateKey, error) { + // TODO (Joerger): rather than requiring a context be passed here, we should + // pre-load the yubikey PIV connection to avoid retry/context logic occurring + // at spontaneous points in the code (anywhere a private key is used). + ctx := context.TODO() + + var keyData yubiKeyPrivateKeyData + if err := json.Unmarshal(keyDataBytes, &keyData); err != nil { + return nil, trace.Wrap(err) + } + + pivSlot, err := parsePIVSlot(keyData.SlotKey) + if err != nil { + return nil, trace.Wrap(err) + } + + y, err := findYubiKey(ctx, keyData.SerialNumber) + if err != nil { + return nil, trace.Wrap(err) + } + + priv, err := y.getPrivateKey(ctx, pivSlot) + if err != nil { + return nil, trace.Wrap(err) + } + + return priv, nil +} + +// Public returns the public key corresponding to this private key. +func (y *YubiKeyPrivateKey) Public() crypto.PublicKey { + return y.pub +} + +// Sign implements crypto.Signer. +func (y *YubiKeyPrivateKey) Sign(rand io.Reader, digest []byte, opts crypto.SignerOpts) (signature []byte, err error) { + yk, err := y.open(y.ctx) + if err != nil { + return nil, trace.Wrap(err) + } + defer yk.Close() + + privateKey, err := yk.PrivateKey(pivSlotNoTouch, y.pub, piv.KeyAuth{}) + if err != nil { + return nil, trace.Wrap(err) + } + + return privateKey.(crypto.Signer).Sign(rand, digest, opts) +} + +func (y *YubiKeyPrivateKey) keyPEM() ([]byte, error) { + keyDataBytes, err := json.Marshal(yubiKeyPrivateKeyData{ + SerialNumber: y.serialNumber, + SlotKey: y.pivSlot.Key, + }) + if err != nil { + return nil, trace.Wrap(err) + } + + return pem.EncodeToMemory(&pem.Block{ + Type: pivYubiKeyPrivateKeyType, + Headers: nil, + Bytes: keyDataBytes, + }), nil +} + +// yubiKey is a specific yubiKey PIV card. +type yubiKey struct { + // card is a reader name used to find and connect to this yubiKey. + // This value may change between OS's, or with other system changes. + card string + // serialNumber is the yubiKey's 8 digit serial number. + serialNumber uint32 +} + +func newYubiKey(ctx context.Context, card string) (*yubiKey, error) { + y := &yubiKey{card: card} + + yk, err := y.open(ctx) + if err != nil { + return nil, trace.Wrap(err) + } + defer yk.Close() + + y.serialNumber, err = yk.Serial() + if err != nil { + return nil, trace.Wrap(err) + } + + return y, nil +} + +// generatePrivateKey generates a new private key from the given PIV slot with the given PIV policies. +func (y *yubiKey) generatePrivateKey(ctx context.Context, slot piv.Slot, touchPolicy piv.TouchPolicy) (*YubiKeyPrivateKey, error) { + yk, err := y.open(ctx) + if err != nil { + return nil, trace.Wrap(err) + } + defer yk.Close() + + opts := piv.Key{ + Algorithm: piv.AlgorithmEC256, + PINPolicy: piv.PINPolicyNever, + TouchPolicy: touchPolicy, + } + pub, err := yk.GenerateKey(piv.DefaultManagementKey, slot, opts) + if err != nil { + return nil, trace.Wrap(err) + } + + // Create a self signed certificate and store it in the PIV slot so that other + // Teleport Clients know to reuse the stored key instead of genearting a new one. + priv, err := yk.PrivateKey(slot, pub, piv.KeyAuth{}) + if err != nil { + return nil, trace.Wrap(err) + } + + cert, err := selfSignedTeleportClientCertificate(priv, pub) + if err != nil { + return nil, trace.Wrap(err) + } + + // Store a self-signed certificate to mark this slot as used by tsh. + if err = yk.SetCertificate(piv.DefaultManagementKey, slot, cert); err != nil { + return nil, trace.Wrap(err) + } + + return newYubiKeyPrivateKey(ctx, y, slot, pub) +} + +// getPrivateKey gets an existing private key from the given PIV slot. +func (y *yubiKey) getPrivateKey(ctx context.Context, slot piv.Slot) (*YubiKeyPrivateKey, error) { + yk, err := y.open(ctx) + if err != nil { + return nil, trace.Wrap(err) + } + defer yk.Close() + + // Check the slot's certificate to see if it contains a self signed Teleport Client cert. + cert, err := yk.Certificate(slot) + if err != nil || cert == nil { + return nil, trace.NotFound("YubiKey certificate slot is empty, expected a Teleport Client cert") + } else if len(cert.Subject.Organization) == 0 || cert.Subject.Organization[0] != certOrgName { + return nil, trace.NotFound("YubiKey certificate slot contained unknown certificate:\n%+v", cert) + } + + // Attest the key to make sure it hasn't been imported. + slotCert, err := yk.Attest(slot) + if err != nil { + return nil, trace.Wrap(err) + } + attestationCert, err := yk.AttestationCertificate() + if err != nil { + return nil, trace.Wrap(err) + } + if _, err = piv.Verify(attestationCert, slotCert); err != nil { + return nil, trace.Wrap(err) + } + + // Verify that the slot's certs have the same public key, otherwise the key + // may have been generated by a non-teleport client. + if pubComparer, ok := cert.PublicKey.(cryptoPublicKeyI); !ok { + return nil, trace.BadParameter("certificate's public key of type %T is not a supported public key", cert.PublicKey) + } else if !pubComparer.Equal(slotCert.PublicKey) { + return nil, trace.NotFound("YubiKey slot contains mismatched certificates and must be regenerated") + } + + return newYubiKeyPrivateKey(ctx, y, slot, slotCert.PublicKey) +} + +// open a connection to YubiKey PIV module. The returned connection should be closed once +// it's been used. The YubiKey PIV module itself takes some additional time to handle closed +// connections, so we use a retry loop to give the PIV module time to close prior connections. +func (y *yubiKey) open(ctx context.Context) (yk *piv.YubiKey, err error) { + linearRetry, err := retryutils.NewLinear(retryutils.LinearConfig{ + // If a PIV connection has just been closed, it take ~5-10 ms to become + // available to new connections. For this reason, we initially wait a + // short 20ms before stepping up to a longer 100ms retry. + First: time.Millisecond * 20, + Step: time.Millisecond * 100, + // Since PIV modules only allow a single connection, it is a bottleneck + // resource. To maximise usage, we use a short 100ms retry to catch the + // connection opening up as soon as possible. + Max: time.Millisecond * 100, + }) + if err != nil { + return nil, trace.Wrap(err) + } + + // Backoff and retry for up to 10 seconds. On login, Teleport Connect tries to open several, + // maybe even hundreds, of connections to the PIV module all at once to load available resources, + // so a long retry period is necessary. + // + // TODO (joerger): Reduce this retry period to something more reasonable, like 1 second, + // and add a way for `tsh` and Teleport Connect to share a single connection to a PIV module. + retryCtx, cancel := context.WithTimeout(ctx, time.Second*10) + defer cancel() + + err = linearRetry.For(retryCtx, func() error { + yk, err = piv.Open(y.card) + if err != nil && !isRetryError(err) { + return retryutils.PermanentRetryError(err) + } + return trace.Wrap(err) + }) + if err != nil { + return nil, trace.Wrap(err) + } + return yk, nil +} + +func isRetryError(err error) bool { + const retryError = "connecting to smart card: the smart card cannot be accessed because of other connections outstanding" + return strings.Contains(err.Error(), retryError) +} + +// findYubiKey finds a yubiKey PIV card by serial number. If no serial +// number is provided, the first yubiKey found will be returned. +func findYubiKey(ctx context.Context, serialNumber uint32) (*yubiKey, error) { + yubiKeyCards, err := findYubiKeyCards() + if err != nil { + return nil, trace.Wrap(err) + } + + if len(yubiKeyCards) == 0 { + return nil, trace.NotFound("no yubiKey devices found") + } + + for _, card := range yubiKeyCards { + y, err := newYubiKey(ctx, card) + if err != nil { + return nil, trace.Wrap(err) + } + + if serialNumber == 0 || y.serialNumber == serialNumber { + return y, nil + } + } + + return nil, trace.NotFound("no yubiKey device found with serial number %q", serialNumber) +} + +// findYubiKeyCards returns a list of connected yubiKey PIV card names. +func findYubiKeyCards() ([]string, error) { + cards, err := piv.Cards() + if err != nil { + return nil, trace.Wrap(err) + } + + var yubiKeyCards []string + for _, card := range cards { + if strings.Contains(strings.ToLower(card), PIVCardTypeYubiKey) { + yubiKeyCards = append(yubiKeyCards, card) + } + } + + return yubiKeyCards, nil +} + +func parsePIVSlot(slotKey uint32) (piv.Slot, error) { + switch slotKey { + case piv.SlotAuthentication.Key: + return piv.SlotAuthentication, nil + case piv.SlotSignature.Key: + return piv.SlotSignature, nil + case piv.SlotCardAuthentication.Key: + return piv.SlotCardAuthentication, nil + case piv.SlotKeyManagement.Key: + return piv.SlotKeyManagement, nil + default: + retiredSlot, ok := piv.RetiredKeyManagementSlot(slotKey) + if !ok { + return piv.Slot{}, trace.BadParameter("slot %X does not exist", slotKey) + } + return retiredSlot, nil + } +} + +// certOrgName is used to identify Teleport Client self-signed certificates stored in yubiKey PIV slots. +const certOrgName = "teleport" + +func selfSignedTeleportClientCertificate(priv crypto.PrivateKey, pub crypto.PublicKey) (*x509.Certificate, error) { + serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128) + serialNumber, err := rand.Int(rand.Reader, serialNumberLimit) // see crypto/tls/generate_cert.go + if err != nil { + return nil, trace.Wrap(err) + } + cert := &x509.Certificate{ + SerialNumber: serialNumber, + PublicKey: pub, + Subject: pkix.Name{ + Organization: []string{certOrgName}, + OrganizationalUnit: []string{api.Version}, + }, + } + if cert.Raw, err = x509.CreateCertificate(rand.Reader, cert, cert, pub, priv); err != nil { + return nil, trace.Wrap(err) + } + return cert, nil +} diff --git a/api/utils/keys/yubikey_other.go b/api/utils/keys/yubikey_other.go new file mode 100644 index 00000000000..0221e975537 --- /dev/null +++ b/api/utils/keys/yubikey_other.go @@ -0,0 +1,35 @@ +//go:build !libpcsclite +// +build !libpcsclite + +/* +Copyright 2022 Gravitational, Inc. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package keys + +import ( + "context" + "crypto" + "errors" + + "github.com/gravitational/trace" +) + +var errPIVUnavailable = errors.New("PIV is unavailable in current build") + +func GetOrGenerateYubiKeyPrivateKey(ctx context.Context, touchRequired bool) (*PrivateKey, error) { + return nil, trace.Wrap(errPIVUnavailable) +} + +func parseYubiKeyPrivateKeyData(keyDataBytes []byte) (crypto.Signer, error) { + return nil, trace.Wrap(errPIVUnavailable) +} diff --git a/api/utils/keys/yubikey_test.go b/api/utils/keys/yubikey_test.go new file mode 100644 index 00000000000..83b17ae399a --- /dev/null +++ b/api/utils/keys/yubikey_test.go @@ -0,0 +1,61 @@ +//go:build libpcsclite +// +build libpcsclite + +/* +Copyright 2022 Gravitational, Inc. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package keys + +import ( + "context" + "os" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestGetOrGenerateYubiKeyPrivateKey tests GetOrGenerateYubiKeyPrivateKey. +func TestGetOrGenerateYubiKeyPrivateKey(t *testing.T) { + // This test expects a yubiKey to be connected with default PIV settings and will overwrite any PIV data on the yubiKey. + if os.Getenv("TELEPORT_TEST_YUBIKEY_PIV") == "" { + t.Skipf("Skipping TestGenerateYubiKeyPrivateKey because TELEPORT_TEST_YUBIKEY_PIV is not set") + } + + ctx := context.Background() + + // Connect to the first yubiKey and reset it. + y, err := findYubiKey(ctx, 0) + require.NoError(t, err) + yk, err := y.open(ctx) + require.NoError(t, err) + require.NoError(t, yk.Reset()) + require.NoError(t, yk.Close()) + + // Generate a new YubiKeyPrivateKey. + priv, err := GetOrGenerateYubiKeyPrivateKey(ctx, false) + require.NoError(t, err) + + // Test creating a self signed certificate with the key. + _, err = selfSignedTeleportClientCertificate(priv, priv.Public()) + require.NoError(t, err) + + // Another call to GetOrGenerateYubiKeyPrivateKey should retrieve the previously generated key. + retrievePriv, err := GetOrGenerateYubiKeyPrivateKey(ctx, false) + require.NoError(t, err) + require.Equal(t, priv, retrievePriv) + + // parsing the key's private key PEM should produce the same key as well. + retrieveKey, err := ParsePrivateKey(priv.PrivateKeyPEM()) + require.NoError(t, err) + require.Equal(t, priv, retrieveKey) +} diff --git a/api/utils/retryutils/jitter.go b/api/utils/retryutils/jitter.go new file mode 100644 index 00000000000..0414bc35e32 --- /dev/null +++ b/api/utils/retryutils/jitter.go @@ -0,0 +1,90 @@ +// Copyright 2022 Gravitational, Inc +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package retryutils + +import ( + "math/rand" + "sync" + "time" + + "github.com/gravitational/trace" +) + +// Jitter is a function which applies random jitter to a +// duration. Used to randomize backoff values. Must be +// safe for concurrent usage. +type Jitter func(time.Duration) time.Duration + +// NewJitter builds a new default jitter (currently jitters on +// the range [d/2,d), but this is subject to change). +func NewJitter() Jitter { + return NewHalfJitter() +} + +// NewFullJitter builds a new jitter on the range [0,d). Most use-cases +// are better served by a jitter with a meaningful minimum value, but if +// the *only* purpose of the jitter is to spread out retries to the greatest +// extent possible (e.g. when retrying a CompareAndSwap operation), a full jitter +// may be appropriate. +func NewFullJitter() Jitter { + rng := rand.New(rand.NewSource(time.Now().UnixNano())) + jitter, _ := newJitter(1, rng) + return jitter +} + +// NewHalfJitter returns a new jitter on the range [d/2,d). This is +// a large range and most suitable for jittering things like backoff +// operations where breaking cycles quickly is a priority. +func NewHalfJitter() Jitter { + rng := rand.New(rand.NewSource(time.Now().UnixNano())) + jitter, _ := newJitter(2, rng) + return jitter +} + +// NewSeventhJitter builds a new jitter on the range [6d/7,d). Prefer smaller +// jitters such as this when jittering periodic operations (e.g. cert rotation +// checks) since large jitters result in significantly increased load. +func NewSeventhJitter() Jitter { + rng := rand.New(rand.NewSource(time.Now().UnixNano())) + jitter, _ := newJitter(7, rng) + return jitter +} + +// rng is an interface implemented by math/rand.Rand. This interface +// is used in testting. +type rng interface { + // Int63n returns, as an int64, a non-negative pseudo-random number + // in the half-open interval [0,n). It panics if n <= 0. + Int63n(n int64) int64 +} + +// newJitter builds a new jitter on the range [d*(n-1)/n,d) +// newJitter only returns an error if n < 1. +func newJitter(n time.Duration, rng rng) (Jitter, error) { + if n < 1 { + return nil, trace.BadParameter("newJitter expects n>=1, but got %v", n) + } + var mu sync.Mutex + return func(d time.Duration) time.Duration { + // values less than 1 cause rng to panic, and some logic + // relies on treating zero duration as non-blocking case. + if d < 1 { + return 0 + } + mu.Lock() + defer mu.Unlock() + return d*(n-1)/n + time.Duration(rng.Int63n(int64(d))/int64(n)) + }, nil +} diff --git a/api/utils/retryutils/jitter_test.go b/api/utils/retryutils/jitter_test.go new file mode 100644 index 00000000000..083a885daaa --- /dev/null +++ b/api/utils/retryutils/jitter_test.go @@ -0,0 +1,114 @@ +/* +Copyright 2021-2022 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package retryutils + +import ( + "fmt" + "testing" + "time" + + "github.com/gravitational/trace" + "github.com/stretchr/testify/require" +) + +func TestNewJitterBadParameter(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + n time.Duration + assertErr require.ErrorAssertionFunc + }{ + { + n: -1, + assertErr: func(t require.TestingT, err error, i ...interface{}) { + require.True(t, trace.IsBadParameter(err), err) + }, + }, + { + n: 0, + assertErr: func(t require.TestingT, err error, i ...interface{}) { + require.True(t, trace.IsBadParameter(err), err) + }, + }, + { + n: 1, + assertErr: require.NoError, + }, + { + n: 7, + assertErr: require.NoError, + }, + } { + t.Run(fmt.Sprintf("n=%v", tc.n), func(t *testing.T) { + _, err := newJitter(tc.n, nil) + tc.assertErr(t, err) + }) + } +} + +func TestNewJitter(t *testing.T) { + t.Parallel() + + baseDuration := time.Second + mockInt63nFloor := mockInt63n(func(n int64) int64 { return 0 }) + mockInt63nCeiling := mockInt63n(func(n int64) int64 { return n - 1 }) + + for _, tc := range []struct { + desc string + n time.Duration + expectFloor time.Duration + expectCeiling time.Duration + }{ + { + desc: "FullJitter", + n: 1, + expectFloor: 0, + expectCeiling: baseDuration - 1, + }, + { + desc: "HalfJitter", + n: 2, + expectFloor: baseDuration / 2, + expectCeiling: baseDuration - 1, + }, + { + desc: "SeventhJitter", + n: 7, + expectFloor: baseDuration * 6 / 7, + expectCeiling: baseDuration - 1, + }, + } { + tc := tc + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + testFloorJitter, err := newJitter(tc.n, mockInt63nFloor) + require.NoError(t, err) + require.Equal(t, tc.expectFloor, testFloorJitter(baseDuration)) + + testCeilingJitter, err := newJitter(tc.n, mockInt63nCeiling) + require.NoError(t, err) + require.Equal(t, tc.expectCeiling, testCeilingJitter(baseDuration)) + }) + } +} + +type mockInt63n func(n int64) int64 + +func (m mockInt63n) Int63n(n int64) int64 { + return m(n) +} diff --git a/api/utils/retryutils/retry.go b/api/utils/retryutils/retry.go new file mode 100644 index 00000000000..70558c763b4 --- /dev/null +++ b/api/utils/retryutils/retry.go @@ -0,0 +1,233 @@ +/* +Copyright 2019-2022 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +// Package retryutils defines common retry and jitter logic. +package retryutils + +import ( + "context" + "fmt" + "time" + + "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" + log "github.com/sirupsen/logrus" +) + +// Retry is an interface that provides retry logic +type Retry interface { + // Reset resets retry state + Reset() + // Inc increments retry attempt + Inc() + // Duration returns retry duration, + // could be 0 + Duration() time.Duration + // After returns time.Time channel + // that fires after Duration delay, + // could fire right away if Duration is 0 + After() <-chan time.Time + // Clone creates a copy of this retry in a + // reset state. + Clone() Retry +} + +// LinearConfig sets up retry configuration +// using arithmetic progression +type LinearConfig struct { + // First is a first element of the progression, + // could be 0 + First time.Duration + // Step is a step of the progression, can't be 0 + Step time.Duration + // Max is a maximum value of the progression, + // can't be 0 + Max time.Duration + // Jitter is an optional jitter function to be applied + // to the delay. Note that supplying a jitter means that + // successive calls to Duration may return different results. + Jitter Jitter `json:"-"` + // AutoReset, if greater than zero, causes the linear retry to automatically + // reset after Max * AutoReset has elapsed since the last call to Incr. + AutoReset int64 + // Clock to override clock in tests + Clock clockwork.Clock +} + +// CheckAndSetDefaults checks and sets defaults +func (c *LinearConfig) CheckAndSetDefaults() error { + if c.Step == 0 { + return trace.BadParameter("missing parameter Step") + } + if c.Max == 0 { + return trace.BadParameter("missing parameter Max") + } + if c.Clock == nil { + c.Clock = clockwork.NewRealClock() + } + return nil +} + +// NewLinear returns a new instance of linear retry +func NewLinear(cfg LinearConfig) (*Linear, error) { + if err := cfg.CheckAndSetDefaults(); err != nil { + return nil, trace.Wrap(err) + } + return newLinear(cfg), nil +} + +// newLinear creates an instance of Linear from a +// previously verified configuration. +func newLinear(cfg LinearConfig) *Linear { + closedChan := make(chan time.Time) + close(closedChan) + return &Linear{LinearConfig: cfg, closedChan: closedChan} +} + +// NewConstant returns a new linear retry with constant interval. +func NewConstant(interval time.Duration) (*Linear, error) { + return NewLinear(LinearConfig{Step: interval, Max: interval}) +} + +// Linear is used to calculate retry period +// that follows the following logic: +// On the first error there is no delay +// on the next error, delay is FastLinear +// on all other errors, delay is SlowLinear +type Linear struct { + // LinearConfig is a linear retry config + LinearConfig + lastUse time.Time + attempt int64 + closedChan chan time.Time +} + +// Reset resets retry period to initial state +func (r *Linear) Reset() { + r.attempt = 0 +} + +// ResetToDelay resets retry period and increments the number of attempts. +func (r *Linear) ResetToDelay() { + r.Reset() + r.Inc() +} + +// Clone creates an identical copy of Linear with fresh state. +func (r *Linear) Clone() Retry { + return newLinear(r.LinearConfig) +} + +// Inc increments attempt counter +func (r *Linear) Inc() { + r.attempt++ +} + +// Duration returns retry duration based on state +func (r *Linear) Duration() time.Duration { + if r.AutoReset > 0 { + now := r.Clock.Now() + if now.After(r.lastUse.Add(r.Max * time.Duration(r.AutoReset))) { + r.Reset() + } + r.lastUse = now + } + + a := r.First + time.Duration(r.attempt)*r.Step + if a < 1 { + return 0 + } + + if a > r.Max { + a = r.Max + } + + if r.Jitter != nil { + a = r.Jitter(a) + } + + return a +} + +// After returns channel that fires with timeout +// defined in Duration method, as a special case +// if Duration is 0 returns a closed channel +func (r *Linear) After() <-chan time.Time { + d := r.Duration() + if d < 1 { + return r.closedChan + } + return r.Clock.After(d) +} + +// String returns user-friendly representation of the LinearPeriod +func (r *Linear) String() string { + return fmt.Sprintf("Linear(attempt=%v, duration=%v)", r.attempt, r.Duration()) +} + +// For retries the provided function until it succeeds or the context expires. +func (r *Linear) For(ctx context.Context, retryFn func() error) error { + for { + err := retryFn() + if err == nil { + return nil + } + if _, ok := trace.Unwrap(err).(*permanentRetryError); ok { + return trace.Wrap(err) + } + log.Debugf("Will retry in %v: %v.", r.Duration(), err) + select { + case <-r.After(): + r.Inc() + case <-ctx.Done(): + return trace.LimitExceeded(ctx.Err().Error()) + } + } +} + +// PermanentRetryError returns a new instance of a permanent retry error. +func PermanentRetryError(err error) error { + return &permanentRetryError{err: err} +} + +// permanentRetryError indicates that retry loop should stop. +type permanentRetryError struct { + err error +} + +// Error returns the original error message. +func (e *permanentRetryError) Error() string { + return e.err.Error() +} + +// RetryFastFor retries a function repeatedly for a set amount of +// time before returning an error. +// +// Intended mostly for tests. +func RetryStaticFor(d time.Duration, w time.Duration, f func() error) error { + start := time.Now() + var err error + + for time.Since(start) < d { + if err = f(); err == nil { + break + } + + time.Sleep(w) + } + + return err +} diff --git a/lib/utils/retry_test.go b/api/utils/retryutils/retry_test.go similarity index 69% rename from lib/utils/retry_test.go rename to api/utils/retryutils/retry_test.go index 7a0874804c4..a8e3ef746e8 100644 --- a/lib/utils/retry_test.go +++ b/api/utils/retryutils/retry_test.go @@ -1,5 +1,5 @@ /* -Copyright 2021 Gravitational, Inc. +Copyright 2021-2022 Gravitational, Inc. Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the License. @@ -14,7 +14,7 @@ See the License for the specific language governing permissions and limitations under the License. */ -package utils +package retryutils import ( "testing" @@ -23,7 +23,29 @@ import ( "github.com/stretchr/testify/require" ) -func Test_LinearRetryMax(t *testing.T) { +// TestLinear tests retry logic +func TestLinear(t *testing.T) { + t.Parallel() + + r, err := NewLinear(LinearConfig{ + Step: time.Second, + Max: 3 * time.Second, + }) + require.NoError(t, err) + require.Equal(t, r.Duration(), time.Duration(0)) + r.Inc() + require.Equal(t, r.Duration(), time.Second) + r.Inc() + require.Equal(t, r.Duration(), 2*time.Second) + r.Inc() + require.Equal(t, r.Duration(), 3*time.Second) + r.Inc() + require.Equal(t, r.Duration(), 3*time.Second) + r.Reset() + require.Equal(t, r.Duration(), time.Duration(0)) +} + +func TestLinearRetryMax(t *testing.T) { t.Parallel() cases := []struct { @@ -31,13 +53,23 @@ func Test_LinearRetryMax(t *testing.T) { config LinearConfig previousCompareFn require.ComparisonAssertionFunc }{ + { + desc: "FullJitter", + config: LinearConfig{ + First: time.Second * 45, + Step: time.Second * 30, + Max: time.Minute, + Jitter: NewFullJitter(), + }, + previousCompareFn: require.NotEqual, + }, { desc: "HalfJitter", config: LinearConfig{ First: time.Second * 45, Step: time.Second * 30, Max: time.Minute, - Jitter: NewJitter(), + Jitter: NewHalfJitter(), }, previousCompareFn: require.NotEqual, }, @@ -51,6 +83,7 @@ func Test_LinearRetryMax(t *testing.T) { }, previousCompareFn: require.NotEqual, }, + { desc: "NoJitter", config: LinearConfig{ @@ -85,7 +118,6 @@ func Test_LinearRetryMax(t *testing.T) { // ensure duration comparison to previous is satisfied tc.previousCompareFn(t, duration, previous) - } }) } diff --git a/api/utils/sshutils/ppk/ppk.go b/api/utils/sshutils/ppk/ppk.go index 2a2f022432d..08b89136990 100644 --- a/api/utils/sshutils/ppk/ppk.go +++ b/api/utils/sshutils/ppk/ppk.go @@ -30,8 +30,9 @@ import ( "fmt" "math/big" - "github.com/gravitational/teleport/api/constants" "github.com/gravitational/trace" + + "github.com/gravitational/teleport/api/constants" ) // ConvertToPPK takes a regular RSA-formatted keypair and converts it into the PPK file format used by the PuTTY SSH client. diff --git a/api/utils/sshutils/ppk/ppk_test.go b/api/utils/sshutils/ppk/ppk_test.go index 5f950353b83..5953787ca42 100644 --- a/api/utils/sshutils/ppk/ppk_test.go +++ b/api/utils/sshutils/ppk/ppk_test.go @@ -21,9 +21,10 @@ import ( "crypto/rsa" "testing" + "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/api/utils/sshutils/ppk" - "github.com/stretchr/testify/require" ) func TestConvertToPPK(t *testing.T) { @@ -217,8 +218,10 @@ Private-MAC: a9b12c6450e46fd7abbaaff5841f8a64f9597c7b2b59bd69d6fd3ceee0ca61ea priv, err := keys.ParsePrivateKey(tc.priv) require.NoError(t, err) - rsaPriv, ok := priv.GetBaseSigner().(*rsa.PrivateKey) + rsaPriv, ok := priv.Signer.(*rsa.PrivateKey) require.True(t, ok) + // Without this line, the linter thinks that "crypto/rsa" is unused... + require.IsType(t, &rsa.PrivateKey{}, rsaPriv) output, err := ppk.ConvertToPPK(rsaPriv, tc.pub) require.NoError(t, err) diff --git a/e b/e index 2137d494a21..39d332a99b3 160000 --- a/e +++ b/e @@ -1 +1 @@ -Subproject commit 2137d494a21411b93149d87d29cc1f0bf885f6cb +Subproject commit 39d332a99b36bbcc46dcd3508fb45526d28732bc diff --git a/go.mod b/go.mod index 17b5f63f3ab..8ce8ba12415 100644 --- a/go.mod +++ b/go.mod @@ -204,6 +204,7 @@ require ( github.com/go-openapi/jsonpointer v0.19.5 // indirect github.com/go-openapi/jsonreference v0.19.5 // indirect github.com/go-openapi/swag v0.19.14 // indirect + github.com/go-piv/piv-go v1.10.0 // indirect github.com/golang-jwt/jwt v3.2.2+incompatible // indirect github.com/golang-jwt/jwt/v4 v4.2.0 // indirect github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe // indirect diff --git a/go.sum b/go.sum index 5c554c40652..655c60ee12d 100644 --- a/go.sum +++ b/go.sum @@ -413,6 +413,8 @@ github.com/go-openapi/jsonreference v0.19.5/go.mod h1:RdybgQwPxbL4UEjuAruzK1x3nE github.com/go-openapi/swag v0.19.5/go.mod h1:POnQmlKehdgb5mhVOsnJFsivZCEZ/vjK9gh66Z9tfKk= github.com/go-openapi/swag v0.19.14 h1:gm3vOOXfiuw5i9p5N9xJvfjvuofpyvLA9Wr6QfK5Fng= github.com/go-openapi/swag v0.19.14/go.mod h1:QYRuS/SOXUCsnplDa677K7+DxSOj6IPNl/eQntq43wQ= +github.com/go-piv/piv-go v1.10.0 h1:P1Y1VjBI5DnXW0+YkKmTuh5opWnMIrKriUaIOblee9Q= +github.com/go-piv/piv-go v1.10.0/go.mod h1:NZ2zmjVkfFaL/CF8cVQ/pXdXtuj110zEKGdJM6fJZZM= github.com/go-sql-driver/mysql v1.3.0/go.mod h1:zAC/RDZ24gD3HViQzih4MyKcchzm+sOG5ZlKdlhCg5w= github.com/go-sql-driver/mysql v1.4.0/go.mod h1:zAC/RDZ24gD3HViQzih4MyKcchzm+sOG5ZlKdlhCg5w= github.com/go-sql-driver/mysql v1.5.0 h1:ozyZYNQW3x3HtqT1jira07DN2PArx2v7/mN66gGcHOs= diff --git a/integration/helpers/helpers.go b/integration/helpers/helpers.go index 6af92d42997..d7291c566e3 100644 --- a/integration/helpers/helpers.go +++ b/integration/helpers/helpers.go @@ -28,17 +28,18 @@ import ( "testing" "time" + "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/constants" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/client" libclient "github.com/gravitational/teleport/lib/client" "github.com/gravitational/teleport/lib/client/identityfile" "github.com/gravitational/teleport/lib/teleagent" - "github.com/gravitational/teleport/lib/utils" - "github.com/stretchr/testify/require" "github.com/gravitational/trace" "golang.org/x/crypto/ssh/agent" @@ -234,7 +235,7 @@ func WaitForProxyCount(t *TeleInstance, clusterName string, count int) error { func WaitForAuditEventTypeWithBackoff(t *testing.T, cli *auth.Server, startTime time.Time, eventType string) []apievents.AuditEvent { max := time.Second timeout := time.After(max) - bf, err := utils.NewLinear(utils.LinearConfig{ + bf, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: max / 10, Max: max, }) diff --git a/integration/helpers/trustedclusters.go b/integration/helpers/trustedclusters.go index 9ad9a1cf70a..8e94dcf1561 100644 --- a/integration/helpers/trustedclusters.go +++ b/integration/helpers/trustedclusters.go @@ -19,14 +19,15 @@ import ( "testing" "time" - "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/reversetunnel" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/require" + + "github.com/gravitational/teleport/api/defaults" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" + "github.com/gravitational/teleport/lib/auth" + "github.com/gravitational/teleport/lib/reversetunnel" ) // WaitForTunnelConnections waits for remote tunnels connections @@ -102,7 +103,7 @@ func WaitForNodeCount(ctx context.Context, t *TeleInstance, clusterName string, iterWaitTime = time.Second ) - err := utils.RetryStaticFor(deadline, iterWaitTime, func() error { + err := retryutils.RetryStaticFor(deadline, iterWaitTime, func() error { remoteSite, err := t.Tunnel.GetSite(clusterName) if err != nil { return trace.Wrap(err) diff --git a/integration/proxy/proxy_helpers.go b/integration/proxy/proxy_helpers.go index fbecf93e59d..2cd2607feee 100644 --- a/integration/proxy/proxy_helpers.go +++ b/integration/proxy/proxy_helpers.go @@ -30,9 +30,18 @@ import ( "time" "github.com/google/uuid" + "github.com/gravitational/trace" + "github.com/jackc/pgconn" + "github.com/stretchr/testify/require" + v1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/client-go/tools/clientcmd" + clientcmdapi "k8s.io/client-go/tools/clientcmd/api" + "github.com/gravitational/teleport/api/breaker" "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/integration/helpers" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/reversetunnel" @@ -43,13 +52,6 @@ import ( "github.com/gravitational/teleport/lib/srv/db/postgres" "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/utils" - "github.com/gravitational/trace" - "github.com/jackc/pgconn" - "github.com/stretchr/testify/require" - v1 "k8s.io/api/core/v1" - metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" - "k8s.io/client-go/tools/clientcmd" - clientcmdapi "k8s.io/client-go/tools/clientcmd/api" ) type Suite struct { @@ -216,7 +218,7 @@ func (p *Suite) mustConnectToClusterAndRunSSHCommand(t *testing.T, config helper require.NoError(t, err) cmd := []string{"echo", "hello world"} - err = utils.RetryStaticFor(deadline, nextIterWaitTime, func() error { + err = retryutils.RetryStaticFor(deadline, nextIterWaitTime, func() error { err = tc.SSH(context.TODO(), cmd, false) return trace.Wrap(err) }) diff --git a/lib/auth/auth.go b/lib/auth/auth.go index bfcfe099d81..b654f97a11b 100644 --- a/lib/auth/auth.go +++ b/lib/auth/auth.go @@ -63,6 +63,7 @@ import ( apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/api/types/wrappers" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth/keystore" "github.com/gravitational/teleport/lib/auth/native" wanlib "github.com/gravitational/teleport/lib/auth/webauthn" @@ -480,7 +481,7 @@ func (a *Server) runPeriodicOperations() { // Create a ticker with jitter heartbeatCheckTicker := interval.New(interval.Config{ Duration: apidefaults.ServerKeepAliveTTL() * 2, - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) promTicker := time.NewTicker(defaults.PrometheusScrapeInterval) missedKeepAliveCount := 0 @@ -502,7 +503,7 @@ func (a *Server) runPeriodicOperations() { releaseCheck := interval.New(interval.Config{ Duration: time.Hour * 24, FirstDuration: firstReleaseCheck, - Jitter: utils.NewFullJitter(), + Jitter: retryutils.NewFullJitter(), }) defer releaseCheck.Stop() for { diff --git a/lib/auth/native/native.go b/lib/auth/native/native.go index 0a75ac8c2f0..241bb91bdd4 100644 --- a/lib/auth/native/native.go +++ b/lib/auth/native/native.go @@ -83,11 +83,17 @@ func GeneratePrivateKey() (*keys.PrivateKey, error) { if err != nil { return nil, trace.Wrap(err) } - rsaSigner, err := keys.NewStandardSigner(rsaKey) + keyDER, err := x509.MarshalPKCS8PrivateKey(rsaKey) if err != nil { return nil, trace.Wrap(err) } - return keys.NewPrivateKey(rsaSigner) + + keyPEM := pem.EncodeToMemory(&pem.Block{ + Type: keys.PKCS8PrivateKeyType, + Headers: nil, + Bytes: keyDER, + }) + return keys.NewPrivateKey(rsaKey, keyPEM) } func getOrGenerateRSAPrivateKey() (*rsa.PrivateKey, error) { diff --git a/lib/backend/dynamo/shards.go b/lib/backend/dynamo/shards.go index 07926a6b169..cbde4e82088 100644 --- a/lib/backend/dynamo/shards.go +++ b/lib/backend/dynamo/shards.go @@ -22,8 +22,8 @@ import ( "time" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" - "github.com/gravitational/teleport/lib/utils" "github.com/aws/aws-sdk-go/aws" "github.com/aws/aws-sdk-go/service/dynamodb" @@ -39,7 +39,7 @@ type shardEvent struct { } func (b *Backend) asyncPollStreams(ctx context.Context) error { - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: b.RetryPeriod / 10, Max: b.RetryPeriod, }) diff --git a/lib/backend/firestore/firestorebk.go b/lib/backend/firestore/firestorebk.go index af46d72a30a..df24fc2659d 100644 --- a/lib/backend/firestore/firestorebk.go +++ b/lib/backend/firestore/firestorebk.go @@ -33,9 +33,9 @@ import ( "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/defaults" - "github.com/gravitational/teleport/lib/utils" "github.com/jonboulle/clockwork" log "github.com/sirupsen/logrus" @@ -305,7 +305,7 @@ func New(ctx context.Context, params backend.Params, options Options) (*Backend, } // kicking off async tasks - linearConfig := utils.LinearConfig{ + linearConfig := retryutils.LinearConfig{ Step: b.RetryPeriod / 10, Max: b.RetryPeriod, } @@ -579,8 +579,8 @@ func (b *Backend) keyToDocumentID(key []byte) string { } // RetryingAsyncFunctionRunner wraps a task target in retry logic -func RetryingAsyncFunctionRunner(ctx context.Context, retryConfig utils.LinearConfig, logger *log.Logger, task func() error, taskName string) { - retry, err := utils.NewLinear(retryConfig) +func RetryingAsyncFunctionRunner(ctx context.Context, retryConfig retryutils.LinearConfig, logger *log.Logger, task func() error, taskName string) { + retry, err := retryutils.NewLinear(retryConfig) if err != nil { logger.WithError(err).Error("Bad retry parameters, returning and not running.") return diff --git a/lib/backend/sqlbk/backend.go b/lib/backend/sqlbk/backend.go index 657ce670e8d..983adb18a09 100644 --- a/lib/backend/sqlbk/backend.go +++ b/lib/backend/sqlbk/backend.go @@ -23,11 +23,12 @@ import ( "sync/atomic" "time" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/backend" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" + "github.com/gravitational/teleport/lib/backend" + "github.com/jonboulle/clockwork" ) @@ -118,7 +119,7 @@ func (b *Backend) retryTx(ctx context.Context, begin func(context.Context) Tx, t ctx, cancel := context.WithTimeout(ctx, b.RetryTimeout) defer cancel() - var delay *utils.Linear + var delay *retryutils.Linear tx := begin(ctx) for { if tx.Err() != nil { @@ -143,11 +144,11 @@ func (b *Backend) retryTx(ctx context.Context, begin func(context.Context) Tx, t if retryDelayPeriod == 0 { // sanity check (0 produces an error in NewLinear) retryDelayPeriod = DefaultRetryDelayPeriod } - delay, err = utils.NewLinear(utils.LinearConfig{ + delay, err = retryutils.NewLinear(retryutils.LinearConfig{ First: retryDelayPeriod, Step: retryDelayPeriod, Max: retryDelayPeriod, - Jitter: utils.NewJitter(), + Jitter: retryutils.NewJitter(), }) if err != nil { return trace.BadParameter("[BUG] invalid retry delay configuration: %v", err) diff --git a/lib/backend/sqlbk/background.go b/lib/backend/sqlbk/background.go index fde94c1e82c..355ed426e60 100644 --- a/lib/backend/sqlbk/background.go +++ b/lib/backend/sqlbk/background.go @@ -21,11 +21,13 @@ import ( "errors" "time" + "github.com/gravitational/trace" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/interval" - "github.com/gravitational/trace" ) // start background goroutine to track expired leases, emit events, and purge records. @@ -66,7 +68,7 @@ func (b *Backend) initLastEventID(ctx context.Context) (lastEventID int64, err e periodic = interval.New(interval.Config{ Duration: b.PollStreamPeriod, FirstDuration: utils.HalfJitter(b.PollStreamPeriod), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer periodic.Stop() } @@ -93,14 +95,14 @@ func (b *Backend) run(eventID int64) { pollPeriodic := interval.New(interval.Config{ Duration: b.PollStreamPeriod, FirstDuration: utils.HalfJitter(b.PollStreamPeriod), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer pollPeriodic.Stop() purgePeriodic := interval.New(interval.Config{ Duration: b.PurgePeriod, FirstDuration: utils.HalfJitter(b.PurgePeriod), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer purgePeriodic.Stop() diff --git a/lib/cache/cache.go b/lib/cache/cache.go index bda59621e3c..dffeed54015 100644 --- a/lib/cache/cache.go +++ b/lib/cache/cache.go @@ -26,6 +26,7 @@ import ( "github.com/gravitational/teleport/api/client/proto" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/observability/metrics" @@ -759,11 +760,11 @@ func New(config Config) (*Cache, error) { // Starts the cache. Should only be called once. func (c *Cache) Start() error { - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(c.MaxRetryPeriod / 10), Step: c.MaxRetryPeriod / 5, Max: c.MaxRetryPeriod, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), Clock: c.Clock, }) if err != nil { @@ -809,7 +810,7 @@ Outer: return c.eventsFanout.NewWatcher(ctx, watch) } -func (c *Cache) update(ctx context.Context, retry utils.Retry) { +func (c *Cache) update(ctx context.Context, retry retryutils.Retry) { defer func() { c.Debugf("Cache is closing, returning from update loop.") // ensure that close operations have been run @@ -893,7 +894,7 @@ func (c *Cache) notify(ctx context.Context, event Event) { // have skipped 1 and 2, but in the absence of such mechanism in Dynamo // we assume that this cache will eventually end up in a correct state // potentially lagging behind the state of the database. -func (c *Cache) fetchAndWatch(ctx context.Context, retry utils.Retry, timer *time.Timer) error { +func (c *Cache) fetchAndWatch(ctx context.Context, retry retryutils.Retry, timer *time.Timer) error { watcher, err := c.Events.NewWatcher(c.ctx, types.Watch{ QueueSize: c.QueueSize, Name: c.Component, @@ -978,7 +979,7 @@ func (c *Cache) fetchAndWatch(ctx context.Context, retry utils.Retry, timer *tim relativeExpiryInterval = interval.New(interval.Config{ Duration: c.Config.RelativeExpiryCheckInterval, FirstDuration: utils.HalfJitter(c.Config.RelativeExpiryCheckInterval), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) break } diff --git a/lib/client/api.go b/lib/client/api.go index a339b00bd71..d93775e73f0 100644 --- a/lib/client/api.go +++ b/lib/client/api.go @@ -52,6 +52,7 @@ import ( "github.com/gravitational/teleport/api/types/wrappers" apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/api/utils/keypaths" + "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/lib/auth" wancli "github.com/gravitational/teleport/lib/auth/webauthncli" "github.com/gravitational/teleport/lib/client/terminal" @@ -3385,6 +3386,32 @@ func (tc *TeleportClient) Login(ctx context.Context) (*Key, error) { ) defer span.End() + // Get a new (or old) key to be signed via proxy on valid login. + var key *Key + var err error + + // TODO (Joerger): Remove this env var check once we can pull server settings to decide whether + // or not to initiate PIV login - https://github.com/gravitational/teleport/pull/15336 + if os.Getenv("PIV_LOGIN") != "" { + priv, err := keys.GetOrGenerateYubiKeyPrivateKey(ctx, false) + if err != nil { + return nil, trace.Wrap(err) + } + key = NewKey(priv) + } else { + // If we find a valid key in the localAgent, reuse it instead of generating a new key. + // This is especially useful if the key is a hardware key, since they need to be reused + // between multiple login sessions and shouldn't be regenerated. + if key, err = tc.localAgent.GetCoreKey(); trace.IsNotFound(err) { + // No core key found, this is the first login for the proxy. Generate a new RSA key. + if key, err = GenerateRSAKey(); err != nil { + return nil, trace.Wrap(err) + } + } else if err != nil { + return nil, trace.Wrap(err) + } + } + // Ping the endpoint to see if it's up and find the type of authentication // supported, also show the message of the day if available. pr, err := tc.PingAndShowMOTD(ctx) @@ -3392,13 +3419,6 @@ func (tc *TeleportClient) Login(ctx context.Context) (*Key, error) { return nil, trace.Wrap(err) } - // generate a new keypair. the public key will be signed via proxy if client's - // password+OTP are valid - key, err := GenerateRSAKey() - if err != nil { - return nil, trace.Wrap(err) - } - var response *auth.SSHLoginResponse var username string switch authType := pr.Auth.Type; { diff --git a/lib/client/keyagent.go b/lib/client/keyagent.go index 2b8c0eee521..69634b85082 100644 --- a/lib/client/keyagent.go +++ b/lib/client/keyagent.go @@ -323,7 +323,7 @@ func (a *LocalKeyAgent) GetKey(clusterName string, opts ...CertOption) (*Key, er } // GetCoreKey returns the key without any cluster-dependent certificates, -// i.e. including only the RSA keypair and the Teleport TLS certificate. +// i.e. including only the private key and the Teleport TLS certificate. func (a *LocalKeyAgent) GetCoreKey() (*Key, error) { return a.GetKey("") } diff --git a/lib/client/keyagent_test.go b/lib/client/keyagent_test.go index e6896163673..ff5627570ca 100644 --- a/lib/client/keyagent_test.go +++ b/lib/client/keyagent_test.go @@ -129,7 +129,7 @@ func TestAddKey(t *testing.T) { // check that we've loaded a cert as well as a private key into the teleport agent // and it's for the user we expected to add a certificate for require.Len(t, teleportAgentKeys, 2) - require.Equal(t, "ssh-rsa-cert-v01@openssh.com", teleportAgentKeys[0].Type()) + require.Equal(t, ssh.CertAlgoRSAv01, teleportAgentKeys[0].Type()) require.Equal(t, "teleport:"+s.username, teleportAgentKeys[0].Comment) require.Equal(t, "ssh-rsa", teleportAgentKeys[1].Type()) require.Equal(t, "teleport:"+s.username, teleportAgentKeys[1].Comment) @@ -144,7 +144,7 @@ func TestAddKey(t *testing.T) { require.True(t, found) found = false for _, sak := range systemAgentKeys { - if sak.Comment == "teleport:"+s.username && sak.Type() == "ssh-rsa-cert-v01@openssh.com" { + if sak.Comment == "teleport:"+s.username && sak.Type() == ssh.CertAlgoRSAv01 { found = true } } diff --git a/lib/events/auditwriter.go b/lib/events/auditwriter.go index 3d9a26474b4..994ad7e0ccd 100644 --- a/lib/events/auditwriter.go +++ b/lib/events/auditwriter.go @@ -24,6 +24,7 @@ import ( apidefaults "github.com/gravitational/teleport/api/defaults" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/utils" @@ -550,7 +551,7 @@ func (a *AuditWriter) completeStream(stream apievents.Stream) { } func (a *AuditWriter) tryResumeStream() (apievents.Stream, error) { - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: defaults.NetworkRetryDuration, Max: defaults.NetworkBackoffDuration, }) diff --git a/lib/events/complete.go b/lib/events/complete.go index 0a804432b5a..acee7be050b 100644 --- a/lib/events/complete.go +++ b/lib/events/complete.go @@ -24,6 +24,7 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types/events" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/utils" @@ -126,7 +127,7 @@ func (u *UploadCompleter) Serve(ctx context.Context) error { periodic := interval.New(interval.Config{ Duration: u.cfg.CheckPeriod, FirstDuration: utils.HalfJitter(u.cfg.CheckPeriod), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer periodic.Stop() diff --git a/lib/events/filesessions/fileasync.go b/lib/events/filesessions/fileasync.go index 0ae9a65044c..d092a5bcee1 100644 --- a/lib/events/filesessions/fileasync.go +++ b/lib/events/filesessions/fileasync.go @@ -28,6 +28,7 @@ import ( "github.com/gravitational/teleport" apidefaults "github.com/gravitational/teleport/api/defaults" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/session" @@ -154,7 +155,7 @@ func (u *Uploader) checkSessionError(sessionID session.ID) (bool, error) { // Serve runs the uploader until stopped func (u *Uploader) Serve(ctx context.Context) error { - backoff, err := utils.NewLinear(utils.LinearConfig{ + backoff, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: u.cfg.ScanPeriod, Max: u.cfg.ScanPeriod * 100, Clock: u.cfg.Clock, diff --git a/lib/events/firestoreevents/firestoreevents.go b/lib/events/firestoreevents/firestoreevents.go index 1323035e068..8317e3270cf 100644 --- a/lib/events/firestoreevents/firestoreevents.go +++ b/lib/events/firestoreevents/firestoreevents.go @@ -28,6 +28,7 @@ import ( "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/observability/metrics" @@ -300,7 +301,7 @@ func New(cfg EventsConfig) (*Log, error) { } } if !b.DisableExpiredDocumentPurge { - go firestorebk.RetryingAsyncFunctionRunner(b.svcContext, utils.LinearConfig{ + go firestorebk.RetryingAsyncFunctionRunner(b.svcContext, retryutils.LinearConfig{ Step: b.RetryPeriod / 10, Max: b.RetryPeriod, }, b.Logger, b.purgeExpiredEvents, "purgeExpiredEvents") diff --git a/lib/events/stream.go b/lib/events/stream.go index 388254a7cb7..27c60d844d5 100644 --- a/lib/events/stream.go +++ b/lib/events/stream.go @@ -28,6 +28,7 @@ import ( "time" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/utils" @@ -710,7 +711,7 @@ func (w *sliceWriter) startUpload(partNumber int64, slice *slice) (*activeUpload <-w.semUploads }() - var retry utils.Retry + var retry retryutils.Retry for i := 0; i < defaults.MaxIterationLimit; i++ { reader, err := slice.reader() if err != nil { @@ -730,7 +731,7 @@ func (w *sliceWriter) startUpload(partNumber int64, slice *slice) (*activeUpload // retry is created on the first upload error if retry == nil { var rerr error - retry, rerr = utils.NewLinear(utils.LinearConfig{ + retry, rerr = retryutils.NewLinear(retryutils.LinearConfig{ Step: defaults.NetworkRetryDuration, Max: defaults.NetworkBackoffDuration, }) diff --git a/lib/events/test/suite.go b/lib/events/test/suite.go index 556119b56d7..681101d0395 100644 --- a/lib/events/test/suite.go +++ b/lib/events/test/suite.go @@ -29,10 +29,10 @@ import ( "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/session" - "github.com/gravitational/teleport/lib/utils" "github.com/google/uuid" "github.com/jonboulle/clockwork" @@ -107,7 +107,7 @@ func (s *EventsSuite) EventPagination(t *testing.T) { var err error var checkpoint string - err = utils.RetryStaticFor(time.Minute*5, time.Second*5, func() error { + err = retryutils.RetryStaticFor(time.Minute*5, time.Second*5, func() error { arr, checkpoint, err = s.Log.SearchEvents(baseTime, toTime, apidefaults.Namespace, nil, 100, types.EventOrderAscending, checkpoint) return err }) @@ -225,7 +225,7 @@ func (s *EventsSuite) SessionEventsCRUD(t *testing.T) { var history []apievents.AuditEvent - err = utils.RetryStaticFor(time.Minute*5, time.Second*5, func() error { + err = retryutils.RetryStaticFor(time.Minute*5, time.Second*5, func() error { history, _, err = s.Log.SearchEvents(s.Clock.Now().Add(-1*time.Hour), s.Clock.Now().Add(time.Hour), apidefaults.Namespace, nil, 100, types.EventOrderAscending, "") return err }) diff --git a/lib/inventory/controller.go b/lib/inventory/controller.go index bb993965377..d574ca9ddb9 100644 --- a/lib/inventory/controller.go +++ b/lib/inventory/controller.go @@ -24,6 +24,7 @@ import ( "github.com/gravitational/teleport/api/client/proto" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/interval" @@ -164,7 +165,7 @@ func (c *Controller) handleControlStream(handle *upstreamHandle) { keepAliveInterval := interval.New(interval.Config{ Duration: c.serverKeepAlive, FirstDuration: utils.HalfJitter(c.serverKeepAlive), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer keepAliveInterval.Stop() diff --git a/lib/inventory/inventory.go b/lib/inventory/inventory.go index b3bd41fc65d..421648ee678 100644 --- a/lib/inventory/inventory.go +++ b/lib/inventory/inventory.go @@ -24,6 +24,7 @@ import ( "github.com/gravitational/teleport/api/client" "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" @@ -83,7 +84,7 @@ func NewDownstreamHandle(fn DownstreamCreateFunc, hello proto.UpstreamInventoryH return handle } -func SendHeartbeat(ctx context.Context, handle DownstreamHandle, hb proto.InventoryHeartbeat, retry utils.Retry) { +func SendHeartbeat(ctx context.Context, handle DownstreamHandle, hb proto.InventoryHeartbeat, retry retryutils.Retry) { for { select { case sender := <-handle.Sender(): diff --git a/lib/restrictedsession/watcher.go b/lib/restrictedsession/watcher.go index caa41e7934b..15dbf0b98cf 100644 --- a/lib/restrictedsession/watcher.go +++ b/lib/restrictedsession/watcher.go @@ -27,6 +27,7 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/utils" @@ -38,11 +39,11 @@ func NewRestrictionsWatcher(cfg RestrictionsWatcherConfig) (*RestrictionsWatcher if err := cfg.CheckAndSetDefaults(); err != nil { return nil, trace.Wrap(err) } - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(cfg.MaxRetryPeriod / 10), Step: cfg.MaxRetryPeriod / 5, Max: cfg.MaxRetryPeriod, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), }) if err != nil { return nil, trace.Wrap(err) @@ -69,7 +70,7 @@ type RestrictionsWatcher struct { resetC chan struct{} // retry is used to manage backoff logic for watches - retry utils.Retry + retry retryutils.Retry wg sync.WaitGroup cancelFn context.CancelFunc diff --git a/lib/reversetunnel/agentpool.go b/lib/reversetunnel/agentpool.go index 36becdc4a54..74d2aa28c76 100644 --- a/lib/reversetunnel/agentpool.go +++ b/lib/reversetunnel/agentpool.go @@ -29,6 +29,7 @@ import ( "github.com/gravitational/teleport/api/client/webclient" "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/api/utils/sshutils" "github.com/gravitational/teleport/lib" "github.com/gravitational/teleport/lib/auth" @@ -84,7 +85,7 @@ type AgentPool struct { cancel context.CancelFunc // backoff limits the rate at which new agents are created. - backoff utils.Retry + backoff retryutils.Retry log logrus.FieldLogger } @@ -177,10 +178,10 @@ func NewAgentPool(ctx context.Context, config AgentPoolConfig) (*AgentPool, erro if err := config.CheckAndSetDefaults(); err != nil { return nil, trace.Wrap(err) } - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: time.Second, Max: maxBackoff, - Jitter: utils.NewJitter(), + Jitter: retryutils.NewJitter(), AutoReset: 4, }) if err != nil { diff --git a/lib/reversetunnel/agentpool_test.go b/lib/reversetunnel/agentpool_test.go index 61e61e857b3..6913fedbf19 100644 --- a/lib/reversetunnel/agentpool_test.go +++ b/lib/reversetunnel/agentpool_test.go @@ -21,14 +21,16 @@ import ( "testing" "time" + "github.com/gravitational/trace" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/reversetunnel/track" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/utils" - "github.com/gravitational/trace" - "github.com/stretchr/testify/require" - "golang.org/x/crypto/ssh" ) type mockAgent struct { @@ -84,7 +86,7 @@ func setupTestAgentPool(t *testing.T) (*AgentPool, *mockClient) { }) require.NoError(t, err) - pool.backoff, err = utils.NewLinear(utils.LinearConfig{ + pool.backoff, err = retryutils.NewLinear(retryutils.LinearConfig{ Step: time.Millisecond, Max: time.Millisecond, }) diff --git a/lib/reversetunnel/remotesite.go b/lib/reversetunnel/remotesite.go index 9444e9fe701..cece8f4dcd8 100644 --- a/lib/reversetunnel/remotesite.go +++ b/lib/reversetunnel/remotesite.go @@ -33,6 +33,7 @@ import ( "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/api/utils/sshutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/services" @@ -455,7 +456,7 @@ func (s *remoteSite) compareAndSwapCertAuthority(ca types.CertAuthority) error { return trace.CompareFailed("remote certificate authority rotation has been updated") } -func (s *remoteSite) updateCertAuthorities(retry utils.Retry, remoteWatcher *services.CertAuthorityWatcher, remoteVersion string) { +func (s *remoteSite) updateCertAuthorities(retry retryutils.Retry, remoteWatcher *services.CertAuthorityWatcher, remoteVersion string) { defer remoteWatcher.Close() for { @@ -652,7 +653,7 @@ func (s *remoteSite) getLocalWatchedCerts(remoteClusterVersion string) (types.Ce }, nil } -func (s *remoteSite) updateLocks(retry utils.Retry) { +func (s *remoteSite) updateLocks(retry retryutils.Retry) { s.Debugf("Watching for remote lock changes.") for { diff --git a/lib/reversetunnel/srv.go b/lib/reversetunnel/srv.go index f9685789fe2..1113eabba56 100644 --- a/lib/reversetunnel/srv.go +++ b/lib/reversetunnel/srv.go @@ -30,6 +30,7 @@ import ( "github.com/gravitational/teleport/api/breaker" "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" apisshutils "github.com/gravitational/teleport/api/utils/sshutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/defaults" @@ -1120,11 +1121,11 @@ func newRemoteSite(srv *server, domainName string, sconn ssh.Conn) (*remoteSite, } remoteSite.certificateCache = certificateCache - caRetry, err := utils.NewLinear(utils.LinearConfig{ + caRetry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(srv.Config.PollingPeriod), Step: srv.Config.PollingPeriod / 5, Max: srv.Config.PollingPeriod, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), Clock: srv.Clock, }) if err != nil { @@ -1148,11 +1149,11 @@ func newRemoteSite(srv *server, domainName string, sconn ssh.Conn) (*remoteSite, remoteSite.updateCertAuthorities(caRetry, remoteWatcher, remoteVersion) }() - lockRetry, err := utils.NewLinear(utils.LinearConfig{ + lockRetry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(srv.Config.PollingPeriod), Step: srv.Config.PollingPeriod / 5, Max: srv.Config.PollingPeriod, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), Clock: srv.Clock, }) if err != nil { diff --git a/lib/service/connect.go b/lib/service/connect.go index 931f4d56ce1..2575666d60f 100644 --- a/lib/service/connect.go +++ b/lib/service/connect.go @@ -36,6 +36,7 @@ import ( apiclient "github.com/gravitational/teleport/api/client" "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/native" @@ -53,12 +54,12 @@ import ( func (process *TeleportProcess) reconnectToAuthService(role types.SystemRole) (*Connector, error) { // TODO(fspmarshall): we should probably have a longer retry period for Instance certs // in order to avoid catastrophic load in the event of an auth server downgrade. - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(process.Config.MaxRetryPeriod / 10), Step: process.Config.MaxRetryPeriod / 5, Max: process.Config.MaxRetryPeriod, Clock: process.Clock, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), }) if err != nil { return nil, trace.Wrap(err) @@ -679,7 +680,7 @@ func (process *TeleportProcess) periodicSyncRotationState() error { periodic := interval.New(interval.Config{ Duration: process.Config.RotationConnectionInterval, FirstDuration: utils.HalfJitter(process.Config.RotationConnectionInterval), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer periodic.Stop() @@ -739,7 +740,7 @@ func (process *TeleportProcess) syncRotationStateCycle() error { periodic := interval.New(interval.Config{ Duration: process.Config.PollingPeriod, FirstDuration: utils.HalfJitter(process.Config.PollingPeriod), - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer periodic.Stop() for { diff --git a/lib/services/local/dynamic_access.go b/lib/services/local/dynamic_access.go index bc4573aaee8..c9ac2b28532 100644 --- a/lib/services/local/dynamic_access.go +++ b/lib/services/local/dynamic_access.go @@ -24,9 +24,9 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" ) @@ -65,7 +65,7 @@ func (s *DynamicAccessService) SetAccessRequestState(ctx context.Context, params return nil, trace.Wrap(err) } retryPeriod := retryPeriodMs * time.Millisecond - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: retryPeriod / 7, Max: retryPeriod, }) @@ -134,7 +134,7 @@ func (s *DynamicAccessService) ApplyAccessReview(ctx context.Context, params typ return nil, trace.Wrap(err) } retryPeriod := retryPeriodMs * time.Millisecond - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: retryPeriod / 7, Max: retryPeriod, }) @@ -352,7 +352,7 @@ func (s *DynamicAccessService) UpdatePluginData(ctx context.Context, params type func (s *DynamicAccessService) updateAccessRequestPluginData(ctx context.Context, params types.PluginDataUpdateParams) error { retryPeriod := retryPeriodMs * time.Millisecond - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: retryPeriod / 7, Max: retryPeriod, }) diff --git a/lib/services/local/presence.go b/lib/services/local/presence.go index bdebeb525f1..4cc1a420a58 100644 --- a/lib/services/local/presence.go +++ b/lib/services/local/presence.go @@ -27,6 +27,7 @@ import ( "github.com/gravitational/teleport/api/constants" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/utils" @@ -40,7 +41,7 @@ import ( // of the cluster - Nodes, Proxies and SSH nodes type PresenceService struct { log *logrus.Entry - jitter utils.Jitter + jitter retryutils.Jitter backend.Backend } @@ -52,7 +53,7 @@ type backendItemToResourceFunc func(item backend.Item) (types.ResourceWithLabels func NewPresenceService(b backend.Backend) *PresenceService { return &PresenceService{ log: logrus.WithFields(logrus.Fields{trace.Component: "Presence"}), - jitter: utils.NewFullJitter(), + jitter: retryutils.NewFullJitter(), Backend: b, } } diff --git a/lib/services/semaphore.go b/lib/services/semaphore.go index 7e359def101..57d9f0341b2 100644 --- a/lib/services/semaphore.go +++ b/lib/services/semaphore.go @@ -20,6 +20,7 @@ import ( "time" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/utils" @@ -71,7 +72,7 @@ func (l *SemaphoreLockConfig) CheckAndSetDefaults() error { type SemaphoreLock struct { cfg SemaphoreLockConfig lease0 types.SemaphoreLease - retry utils.Retry + retry retryutils.Retry ticker *time.Ticker doneC chan struct{} closeOnce sync.Once @@ -213,13 +214,13 @@ Outer: type AcquireSemaphoreWithRetryConfig struct { Service types.Semaphores Request types.AcquireSemaphoreRequest - Retry utils.LinearConfig + Retry retryutils.LinearConfig } // AcquireSemaphoreWithRetry tries to acquire the semaphore according to the // retry schedule until it succeeds or context expires. func AcquireSemaphoreWithRetry(ctx context.Context, req AcquireSemaphoreWithRetryConfig) (*types.SemaphoreLease, error) { - retry, err := utils.NewLinear(req.Retry) + retry, err := retryutils.NewLinear(req.Retry) if err != nil { return nil, trace.Wrap(err) } @@ -242,10 +243,10 @@ func AcquireSemaphoreLock(ctx context.Context, cfg SemaphoreLockConfig) (*Semaph return nil, trace.Wrap(err) } // set up retry with a ratio which will result in 3-4 retries before the lease expires - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Max: cfg.Expiry / 4, Step: cfg.Expiry / 16, - Jitter: utils.NewJitter(), + Jitter: retryutils.NewJitter(), }) if err != nil { return nil, trace.Wrap(err) diff --git a/lib/services/watcher.go b/lib/services/watcher.go index 71e83f4831b..6c08ded54c0 100644 --- a/lib/services/watcher.go +++ b/lib/services/watcher.go @@ -24,6 +24,7 @@ import ( "github.com/gravitational/teleport/api/constants" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/utils" @@ -96,11 +97,11 @@ func (cfg *ResourceWatcherConfig) CheckAndSetDefaults() error { // It is the caller's responsibility to verify the inputs' validity // incl. cfg.CheckAndSetDefaults. func newResourceWatcher(ctx context.Context, collector resourceCollector, cfg ResourceWatcherConfig) (*resourceWatcher, error) { - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: utils.HalfJitter(cfg.MaxRetryPeriod / 10), Step: cfg.MaxRetryPeriod / 5, Max: cfg.MaxRetryPeriod, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), Clock: cfg.Clock, }) if err != nil { @@ -133,7 +134,7 @@ type resourceWatcher struct { cancel context.CancelFunc // retry is used to manage backoff logic for watchers. - retry utils.Retry + retry retryutils.Retry // failureStartedAt records when the current sync failures were first // detected, zero if there are no failures present. diff --git a/lib/srv/db/cloud/iam.go b/lib/srv/db/cloud/iam.go index e14ca7f32ea..2b1c9eac9ca 100644 --- a/lib/srv/db/cloud/iam.go +++ b/lib/srv/db/cloud/iam.go @@ -27,11 +27,11 @@ import ( "github.com/jonboulle/clockwork" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/cloud" awslib "github.com/gravitational/teleport/lib/cloud/aws" "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" "github.com/sirupsen/logrus" @@ -241,10 +241,10 @@ func (c *IAM) processTask(ctx context.Context, task iamTask) error { }, // Retry with some jitters up to twice of the semaphore expire time. - Retry: utils.LinearConfig{ + Retry: retryutils.LinearConfig{ Step: 10 * time.Second, Max: 2 * time.Minute, - Jitter: utils.NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), }, }) if err != nil { diff --git a/lib/srv/db/cloud/users/users.go b/lib/srv/db/cloud/users/users.go index 5f204a7fa43..df386d36a54 100644 --- a/lib/srv/db/cloud/users/users.go +++ b/lib/srv/db/cloud/users/users.go @@ -20,8 +20,8 @@ import ( "time" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/cloud" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/interval" "github.com/gravitational/trace" @@ -149,7 +149,7 @@ func (u *Users) Start(ctx context.Context, getAllDatabases func() types.Database ticker := interval.New(interval.Config{ // Use jitter for HA setups. - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), // NewSeventhJitter builds a new jitter on the range [6n/7,n). // Use n = cfg.Interval*7/6 gives an effective duration range of diff --git a/lib/srv/db/common/auth.go b/lib/srv/db/common/auth.go index e4b3d8b275d..f3e756ff5a3 100644 --- a/lib/srv/db/common/auth.go +++ b/lib/srv/db/common/auth.go @@ -27,6 +27,7 @@ import ( "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" libauth "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/native" "github.com/gravitational/teleport/lib/cloud" @@ -238,7 +239,7 @@ func (a *dbAuth) GetCloudSQLPassword(ctx context.Context, sessionCtx *Session) ( // Cloud SQL will return 409 to a user update operation if there is another // one in progress, so retry upon encountering it. Also, be nice to the API // and retry with a backoff. - retry, err := utils.NewConstant(time.Second) + retry, err := retryutils.NewConstant(time.Second) if err != nil { return "", trace.Wrap(err) } @@ -249,7 +250,7 @@ func (a *dbAuth) GetCloudSQLPassword(ctx context.Context, sessionCtx *Session) ( Password: token, }) if err != nil && !trace.IsCompareFailed(ConvertError(err)) { // We only want to retry on 409. - return utils.PermanentRetryError(err) + return retryutils.PermanentRetryError(err) } return trace.Wrap(err) }) diff --git a/lib/srv/db/mysql/engine.go b/lib/srv/db/mysql/engine.go index a699d39fd25..f0c5b2816b3 100644 --- a/lib/srv/db/mysql/engine.go +++ b/lib/srv/db/mysql/engine.go @@ -29,6 +29,7 @@ import ( "github.com/go-mysql-org/go-mysql/server" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/srv/db/cloud" @@ -416,7 +417,7 @@ func (e *Engine) makeAcquireSemaphoreConfig(sessionCtx *common.Session) services }, // If multiple connections are being established simultaneously to the // same database as the same user, retry for a few seconds. - Retry: utils.LinearConfig{ + Retry: retryutils.LinearConfig{ Step: time.Second, Max: time.Second, Clock: e.Clock, diff --git a/lib/srv/heartbeatv2.go b/lib/srv/heartbeatv2.go index 31c7d01f33f..da4e19cc830 100644 --- a/lib/srv/heartbeatv2.go +++ b/lib/srv/heartbeatv2.go @@ -23,6 +23,7 @@ import ( "github.com/gravitational/teleport/api/client/proto" apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/inventory" @@ -204,7 +205,7 @@ func (h *HeartbeatV2) run() { h.announce = interval.New(interval.Config{ FirstDuration: utils.HalfJitter(h.announceInterval), Duration: h.announceInterval, - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer h.announce.Stop() @@ -212,7 +213,7 @@ func (h *HeartbeatV2) run() { h.poll = interval.New(interval.Config{ FirstDuration: utils.HalfJitter(h.pollInterval), Duration: h.pollInterval, - Jitter: utils.NewSeventhJitter(), + Jitter: retryutils.NewSeventhJitter(), }) defer h.poll.Stop() diff --git a/lib/srv/usermgmt.go b/lib/srv/usermgmt.go index ba9feb0eace..6083a1c88f7 100644 --- a/lib/srv/usermgmt.go +++ b/lib/srv/usermgmt.go @@ -30,9 +30,9 @@ import ( log "github.com/sirupsen/logrus" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/services/local" - "github.com/gravitational/teleport/lib/utils" ) // NewHostUsers initialize a new HostUsers object @@ -259,7 +259,7 @@ func (u *HostUserManagement) doWithUserLock(f func(types.SemaphoreLease) error) MaxLeases: 1, Expires: time.Now().Add(time.Second * 20), }, - Retry: utils.LinearConfig{ + Retry: retryutils.LinearConfig{ Step: time.Second * 5, Max: time.Minute, }, diff --git a/lib/tbot/ca_rotation.go b/lib/tbot/ca_rotation.go index bc4f7bc47da..3bc09158a39 100644 --- a/lib/tbot/ca_rotation.go +++ b/lib/tbot/ca_rotation.go @@ -22,11 +22,12 @@ import ( "sync" "time" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils" - libUtils "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" "github.com/sirupsen/logrus" + + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/api/utils/retryutils" ) // See https://github.com/gravitational/teleport/blob/1aa38f4bc56997ba13b26a1ef1b4da7a3a078930/lib/auth/rotate.go#L135 @@ -105,7 +106,7 @@ func (b *Bot) caRotationLoop(ctx context.Context) error { }, debouncePeriod: time.Second * 10, } - jitter := libUtils.NewJitter() + jitter := retryutils.NewJitter() for { err := b.watchCARotations(ctx, rd.attempt) diff --git a/lib/tbot/renew.go b/lib/tbot/renew.go index c17c9f71c4f..8945af83db2 100644 --- a/lib/tbot/renew.go +++ b/lib/tbot/renew.go @@ -25,10 +25,14 @@ import ( "strings" "time" + "github.com/gravitational/trace" + "golang.org/x/crypto/ssh" + "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/authclient" "github.com/gravitational/teleport/lib/auth/native" @@ -40,8 +44,6 @@ import ( "github.com/gravitational/teleport/lib/tbot/identity" "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/utils" - "github.com/gravitational/trace" - "golang.org/x/crypto/ssh" ) // generateKeys generates TLS and SSH keypairs. @@ -674,7 +676,7 @@ func (b *Bot) renewLoop(ctx context.Context) error { } ticker := time.NewTicker(b.cfg.RenewalInterval) - jitter := utils.NewJitter() + jitter := retryutils.NewJitter() defer ticker.Stop() for { var err error diff --git a/lib/utils/interval/interval.go b/lib/utils/interval/interval.go index d3fad8e67dc..04c5eab27e4 100644 --- a/lib/utils/interval/interval.go +++ b/lib/utils/interval/interval.go @@ -21,7 +21,7 @@ import ( "sync" "time" - "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/api/utils/retryutils" ) // Interval functions similarly to time.Ticker, with the added benefit of being @@ -54,7 +54,7 @@ type Config struct { // It is usually preferable to use a smaller jitter (e.g. NewSeventhJitter()) // for this parameter, since periodic operations are typically costly and the // effect of the jitter is cumulative. - Jitter utils.Jitter + Jitter retryutils.Jitter } // NewNoop creates a new interval that will never fire. diff --git a/lib/utils/retry.go b/lib/utils/retry.go index feb2d612cd4..fa98cea920d 100644 --- a/lib/utils/retry.go +++ b/lib/utils/retry.go @@ -17,173 +17,25 @@ limitations under the License. package utils import ( - "context" - "fmt" - "math/rand" - "sync" "time" - "github.com/gravitational/trace" - "github.com/jonboulle/clockwork" - log "github.com/sirupsen/logrus" + "github.com/gravitational/teleport/api/utils/retryutils" ) // HalfJitter is a global jitter instance used for one-off jitters. // Prefer instantiating a new jitter instance for operations that require // repeated calls. -var HalfJitter = NewHalfJitter() +var HalfJitter = retryutils.NewHalfJitter() // SeventhJitter is a global jitter instance used for one-off jitters. // Prefer instantiating a new jitter instance for operations that require // repeated calls. -var SeventhJitter = NewSeventhJitter() +var SeventhJitter = retryutils.NewSeventhJitter() // FullJitter is a global jitter instance used for one-off jitters. // Prefer instantiating a new jitter instance for operations that require // repeated calls -var FullJitter = NewFullJitter() - -// Jitter is a function which applies random jitter to a -// duration. Used to randomize backoff values. Must be -// safe for concurrent usage. -type Jitter func(time.Duration) time.Duration - -// NewJitter builds a new default jitter (currently jitters on -// the range [n/2,n), but this is subject to change). -func NewJitter() Jitter { - return NewHalfJitter() -} - -// NewHalfJitter returns a new jitter on the range [n/2,n). This is -// a large range and most suitable for jittering things like backoff -// operations where breaking cycles quickly is a priority. -func NewHalfJitter() Jitter { - var mu sync.Mutex - rng := rand.New(rand.NewSource(time.Now().UnixNano())) - return func(d time.Duration) time.Duration { - // values less than 1 cause rng to panic, and some logic - // relies on treating zero duration as non-blocking case. - if d < 1 { - return 0 - } - mu.Lock() - defer mu.Unlock() - return (d / 2) + time.Duration(rng.Int63n(int64(d))/2) - } -} - -// NewSeventhJitter builds a new jitter on the range [6n/7,n). Prefer smaller -// jitters such as this when jittering periodic operations (e.g. cert rotation -// checks) since large jitters result in significantly increased load. -func NewSeventhJitter() Jitter { - var mu sync.Mutex - rng := rand.New(rand.NewSource(time.Now().UnixNano())) - return func(d time.Duration) time.Duration { - // values less than 1 cause rng to panic, and some logic - // relies on treating zero duration as non-blocking case. - if d < 1 { - return 0 - } - mu.Lock() - defer mu.Unlock() - return (6 * d / 7) + time.Duration(rng.Int63n(int64(d))/7) - } -} - -// NewFullJitter builds a new jitter on the range (0,n]. Most use-cases -// are better served by a jitter with a meaningful minimum value, but if -// the *only* purpose of the jitter is to spread out retries to the greatest -// extent possible (e.g. when retrying a CompareAndSwap operation), a full jitter -// may be appropriate. -func NewFullJitter() Jitter { - var mu sync.Mutex - rng := rand.New(rand.NewSource(time.Now().UnixNano())) - return func(d time.Duration) time.Duration { - // values less than 1 cause rng to panic, and some logic - // relies on treating zero duration as non-blocking case. - if d < 1 { - return 0 - } - mu.Lock() - defer mu.Unlock() - return time.Duration(1) + time.Duration(rng.Int63n(int64(d))) - } -} - -// Retry is an interface that provides retry logic -type Retry interface { - // Reset resets retry state - Reset() - // Inc increments retry attempt - Inc() - // Duration returns retry duration, - // could be 0 - Duration() time.Duration - // After returns time.Time channel - // that fires after Duration delay, - // could fire right away if Duration is 0 - After() <-chan time.Time - // Clone creates a copy of this retry in a - // reset state. - Clone() Retry -} - -// LinearConfig sets up retry configuration -// using arithmetic progression -type LinearConfig struct { - // First is a first element of the progression, - // could be 0 - First time.Duration - // Step is a step of the progression, can't be 0 - Step time.Duration - // Max is a maximum value of the progression, - // can't be 0 - Max time.Duration - // Jitter is an optional jitter function to be applied - // to the delay. Note that supplying a jitter means that - // successive calls to Duration may return different results. - Jitter Jitter `json:"-"` - // AutoReset, if greater than zero, causes the linear retry to automatically - // reset after Max * AutoReset has elapsed since the last call to Incr. - AutoReset int64 - // Clock to override clock in tests - Clock clockwork.Clock -} - -// CheckAndSetDefaults checks and sets defaults -func (c *LinearConfig) CheckAndSetDefaults() error { - if c.Step == 0 { - return trace.BadParameter("missing parameter Step") - } - if c.Max == 0 { - return trace.BadParameter("missing parameter Max") - } - if c.Clock == nil { - c.Clock = clockwork.NewRealClock() - } - return nil -} - -// NewLinear returns a new instance of linear retry -func NewLinear(cfg LinearConfig) (*Linear, error) { - if err := cfg.CheckAndSetDefaults(); err != nil { - return nil, trace.Wrap(err) - } - return newLinear(cfg), nil -} - -// newLinear creates an instance of Linear from a -// previously verified configuration. -func newLinear(cfg LinearConfig) *Linear { - closedChan := make(chan time.Time) - close(closedChan) - return &Linear{LinearConfig: cfg, closedChan: closedChan} -} - -// NewConstant returns a new linear retry with constant interval. -func NewConstant(interval time.Duration) (*Linear, error) { - return NewLinear(LinearConfig{Step: interval, Max: interval}) -} +var FullJitter = retryutils.NewFullJitter() // NewDefaultLinear creates a linear retry using a half jitter, 10s step, and maxing out // at 1 minute. These values were selected by reviewing commonly used parameters elsewhere @@ -192,12 +44,12 @@ func NewConstant(interval time.Duration) (*Linear, error) { // registration and auth connector setup. It also includes an auto-reset value of 5m. Auto-reset // is less commonly used, and if used should probably be shorter, but 5m is a reasonable // safety net to reduce the impact of accidental misuse. -func NewDefaultLinear() *Linear { - retry, err := NewLinear(LinearConfig{ +func NewDefaultLinear() *retryutils.Linear { + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ First: HalfJitter(time.Second * 5), Step: time.Second * 10, Max: time.Minute, - Jitter: NewHalfJitter(), + Jitter: retryutils.NewHalfJitter(), AutoReset: 5, }) if err != nil { @@ -205,133 +57,3 @@ func NewDefaultLinear() *Linear { } return retry } - -// Linear is used to calculate retry period -// that follows the following logic: -// On the first error there is no delay -// on the next error, delay is FastLinear -// on all other errors, delay is SlowLinear -type Linear struct { - // LinearConfig is a linear retry config - LinearConfig - lastUse time.Time - attempt int64 - closedChan chan time.Time -} - -// Reset resets retry period to initial state -func (r *Linear) Reset() { - r.attempt = 0 -} - -// ResetToDelay resets retry period and increments the number of attempts. -func (r *Linear) ResetToDelay() { - r.Reset() - r.Inc() -} - -// Clone creates an identical copy of Linear with fresh state. -func (r *Linear) Clone() Retry { - return newLinear(r.LinearConfig) -} - -// Inc increments attempt counter -func (r *Linear) Inc() { - r.attempt++ -} - -// Duration returns retry duration based on state -func (r *Linear) Duration() time.Duration { - if r.AutoReset > 0 { - now := r.Clock.Now() - if now.After(r.lastUse.Add(r.Max * time.Duration(r.AutoReset))) { - r.Reset() - } - r.lastUse = now - } - - a := r.First + time.Duration(r.attempt)*r.Step - if a < 1 { - return 0 - } - - if a > r.Max { - a = r.Max - } - - if r.Jitter != nil { - a = r.Jitter(a) - } - - return a -} - -// After returns channel that fires with timeout -// defined in Duration method, as a special case -// if Duration is 0 returns a closed channel -func (r *Linear) After() <-chan time.Time { - d := r.Duration() - if d < 1 { - return r.closedChan - } - return r.Clock.After(d) -} - -// String returns user-friendly representation of the LinearPeriod -func (r *Linear) String() string { - return fmt.Sprintf("Linear(attempt=%v, duration=%v)", r.attempt, r.Duration()) -} - -// For retries the provided function until it succeeds or the context expires. -func (r *Linear) For(ctx context.Context, retryFn func() error) error { - for { - err := retryFn() - if err == nil { - return nil - } - if _, ok := trace.Unwrap(err).(*permanentRetryError); ok { - return trace.Wrap(err) - } - log.Debugf("Will retry in %v: %v.", r.Duration(), err) - select { - case <-r.After(): - r.Inc() - case <-ctx.Done(): - return trace.LimitExceeded(ctx.Err().Error()) - } - } -} - -// PermanentRetryError returns a new instance of a permanent retry error. -func PermanentRetryError(err error) error { - return &permanentRetryError{err: err} -} - -// permanentRetryError indicates that retry loop should stop. -type permanentRetryError struct { - err error -} - -// Error returns the original error message. -func (e *permanentRetryError) Error() string { - return e.err.Error() -} - -// RetryFastFor retries a function repeatedly for a set amount of -// time before returning an error. -// -// Intended mostly for tests. -func RetryStaticFor(d time.Duration, w time.Duration, f func() error) error { - start := time.Now() - var err error - - for time.Since(start) < d { - if err = f(); err == nil { - break - } - - time.Sleep(w) - } - - return err -} diff --git a/lib/utils/utils_test.go b/lib/utils/utils_test.go index 2cddf425de6..020622828ad 100644 --- a/lib/utils/utils_test.go +++ b/lib/utils/utils_test.go @@ -27,6 +27,7 @@ import ( "time" "github.com/google/uuid" + "github.com/gravitational/teleport/lib/fixtures" "github.com/stretchr/testify/require" @@ -39,28 +40,6 @@ func TestMain(m *testing.M) { os.Exit(m.Run()) } -// TestLinear tests retry logic -func TestLinear(t *testing.T) { - t.Parallel() - - r, err := NewLinear(LinearConfig{ - Step: time.Second, - Max: 3 * time.Second, - }) - require.NoError(t, err) - require.Equal(t, r.Duration(), time.Duration(0)) - r.Inc() - require.Equal(t, r.Duration(), time.Second) - r.Inc() - require.Equal(t, r.Duration(), 2*time.Second) - r.Inc() - require.Equal(t, r.Duration(), 3*time.Second) - r.Inc() - require.Equal(t, r.Duration(), 3*time.Second) - r.Reset() - require.Equal(t, r.Duration(), time.Duration(0)) -} - func TestHostUUIDIdempotent(t *testing.T) { t.Parallel() diff --git a/operator/main.go b/operator/main.go index 5a611d6f70e..255edcdc829 100644 --- a/operator/main.go +++ b/operator/main.go @@ -35,7 +35,7 @@ import ( "sigs.k8s.io/controller-runtime/pkg/healthz" "sigs.k8s.io/controller-runtime/pkg/log/zap" - "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/api/utils/retryutils" resourcesv2 "github.com/gravitational/teleport/operator/apis/resources/v2" resourcesv5 "github.com/gravitational/teleport/operator/apis/resources/v5" resourcescontrollers "github.com/gravitational/teleport/operator/controllers/resources" @@ -109,7 +109,7 @@ func main() { var bot *sidecar.Bot - retry, err := utils.NewLinear(utils.LinearConfig{ + retry, err := retryutils.NewLinear(retryutils.LinearConfig{ Step: 100 * time.Millisecond, Max: time.Second, }) diff --git a/tool/tbot/kube_test.go b/tool/tbot/kube_test.go index eb78b85938f..a614c45a8ab 100644 --- a/tool/tbot/kube_test.go +++ b/tool/tbot/kube_test.go @@ -24,12 +24,14 @@ import ( "testing" "time" + "github.com/jonboulle/clockwork" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" + "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/identityfile" "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/tlsca" - "github.com/jonboulle/clockwork" - "github.com/stretchr/testify/require" ) func TestGetKubeCredentialData(t *testing.T) { @@ -54,7 +56,7 @@ func TestGetKubeCredentialData(t *testing.T) { idFile := &identityfile.IdentityFile{ PrivateKey: privateKeyBytes, Certs: identityfile.Certs{ - SSH: []byte("ssh ssh-cert"), // dummy value + SSH: []byte(ssh.CertAlgoRSAv01), // dummy value TLS: certBytes, }, CACerts: identityfile.CACerts{ diff --git a/tool/tsh/proxy_test.go b/tool/tsh/proxy_test.go index d7eac397933..fd9268292c0 100644 --- a/tool/tsh/proxy_test.go +++ b/tool/tsh/proxy_test.go @@ -37,6 +37,7 @@ import ( apidefaults "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/defaults" @@ -550,7 +551,7 @@ func runOpenSSHCommand(t *testing.T, configFile string, sshConnString string, po } func mustRunOpenSSHCommand(t *testing.T, configFile string, sshConnString string, port int, args ...string) { - err := utils.RetryStaticFor(time.Second*10, time.Millisecond*500, func() error { + err := retryutils.RetryStaticFor(time.Second*10, time.Millisecond*500, func() error { err := runOpenSSHCommand(t, configFile, sshConnString, port, args...) return trace.Wrap(err) }) diff --git a/tool/tsh/tsh.go b/tool/tsh/tsh.go index 69bc86b4fd4..e7813714dcb 100644 --- a/tool/tsh/tsh.go +++ b/tool/tsh/tsh.go @@ -1458,6 +1458,7 @@ func onLogin(cf *CLIConf) error { } return trace.Wrap(err) } + tc.AllowStdinHijack = false // the login operation may update the username and should be considered the more @@ -1466,7 +1467,6 @@ func onLogin(cf *CLIConf) error { // TODO(fspmarshall): Refactor access request & cert reissue logic to allow // access requests to be applied to identity files. - if makeIdentityFile { if err := setupNoninteractiveClient(tc, key); err != nil { return trace.Wrap(err) @@ -3269,7 +3269,7 @@ func onShow(cf *CLIConf) error { return trace.Wrap(err) } - fmt.Printf("Cert: %#v\nPriv: %#v\nPub: %#v\n", cert, key.GetBaseSigner(), key.MarshalSSHPublicKey()) + fmt.Printf("Cert: %#v\nPriv: %#v\nPub: %#v\n", cert, key.Signer, key.MarshalSSHPublicKey()) fmt.Printf("Fingerprint: %s\n", ssh.FingerprintSHA256(key.SSHPublicKey())) return nil }