diff --git a/lib/auth/authclient/clt.go b/lib/auth/authclient/clt.go index ede474ac3f4..184e6b667e3 100644 --- a/lib/auth/authclient/clt.go +++ b/lib/auth/authclient/clt.go @@ -180,6 +180,11 @@ func NewClient(cfg client.Config, params ...roundtrip.ClientParam) (*Client, err TLS: httpTLS, Dialer: httpDialer, ALPNSNIAuthDialClusterName: cfg.ALPNSNIAuthDialClusterName, + // we are ok with the HTTP client using a separate circuit breaker with + // the same configuration, since there's so few auth API calls using + // HTTP and most of them are used by the Proxy, which is going to have a + // no-op circuit breaker to begin with + CircuitBreakerConfig: cfg.CircuitBreakerConfig, } httpClient, err := NewHTTPClient(httpClientCfg, params...) if err != nil { diff --git a/lib/auth/authclient/clt_test.go b/lib/auth/authclient/clt_test.go index 9f7fa33e288..4b055182e8f 100644 --- a/lib/auth/authclient/clt_test.go +++ b/lib/auth/authclient/clt_test.go @@ -19,11 +19,25 @@ package authclient import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "net" + "net/http" "testing" + "testing/synctest" + "time" "github.com/google/go-cmp/cmp" + "github.com/gravitational/trace" "github.com/stretchr/testify/require" + "google.golang.org/grpc/test/bufconn" + "github.com/gravitational/teleport/api/breaker" + "github.com/gravitational/teleport/api/client" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/fixtures" ) @@ -84,3 +98,70 @@ func TestValidateTrustedClusterResponseProto(t *testing.T) { backToNative := ValidateTrustedClusterResponseFromProto(proto) require.Empty(t, cmp.Diff(native, backToNative)) } + +func TestHTTPCircuitBreaker(t *testing.T) { + synctest.Test(t, synctestHTTPCircuitBreaker) +} +func synctestHTTPCircuitBreaker(t *testing.T) { + srv := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "this hit the server", 500) + }), + TLSConfig: &tls.Config{ + Certificates: []tls.Certificate{requireSnakeoilCert(t)}, + }, + } + listener := bufconn.Listen(100) + + serveErr := make(chan error, 1) + go func() { + serveErr <- srv.ServeTLS(listener, "", "") + }() + defer func() { + require.ErrorIs(t, <-serveErr, http.ErrServerClosed) + }() + + defer func() { + require.NoError(t, srv.Close()) + }() + + clt, err := NewClient(client.Config{ + Dialer: client.ContextDialerFunc(func(ctx context.Context, network, addr string) (net.Conn, error) { + return listener.DialContext(ctx) + }), + Credentials: []client.Credentials{ + client.LoadTLS(&tls.Config{InsecureSkipVerify: true}), + }, + CircuitBreakerConfig: breaker.Config{ + Interval: time.Hour, + Trip: breaker.StaticTripper(true), + Recover: breaker.StaticTripper(true), + IsSuccessful: func(v any, err error) bool { return false }, + }, + DialInBackground: true, + }) + require.NoError(t, err) + defer clt.Close() + + _, err = clt.HTTPClient.Get(t.Context(), clt.HTTPClient.Endpoint(), nil) + require.ErrorContains(t, err, "this hit the server") + + // the breaker should be tripped now, unlike what a default circuit breaker would do + + _, err = clt.HTTPClient.Get(t.Context(), clt.HTTPClient.Endpoint(), nil) + require.ErrorAs(t, err, new(*trace.ConnectionProblemError)) + require.ErrorContains(t, err, "Unable to communicate with the Teleport Auth Service") +} + +func requireSnakeoilCert(t testing.TB) tls.Certificate { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + cert := &x509.Certificate{} + certDER, err := x509.CreateCertificate(rand.Reader, cert, cert, key.Public(), key) + require.NoError(t, err) + return tls.Certificate{ + Certificate: [][]byte{certDER}, + PrivateKey: key, + } +}