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 } }