Skip to content
Open
Show file tree
Hide file tree
Changes from 6 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 19 additions & 4 deletions pkg/yum/module_stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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
}
Expand All @@ -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)
Expand Down
19 changes: 16 additions & 3 deletions pkg/yum/module_stream_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 TestParseModuleMDsMaxLimit(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
Expand All @@ -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)
Expand Down
27 changes: 19 additions & 8 deletions pkg/yum/repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"compress/gzip"
"context"
"encoding/xml"
"errors"
"fmt"
"io"
"net/http"
Expand Down Expand Up @@ -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)
}

Expand Down Expand Up @@ -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{}
Expand All @@ -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 {
Comment thread
xbhouse marked this conversation as resolved.
break
}

switch elType := t.(type) {
Expand All @@ -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)
Expand Down
17 changes: 16 additions & 1 deletion pkg/yum/repository_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -303,12 +303,27 @@ 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 TestParseCompsXMLMaxLimit(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) {
xmlFile, err := os.Open("mocks/primary.xml.gz")
Expand Down
10 changes: 10 additions & 0 deletions pkg/yum/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package yum

import (
"bufio"
"fmt"
"io"

"github.com/h2non/filetype"
Expand Down Expand Up @@ -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
}
Loading