From a8ddc894fff3e5dcb8b163577e1d1c92f35d588b Mon Sep 17 00:00:00 2001 From: Tao Tong Date: Tue, 11 Aug 2026 21:42:03 +0000 Subject: [PATCH 1/2] Set Content-Length when the request body is set instead of via middleware Add Request.SetStreamWithLength, which sets ContentLength when the stream length is known. The protocol serializers and request compression now use it, so the SDK no longer needs the standalone ComputeContentLength middleware. That middleware is kept but deprecated. --- .../HttpBindingProtocolGenerator.java | 2 +- .../rpc2/cbor/SerializeMiddleware.java | 2 +- .../requestcompression/request_compression.go | 2 + .../request_compression_test.go | 42 +++++++++++++++++++ transport/http/middleware_content_length.go | 3 ++ transport/http/protocol/awsjson/awsjson.go | 4 +- transport/http/protocol/awsquery/awsquery.go | 2 +- transport/http/protocol/ec2query/ec2query.go | 2 +- .../internal/httpbinding/serializer.go | 2 +- transport/http/protocol/rpcv2/rpcv2.go | 2 +- transport/http/request.go | 19 +++++++++ transport/http/request_test.go | 42 +++++++++++++++++++ 12 files changed, 116 insertions(+), 8 deletions(-) diff --git a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java index 913b5f889..746fb9777 100644 --- a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java +++ b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java @@ -565,7 +565,7 @@ protected void writeSetPayloadShapeHeader(GoWriter writer, Shape payloadShape) { */ protected void writeSetStream(GoWriter writer, String operand) { writer.write(""" - if request, err = request.SetStream($L); err != nil { + if request, err = request.SetStreamWithLength($L); err != nil { return out, metadata, &smithy.SerializationError{Err: err} }""", operand); } diff --git a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java index 0e303ef77..324e64347 100644 --- a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java +++ b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java @@ -60,7 +60,7 @@ public Writable generateSerialize() { } payload := $reader:T($encode:T(cv)) - if req, err = req.SetStream(payload); err != nil { + if req, err = req.SetStreamWithLength(payload); err != nil { return out, metadata, &$error:T{Err: err} } diff --git a/private/requestcompression/request_compression.go b/private/requestcompression/request_compression.go index 7c4147603..304c30319 100644 --- a/private/requestcompression/request_compression.go +++ b/private/requestcompression/request_compression.go @@ -89,6 +89,8 @@ func (m requestCompression) HandleSerialize( } *req = *newReq + req.ContentLength = int64(len(compressedBytes)) + if val := req.Header.Get("Content-Encoding"); val != "" { req.Header.Set("Content-Encoding", fmt.Sprintf("%s, %s", val, algorithm)) } else { 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..859c85de8 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.SetStreamWithLength, 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/protocol/awsjson/awsjson.go b/transport/http/protocol/awsjson/awsjson.go index 6aa7f97ac..a107b189f 100644 --- a/transport/http/protocol/awsjson/awsjson.go +++ b/transport/http/protocol/awsjson/awsjson.go @@ -95,7 +95,7 @@ func (p *Protocol) SerializeRequest( } if schema.Input == nil { - sreq, err := req.SetStream(strings.NewReader("{}")) + sreq, err := req.SetStreamWithLength(strings.NewReader("{}")) if err != nil { return fmt.Errorf("set stream: %w", err) } @@ -106,7 +106,7 @@ func (p *Protocol) SerializeRequest( ss := internaljson.NewShapeSerializer() in.Serialize(ss) - sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/awsquery/awsquery.go b/transport/http/protocol/awsquery/awsquery.go index a630f0653..b7eaffdeb 100644 --- a/transport/http/protocol/awsquery/awsquery.go +++ b/transport/http/protocol/awsquery/awsquery.go @@ -88,7 +88,7 @@ func (p *Protocol) SerializeRequest( in.Serialize(ss) } - sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/ec2query/ec2query.go b/transport/http/protocol/ec2query/ec2query.go index ddd2d9c2a..1b92443b5 100644 --- a/transport/http/protocol/ec2query/ec2query.go +++ b/transport/http/protocol/ec2query/ec2query.go @@ -89,7 +89,7 @@ func (p *Protocol) SerializeRequest( in.Serialize(ss) } - sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/internal/httpbinding/serializer.go b/transport/http/protocol/internal/httpbinding/serializer.go index 33eee224c..6273f724d 100644 --- a/transport/http/protocol/internal/httpbinding/serializer.go +++ b/transport/http/protocol/internal/httpbinding/serializer.go @@ -137,7 +137,7 @@ func (s *ShapeSerializer) setBody(body io.Reader, contentType string) error { if s.request.Header.Get("Content-Type") == "" { s.request.Header.Set("Content-Type", contentType) } - sreq, err := s.request.SetStream(body) + sreq, err := s.request.SetStreamWithLength(body) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/rpcv2/rpcv2.go b/transport/http/protocol/rpcv2/rpcv2.go index 48b65b71e..767443293 100644 --- a/transport/http/protocol/rpcv2/rpcv2.go +++ b/transport/http/protocol/rpcv2/rpcv2.go @@ -98,7 +98,7 @@ func (p *Protocol) SerializeRequest( req.Header.Set("Content-Type", "application/cbor") - sreq, err := req.SetStream(bytes.NewReader(payload)) + sreq, err := req.SetStreamWithLength(bytes.NewReader(payload)) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/request.go b/transport/http/request.go index 5cbf6f10a..de5aa003f 100644 --- a/transport/http/request.go +++ b/transport/http/request.go @@ -154,6 +154,25 @@ func (r *Request) SetStream(reader io.Reader) (rc *Request, err error) { return rc, err } +// SetStreamWithLength sets the request stream like SetStream, and additionally +// sets ContentLength when the stream's length can be determined. If it cannot +// (e.g. a non-seekable stream), ContentLength is left unchanged so an +// explicitly provided value is preserved. +func (r *Request) SetStreamWithLength(reader io.Reader) (rc *Request, err error) { + rc, err = r.SetStream(reader) + if err != nil { + return rc, err + } + + if n, ok, err := rc.StreamLength(); err != nil { + return rc, err + } else if ok { + rc.ContentLength = n + } + + return rc, nil +} + // Build returns a build standard HTTP request value from the Smithy request. // The request's stream is wrapped in a safe container that allows it to be // reused for subsequent attempts. diff --git a/transport/http/request_test.go b/transport/http/request_test.go index fafd50f95..8dbc383a3 100644 --- a/transport/http/request_test.go +++ b/transport/http/request_test.go @@ -218,3 +218,45 @@ func TestRequestSetStream(t *testing.T) { }) } } + +func TestRequestSetStreamWithLength(t *testing.T) { + cases := map[string]struct { + reader io.Reader + expectContentLength int64 + }{ + "nil stream": { + expectContentLength: 0, + }, + "empty seekable stream": { + reader: bytes.NewReader([]byte{}), + expectContentLength: 0, + }, + "unseekable stream with len": { + reader: bytes.NewBuffer([]byte("abc123")), + expectContentLength: 6, + }, + "seekable stream": { + reader: bytes.NewReader([]byte("abc123")), + expectContentLength: 6, + }, + // Length not determinable: left at default (-1) so a user-set value survives. + "unseekable no len stream": { + reader: io.NopCloser(bytes.NewBuffer([]byte("abc123"))), + expectContentLength: -1, + }, + } + + for name, c := range cases { + t.Run(name, func(t *testing.T) { + req := NewStackRequest().(*Request) + req, err := req.SetStreamWithLength(c.reader) + if err != nil { + t.Fatalf("expect no error, got %v", err) + } + + if e, a := c.expectContentLength, req.ContentLength; e != a { + t.Errorf("expect content-length %v, got %v", e, a) + } + }) + } +} From 26ed9fb45dc77805e58475919c921663640a237a Mon Sep 17 00:00:00 2001 From: Tao Tong Date: Wed, 12 Aug 2026 20:51:03 +0000 Subject: [PATCH 2/2] Fold length computation into SetStream instead of a new method Per review, SetStream now sets ContentLength when the stream length can be determined, instead of adding a separate SetStreamWithLength method. --- .../HttpBindingProtocolGenerator.java | 2 +- .../rpc2/cbor/SerializeMiddleware.java | 2 +- .../middleware_capture_request_test.go | 3 +- .../requestcompression/request_compression.go | 2 - transport/http/middleware_content_length.go | 2 +- .../http/middleware_content_length_test.go | 128 ------------------ transport/http/protocol/awsjson/awsjson.go | 4 +- transport/http/protocol/awsquery/awsquery.go | 2 +- transport/http/protocol/ec2query/ec2query.go | 2 +- .../internal/httpbinding/serializer.go | 2 +- transport/http/protocol/rpcv2/rpcv2.go | 2 +- transport/http/request.go | 18 +-- transport/http/request_test.go | 64 +++------ 13 files changed, 30 insertions(+), 203 deletions(-) diff --git a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java index 746fb9777..913b5f889 100644 --- a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java +++ b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/integration/HttpBindingProtocolGenerator.java @@ -565,7 +565,7 @@ protected void writeSetPayloadShapeHeader(GoWriter writer, Shape payloadShape) { */ protected void writeSetStream(GoWriter writer, String operand) { writer.write(""" - if request, err = request.SetStreamWithLength($L); err != nil { + if request, err = request.SetStream($L); err != nil { return out, metadata, &smithy.SerializationError{Err: err} }""", operand); } diff --git a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java index 324e64347..0e303ef77 100644 --- a/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java +++ b/codegen/smithy-go-codegen/src/main/java/software/amazon/smithy/go/codegen/protocol/rpc2/cbor/SerializeMiddleware.java @@ -60,7 +60,7 @@ public Writable generateSerialize() { } payload := $reader:T($encode:T(cv)) - if req, err = req.SetStreamWithLength(payload); err != nil { + if req, err = req.SetStream(payload); err != nil { return out, metadata, &$error:T{Err: err} } 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.go b/private/requestcompression/request_compression.go index 304c30319..7c4147603 100644 --- a/private/requestcompression/request_compression.go +++ b/private/requestcompression/request_compression.go @@ -89,8 +89,6 @@ func (m requestCompression) HandleSerialize( } *req = *newReq - req.ContentLength = int64(len(compressedBytes)) - if val := req.Header.Get("Content-Encoding"); val != "" { req.Header.Set("Content-Encoding", fmt.Sprintf("%s, %s", val, algorithm)) } else { diff --git a/transport/http/middleware_content_length.go b/transport/http/middleware_content_length.go index 859c85de8..7a1c47fe3 100644 --- a/transport/http/middleware_content_length.go +++ b/transport/http/middleware_content_length.go @@ -16,7 +16,7 @@ type ComputeContentLength struct { // stack's Build step. // // Deprecated: Content-Length is now set when the request body is set via -// Request.SetStreamWithLength, so this middleware is no longer used. +// 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/protocol/awsjson/awsjson.go b/transport/http/protocol/awsjson/awsjson.go index a107b189f..6aa7f97ac 100644 --- a/transport/http/protocol/awsjson/awsjson.go +++ b/transport/http/protocol/awsjson/awsjson.go @@ -95,7 +95,7 @@ func (p *Protocol) SerializeRequest( } if schema.Input == nil { - sreq, err := req.SetStreamWithLength(strings.NewReader("{}")) + sreq, err := req.SetStream(strings.NewReader("{}")) if err != nil { return fmt.Errorf("set stream: %w", err) } @@ -106,7 +106,7 @@ func (p *Protocol) SerializeRequest( ss := internaljson.NewShapeSerializer() in.Serialize(ss) - sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/awsquery/awsquery.go b/transport/http/protocol/awsquery/awsquery.go index b7eaffdeb..a630f0653 100644 --- a/transport/http/protocol/awsquery/awsquery.go +++ b/transport/http/protocol/awsquery/awsquery.go @@ -88,7 +88,7 @@ func (p *Protocol) SerializeRequest( in.Serialize(ss) } - sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/ec2query/ec2query.go b/transport/http/protocol/ec2query/ec2query.go index 1b92443b5..ddd2d9c2a 100644 --- a/transport/http/protocol/ec2query/ec2query.go +++ b/transport/http/protocol/ec2query/ec2query.go @@ -89,7 +89,7 @@ func (p *Protocol) SerializeRequest( in.Serialize(ss) } - sreq, err := req.SetStreamWithLength(bytes.NewReader(ss.Bytes())) + sreq, err := req.SetStream(bytes.NewReader(ss.Bytes())) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/internal/httpbinding/serializer.go b/transport/http/protocol/internal/httpbinding/serializer.go index 6273f724d..33eee224c 100644 --- a/transport/http/protocol/internal/httpbinding/serializer.go +++ b/transport/http/protocol/internal/httpbinding/serializer.go @@ -137,7 +137,7 @@ func (s *ShapeSerializer) setBody(body io.Reader, contentType string) error { if s.request.Header.Get("Content-Type") == "" { s.request.Header.Set("Content-Type", contentType) } - sreq, err := s.request.SetStreamWithLength(body) + sreq, err := s.request.SetStream(body) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/protocol/rpcv2/rpcv2.go b/transport/http/protocol/rpcv2/rpcv2.go index 767443293..48b65b71e 100644 --- a/transport/http/protocol/rpcv2/rpcv2.go +++ b/transport/http/protocol/rpcv2/rpcv2.go @@ -98,7 +98,7 @@ func (p *Protocol) SerializeRequest( req.Header.Set("Content-Type", "application/cbor") - sreq, err := req.SetStreamWithLength(bytes.NewReader(payload)) + sreq, err := req.SetStream(bytes.NewReader(payload)) if err != nil { return fmt.Errorf("set stream: %w", err) } diff --git a/transport/http/request.go b/transport/http/request.go index de5aa003f..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,26 +154,13 @@ func (r *Request) SetStream(reader io.Reader) (rc *Request, err error) { rc.isStreamSeekable = isStreamSeekable rc.streamStartPos = streamStartPos - return rc, err -} - -// SetStreamWithLength sets the request stream like SetStream, and additionally -// sets ContentLength when the stream's length can be determined. If it cannot -// (e.g. a non-seekable stream), ContentLength is left unchanged so an -// explicitly provided value is preserved. -func (r *Request) SetStreamWithLength(reader io.Reader) (rc *Request, err error) { - rc, err = r.SetStream(reader) - if err != nil { - return rc, err - } - if n, ok, err := rc.StreamLength(); err != nil { return rc, err } else if ok { rc.ContentLength = n } - return rc, nil + return rc, err } // Build returns a build standard HTTP request value from the Smithy request. diff --git a/transport/http/request_test.go b/transport/http/request_test.go index 8dbc383a3..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) } @@ -218,45 +228,3 @@ func TestRequestSetStream(t *testing.T) { }) } } - -func TestRequestSetStreamWithLength(t *testing.T) { - cases := map[string]struct { - reader io.Reader - expectContentLength int64 - }{ - "nil stream": { - expectContentLength: 0, - }, - "empty seekable stream": { - reader: bytes.NewReader([]byte{}), - expectContentLength: 0, - }, - "unseekable stream with len": { - reader: bytes.NewBuffer([]byte("abc123")), - expectContentLength: 6, - }, - "seekable stream": { - reader: bytes.NewReader([]byte("abc123")), - expectContentLength: 6, - }, - // Length not determinable: left at default (-1) so a user-set value survives. - "unseekable no len stream": { - reader: io.NopCloser(bytes.NewBuffer([]byte("abc123"))), - expectContentLength: -1, - }, - } - - for name, c := range cases { - t.Run(name, func(t *testing.T) { - req := NewStackRequest().(*Request) - req, err := req.SetStreamWithLength(c.reader) - if err != nil { - t.Fatalf("expect no error, got %v", err) - } - - if e, a := c.expectContentLength, req.ContentLength; e != a { - t.Errorf("expect content-length %v, got %v", e, a) - } - }) - } -}