diff --git a/pkg/yum/module_stream.go b/pkg/yum/module_stream.go index 35d0d1d..b85e7ab 100644 --- a/pkg/yum/module_stream.go +++ b/pkg/yum/module_stream.go @@ -94,8 +94,8 @@ func (r *Repository) ModuleMDs(ctx context.Context) ([]ModuleMD, int, error) { } defer resp.Body.Close() - if moduleMDs, err = parseModuleMDs(resp.Body); err != nil { - return nil, resp.StatusCode, fmt.Errorf("error parsing comps.xml: %w", err) + if moduleMDs, err = parseModuleMDs(resp.Body, *r.settings.MaxXmlSize); err != nil { + return nil, resp.StatusCode, fmt.Errorf("error parsing modulemds: %w", err) } return moduleMDs, resp.StatusCode, nil @@ -108,7 +108,7 @@ func (r *Repository) ModuleMDs(ctx context.Context) ([]ModuleMD, int, error) { // this breaks parsing into two parts: // 1. use node to read the document type // 2. if the document type is modulemd, fully decode the value -func parseModuleMDs(body io.ReadCloser) ([]ModuleMD, error) { +func parseModuleMDs(body io.ReadCloser, maxSize int64) ([]ModuleMD, error) { moduleMDs := make([]ModuleMD, 0) reader, err := ExtractIfCompressed(body) @@ -118,11 +118,19 @@ func parseModuleMDs(body io.ReadCloser) ([]ModuleMD, error) { yaml.RegisterCustomUnmarshaler[StreamVersion](unmarshalStreamVersion) - decoder := yaml.NewDecoder(reader) + // Wrap with maxSize + 1 so limit error only triggers when limit is exceeded + limitedReader := io.LimitReader(reader, maxSize+1) + decoder := yaml.NewDecoder(limitedReader) + for { var node ast.Node err := decoder.Decode(&node) + if err != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return nil, limitErr + } + if errors.Is(err, io.EOF) { break } @@ -133,12 +141,19 @@ func parseModuleMDs(body io.ReadCloser) ([]ModuleMD, error) { Document string `yaml:"document"` } if err := yaml.NodeToValue(node, &docType); err != nil { + // Check limit if NodeToValue fails due to an incomplete/truncated AST + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return nil, limitErr + } return nil, fmt.Errorf("error decoding document type: %w", err) } if docType.Document == "modulemd" { var module ModuleMD if err := yaml.NodeToValue(node, &module); err != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return nil, limitErr + } return nil, fmt.Errorf("error decoding modulemd: %w", err) } moduleMDs = append(moduleMDs, module) diff --git a/pkg/yum/module_stream_test.go b/pkg/yum/module_stream_test.go index 5f622f6..1a1f030 100644 --- a/pkg/yum/module_stream_test.go +++ b/pkg/yum/module_stream_test.go @@ -13,18 +13,31 @@ func TestParseModuleMDs(t *testing.T) { f, err := os.Open("mocks/module.yaml.zst") assert.NoError(t, err) - parsed, err := parseModuleMDs(f) + parsed, err := parseModuleMDs(f, DefaultMaxXmlSize) assert.NoError(t, err) assert.Equal(t, 13, len(parsed)) assert.NotEmpty(t, parsed[0].Data.Name) assert.NotEmpty(t, parsed[0].Data.Artifacts.Rpms) } +// A maxSize that's smaller than the decompressed modules.yaml must bound how much is read, +// rather than fully decompressing/parsing the payload (decompression-bomb protection). +func TestParseModuleMDsMaxLimitError(t *testing.T) { + f, err := os.Open("mocks/module.yaml.zst") + assert.NoError(t, err) + defer f.Close() + + parsed, err := parseModuleMDs(f, 10) + assert.Error(t, err) + assert.ErrorContains(t, err, "decompression limit of 10 bytes exceeded") + assert.Empty(t, parsed) +} + func TestStreamVersionPrecision(t *testing.T) { f, err := os.Open("mocks/module.yaml.zst") assert.NoError(t, err) - parsed, err := parseModuleMDs(f) + parsed, err := parseModuleMDs(f, DefaultMaxXmlSize) assert.NoError(t, err) handlesFloatFound, handlesStringFound := false, false @@ -48,7 +61,7 @@ func TestParseRhel8Modules(t *testing.T) { defer f.Close() require.NoError(t, err) - modules, err := parseModuleMDs(f) + modules, err := parseModuleMDs(f, DefaultMaxXmlSize) require.NoError(t, err) assert.Len(t, modules, 961) diff --git a/pkg/yum/repository.go b/pkg/yum/repository.go index d52ce08..5e7c9f2 100644 --- a/pkg/yum/repository.go +++ b/pkg/yum/repository.go @@ -5,6 +5,7 @@ import ( "compress/gzip" "context" "encoding/xml" + "errors" "fmt" "io" "net/http" @@ -220,7 +221,7 @@ func (r *Repository) Comps(ctx context.Context) (*Comps, int, error) { defer resp.Body.Close() - if comps, err = ParseCompsXML(resp.Body, compsURL); err != nil { + if comps, err = ParseCompsXML(resp.Body, compsURL, *r.settings.MaxXmlSize); err != nil { return nil, resp.StatusCode, fmt.Errorf("error parsing comps.xml: %w", err) } @@ -462,7 +463,7 @@ func ParseRepomdXML(body io.ReadCloser) (Repomd, error) { } // ParseCompsXML creates PackageGroup array and Environment array from comps.xml body response -func ParseCompsXML(body io.ReadCloser, url *string) (Comps, error) { +func ParseCompsXML(body io.ReadCloser, url *string, maxSize int64) (Comps, error) { var reader io.Reader var comps Comps packageGroups := []PackageGroup{} @@ -474,17 +475,21 @@ func ParseCompsXML(body io.ReadCloser, url *string) (Comps, error) { return comps, err } - decoder := xml.NewDecoder(reader) + // Wrap with maxSize + 1 so limit error only triggers when limit is exceeded + limitedReader := io.LimitReader(reader, maxSize+1) + decoder := xml.NewDecoder(limitedReader) for { t, decodeError := decoder.Token() - if decodeError == io.EOF { - break - } else if decodeError != nil { + if decodeError != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return comps, limitErr + } + if errors.Is(decodeError, io.EOF) { + break + } return comps, fmt.Errorf("error decoding token: %w", decodeError) - } else if t == nil { - break } switch elType := t.(type) { @@ -493,12 +498,18 @@ func ParseCompsXML(body io.ReadCloser, url *string) (Comps, error) { case "group": var packageGroup PackageGroup if decodeElementError := decoder.DecodeElement(&packageGroup, &elType); decodeElementError != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return comps, limitErr + } return comps, decodeElementError } packageGroups = append(packageGroups, packageGroup) case "environment": var environment Environment if decodeElementError := decoder.DecodeElement(&environment, &elType); decodeElementError != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return comps, limitErr + } return comps, decodeElementError } environments = append(environments, environment) @@ -569,20 +580,22 @@ func ParseCompressedXMLData(body io.Reader, maxSize int64) ([]Package, error) { return []Package{}, fmt.Errorf("error unzipping response body: %w", err) } - limitedReader := io.LimitReader(reader, maxSize) + // Wrap with maxSize + 1 so limit error only triggers when limit is exceeded + limitedReader := io.LimitReader(reader, maxSize+1) decoder := xml.NewDecoder(limitedReader) for { // Read tokens from the XML document in a stream. t, decodeError := decoder.Token() - // If we are at the end of the file, we are done - if decodeError == io.EOF { - break - } else if decodeError != nil { + if decodeError != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return []Package{}, limitErr + } + if errors.Is(decodeError, io.EOF) { + break + } return []Package{}, fmt.Errorf("error decoding token: %w", decodeError) - } else if t == nil { - break } // Here, we inspect the token @@ -593,6 +606,9 @@ func ParseCompressedXMLData(body io.Reader, maxSize int64) ([]Package, error) { case "package": var pkg Package if decodeElementError := decoder.DecodeElement(&pkg, &elType); decodeElementError != nil { + if limitErr := CheckLimit(limitedReader, maxSize); limitErr != nil { + return []Package{}, limitErr + } return result, decodeElementError } // Ensure that the type is "rpm" before pushing our array diff --git a/pkg/yum/repository_test.go b/pkg/yum/repository_test.go index bf61eb2..5c4c49f 100644 --- a/pkg/yum/repository_test.go +++ b/pkg/yum/repository_test.go @@ -303,28 +303,44 @@ func TestParseCompsXML(t *testing.T) { xmlFile, err := os.Open(path) assert.NoError(t, err) defer xmlFile.Close() - comps, err := ParseCompsXML(xmlFile, &path) + comps, err := ParseCompsXML(xmlFile, &path, DefaultMaxXmlSize) assert.NoError(t, err) assert.NotEmpty(t, comps) } } +// A maxSize that's smaller than the decompressed comps.xml must bound how much is read, +// rather than fully decompressing/parsing the payload (decompression-bomb protection). +func TestParseCompsXMLMaxLimitError(t *testing.T) { + path := "mocks/comps.xml.gz" + xmlFile, err := os.Open(path) + assert.NoError(t, err) + defer xmlFile.Close() + + comps, err := ParseCompsXML(xmlFile, &path, 10) + assert.Error(t, err) + assert.ErrorContains(t, err, "decompression limit of 10 bytes exceeded") + assert.Empty(t, comps.PackageGroups) + assert.Empty(t, comps.Environments) +} + // if the xml is half complete, you get a parse error -func TestParseCompressedXMLDataWithError(t *testing.T) { +func TestParseCompressedXMLDataMaxLimitError(t *testing.T) { xmlFile, err := os.Open("mocks/primary.xml.gz") assert.NoError(t, err) defer xmlFile.Close() result, err := ParseCompressedXMLData(xmlFile, 200) assert.Error(t, err) + assert.ErrorContains(t, err, "decompression limit of 200 bytes exceeded") assert.Empty(t, result) } // If no elements are parsed, no error is thrown, but you get empty results -func TestParseCompressedXMLDataMaxLimit(t *testing.T) { +func TestParseCompressedXMLDataNoXMLElements(t *testing.T) { xmlFile, err := os.Open("mocks/aaaa.xml.gz") assert.NoError(t, err) defer xmlFile.Close() - result, err := ParseCompressedXMLData(xmlFile, 10) + result, err := ParseCompressedXMLData(xmlFile, DefaultMaxXmlSize) assert.NoError(t, err) assert.Empty(t, result) } diff --git a/pkg/yum/utils.go b/pkg/yum/utils.go index dcb4213..4058290 100644 --- a/pkg/yum/utils.go +++ b/pkg/yum/utils.go @@ -2,6 +2,7 @@ package yum import ( "bufio" + "fmt" "io" "github.com/h2non/filetype" @@ -36,3 +37,12 @@ func ExtractIfCompressed(reader io.ReadCloser) (extractedReader io.Reader, err e return bufferedReader, nil } } + +// CheckLimit inspects an io.Reader (typically returned from io.LimitReader) +// to see if the byte limit has been exceeded. +func CheckLimit(r io.Reader, maxSize int64) error { + if lr, ok := r.(*io.LimitedReader); ok && lr.N == 0 { + return fmt.Errorf("decompression limit of %d bytes exceeded", maxSize) + } + return nil +}