Files
rvbox/internal/agentproto/compression.go
T

47 lines
1.6 KiB
Go

package agentproto
import (
"bytes"
"errors"
"fmt"
"io"
"github.com/klauspost/compress/zstd"
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
)
var ErrInvalidOutputChunk = errors.New("invalid output chunk")
// DecodeOutputChunk validates declared sizes before bounded decompression.
func DecodeOutputChunk(chunk *rvboxv1.OutputChunk, maxRawBytes uint64) ([]byte, error) {
if chunk == nil || (chunk.Stream != rvboxv1.StreamKind_STREAM_STDOUT && chunk.Stream != rvboxv1.StreamKind_STREAM_STDERR) || chunk.CompressedSize != uint64(len(chunk.Data)) || chunk.UncompressedSize > maxRawBytes {
return nil, ErrInvalidOutputChunk
}
switch chunk.Compression {
case rvboxv1.Compression_COMPRESSION_NONE:
if chunk.UncompressedSize != uint64(len(chunk.Data)) {
return nil, ErrInvalidOutputChunk
}
return bytes.Clone(chunk.Data), nil
case rvboxv1.Compression_COMPRESSION_ZSTD:
if chunk.UncompressedSize == 0 && len(chunk.Data) != 0 {
return nil, ErrInvalidOutputChunk
}
decoder, err := zstd.NewReader(bytes.NewReader(chunk.Data), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(maxRawBytes+1))
if err != nil {
return nil, fmt.Errorf("%w: zstd header: %v", ErrInvalidOutputChunk, err)
}
defer decoder.Close()
decoded, err := io.ReadAll(io.LimitReader(decoder, int64(maxRawBytes)+1))
if err != nil {
return nil, fmt.Errorf("%w: zstd decode: %v", ErrInvalidOutputChunk, err)
}
if uint64(len(decoded)) != chunk.UncompressedSize || uint64(len(decoded)) > maxRawBytes {
return nil, ErrInvalidOutputChunk
}
return decoded, nil
default:
return nil, ErrInvalidOutputChunk
}
}