diff --git a/private/protocol/middleware_capture_request_test.go b/private/protocol/middleware_capture_request_test.go index 2777f720a..66cc6f84d 100644 --- a/private/protocol/middleware_capture_request_test.go +++ b/private/protocol/middleware_capture_request_test.go @@ -32,14 +32,13 @@ func TestAddCaptureRequestMiddleware(t *testing.T) { Path: "test/path", RawQuery: "language=us®ion=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", diff --git a/private/requestcompression/request_compression_test.go b/private/requestcompression/request_compression_test.go index b29947dfd..a6b3b2a17 100644 --- a/private/requestcompression/request_compression_test.go +++ b/private/requestcompression/request_compression_test.go @@ -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) diff --git a/transport/http/middleware_content_length.go b/transport/http/middleware_content_length.go index 9969389bb..7a1c47fe3 100644 --- a/transport/http/middleware_content_length.go +++ b/transport/http/middleware_content_length.go @@ -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) } diff --git a/transport/http/middleware_content_length_test.go b/transport/http/middleware_content_length_test.go index 16a1f265c..607c36908 100644 --- a/transport/http/middleware_content_length_test.go +++ b/transport/http/middleware_content_length_test.go @@ -1,9 +1,7 @@ package http import ( - "bytes" "context" - "fmt" "io" "strings" "testing" @@ -11,137 +9,11 @@ import ( "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 diff --git a/transport/http/request.go b/transport/http/request.go index 5cbf6f10a..87acdfe57 100644 --- a/transport/http/request.go +++ b/transport/http/request.go @@ -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() @@ -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 } diff --git a/transport/http/request_test.go b/transport/http/request_test.go index fafd50f95..434b9bd4f 100644 --- a/transport/http/request_test.go +++ b/transport/http/request_test.go @@ -3,6 +3,7 @@ package http import ( "bytes" "context" + "fmt" "io" "net/http" "os" @@ -117,6 +118,7 @@ func TestRequestSetStream(t *testing.T) { expectNilStream bool expectNilBody bool expectReqContentLength int64 + expectErr string }{ "nil stream": { expectNilStream: true, @@ -174,6 +176,11 @@ 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 { @@ -181,6 +188,15 @@ func TestRequestSetStream(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) } @@ -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) }