From 1d804737f9f112768a50041571ebdbcf2cffe209 Mon Sep 17 00:00:00 2001 From: Anunay Maheshwari Date: Thu, 16 Apr 2026 01:12:33 +0530 Subject: [PATCH] fix: limit join request and WHIP request body to http.DefaultMaxHeaderBytes (#4450) * fix: CS-1665 * cleanup * cleanup and testes * updates --- pkg/config/config.go | 2 +- pkg/service/rtcservice.go | 27 ++++++++++++-------- pkg/service/utils.go | 29 +++++++++++++++++++++ pkg/service/utils_test.go | 52 ++++++++++++++++++++++++++++++++++++++ pkg/service/whipservice.go | 17 ++++++++++--- 5 files changed, 111 insertions(+), 16 deletions(-) diff --git a/pkg/config/config.go b/pkg/config/config.go index b8806c02f..1da51fd5d 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -439,7 +439,7 @@ var DefaultConfig = Config{ Agents: agent.Config{ TargetLoad: agent.DefaultTargetLoad, }, - PSRPC: rpc.DefaultPSRPCConfig, + PSRPC: rpc.DefaultPSRPCConfig, Keys: map[string]string{}, Metric: metric.DefaultMetricConfig, WebHook: webhook.DefaultWebHookConfig, diff --git a/pkg/service/rtcservice.go b/pkg/service/rtcservice.go index 0747918ad..04685b1e9 100644 --- a/pkg/service/rtcservice.go +++ b/pkg/service/rtcservice.go @@ -15,14 +15,11 @@ package service import ( - "bytes" - "compress/gzip" "context" "encoding/base64" "encoding/json" "errors" "fmt" - "io" "maps" "math/rand" "net/http" @@ -128,9 +125,7 @@ func decodeAttributes(str string) (map[string]string, error) { return attrs, nil } -var gzipReaderPool = sync.Pool{ - New: func() any { return &gzip.Reader{} }, -} +var errJoinRequestTooLarge = errors.New("join request too large") func (s *RTCService) validateInternal( lgr logger.Logger, @@ -138,6 +133,10 @@ func (s *RTCService) validateInternal( needsJoinRequest bool, strict bool, ) (livekit.RoomName, routing.ParticipantInit, int, error) { + if claims := GetGrants(r.Context()); claims == nil || claims.Video == nil { + return "", routing.ParticipantInit{}, http.StatusUnauthorized, rtc.ErrPermissionDenied + } + var params ValidateConnectRequestParams useSinglePeerConnection := false joinRequest := &livekit.JoinRequest{} @@ -174,17 +173,23 @@ func (s *RTCService) validateInternal( switch wrappedJoinRequest.Compression { case livekit.WrappedJoinRequest_NONE: + if len(wrappedJoinRequest.JoinRequest) > http.DefaultMaxHeaderBytes { + return "", routing.ParticipantInit{}, http.StatusBadRequest, errJoinRequestTooLarge + } if err := proto.Unmarshal(wrappedJoinRequest.JoinRequest, joinRequest); err != nil { return "", routing.ParticipantInit{}, http.StatusBadRequest, errors.New("cannot unmarshal join request") } case livekit.WrappedJoinRequest_GZIP: - reader := gzipReaderPool.Get().(*gzip.Reader) - defer gzipReaderPool.Put(reader) - reader.Reset(bytes.NewReader(wrappedJoinRequest.JoinRequest)) - protoBytes, err := io.ReadAll(reader) + protoBytes, err := DecompressGzip(wrappedJoinRequest.JoinRequest) if err != nil { - return "", routing.ParticipantInit{}, http.StatusBadRequest, errors.New("cannot read decompressed join request") + switch { + case errors.Is(err, ErrGzipTooLarge): + err = errJoinRequestTooLarge + case errors.Is(err, ErrGzipReadFailed): + err = errors.New("cannot read decompressed join request") + } + return "", routing.ParticipantInit{}, http.StatusBadRequest, err } if err := proto.Unmarshal(protoBytes, joinRequest); err != nil { diff --git a/pkg/service/utils.go b/pkg/service/utils.go index 7d2ef75d7..92189710e 100644 --- a/pkg/service/utils.go +++ b/pkg/service/utils.go @@ -15,10 +15,13 @@ package service import ( + "bytes" + "compress/gzip" "context" "encoding/json" "errors" "fmt" + "io" "net" "net/http" "regexp" @@ -39,6 +42,32 @@ import ( "github.com/livekit/protocol/logger" ) +var ( + ErrGzipReadFailed = errors.New("cannot read decompressed data") + ErrGzipTooLarge = errors.New("decompressed data too large") +) + +var gzipReaderPool = sync.Pool{ + New: func() any { return &gzip.Reader{} }, +} + +func DecompressGzip(compressed []byte) ([]byte, error) { + reader := gzipReaderPool.Get().(*gzip.Reader) + defer gzipReaderPool.Put(reader) + if err := reader.Reset(bytes.NewReader(compressed)); err != nil { + return nil, fmt.Errorf("%w: %w", ErrGzipReadFailed, err) + } + + out, err := io.ReadAll(io.LimitReader(reader, http.DefaultMaxHeaderBytes+1)) + if err != nil { + return nil, fmt.Errorf("%w: %w", ErrGzipReadFailed, err) + } + if len(out) > http.DefaultMaxHeaderBytes { + return nil, ErrGzipTooLarge + } + return out, nil +} + func handleError(w http.ResponseWriter, r *http.Request, status int, err error, keysAndValues ...any) { keysAndValues = append(keysAndValues, "status", status) if r != nil && r.URL != nil { diff --git a/pkg/service/utils_test.go b/pkg/service/utils_test.go index 5d8cdcaf0..1f675c6f8 100644 --- a/pkg/service/utils_test.go +++ b/pkg/service/utils_test.go @@ -15,7 +15,10 @@ package service_test import ( + "bytes" + "compress/gzip" "context" + "net/http" "testing" "time" @@ -75,3 +78,52 @@ func TestIsValidDomain(t *testing.T) { require.Equal(t, service.IsValidDomain(key), result) } } + +func compress(t *testing.T, payload []byte) []byte { + t.Helper() + var buf bytes.Buffer + gw, err := gzip.NewWriterLevel(&buf, gzip.BestCompression) + require.NoError(t, err) + _, err = gw.Write(payload) + require.NoError(t, err) + require.NoError(t, gw.Close()) + return buf.Bytes() +} + +func TestDecompressGzip(t *testing.T) { + t.Run("small payload", func(t *testing.T) { + out, err := service.DecompressGzip(compress(t, []byte("hello world"))) + require.NoError(t, err) + require.Equal(t, []byte("hello world"), out) + }) + + t.Run("payload exactly at cap", func(t *testing.T) { + raw := make([]byte, http.DefaultMaxHeaderBytes) + out, err := service.DecompressGzip(compress(t, raw)) + require.NoError(t, err) + require.Len(t, out, http.DefaultMaxHeaderBytes) + }) + + t.Run("payload one byte over capd", func(t *testing.T) { + raw := make([]byte, http.DefaultMaxHeaderBytes+1) + _, err := service.DecompressGzip(compress(t, raw)) + require.ErrorIs(t, err, service.ErrGzipTooLarge) + }) + + t.Run("gzip decompression bomb", func(t *testing.T) { + // 100 MB of zeros + raw := make([]byte, 100<<20) + compressed := compress(t, raw) + require.Less(t, len(compressed), 1<<20, + "sanity: bomb input should compress dramatically") + + _, err := service.DecompressGzip(compressed) + require.ErrorIs(t, err, service.ErrGzipTooLarge) + }) + + t.Run("malformed gzip compression", func(t *testing.T) { + _, err := service.DecompressGzip([]byte("not gzip data")) + require.Error(t, err) + require.Contains(t, err.Error(), "cannot read decompressed") + }) +} diff --git a/pkg/service/whipservice.go b/pkg/service/whipservice.go index 244bcc85e..7f5f20fd9 100644 --- a/pkg/service/whipservice.go +++ b/pkg/service/whipservice.go @@ -119,7 +119,7 @@ type createRequest struct { FromIngress bool } -func (s *WHIPService) validateCreate(r *http.Request) (*createRequest, int, error) { +func (s *WHIPService) validateCreate(w http.ResponseWriter, r *http.Request) (*createRequest, int, error) { claims := GetGrants(r.Context()) if claims == nil || claims.Video == nil { return nil, http.StatusUnauthorized, rtc.ErrPermissionDenied @@ -156,8 +156,12 @@ func (s *WHIPService) validateCreate(r *http.Request) (*createRequest, int, erro fromIngress := r.Header.Get("X-Livekit-Ingress") - offerSDPBytes, err := io.ReadAll(r.Body) + offerSDPBytes, err := io.ReadAll(http.MaxBytesReader(w, r.Body, http.DefaultMaxHeaderBytes)) if err != nil { + var maxErr *http.MaxBytesError + if errors.As(err, &maxErr) { + return nil, http.StatusRequestEntityTooLarge, fmt.Errorf("request body exceeds %d bytes", maxErr.Limit) + } return nil, http.StatusBadRequest, fmt.Errorf("body does not have SDP offer: %s", err) } if len(offerSDPBytes) == 0 { @@ -212,7 +216,7 @@ func (s *WHIPService) handleCreate(w http.ResponseWriter, r *http.Request) { w.Header().Add("Content-type", "application/sdp") - req, status, err := s.validateCreate(r) + req, status, err := s.validateCreate(w, r) if err != nil { s.handleError("Create", w, r, status, err) return @@ -462,8 +466,13 @@ func (s *WHIPService) handleParticipantPatch(w http.ResponseWriter, r *http.Requ return } - sdpFragmentBytes, err := io.ReadAll(r.Body) + sdpFragmentBytes, err := io.ReadAll(http.MaxBytesReader(w, r.Body, http.DefaultMaxHeaderBytes)) if err != nil { + var maxErr *http.MaxBytesError + if errors.As(err, &maxErr) { + s.handleError("Patch", w, r, http.StatusRequestEntityTooLarge, fmt.Errorf("request body exceeds %d bytes", maxErr.Limit)) + return + } s.handleError("Patch", w, r, http.StatusBadRequest, fmt.Errorf("body does not have SDP fragment: %s", err)) return }