mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user