diff --git a/adapters/core2sn/class.go b/adapters/core2sn/class.go index 3235177619..3aabb642a8 100644 --- a/adapters/core2sn/class.go +++ b/adapters/core2sn/class.go @@ -116,7 +116,9 @@ func AdaptDeprecatedEntryPoint(ep *core.DeprecatedEntryPoint) starknet.EntryPoin func AdaptDeprecatedCairoClass( class *core.DeprecatedCairoClass, ) (starknet.DeprecatedCairoClass, error) { - decompressedProgram, err := compression.Gzip64Decode(class.Program) + decompressedProgram, err := compression.Gzip64Decode( + class.Program, core.MaxDeprecatedClassProgramSize, + ) if err != nil { return starknet.DeprecatedCairoClass{}, err } diff --git a/core/class.go b/core/class.go index 6a5be345a6..617565f6cc 100644 --- a/core/class.go +++ b/core/class.go @@ -22,6 +22,10 @@ var ( const minDeclaredClassSize = 8 +// MaxDeprecatedClassProgramSize bounds the decompressed size of a deprecated +// Cairo 0 class program +const MaxDeprecatedClassProgramSize = 16 * db.Megabyte + // Single felt identifying the number "0.1.0" as a short string var SierraVersion010 felt.Felt = felt.Felt( [4]uint64{ diff --git a/core/class_hash.go b/core/class_hash.go index 630273a08a..8c8644667d 100644 --- a/core/class_hash.go +++ b/core/class_hash.go @@ -1,6 +1,7 @@ package core import ( + "fmt" "sync" "github.com/NethermindEth/juno/core/crypto" @@ -9,14 +10,16 @@ import ( ) func deprecatedCairoClassHash(class *DeprecatedCairoClass) (felt.Felt, error) { - decompressedProgram, err := compression.Gzip64Decode(class.Program) + decompressedProgram, err := compression.Gzip64Decode( + class.Program, MaxDeprecatedClassProgramSize, + ) if err != nil { - return felt.Felt{}, err + return felt.Felt{}, fmt.Errorf("decompressing Cairo Zero class: %w", err) } program, err := unmarshalDeprecatedCairoProgram(decompressedProgram) if err != nil { - return felt.Felt{}, err + return felt.Felt{}, fmt.Errorf("unmarshalling Cairo Zero class: %w", err) } var ( diff --git a/rpc/v10/transaction_test.go b/rpc/v10/transaction_test.go index 91dd9010a0..1274472eec 100644 --- a/rpc/v10/transaction_test.go +++ b/rpc/v10/transaction_test.go @@ -2285,7 +2285,7 @@ func TestContractClassToGatewayPayload(t *testing.T) { require.Equal(t, class.EntryPoints, decoded.EntryPoints) require.Equal(t, class.ABI, decoded.ABI) - sierraJSON, err := compression.Gzip64Decode(decoded.SierraProgram) + sierraJSON, err := compression.Gzip64Decode(decoded.SierraProgram, compression.NoLimit) require.NoError(t, err, "sierra_program must be gzip+base64 encoded") var roundTripped []felt.Felt diff --git a/rpc/v8/class.go b/rpc/v8/class.go index 10ddaf0c61..e4e2ebd5b6 100644 --- a/rpc/v8/class.go +++ b/rpc/v8/class.go @@ -67,7 +67,10 @@ func adaptDeclaredClass( } base64Program := string(program[1 : len(program)-1]) - feederClass.DeprecatedCairo.Program, err = compression.Gzip64Decode(base64Program) + feederClass.DeprecatedCairo.Program, err = compression.Gzip64Decode( + base64Program, + core.MaxDeprecatedClassProgramSize, + ) if err != nil { return nil, err } diff --git a/rpc/v9/class.go b/rpc/v9/class.go index 5803cab3ad..b3d21640e5 100644 --- a/rpc/v9/class.go +++ b/rpc/v9/class.go @@ -73,7 +73,10 @@ func AdaptDeclaredClass( } base64Program := string(program[1 : len(program)-1]) - feederClass.DeprecatedCairo.Program, err = compression.Gzip64Decode(base64Program) + feederClass.DeprecatedCairo.Program, err = compression.Gzip64Decode( + base64Program, + core.MaxDeprecatedClassProgramSize, + ) if err != nil { return nil, err } diff --git a/utils/compression/compression.go b/utils/compression/compression.go index 6d6902bf2f..d606a93905 100644 --- a/utils/compression/compression.go +++ b/utils/compression/compression.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "io" + "math" "sync" "github.com/klauspost/compress/gzip" @@ -27,6 +28,9 @@ const ( levelCount = maxLevel - minLevel + 1 ) +// NoLimit disables the decompressed-size bound in Gzip64Decode. +const NoLimit int64 = math.MaxInt64 + var ErrWriterNotAcquired = errors.New("using writer after release") // gzipWriterPools holds one pool per compression level. @@ -134,6 +138,7 @@ func GzipWriterLevel(dst io.Writer, level int) *Writer { return writer } +// Gzip64Encode encodes data with default compression func Gzip64Encode(data []byte) (string, error) { var compressedBuffer bytes.Buffer gzipWriter := GzipWriter(&compressedBuffer) @@ -147,7 +152,10 @@ func Gzip64Encode(data []byte) (string, error) { return base64.StdEncoding.EncodeToString(compressedBuffer.Bytes()), nil } -func Gzip64Decode(data string) ([]byte, error) { +// Gzip64Decode decompress data with a size limit of `maxDecompressedSize`. +// If decoded data turns out to be bigger an error is retured. Use +// [NoLimit] for unbounded decompression. +func Gzip64Decode(data string, maxDecompressedSize int64) ([]byte, error) { decodedBytes, err := base64.StdEncoding.DecodeString(data) if err != nil { return nil, err @@ -156,13 +164,30 @@ func Gzip64Decode(data string) ([]byte, error) { if err != nil { return nil, err } - decompressedBytes, err := io.ReadAll(gzipReader) + + // Read one byte more than the limit to differentiate between the decompressed + // size fitting (<= maxDecompressedSize) and overflowing (> maxDecompressedSize). + // The guard keeps NoLimit from overflowing. + readLimit := maxDecompressedSize + if readLimit < NoLimit { + readLimit++ + } + + limited := io.LimitReader(gzipReader, readLimit) + decompressedBytes, err := io.ReadAll(limited) if err != nil { return nil, err } + if int64(len(decompressedBytes)) > maxDecompressedSize { + return nil, fmt.Errorf( + "decompressed data exceeded the maximum byte size: %d", maxDecompressedSize, + ) + } + err = gzipReader.Close() if err != nil { return nil, err } + return decompressedBytes, nil } diff --git a/utils/compression/compression_bench_test.go b/utils/compression/compression_bench_test.go index dd0919a413..32e476060c 100644 --- a/utils/compression/compression_bench_test.go +++ b/utils/compression/compression_bench_test.go @@ -51,7 +51,7 @@ func BenchmarkGzip64Decode(b *testing.B) { b.ReportAllocs() b.SetBytes(int64(size)) for b.Loop() { - if _, err := compression.Gzip64Decode(encoded); err != nil { + if _, err := compression.Gzip64Decode(encoded, compression.NoLimit); err != nil { b.Fatal(err) } } diff --git a/utils/compression/compression_test.go b/utils/compression/compression_test.go index 4f865c07f4..4b95865a6a 100644 --- a/utils/compression/compression_test.go +++ b/utils/compression/compression_test.go @@ -2,6 +2,7 @@ package compression_test import ( "bytes" + "encoding/base64" "errors" "io" "runtime" @@ -24,16 +25,135 @@ func TestGzip64(t *testing.T) { require.NoError(t, err) assert.Equal(t, expectedComBytes, comBytes) - decompBytes, err := compression.Gzip64Decode(comBytes) + decompBytes, err := compression.Gzip64Decode(comBytes, compression.NoLimit) require.NoError(t, err) assert.Equal(t, bytes, decompBytes) } +func TestGzip64Decode(t *testing.T) { + const limit = 1024 + tests := []struct { + name string + payload []byte + limit int64 + wantErrContains string // empty means the payload is expected to round-trip + }{ + { + name: "within limit", + payload: bytes.Repeat([]byte("a"), limit/2), + limit: limit, + }, + { + name: "unbounded limit", + payload: []byte{0}, + limit: compression.NoLimit, + }, + { + name: "empty payload", + payload: []byte{}, + limit: limit, + }, + { + // The budget is inclusive: a payload that exactly fills it is valid. + name: "exactly at limit", + payload: bytes.Repeat([]byte("a"), limit), + limit: limit, + }, + { + // One byte over is the smallest overflow the +1 read must catch. + name: "one byte over limit", + payload: bytes.Repeat([]byte("a"), limit+1), + limit: limit, + wantErrContains: "decompressed data exceeded the maximum byte size:", + }, + { + name: "zero limit rejects any content", + payload: []byte("x"), + limit: 0, + wantErrContains: "decompressed data exceeded the maximum byte size:", + }, + { + name: "zero limit accepts empty payload", + payload: []byte{}, + limit: 0, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + encoded, err := compression.Gzip64Encode(test.payload) + require.NoError(t, err) + + decoded, err := compression.Gzip64Decode(encoded, test.limit) + if test.wantErrContains != "" { + assert.ErrorContains(t, err, test.wantErrContains) + assert.Nil(t, decoded) + return + } + + require.NoError(t, err) + assert.Equal(t, test.payload, decoded) + }) + } +} + +// A compression bomb — a few KiB of input inflating to 64 MiB — must be rejected +// cheaply. +func TestGzip64DecodeCompressionBomb(t *testing.T) { + const limit = 1024 + bomb := make([]byte, 64*1024*1024) // zeros compress ~1000:1 + encoded, err := compression.Gzip64Encode(bomb) + require.NoError(t, err) + + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + decoded, err := compression.Gzip64Decode(encoded, limit) + runtime.ReadMemStats(&after) + + assert.ErrorContains(t, err, "decompressed data exceeded the maximum byte size:") + assert.Nil(t, decoded) + assert.Less(t, after.TotalAlloc-before.TotalAlloc, uint64(8*1024*1024), + "rejecting an over-limit stream must not materialise the decompressed payload") +} + +// Check stream corruption is properly shown. +func TestGzip64DecodeCorruptStream(t *testing.T) { + const limit = 1024 + payload := bytes.Repeat([]byte("a"), limit) + encoded, err := compression.Gzip64Encode(payload) + require.NoError(t, err) + raw, err := base64.StdEncoding.DecodeString(encoded) + require.NoError(t, err) + + // Corrupt the gzip footer + t.Run("corrupt checksum at exactly the limit", func(t *testing.T) { + corrupted := bytes.Clone(raw) + corrupted[len(corrupted)-1] ^= 0xff // corrupt the footer + + decoded, err := compression.Gzip64Decode( + base64.StdEncoding.EncodeToString(corrupted), limit, + ) + assert.ErrorIs(t, err, gzip.ErrChecksum) + assert.Nil(t, decoded) + }) + + t.Run("truncated stream", func(t *testing.T) { + decoded, err := compression.Gzip64Decode( + // remove the footer and part of the data. + base64.StdEncoding.EncodeToString(raw[:len(raw)/2]), + limit, + ) + + assert.ErrorIs(t, err, io.ErrUnexpectedEOF) + assert.Nil(t, decoded) + }) +} + func FuzzGzip64(f *testing.F) { f.Fuzz(func(t *testing.T, data []byte) { compressed, err := compression.Gzip64Encode(data) require.NoError(t, err) - decompressed, err := compression.Gzip64Decode(compressed) + decompressed, err := compression.Gzip64Decode(compressed, compression.NoLimit) require.NoError(t, err) assert.Equal(t, data, decompressed) }) @@ -50,7 +170,7 @@ func TestGzip64EncodeAcrossSuccessiveCalls(t *testing.T) { encoded, err := compression.Gzip64Encode(payload) require.NoError(t, err) - decoded, err := compression.Gzip64Decode(encoded) + decoded, err := compression.Gzip64Decode(encoded, compression.NoLimit) require.NoError(t, err) assert.Equal(t, payload, decoded) }) @@ -204,7 +324,7 @@ func TestGzip64EncodeConcurrent(t *testing.T) { payload := bytes.Repeat([]byte{byte('a' + i)}, chunk*(i+1)) encoded, err := compression.Gzip64Encode(payload) assert.NoError(t, err) - decoded, err := compression.Gzip64Decode(encoded) + decoded, err := compression.Gzip64Decode(encoded, compression.NoLimit) assert.NoError(t, err) assert.Equal(t, payload, decoded) })