diff --git a/internal/api/api_config_global.go b/internal/api/api_config_global.go index b43e0fa3..a5e56f90 100644 --- a/internal/api/api_config_global.go +++ b/internal/api/api_config_global.go @@ -1,7 +1,6 @@ package api //nolint:revive import ( - "io" "net/http" "github.com/bluenviron/mediamtx/internal/conf" @@ -19,7 +18,7 @@ func (a *API) onConfigGlobalGet(ctx *gin.Context) { func (a *API) onConfigGlobalPatch(ctx *gin.Context) { var c conf.OptionalGlobal - err := jsonwrapper.Decode(io.LimitReader(ctx.Request.Body, maxInboundConfigSize), &c) + err := jsonwrapper.Decode(&customLimitReader{ctx.Request.Body, maxInboundConfigSize}, &c) if err != nil { a.writeError(ctx, http.StatusBadRequest, err) return diff --git a/internal/api/api_config_pathdefaults.go b/internal/api/api_config_pathdefaults.go index d0675720..0c70fbbe 100644 --- a/internal/api/api_config_pathdefaults.go +++ b/internal/api/api_config_pathdefaults.go @@ -1,7 +1,6 @@ package api //nolint:revive import ( - "io" "net/http" "github.com/bluenviron/mediamtx/internal/conf" @@ -19,7 +18,7 @@ func (a *API) onConfigPathDefaultsGet(ctx *gin.Context) { func (a *API) onConfigPathDefaultsPatch(ctx *gin.Context) { var p conf.OptionalPath - err := jsonwrapper.Decode(io.LimitReader(ctx.Request.Body, maxInboundConfigSize), &p) + err := jsonwrapper.Decode(&customLimitReader{ctx.Request.Body, maxInboundConfigSize}, &p) if err != nil { a.writeError(ctx, http.StatusBadRequest, err) return diff --git a/internal/api/api_config_paths.go b/internal/api/api_config_paths.go index 6ac5b9e2..be4c297f 100644 --- a/internal/api/api_config_paths.go +++ b/internal/api/api_config_paths.go @@ -3,7 +3,6 @@ package api //nolint:revive import ( "errors" "fmt" - "io" "net/http" "github.com/bluenviron/mediamtx/internal/conf" @@ -64,7 +63,7 @@ func (a *API) onConfigPathsAdd(ctx *gin.Context) { //nolint:dupl } var p conf.OptionalPath - err := jsonwrapper.Decode(io.LimitReader(ctx.Request.Body, maxInboundConfigSize), &p) + err := jsonwrapper.Decode(&customLimitReader{ctx.Request.Body, maxInboundConfigSize}, &p) if err != nil { a.writeError(ctx, http.StatusBadRequest, err) return @@ -101,7 +100,7 @@ func (a *API) onConfigPathsPatch(ctx *gin.Context) { //nolint:dupl } var p conf.OptionalPath - err := jsonwrapper.Decode(io.LimitReader(ctx.Request.Body, maxInboundConfigSize), &p) + err := jsonwrapper.Decode(&customLimitReader{ctx.Request.Body, maxInboundConfigSize}, &p) if err != nil { a.writeError(ctx, http.StatusBadRequest, err) return @@ -142,7 +141,7 @@ func (a *API) onConfigPathsReplace(ctx *gin.Context) { //nolint:dupl } var p conf.OptionalPath - err := jsonwrapper.Decode(io.LimitReader(ctx.Request.Body, maxInboundConfigSize), &p) + err := jsonwrapper.Decode(&customLimitReader{ctx.Request.Body, maxInboundConfigSize}, &p) if err != nil { a.writeError(ctx, http.StatusBadRequest, err) return diff --git a/internal/api/custom_limit_reader.go b/internal/api/custom_limit_reader.go new file mode 100644 index 00000000..d05b2b21 --- /dev/null +++ b/internal/api/custom_limit_reader.go @@ -0,0 +1,26 @@ +package api + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/auth/custom_limit_reader.go b/internal/auth/custom_limit_reader.go new file mode 100644 index 00000000..1b55d918 --- /dev/null +++ b/internal/auth/custom_limit_reader.go @@ -0,0 +1,26 @@ +package auth + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/auth/manager.go b/internal/auth/manager.go index 938dcc64..9c5adecb 100644 --- a/internal/auth/manager.go +++ b/internal/auth/manager.go @@ -235,7 +235,7 @@ func (m *Manager) authenticateHTTP(req *Request, token string) (string, error) { defer res.Body.Close() if res.StatusCode < 200 || res.StatusCode > 299 { - resBody, err2 := io.ReadAll(io.LimitReader(res.Body, maxInboundBodySize)) + resBody, err2 := io.ReadAll(&customLimitReader{res.Body, maxInboundBodySize}) if err2 == nil && len(resBody) != 0 { return "", fmt.Errorf("server replied with code %d: %s", res.StatusCode, string(resBody)) } @@ -306,7 +306,7 @@ func (m *Manager) pullJWTJWKS() (jwt.Keyfunc, error) { defer res.Body.Close() var raw json.RawMessage - err = json.NewDecoder(io.LimitReader(res.Body, maxInboundBodySize)).Decode(&raw) + err = json.NewDecoder(&customLimitReader{res.Body, maxInboundBodySize}).Decode(&raw) if err != nil { return nil, err } diff --git a/internal/protocols/whip/client.go b/internal/protocols/whip/client.go index f978a38b..2fcdcf0b 100644 --- a/internal/protocols/whip/client.go +++ b/internal/protocols/whip/client.go @@ -264,7 +264,7 @@ func (c *Client) postOffer( return nil, fmt.Errorf("ETag is missing") } - sdp, err := io.ReadAll(io.LimitReader(res.Body, maxInboundSDPSize)) + sdp, err := io.ReadAll(&customLimitReader{res.Body, maxInboundSDPSize}) if err != nil { return nil, err } diff --git a/internal/protocols/whip/custom_limit_reader.go b/internal/protocols/whip/custom_limit_reader.go new file mode 100644 index 00000000..6b4cb119 --- /dev/null +++ b/internal/protocols/whip/custom_limit_reader.go @@ -0,0 +1,26 @@ +package whip + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/servers/hls/hlsjsdownloader/custom_limit_reader.go b/internal/servers/hls/hlsjsdownloader/custom_limit_reader.go new file mode 100644 index 00000000..43540fd8 --- /dev/null +++ b/internal/servers/hls/hlsjsdownloader/custom_limit_reader.go @@ -0,0 +1,26 @@ +package main + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/servers/hls/hlsjsdownloader/main.go b/internal/servers/hls/hlsjsdownloader/main.go index 1e332ba5..404d7b88 100644 --- a/internal/servers/hls/hlsjsdownloader/main.go +++ b/internal/servers/hls/hlsjsdownloader/main.go @@ -38,7 +38,7 @@ func do() error { return fmt.Errorf("bad status code: %v", res.StatusCode) } - zipBuf, err := io.ReadAll(io.LimitReader(res.Body, maxInboundHLSJSSize)) + zipBuf, err := io.ReadAll(&customLimitReader{res.Body, maxInboundHLSJSSize}) if err != nil { return err } diff --git a/internal/servers/webrtc/custom_limit_reader.go b/internal/servers/webrtc/custom_limit_reader.go new file mode 100644 index 00000000..82da9ca3 --- /dev/null +++ b/internal/servers/webrtc/custom_limit_reader.go @@ -0,0 +1,26 @@ +package webrtc + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/servers/webrtc/http_server.go b/internal/servers/webrtc/http_server.go index f78d5047..022e6482 100644 --- a/internal/servers/webrtc/http_server.go +++ b/internal/servers/webrtc/http_server.go @@ -194,7 +194,7 @@ func (s *httpServer) onWHIPPost(ctx *gin.Context, pathName string, publish bool) return } - offer, err := io.ReadAll(io.LimitReader(ctx.Request.Body, maxInboundSDPSize)) + offer, err := io.ReadAll(&customLimitReader{ctx.Request.Body, maxInboundSDPSize}) if err != nil { return } @@ -260,7 +260,7 @@ func (s *httpServer) onWHIPPatch(ctx *gin.Context, pathName string, rawSecret st return } - byts, err := io.ReadAll(io.LimitReader(ctx.Request.Body, maxInboundSDPSize)) + byts, err := io.ReadAll(&customLimitReader{ctx.Request.Body, maxInboundSDPSize}) if err != nil { return } diff --git a/internal/staticsources/rpicamera/mtxrpicamdownloader/custom_limit_reader.go b/internal/staticsources/rpicamera/mtxrpicamdownloader/custom_limit_reader.go new file mode 100644 index 00000000..43540fd8 --- /dev/null +++ b/internal/staticsources/rpicamera/mtxrpicamdownloader/custom_limit_reader.go @@ -0,0 +1,26 @@ +package main + +import ( + "fmt" + "io" +) + +var errSizeExceeded = fmt.Errorf("size exceeds maximum allowed") + +// like io.LimitReader, but returns a dedicated error if the limit is exceeded. +type customLimitReader struct { + r io.Reader + n int64 +} + +func (l *customLimitReader) Read(p []byte) (n int, err error) { + if l.n <= 0 { + return 0, errSizeExceeded + } + if int64(len(p)) > l.n { + p = p[0:l.n] + } + n, err = l.r.Read(p) + l.n -= int64(n) + return +} diff --git a/internal/staticsources/rpicamera/mtxrpicamdownloader/main.go b/internal/staticsources/rpicamera/mtxrpicamdownloader/main.go index dea9ee05..d311ffbb 100644 --- a/internal/staticsources/rpicamera/mtxrpicamdownloader/main.go +++ b/internal/staticsources/rpicamera/mtxrpicamdownloader/main.go @@ -79,7 +79,7 @@ func doSingle(version string, f string) error { return fmt.Errorf("bad status code: %v", res.StatusCode) } - buf, err := io.ReadAll(io.LimitReader(res.Body, maxInboundRPICameraSize)) + buf, err := io.ReadAll(&customLimitReader{res.Body, maxInboundRPICameraSize}) if err != nil { return err }