From 6413a5ae2ee7dfbb36c2fa4dd962d5783271f229 Mon Sep 17 00:00:00 2001 From: Edoardo Spadolini Date: Sat, 23 May 2026 01:03:52 +0200 Subject: [PATCH] Add an optional unstable rate limit for ResolveSSHTarget (#66952) * Move CreateAuditStream unstable envvar to lib/service * Add an optional rate limit for ResolveSSHTarget * Fix logging oddities * Actually set the create audit stream limit metric * Fix and synctest TestCreateAuditStreamLimit --- lib/auth/apiserver.go | 10 ++++ lib/auth/grpcserver.go | 65 +++++++++++++---------- lib/auth/grpcserver_test.go | 101 ++++++++++++++++++++++++++++++++---- lib/service/service.go | 30 +++++++++++ 4 files changed, 170 insertions(+), 36 deletions(-) diff --git a/lib/auth/apiserver.go b/lib/auth/apiserver.go index be65b7b4489..e7f9e937c41 100644 --- a/lib/auth/apiserver.go +++ b/lib/auth/apiserver.go @@ -68,6 +68,16 @@ type APIConfig struct { // DisableJoinV1 disables registration of the new join gRPC service. // Intended for tests that need to exercise legacy join fallback paths. DisableJoinV1 bool + // CreateAuditStreamInflightLimit, if set, is the maximum amount of allowed + // in-flight CreateAuditStream rpc calls. Calls beyond the limit will + // immediately return with an error. A non-positive limit means that no + // calls will be allowed. + CreateAuditStreamInflightLimit *int + // ResolveSSHTargetRateLimit, if set, is the (server-wide) rate limit for + // the ResolveSSHTarget rpc (i.e. the number of allowed calls per second), + // with an allowed burst rate equal to the rate per second (rounded up). + // Calls beyond the limit will block and wait for their turn. + ResolveSSHTargetRateLimit *float64 } // CheckAndSetDefaults checks and sets default values diff --git a/lib/auth/grpcserver.go b/lib/auth/grpcserver.go index 36590e3f80f..ac5c3bf4d17 100644 --- a/lib/auth/grpcserver.go +++ b/lib/auth/grpcserver.go @@ -26,10 +26,9 @@ import ( "io" "iter" "log/slog" + "math" "net" - "os" goslices "slices" - "strconv" "strings" "time" @@ -39,6 +38,7 @@ import ( "github.com/prometheus/client_golang/prometheus" "go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc" collectortracepb "go.opentelemetry.io/proto/otlp/collector/trace/v1" + "golang.org/x/time/rate" "google.golang.org/grpc" "google.golang.org/grpc/codes" _ "google.golang.org/grpc/encoding/gzip" // gzip compressor for gRPC. @@ -232,17 +232,23 @@ type GRPCServer struct { // collect and forward spans collectortracepb.TraceServiceServer - // createAuditStreamSemaphore, if not nil, is used to limit the amount of - // in-flight CreateAuditStream RPCs, by sending a value in at the beginning - // of the RPC and pulling one out before returning. - createAuditStreamSemaphore chan struct{} - // createAuthenticateChallengeLimiter is a rate limiter for invocations of // /proto.AuthService/CreateAuthenticateChallenge that don't rely on a user // context and thus warrant additional rate limiting since they are // unauthenticated (either through direct API connections or coming from the // proxy on behalf of a remote unauthenticated user). createAuthenticateChallengeLimiter *limiter.RateLimiter + + // createAuditStreamSemaphore, if not nil, is used to limit the amount of + // in-flight CreateAuditStream RPCs, by sending a value in at the beginning + // of the RPC and pulling one out before returning. + createAuditStreamSemaphore chan struct{} + + // resolveSSHTargetRateLimiter is an optional (server-wide) rate limiter for + // calls to ResolveSSHTarget, since those might end up iterating over the + // whole inventory of nodes and that can be quite heavy. Calls beyond the + // rate limit will block until their execution would be allowed. + resolveSSHTargetRateLimiter *rate.Limiter } func (g *GRPCServer) SetServingStatus(service string, servingStatus grpc_health_v1.HealthCheckResponse_ServingStatus) { @@ -5104,6 +5110,12 @@ func (g *GRPCServer) ResolveSSHTarget(ctx context.Context, req *authpb.ResolveSS return nil, trace.Wrap(err) } + if l := g.resolveSSHTargetRateLimiter; l != nil { + if err := l.Wait(ctx); err != nil { + return nil, trace.Wrap(err) + } + } + rsp, err := auth.ServerWithRoles.ResolveSSHTarget(ctx, req) if err != nil { return nil, trace.Wrap(err) @@ -6386,6 +6398,23 @@ func NewGRPCServer(cfg GRPCServerConfig) (*GRPCServer, error) { return nil, trace.Wrap(err) } + var createAuditStreamSemaphore chan struct{} + if cfg.CreateAuditStreamInflightLimit != nil { + metrics.RegisterPrometheusCollectors( + createAuditStreamAcceptedTotalMetric, + createAuditStreamRejectedTotalMetric, + createAuditStreamLimitMetric, + ) + limit := max(0, *cfg.CreateAuditStreamInflightLimit) + createAuditStreamLimitMetric.Set(float64(limit)) + createAuditStreamSemaphore = make(chan struct{}, limit) + } + + var resolveSSHTargetRateLimiter *rate.Limiter + if cfg.ResolveSSHTargetRateLimit != nil { + resolveSSHTargetRateLimiter = rate.NewLimiter(rate.Limit(*cfg.ResolveSSHTargetRateLimit), int(math.Ceil(*cfg.ResolveSSHTargetRateLimit))) + } + authServer := &GRPCServer{ APIConfig: cfg.APIConfig, logger: logger, @@ -6393,26 +6422,8 @@ func NewGRPCServer(cfg GRPCServerConfig) (*GRPCServer, error) { healthcheck: health.NewServer(), createAuthenticateChallengeLimiter: createAuthenticateChallengeLimiter, - } - - if en := os.Getenv("TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT"); en != "" { - inflightLimit, err := strconv.ParseInt(en, 10, 64) - if err != nil { - logger.ErrorContext(context.Background(), "Failed to parse the TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT envvar, limit will not be enforced") - inflightLimit = -1 - } - if inflightLimit == 0 { - logger.WarnContext(context.Background(), "TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT is set to 0, no CreateAuditStream RPCs will be allowed") - } - metrics.RegisterPrometheusCollectors( - createAuditStreamAcceptedTotalMetric, - createAuditStreamRejectedTotalMetric, - createAuditStreamLimitMetric, - ) - createAuditStreamLimitMetric.Set(float64(inflightLimit)) - if inflightLimit >= 0 { - authServer.createAuditStreamSemaphore = make(chan struct{}, inflightLimit) - } + createAuditStreamSemaphore: createAuditStreamSemaphore, + resolveSSHTargetRateLimiter: resolveSSHTargetRateLimiter, } authpb.RegisterAuthServiceServer(server, authServer) diff --git a/lib/auth/grpcserver_test.go b/lib/auth/grpcserver_test.go index 6b12c5a4956..285b7690304 100644 --- a/lib/auth/grpcserver_test.go +++ b/lib/auth/grpcserver_test.go @@ -4810,6 +4810,10 @@ func TestCustomRateLimiting(t *testing.T) { t.Run("unauthenticated CreateAuthenticateChallenge", func(t *testing.T) { synctest.Test(t, synctestCustomRateLimitingUnauthenticatedCreateAuthenticateChallenge) }) + + t.Run("ResolveSSHTarget", func(t *testing.T) { + synctest.Test(t, synctestCustomRateLimitingResolveSSHTarget) + }) } func synctestCustomRateLimitingUnauthenticatedCreateAuthenticateChallenge(t *testing.T) { @@ -4875,6 +4879,68 @@ func synctestCustomRateLimitingUnauthenticatedCreateAuthenticateChallenge(t *tes } } +func synctestCustomRateLimitingResolveSSHTarget(t *testing.T) { + ctx := t.Context() + + as, err := authtest.NewAuthServer(authtest.AuthServerConfig{ + Dir: t.TempDir(), + }) + require.NoError(t, err) + defer as.Close() + + unlimitedSrv, err := as.NewTestTLSServer( + authtest.WithBufconnListener(), + ) + require.NoError(t, err) + defer unlimitedSrv.Close() + + unlimitedClt, err := unlimitedSrv.NewClient(authtest.TestAdmin()) + require.NoError(t, err) + defer unlimitedClt.Close() + + _, err = unlimitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + _, err = unlimitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + _, err = unlimitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + _, err = unlimitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + + limitedSrv, err := as.NewTestTLSServer( + authtest.WithBufconnListener(), + func(c *authtest.TLSServerConfig) { + l := 2.0 + c.APIConfig.ResolveSSHTargetRateLimit = &l + }, + ) + require.NoError(t, err) + defer limitedSrv.Close() + + limitedClt, err := limitedSrv.NewClient(authtest.TestAdmin()) + require.NoError(t, err) + defer limitedClt.Close() + + _, err = limitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + _, err = limitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + require.ErrorAs(t, err, new(*trace.BadParameterError)) + errC := make(chan error, 1) + go func() { + _, err := limitedClt.ResolveSSHTarget(ctx, new(proto.ResolveSSHTargetRequest)) + errC <- err + }() + + synctest.Wait() + + require.Empty(t, errC) + + time.Sleep(time.Second) + synctest.Wait() + + require.ErrorAs(t, <-errC, new(*trace.BadParameterError)) +} + type mockAuthorizer struct { ctx *authz.Context err error @@ -5962,17 +6028,35 @@ func TestGetVnetConfig(t *testing.T) { } func TestCreateAuditStreamLimit(t *testing.T) { - const N = 5 - t.Setenv("TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT", fmt.Sprintf("%d", N)) + synctest.Test(t, synctestCreateAuditStreamLimit) +} +func synctestCreateAuditStreamLimit(t *testing.T) { + const inflightLimit = 5 ctx := t.Context() - server := newTestTLSServer(t) + as, err := authtest.NewAuthServer(authtest.AuthServerConfig{ + Dir: t.TempDir(), + }) + require.NoError(t, err) + defer as.Close() + + server, err := as.NewTestTLSServer( + authtest.WithBufconnListener(), + func(c *authtest.TLSServerConfig) { + n := inflightLimit + c.APIConfig.CreateAuditStreamInflightLimit = &n + }, + ) + require.NoError(t, err) + defer server.Close() + clt, err := server.NewClient(authtest.TestServerID(types.RoleNode, uuid.NewString())) require.NoError(t, err) + defer clt.Close() // HACK(espadolini): we're piggybacking on the prometheus counter which - // can't change while this test is running (we set an envvar, so we can't be + // can't change while this test is running (this is a synctest test that's not // running in parallel with other tests) but it's still pretty awful, and // it'd be much better to actually check that the streams were accepted by // the server; unfortunately, the CreateAuditStream stream doesn't actually @@ -5985,15 +6069,14 @@ func TestCreateAuditStreamLimit(t *testing.T) { } currentAcceptedTotal := getAcceptedTotal() - for range N { + for range inflightLimit { stream, err := clt.CreateAuditStream(ctx, session.NewID()) require.NoError(t, err) - t.Cleanup(func() { stream.Close(ctx) }) + defer stream.Close(ctx) } - require.EventuallyWithT(t, func(t *assert.CollectT) { - require.EqualValues(t, currentAcceptedTotal+N, getAcceptedTotal()) - }, time.Second, 100*time.Millisecond) + synctest.Wait() + require.EqualValues(t, currentAcceptedTotal+inflightLimit, getAcceptedTotal()) ac := proto.NewAuthServiceClient(clt.APIClient.GetConnection()) stream, err := ac.CreateAuditStream(ctx) diff --git a/lib/service/service.go b/lib/service/service.go index 9c20a17f331..85bc1ed376e 100644 --- a/lib/service/service.go +++ b/lib/service/service.go @@ -2666,6 +2666,34 @@ func (process *TeleportProcess) initAuthService() error { return trace.Wrap(err) } } + + var createAuditStreamInflightLimit *int + if en := os.Getenv("TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT"); en != "" { + limit, err := strconv.ParseInt(en, 10, 0) + if err != nil { + logger.ErrorContext(process.ExitContext(), "Failed to parse the TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT envvar, limit will not be enforced", "error", err) + } else if limit >= 0 { + l := int(limit) + createAuditStreamInflightLimit = &l + if *createAuditStreamInflightLimit == 0 { + logger.WarnContext(process.ExitContext(), "TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT is set to 0, no CreateAuditStream RPCs will be allowed") + } else { + logger.DebugContext(process.ExitContext(), "TELEPORT_UNSTABLE_CREATEAUDITSTREAM_INFLIGHT_LIMIT is set, enabling in-flight limit for CreateAuditStream", "limit", *createAuditStreamInflightLimit) + } + } + } + + var resolveSSHTargetRateLimit *float64 + if en := os.Getenv("TELEPORT_UNSTABLE_RESOLVESSHTARGET_RATE_LIMIT"); en != "" { + limit, err := strconv.ParseFloat(en, 64) + if err != nil { + logger.ErrorContext(process.ExitContext(), "Failed to parse the TELEPORT_UNSTABLE_RESOLVESSHTARGET_RATE_LIMIT envvar, limit will not be enforced", "error", err) + } else { + resolveSSHTargetRateLimit = &limit + logger.DebugContext(process.ExitContext(), "TELEPORT_UNSTABLE_RESOLVESSHTARGET_RATE_LIMIT is set, enabling rate limit for ResolveSSHTarget", "limit", *resolveSSHTargetRateLimit) + } + } + apiConf := &auth.APIConfig{ AuthServer: authServer, Authorizer: authorizer, @@ -2680,6 +2708,8 @@ func (process *TeleportProcess) initAuthService() error { CA: accessGraphCAData, Insecure: cfg.AccessGraph.Insecure, }, + CreateAuditStreamInflightLimit: createAuditStreamInflightLimit, + ResolveSSHTargetRateLimit: resolveSSHTargetRateLimit, } // Auth initialization is done (including creation/updating of all singleton