middleware_capture_request_compression.go (1658B)
1 package requestcompression 2 3 import ( 4 "bytes" 5 "context" 6 "fmt" 7 "io" 8 "net/http" 9 10 "github.com/aws/smithy-go/middleware" 11 smithyhttp "github.com/aws/smithy-go/transport/http" 12 ) 13 14 const captureUncompressedRequestID = "CaptureUncompressedRequest" 15 16 // AddCaptureUncompressedRequestMiddleware captures http request before compress encoding for check 17 func AddCaptureUncompressedRequestMiddleware(stack *middleware.Stack, buf *bytes.Buffer) error { 18 return stack.Serialize.Insert(&captureUncompressedRequestMiddleware{ 19 buf: buf, 20 }, "RequestCompression", middleware.Before) 21 } 22 23 type captureUncompressedRequestMiddleware struct { 24 req *http.Request 25 buf *bytes.Buffer 26 bytes []byte 27 } 28 29 // ID returns id of the captureUncompressedRequestMiddleware 30 func (*captureUncompressedRequestMiddleware) ID() string { 31 return captureUncompressedRequestID 32 } 33 34 // HandleSerialize captures request payload before it is compressed by request compression middleware 35 func (m *captureUncompressedRequestMiddleware) HandleSerialize(ctx context.Context, input middleware.SerializeInput, next middleware.SerializeHandler, 36 ) ( 37 output middleware.SerializeOutput, metadata middleware.Metadata, err error, 38 ) { 39 request, ok := input.Request.(*smithyhttp.Request) 40 if !ok { 41 return output, metadata, fmt.Errorf("error when retrieving http request") 42 } 43 44 _, err = io.Copy(m.buf, request.GetStream()) 45 if err != nil { 46 return output, metadata, fmt.Errorf("error when copying http request stream: %q", err) 47 } 48 if err = request.RewindStream(); err != nil { 49 return output, metadata, fmt.Errorf("error when rewinding request stream: %q", err) 50 } 51 52 return next.HandleSerialize(ctx, input) 53 }