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
This commit is contained in:
Edoardo Spadolini
2026-05-22 23:03:52 +00:00
committed by GitHub
parent 57ee7542c8
commit 6413a5ae2e
4 changed files with 170 additions and 36 deletions
+10
View File
@@ -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
+38 -27
View File
@@ -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)
+92 -9
View File
@@ -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)
+30
View File
@@ -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