47 lines
1.6 KiB
Go
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
|
|
}
|
|
}
|