middleware.go (3969B)
1 package v4a 2 3 import ( 4 "context" 5 "fmt" 6 "net/http" 7 "time" 8 9 awsmiddleware "github.com/aws/aws-sdk-go-v2/aws/middleware" 10 v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" 11 internalauth "github.com/aws/aws-sdk-go-v2/internal/auth" 12 "github.com/aws/smithy-go/middleware" 13 smithyhttp "github.com/aws/smithy-go/transport/http" 14 ) 15 16 // HTTPSigner is SigV4a HTTP signer implementation 17 type HTTPSigner interface { 18 SignHTTP(ctx context.Context, credentials Credentials, r *http.Request, payloadHash string, service string, regionSet []string, signingTime time.Time, optfns ...func(*SignerOptions)) error 19 } 20 21 // SignHTTPRequestMiddlewareOptions is the middleware options for constructing a SignHTTPRequestMiddleware. 22 type SignHTTPRequestMiddlewareOptions struct { 23 Credentials CredentialsProvider 24 Signer HTTPSigner 25 LogSigning bool 26 } 27 28 // SignHTTPRequestMiddleware is a middleware for signing an HTTP request using SigV4a. 29 type SignHTTPRequestMiddleware struct { 30 credentials CredentialsProvider 31 signer HTTPSigner 32 logSigning bool 33 } 34 35 // NewSignHTTPRequestMiddleware constructs a SignHTTPRequestMiddleware using the given SignHTTPRequestMiddlewareOptions. 36 func NewSignHTTPRequestMiddleware(options SignHTTPRequestMiddlewareOptions) *SignHTTPRequestMiddleware { 37 return &SignHTTPRequestMiddleware{ 38 credentials: options.Credentials, 39 signer: options.Signer, 40 logSigning: options.LogSigning, 41 } 42 } 43 44 // ID the middleware identifier. 45 func (s *SignHTTPRequestMiddleware) ID() string { 46 return "Signing" 47 } 48 49 // HandleFinalize signs an HTTP request using SigV4a. 50 func (s *SignHTTPRequestMiddleware) HandleFinalize( 51 ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler, 52 ) ( 53 out middleware.FinalizeOutput, metadata middleware.Metadata, err error, 54 ) { 55 if !hasCredentialProvider(s.credentials) { 56 return next.HandleFinalize(ctx, in) 57 } 58 59 req, ok := in.Request.(*smithyhttp.Request) 60 if !ok { 61 return out, metadata, fmt.Errorf("unexpected request middleware type %T", in.Request) 62 } 63 64 signingName, signingRegion := awsmiddleware.GetSigningName(ctx), awsmiddleware.GetSigningRegion(ctx) 65 payloadHash := v4.GetPayloadHash(ctx) 66 if len(payloadHash) == 0 { 67 return out, metadata, &SigningError{Err: fmt.Errorf("computed payload hash missing from context")} 68 } 69 70 credentials, err := s.credentials.RetrievePrivateKey(ctx) 71 if err != nil { 72 return out, metadata, &SigningError{Err: fmt.Errorf("failed to retrieve credentials: %w", err)} 73 } 74 75 signerOptions := []func(o *SignerOptions){ 76 func(o *SignerOptions) { 77 o.Logger = middleware.GetLogger(ctx) 78 o.LogSigning = s.logSigning 79 }, 80 } 81 82 // existing DisableURIPathEscaping is equivalent in purpose 83 // to authentication scheme property DisableDoubleEncoding 84 disableDoubleEncoding, overridden := internalauth.GetDisableDoubleEncoding(ctx) 85 if overridden { 86 signerOptions = append(signerOptions, func(o *SignerOptions) { 87 o.DisableURIPathEscaping = disableDoubleEncoding 88 }) 89 } 90 91 err = s.signer.SignHTTP(ctx, credentials, req.Request, payloadHash, signingName, []string{signingRegion}, time.Now().UTC(), signerOptions...) 92 if err != nil { 93 return out, metadata, &SigningError{Err: fmt.Errorf("failed to sign http request, %w", err)} 94 } 95 96 return next.HandleFinalize(ctx, in) 97 } 98 99 func hasCredentialProvider(p CredentialsProvider) bool { 100 if p == nil { 101 return false 102 } 103 104 return true 105 } 106 107 // RegisterSigningMiddleware registers the SigV4a signing middleware to the stack. If a signing middleware is already 108 // present, this provided middleware will be swapped. Otherwise the middleware will be added at the tail of the 109 // finalize step. 110 func RegisterSigningMiddleware(stack *middleware.Stack, signingMiddleware *SignHTTPRequestMiddleware) (err error) { 111 const signedID = "Signing" 112 _, present := stack.Finalize.Get(signedID) 113 if present { 114 _, err = stack.Finalize.Swap(signedID, signingMiddleware) 115 } else { 116 err = stack.Finalize.Add(signingMiddleware, middleware.After) 117 } 118 return err 119 }