Skip to content
Open
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
3 changes: 1 addition & 2 deletions private/protocol/middleware_capture_request_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,14 +32,13 @@ func TestAddCaptureRequestMiddleware(t *testing.T) {
Path: "test/path",
RawQuery: "language=us&region=us-west+east",
},
ContentLength: 100,
},
ExpectRequest: &http.Request{
Method: "PUT",
Header: map[string][]string{
"Foo": {"bar", "too"},
"Checksum": {"SHA256"},
"Content-Length": {"100"},
"Content-Length": {"12"},
},
URL: &url.URL{
Path: "test/path",
Expand Down
42 changes: 42 additions & 0 deletions private/requestcompression/request_compression_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,48 @@ func TestRequestCompression(t *testing.T) {
}
}

// TestRequestCompressionSetsContentLength verifies that, once the body is
// compressed, ContentLength is updated to the compressed size (overwriting the
// original length) so the correct Content-Length is sent.
func TestRequestCompressionSetsContentLength(t *testing.T) {
original := strings.Repeat("Hi, world! ", 100)

req := http.NewStackRequest().(*http.Request)
req, _ = req.SetStream(strings.NewReader(original))

m := requestCompression{
compressAlgorithms: []string{GZIP},
}

var updatedRequest *http.Request
_, _, err := m.HandleSerialize(context.Background(),
middleware.SerializeInput{Request: req},
middleware.SerializeHandlerFunc(func(ctx context.Context, input middleware.SerializeInput) (
out middleware.SerializeOutput, metadata middleware.Metadata, err error) {
updatedRequest = input.Request.(*http.Request)
return out, metadata, nil
}),
)
if err != nil {
t.Fatalf("expect no error, got %v", err)
}

size, ok, err := updatedRequest.StreamLength()
if err != nil {
t.Fatalf("expect no error getting stream length, got %v", err)
}
if !ok {
t.Fatal("expect compressed stream length to be known")
}
if e, a := size, updatedRequest.ContentLength; e != a {
t.Errorf("expect ContentLength %d, got %d", e, a)
}
// And it should reflect compression, i.e. differ from the original length.
if updatedRequest.ContentLength == int64(len(original)) {
t.Errorf("expect ContentLength to be updated to compressed size, still %d", updatedRequest.ContentLength)
}
}

func testUnzipContent(content io.Reader, expect []byte, disableRequestCompression bool, requestMinCompressionSizeBytes int64) error {
if disableRequestCompression || int64(len(expect)) < requestMinCompressionSizeBytes {
b, err := io.ReadAll(content)
Expand Down
3 changes: 3 additions & 0 deletions transport/http/middleware_content_length.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@ type ComputeContentLength struct {

// AddComputeContentLengthMiddleware adds ComputeContentLength to the middleware
// stack's Build step.
//
// Deprecated: Content-Length is now set when the request body is set via
// Request.SetStream, so this middleware is no longer used.
func AddComputeContentLengthMiddleware(stack *middleware.Stack) error {
return stack.Build.Add(&ComputeContentLength{}, middleware.After)
}
Expand Down
128 changes: 0 additions & 128 deletions transport/http/middleware_content_length_test.go
Original file line number Diff line number Diff line change
@@ -1,147 +1,19 @@
package http

import (
"bytes"
"context"
"fmt"
"io"
"strings"
"testing"

"github.com/aws/smithy-go/middleware"
)

func TestContentLengthMiddleware(t *testing.T) {
cases := map[string]struct {
Stream io.Reader
ExpectNilStream bool
ExpectLen int64
ExpectErr string
}{
// Cases
"bytes.Reader": {
Stream: bytes.NewReader(make([]byte, 10)),
ExpectLen: 10,
ExpectNilStream: false,
},
"bytes.Buffer": {
Stream: bytes.NewBuffer(make([]byte, 10)),
ExpectLen: 10,
ExpectNilStream: false,
},
"strings.Reader": {
Stream: strings.NewReader("hello"),
ExpectLen: 5,
ExpectNilStream: false,
},
"empty stream": {
Stream: strings.NewReader(""),
ExpectLen: 0,
ExpectNilStream: false,
},
"empty stream bytes": {
Stream: bytes.NewReader([]byte{}),
ExpectLen: 0,
ExpectNilStream: false,
},
"nil stream": {
ExpectLen: 0,
ExpectNilStream: true,
},
"un-seekable and no length": {
Stream: &basicReader{buf: make([]byte, 10)},
ExpectLen: -1,
ExpectNilStream: false,
},
"with error": {
Stream: &errorSecondSeekableReader{err: fmt.Errorf("seek failed")},
ExpectErr: "seek failed",
ExpectLen: -1,
ExpectNilStream: false,
},
}

for name, c := range cases {
t.Run(name, func(t *testing.T) {
var err error
req := NewStackRequest().(*Request)
req, err = req.SetStream(c.Stream)
if err != nil {
t.Fatalf("expect to set stream, %v", err)
}

var updatedRequest *Request
var m ComputeContentLength
_, _, err = m.HandleBuild(context.Background(),
middleware.BuildInput{Request: req},
middleware.BuildHandlerFunc(func(ctx context.Context, input middleware.BuildInput) (
out middleware.BuildOutput, metadata middleware.Metadata, err error) {
updatedRequest = input.Request.(*Request)
return out, metadata, nil
}),
)
if len(c.ExpectErr) != 0 {
if err == nil {
t.Fatalf("expect error, got none")
}
if e, a := c.ExpectErr, err.Error(); !strings.Contains(a, e) {
t.Fatalf("expect error to contain %q, got %v", e, a)
}
return
} else if err != nil {
t.Fatalf("expect no error, got %v", err)
}

if e, a := c.ExpectLen, updatedRequest.ContentLength; e != a {
t.Errorf("expect %v content-length, got %v", e, a)
}

if e, a := c.ExpectNilStream, updatedRequest.stream == nil; e != a {
t.Errorf("expect %v nil stream, got %v", e, a)
}
})
}
}

func TestContentLengthMiddleware_HeaderSet(t *testing.T) {
req := NewStackRequest().(*Request)
req.Header.Set("Content-Length", "1234")

var err error
req, err = req.SetStream(strings.NewReader("hello"))
if err != nil {
t.Fatalf("expect to set stream, %v", err)
}

var m ComputeContentLength
_, _, err = m.HandleBuild(context.Background(),
middleware.BuildInput{Request: req},
nopBuildHandler,
)
if err != nil {
t.Fatalf("expect middleware to run, %v", err)
}

if e, a := "1234", req.Header.Get("Content-Length"); e != a {
t.Errorf("expect Content-Length not to change, got %v", a)
}
}

var nopBuildHandler = middleware.BuildHandlerFunc(func(ctx context.Context, input middleware.BuildInput) (
out middleware.BuildOutput, metadata middleware.Metadata, err error) {
return out, metadata, nil
})

type basicReader struct {
buf []byte
}

func (r *basicReader) Read(p []byte) (int, error) {
n := copy(p, r.buf)
r.buf = r.buf[n:]
return n, nil
}

type errorSecondSeekableReader struct {
err error
count int
Expand Down
9 changes: 9 additions & 0 deletions transport/http/request.go
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,9 @@ func (r *Request) IsStreamSeekable() bool {
// SetStream returns a clone of the request with the stream set to the provided
// reader. May return an error if the provided reader is seekable but returns
// an error.
//
// ContentLength is set to the stream's length when it can be determined, and
// left unchanged otherwise.
func (r *Request) SetStream(reader io.Reader) (rc *Request, err error) {
rc = r.Clone()

Expand Down Expand Up @@ -151,6 +154,12 @@ func (r *Request) SetStream(reader io.Reader) (rc *Request, err error) {
rc.isStreamSeekable = isStreamSeekable
rc.streamStartPos = streamStartPos

if n, ok, err := rc.StreamLength(); err != nil {
return rc, err
} else if ok {
rc.ContentLength = n
}

return rc, err
}

Expand Down
22 changes: 16 additions & 6 deletions transport/http/request_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package http
import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"os"
Expand Down Expand Up @@ -117,6 +118,7 @@ func TestRequestSetStream(t *testing.T) {
expectNilStream bool
expectNilBody bool
expectReqContentLength int64
expectErr string
}{
"nil stream": {
expectNilStream: true,
Expand Down Expand Up @@ -174,13 +176,27 @@ func TestRequestSetStream(t *testing.T) {
expectNilStream: true,
expectNilBody: true,
},
// Seeking to compute the length fails.
"seek error": {
reader: &errorSecondSeekableReader{err: fmt.Errorf("seek failed")},
expectErr: "seek failed",
},
}

for name, c := range cases {
t.Run(name, func(t *testing.T) {
var err error
req := NewStackRequest().(*Request)
req, err = req.SetStream(c.reader)
if len(c.expectErr) != 0 {
if err == nil {
t.Fatalf("expect error, got none")
}
if e, a := c.expectErr, err.Error(); !strings.Contains(a, e) {
t.Fatalf("expect error to contain %q, got %v", e, a)
}
return
}
if err != nil {
t.Fatalf("expect not error, got %v", err)
}
Expand All @@ -195,12 +211,6 @@ func TestRequestSetStream(t *testing.T) {
t.Errorf("expect %v nil stream, got %v", e, a)
}

if l, ok, err := req.StreamLength(); err != nil {
t.Fatalf("expect no stream length error, got %v", err)
} else if ok {
req.ContentLength = l
}

if e, a := c.expectContentLength, req.ContentLength; e != a {
t.Errorf("expect %v content-length, got %v", e, a)
}
Expand Down
Loading