Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
8 changes: 6 additions & 2 deletions cli/cmd/model_pull.go
Original file line number Diff line number Diff line change
Expand Up @@ -109,11 +109,15 @@ func (r *runners) pullModel(cmd *cobra.Command, args []string) error {
}
defer manifestContent.Close()

// Read the manifest content into a byte slice
manifestBytes, err := io.ReadAll(manifestContent)
// Read the manifest content into a byte slice (OCI manifests are small JSON).
const maxManifestBody = 16 << 20 // 16 MiB
manifestBytes, err := io.ReadAll(io.LimitReader(manifestContent, maxManifestBody+1))
if err != nil {
return err
}
if len(manifestBytes) > maxManifestBody {
return fmt.Errorf("OCI manifest exceeds %d bytes", maxManifestBody)
}

var manifest v1.Manifest
if err := json.Unmarshal(manifestBytes, &manifest); err != nil {
Expand Down
17 changes: 15 additions & 2 deletions pkg/cmxmetadata/cmxmetadata.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@ const (
mmdsPath = "/latest/vendor-api"
mmdsTimeout = 500 * time.Millisecond // fail fast if not in CMX
tokenLeeway = 60 * time.Second // refresh token this early before expiry
// maxMetadataBody bounds MMDS / token JSON responses (credentials, not archives).
maxMetadataBody = 1 << 20 // 1 MiB
)

// ErrNotAvailable is returned when the CMX metadata service is not reachable.
Expand Down Expand Up @@ -61,7 +63,7 @@ func GetVMMetadata() (*VMMetadata, error) {
return nil, ErrNotAvailable
}

body, err := io.ReadAll(resp.Body)
body, err := readAllLimited(resp.Body, maxMetadataBody)
if err != nil {
return nil, ErrNotAvailable
}
Expand Down Expand Up @@ -140,7 +142,7 @@ func exchangeCredentials(meta *VMMetadata) (string, *time.Time, error) {
}
defer resp.Body.Close()

body, err := io.ReadAll(resp.Body)
body, err := readAllLimited(resp.Body, maxMetadataBody)
if err != nil {
return "", nil, fmt.Errorf("reading token response body: %w", err)
}
Expand All @@ -166,3 +168,14 @@ func exchangeCredentials(meta *VMMetadata) (string, *time.Time, error) {

return tokenResp.AccessToken, expiresAt, nil
}

func readAllLimited(r io.Reader, limit int64) ([]byte, error) {
body, err := io.ReadAll(io.LimitReader(r, limit+1))
if err != nil {
return nil, err
}
if int64(len(body)) > limit {
return nil, fmt.Errorf("response body exceeds %d bytes", limit)
}
return body, nil
}
7 changes: 6 additions & 1 deletion pkg/credentials/fetch.go
Original file line number Diff line number Diff line change
Expand Up @@ -122,10 +122,15 @@ func exchangeNonceForToken(uri string, nonce string) (string, error) {
return "", fmt.Errorf("unexpected status code: %d", resp.StatusCode)
}

b, err := io.ReadAll(resp.Body)
// Token JSON is small; bound allocation against a hostile exchange endpoint.
const maxTokenResponse = 1 << 20 // 1 MiB
b, err := io.ReadAll(io.LimitReader(resp.Body, maxTokenResponse+1))
if err != nil {
return "", err
}
if len(b) > maxTokenResponse {
return "", fmt.Errorf("token response exceeds %d bytes", maxTokenResponse)
}

type tokenResponse struct {
Token string `json:"token"`
Expand Down
18 changes: 16 additions & 2 deletions pkg/tools/checksum.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,9 @@ var checksumHTTPClient = &http.Client{
Timeout: 30 * time.Second,
}

// maxChecksumBody bounds checksum text files (not the archives themselves).
const maxChecksumBody = 1 << 20 // 1 MiB

// VerifyHelmChecksum verifies a Helm binary against its .sha256sum file
func VerifyHelmChecksum(data []byte, archiveURL string) error {
// Helm provides per-file checksums: <url>.sha256sum
Expand All @@ -31,7 +34,7 @@ func VerifyHelmChecksum(data []byte, archiveURL string) error {
return fmt.Errorf("checksum file not found (HTTP %d): %s", resp.StatusCode, checksumURL)
}

checksumData, err := io.ReadAll(resp.Body)
checksumData, err := readAllLimited(resp.Body, maxChecksumBody)
if err != nil {
return fmt.Errorf("reading checksum file: %w", err)
}
Expand Down Expand Up @@ -71,7 +74,7 @@ func VerifyTroubleshootChecksum(data []byte, version, filename string) error {
return fmt.Errorf("checksums file not found (HTTP %d): %s", resp.StatusCode, checksumURL)
}

checksumData, err := io.ReadAll(resp.Body)
checksumData, err := readAllLimited(resp.Body, maxChecksumBody)
if err != nil {
return fmt.Errorf("reading checksums file: %w", err)
}
Expand Down Expand Up @@ -102,3 +105,14 @@ func VerifyTroubleshootChecksum(data []byte, version, filename string) error {

return nil
}

func readAllLimited(r io.Reader, limit int64) ([]byte, error) {
body, err := io.ReadAll(io.LimitReader(r, limit+1))
if err != nil {
return nil, err
}
if int64(len(body)) > limit {
return nil, fmt.Errorf("response body exceeds %d bytes", limit)
}
return body, nil
}
23 changes: 23 additions & 0 deletions pkg/tools/checksum_limit_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
package tools

import (
"strings"
"testing"
)

func TestReadAllLimited(t *testing.T) {
t.Parallel()

got, err := readAllLimited(strings.NewReader("abc"), 10)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if string(got) != "abc" {
t.Fatalf("got %q, want abc", got)
}

_, err = readAllLimited(strings.NewReader(strings.Repeat("x", 5)), 4)
if err == nil {
t.Fatal("expected error when body exceeds limit")
}
}