From ba501e482d787f43bd500be872a3c4b6c6b2eb92 Mon Sep 17 00:00:00 2001 From: Ben McClelland Date: Sat, 25 Nov 2023 17:25:30 -0700 Subject: [PATCH] feat: steaming requests for put object and put part This builds on the previous work that sets up the body streaming for the put object and put part requests. This adds the auth and checksum readers to postpone the v4auth checks and the content checksum until the end of the body stream. This means that the backend with start reading the data from the body stream before the request is fully validated and signatures checked. So the backend must check the error returned from the body reader for the final auth and content checks. The backend is expected to discard the data upon error. This should increase performance and reduce memory utilization to no longer require caching the entire request body in memory for put object and put part. --- backend/posix/posix.go | 11 -- backend/s3proxy/s3.go | 12 +- integration/group-tests.go | 2 + integration/tests.go | 13 +- runtests.sh | 1 + s3api/controllers/base.go | 29 +++- s3api/middlewares/authentication.go | 196 +++++++-------------- s3api/middlewares/body-reader.go | 31 ++++ s3api/middlewares/md5.go | 22 +-- s3api/utils/auth-reader.go | 261 ++++++++++++++++++++++++++++ s3api/utils/csum-reader.go | 130 ++++++++++++++ s3api/utils/reader.go | 72 -------- s3api/utils/utils.go | 5 +- s3api/utils/utils_test.go | 2 +- 14 files changed, 550 insertions(+), 237 deletions(-) create mode 100644 s3api/middlewares/body-reader.go create mode 100644 s3api/utils/auth-reader.go create mode 100644 s3api/utils/csum-reader.go delete mode 100644 s3api/utils/reader.go diff --git a/backend/posix/posix.go b/backend/posix/posix.go index f40af9bd..a5c88da7 100644 --- a/backend/posix/posix.go +++ b/backend/posix/posix.go @@ -37,7 +37,6 @@ import ( "github.com/pkg/xattr" "github.com/versity/versitygw/auth" "github.com/versity/versitygw/backend" - "github.com/versity/versitygw/s3api/utils" "github.com/versity/versitygw/s3err" "github.com/versity/versitygw/s3response" ) @@ -895,11 +894,6 @@ func (p *Posix) UploadPart(_ context.Context, input *s3.UploadPartInput) (string return "", fmt.Errorf("write part data: %w", err) } - rdr, ok := r.(*utils.HashReader) - if ok && rdr.Err() != nil { - return "", err - } - err = f.link() if err != nil { return "", fmt.Errorf("link object in namespace: %w", err) @@ -1107,11 +1101,6 @@ func (p *Posix) PutObject(ctx context.Context, po *s3.PutObjectInput) (string, e if err != nil { return "", fmt.Errorf("write object data: %w", err) } - - r, ok := po.Body.(*utils.HashReader) - if ok && r.Err() != nil { - return "", r.Err() - } dir := filepath.Dir(name) if dir != "" { err = mkdirAll(dir, os.FileMode(0755), *po.Bucket, *po.Key) diff --git a/backend/s3proxy/s3.go b/backend/s3proxy/s3.go index 14fa7fcb..3952b9ef 100644 --- a/backend/s3proxy/s3.go +++ b/backend/s3proxy/s3.go @@ -270,7 +270,11 @@ func (s *S3be) UploadPart(ctx context.Context, input *s3.UploadPartInput) (etag return "", err } - output, err := client.UploadPart(ctx, input) + // streaming backend is not seekable, + // use unsigned payload for streaming ops + output, err := client.UploadPart(ctx, input, s3.WithAPIOptions( + v4.SwapComputePayloadSHA256ForUnsignedPayloadMiddleware, + )) err = handleError(err) if err != nil { return "", err @@ -303,7 +307,11 @@ func (s *S3be) PutObject(ctx context.Context, input *s3.PutObjectInput) (string, return "", err } - output, err := client.PutObject(ctx, input) + // streaming backend is not seekable, + // use unsigned payload for streaming ops + output, err := client.PutObject(ctx, input, s3.WithAPIOptions( + v4.SwapComputePayloadSHA256ForUnsignedPayloadMiddleware, + )) err = handleError(err) if err != nil { return "", err diff --git a/integration/group-tests.go b/integration/group-tests.go index e9dd1189..ab6899e3 100644 --- a/integration/group-tests.go +++ b/integration/group-tests.go @@ -51,6 +51,7 @@ func TestPutObject(s *S3Conf) { PutObject_special_chars(s) PutObject_invalid_long_tags(s) PutObject_success(s) + PutObject_invalid_credentials(s) } func TestHeadObject(s *S3Conf) { @@ -262,6 +263,7 @@ func GetIntTests() IntTests { "PutObject_special_chars": PutObject_special_chars, "PutObject_invalid_long_tags": PutObject_invalid_long_tags, "PutObject_success": PutObject_success, + "PutObject_invalid_credentials": PutObject_invalid_credentials, "HeadObject_non_existing_object": HeadObject_non_existing_object, "HeadObject_success": HeadObject_success, "GetObject_non_existing_key": GetObject_non_existing_key, diff --git a/integration/tests.go b/integration/tests.go index a62fcfb5..e8d9c76d 100644 --- a/integration/tests.go +++ b/integration/tests.go @@ -447,7 +447,7 @@ func Authentication_invalid_signed_headers(s *S3Conf) error { return err } defer resp.Body.Close() - if err := checkAuthErr(resp, s3err.GetAPIError(s3err.ErrCredMalformed)); err != nil { + if err := checkAuthErr(resp, s3err.GetAPIError(s3err.ErrInvalidQueryParams)); err != nil { return err } @@ -1088,6 +1088,17 @@ func PutObject_success(s *S3Conf) error { }) } +func PutObject_invalid_credentials(s *S3Conf) error { + testName := "PutObject_invalid_credentials" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + newconf := *s + newconf.awsSecret = newconf.awsSecret + "badpassword" + client := s3.NewFromConfig(newconf.Config()) + err := putObjects(client, []string{"my-obj"}, bucket) + return checkApiErr(err, s3err.GetAPIError(s3err.ErrSignatureDoesNotMatch)) + }) +} + func HeadObject_non_existing_object(s *S3Conf) error { testName := "HeadObject_non_existing_object" return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { diff --git a/runtests.sh b/runtests.sh index 178f223a..2a219010 100755 --- a/runtests.sh +++ b/runtests.sh @@ -1,6 +1,7 @@ #!/bin/bash # make temp dirs +rm -rf /tmp/gw mkdir /tmp/gw rm -rf /tmp/covdata mkdir /tmp/covdata diff --git a/s3api/controllers/base.go b/s3api/controllers/base.go index d5a5b32b..d7cc0c42 100644 --- a/s3api/controllers/base.go +++ b/s3api/controllers/base.go @@ -15,7 +15,9 @@ package controllers import ( + "bytes" "encoding/xml" + "errors" "fmt" "io" "log" @@ -515,7 +517,14 @@ func (c S3ApiController) PutActions(ctx *fiber.Ctx) error { return SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidRequest), &MetaOpts{Logger: c.logger, Action: "UploadPart", BucketOwner: parsedAcl.Owner}) } - body := ctx.Locals("body-reader").(io.Reader) + var body io.Reader + bodyi := ctx.Locals("body-reader") + if bodyi != nil { + body = bodyi.(io.Reader) + } else { + body = bytes.NewReader([]byte{}) + } + ctx.Locals("logReqBody", false) etag, err := c.be.UploadPart(ctx.Context(), &s3.UploadPartInput{ Bucket: &bucket, @@ -654,7 +663,13 @@ func (c S3ApiController) PutActions(ctx *fiber.Ctx) error { return SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidRequest), &MetaOpts{Logger: c.logger, Action: "PutObject", BucketOwner: parsedAcl.Owner}) } - rdr := ctx.Locals("body-reader").(io.Reader) + var body io.Reader + bodyi := ctx.Locals("body-reader") + if bodyi != nil { + body = bodyi.(io.Reader) + } else { + body = bytes.NewReader([]byte{}) + } ctx.Locals("logReqBody", false) etag, err := c.be.PutObject(ctx.Context(), &s3.PutObjectInput{ @@ -662,7 +677,7 @@ func (c S3ApiController) PutActions(ctx *fiber.Ctx) error { Key: &keyStart, ContentLength: &contentLength, Metadata: metadata, - Body: rdr, + Body: body, Tagging: &tagging, }) ctx.Response().Header.Set("ETag", etag) @@ -1002,10 +1017,10 @@ func SendResponse(ctx *fiber.Ctx, err error, l *MetaOpts) error { }) } if err != nil { - serr, ok := err.(s3err.APIError) - if ok { - ctx.Status(serr.HTTPStatusCode) - return ctx.Send(s3err.GetAPIErrorResponse(serr, "", "", "")) + var apierr s3err.APIError + if errors.As(err, &apierr) { + ctx.Status(apierr.HTTPStatusCode) + return ctx.Send(s3err.GetAPIErrorResponse(apierr, "", "", "")) } log.Printf("Internal Error, %v", err) diff --git a/s3api/middlewares/authentication.go b/s3api/middlewares/authentication.go index b582860c..e0ac1ccb 100644 --- a/s3api/middlewares/authentication.go +++ b/s3api/middlewares/authentication.go @@ -18,15 +18,11 @@ import ( "crypto/sha256" "encoding/hex" "fmt" - "math" + "io" "net/http" - "os" - "strings" + "strconv" "time" - "github.com/aws/aws-sdk-go-v2/aws" - v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" - "github.com/aws/smithy-go/logging" "github.com/gofiber/fiber/v2" "github.com/versity/versitygw/auth" "github.com/versity/versitygw/s3api/controllers" @@ -37,7 +33,6 @@ import ( const ( iso8601Format = "20060102T150405Z" - YYYYMMDD = "20060102" ) type RootUserConfig struct { @@ -53,144 +48,93 @@ func VerifyV4Signature(root RootUserConfig, iam auth.IAMService, logger s3log.Au ctx.Locals("startTime", time.Now()) authorization := ctx.Get("Authorization") if authorization == "" { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrAuthHeaderEmpty), &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrAuthHeaderEmpty), logger) } - // Check the signature version - authParts := strings.Split(authorization, ",") - for i, el := range authParts { - authParts[i] = strings.TrimSpace(el) + authData, err := utils.ParseAuthorization(authorization) + if err != nil { + return sendResponse(ctx, err, logger) } - if len(authParts) != 3 { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrMissingFields), &controllers.MetaOpts{Logger: logger}) + if authData.Algorithm != "AWS4-HMAC-SHA256" { + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureVersionNotSupported), logger) } - startParts := strings.Split(authParts[0], " ") - - if startParts[0] != "AWS4-HMAC-SHA256" { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureVersionNotSupported), &controllers.MetaOpts{Logger: logger}) - } - - credKv := strings.Split(startParts[1], "=") - if len(credKv) != 2 { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrCredMalformed), &controllers.MetaOpts{Logger: logger}) - } - // Credential variables validation - creds := strings.Split(credKv[1], "/") - if len(creds) != 5 { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrCredMalformed), &controllers.MetaOpts{Logger: logger}) - } - if creds[4] != "aws4_request" { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureTerminationStr), &controllers.MetaOpts{Logger: logger}) - } - if creds[3] != "s3" { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureIncorrService), &controllers.MetaOpts{Logger: logger}) - } - if creds[2] != region { - return controllers.SendResponse(ctx, s3err.APIError{ + if authData.Region != region { + return sendResponse(ctx, s3err.APIError{ Code: "SignatureDoesNotMatch", - Description: fmt.Sprintf("Credential should be scoped to a valid Region, not %v", creds[2]), + Description: fmt.Sprintf("Credential should be scoped to a valid Region, not %v", authData.Region), HTTPStatusCode: http.StatusForbidden, - }, &controllers.MetaOpts{Logger: logger}) + }, logger) } - ctx.Locals("isRoot", creds[0] == root.Access) + ctx.Locals("isRoot", authData.Access == root.Access) - _, err := time.Parse(YYYYMMDD, creds[1]) - if err != nil { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureDateDoesNotMatch), &controllers.MetaOpts{Logger: logger}) - } - - signHdrKv := strings.Split(authParts[1], "=") - if len(signHdrKv) != 2 { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrCredMalformed), &controllers.MetaOpts{Logger: logger}) - } - signedHdrs := strings.Split(signHdrKv[1], ";") - - account, err := acct.getAccount(creds[0]) + account, err := acct.getAccount(authData.Access) if err == auth.ErrNoSuchUser { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidAccessKeyID), &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidAccessKeyID), logger) } if err != nil { - return controllers.SendResponse(ctx, err, &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, err, logger) } ctx.Locals("account", account) // Check X-Amz-Date header date := ctx.Get("X-Amz-Date") if date == "" { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrMissingDateHeader), &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrMissingDateHeader), logger) } // Parse the date and check the date validity tdate, err := time.Parse(iso8601Format, date) if err != nil { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrMalformedDate), &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrMalformedDate), logger) } - if date[:8] != creds[1] { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureDateDoesNotMatch), &controllers.MetaOpts{Logger: logger}) + if date[:8] != authData.Date { + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureDateDoesNotMatch), logger) } // Validate the dates difference err = validateDate(tdate) if err != nil { - return controllers.SendResponse(ctx, err, &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, err, logger) } - hashPayloadHeader := ctx.Get("X-Amz-Content-Sha256") - ok := isSpecialPayload(hashPayloadHeader) + if utils.IsBigDataAction(ctx) { + // for streaming PUT actions, authorization is deferred + // until end of stream due to need to get length and + // checksum of the stream to validate authorization + wrapBodyReader(ctx, func(r io.Reader) io.Reader { + return utils.NewAuthReader(ctx, r, authData, account.Secret, debug) + }) + return ctx.Next() + } - if !ok { - if utils.IsBigDataAction(ctx) { - rdr := ctx.Request().BodyStream() - r := utils.NewHashReader(rdr, sha256.New(), hashPayloadHeader, utils.HashTypeSha256) - ctx.Locals("body-reader", r) - } else { - // Calculate the hash of the request payload - hashedPayload := sha256.Sum256(ctx.Body()) - hexPayload := hex.EncodeToString(hashedPayload[:]) + hashPayload := ctx.Get("X-Amz-Content-Sha256") + if !utils.IsSpecialPayload(hashPayload) { + // Calculate the hash of the request payload + hashedPayload := sha256.Sum256(ctx.Body()) + hexPayload := hex.EncodeToString(hashedPayload[:]) - // Compare the calculated hash with the hash provided - if hashPayloadHeader != hexPayload { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrContentSHA256Mismatch), &controllers.MetaOpts{Logger: logger}) - } + // Compare the calculated hash with the hash provided + if hashPayload != hexPayload { + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrContentSHA256Mismatch), logger) } } - // Create a new http request instance from fasthttp request - req, err := utils.CreateHttpRequestFromCtx(ctx, signedHdrs) + var contentLength int64 + contentLengthStr := ctx.Get("Content-Length") + if contentLengthStr != "" { + contentLength, err = strconv.ParseInt(contentLengthStr, 10, 64) + if err != nil { + return sendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidRequest), logger) + } + } + + err = utils.CheckValidSignature(ctx, authData, account.Secret, hashPayload, tdate, contentLength, debug) if err != nil { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInternalError), &controllers.MetaOpts{Logger: logger}) - } - - signer := v4.NewSigner() - - signErr := signer.SignHTTP(req.Context(), aws.Credentials{ - AccessKeyID: creds[0], - SecretAccessKey: account.Secret, - }, req, hashPayloadHeader, creds[3], region, tdate, func(options *v4.SignerOptions) { - options.DisableURIPathEscaping = true - if true { - options.LogSigning = true - options.Logger = logging.NewStandardLogger(os.Stderr) - } - }) - if signErr != nil { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInternalError), &controllers.MetaOpts{Logger: logger}) - } - - parts := strings.Split(req.Header.Get("Authorization"), " ") - if len(parts) < 4 { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrMissingFields), &controllers.MetaOpts{Logger: logger}) - } - calculatedSign := strings.Split(parts[3], "=")[1] - expectedSign := strings.Split(authParts[2], "=")[1] - fmt.Println(calculatedSign, expectedSign) - - if expectedSign != calculatedSign { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrSignatureDoesNotMatch), &controllers.MetaOpts{Logger: logger}) + return sendResponse(ctx, err, logger) } return ctx.Next() @@ -214,39 +158,29 @@ func (a accounts) getAccount(access string) (auth.Account, error) { return a.iam.GetUserAccount(access) } -func isSpecialPayload(str string) bool { - specialValues := map[string]bool{ - "UNSIGNED-PAYLOAD": true, - "STREAMING-UNSIGNED-PAYLOAD-TRAILER": true, - "STREAMING-AWS4-HMAC-SHA256-PAYLOAD": true, - "STREAMING-AWS4-HMAC-SHA256-PAYLOAD-TRAILER": true, - "STREAMING-AWS4-ECDSA-P256-SHA256-PAYLOAD": true, - "STREAMING-AWS4-ECDSA-P256-SHA256-PAYLOAD-TRAILER": true, - } - - return specialValues[str] -} - func validateDate(date time.Time) error { now := time.Now().UTC() diff := date.Unix() - now.Unix() // Checks the dates difference to be less than a minute - if math.Abs(float64(diff)) > 60 { - if diff > 0 { - return s3err.APIError{ - Code: "SignatureDoesNotMatch", - Description: fmt.Sprintf("Signature not yet current: %s is still later than %s", date.Format(iso8601Format), now.Format(iso8601Format)), - HTTPStatusCode: http.StatusForbidden, - } - } else { - return s3err.APIError{ - Code: "SignatureDoesNotMatch", - Description: fmt.Sprintf("Signature expired: %s is now earlier than %s", date.Format(iso8601Format), now.Format(iso8601Format)), - HTTPStatusCode: http.StatusForbidden, - } + if diff > 60 { + return s3err.APIError{ + Code: "SignatureDoesNotMatch", + Description: fmt.Sprintf("Signature not yet current: %s is still later than %s", date.Format(iso8601Format), now.Format(iso8601Format)), + HTTPStatusCode: http.StatusForbidden, + } + } + if diff < -60 { + return s3err.APIError{ + Code: "SignatureDoesNotMatch", + Description: fmt.Sprintf("Signature expired: %s is now earlier than %s", date.Format(iso8601Format), now.Format(iso8601Format)), + HTTPStatusCode: http.StatusForbidden, } } return nil } + +func sendResponse(ctx *fiber.Ctx, err error, logger s3log.AuditLogger) error { + return controllers.SendResponse(ctx, err, &controllers.MetaOpts{Logger: logger}) +} diff --git a/s3api/middlewares/body-reader.go b/s3api/middlewares/body-reader.go new file mode 100644 index 00000000..aefd4fe8 --- /dev/null +++ b/s3api/middlewares/body-reader.go @@ -0,0 +1,31 @@ +// Copyright 2023 Versity Software +// This file is licensed under the Apache License, Version 2.0 +// (the "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package middlewares + +import ( + "io" + + "github.com/gofiber/fiber/v2" +) + +func wrapBodyReader(ctx *fiber.Ctx, wr func(io.Reader) io.Reader) { + r, ok := ctx.Locals("body-reader").(io.Reader) + if !ok { + r = ctx.Request().BodyStream() + } + + r = wr(r) + ctx.Locals("body-reader", r) +} diff --git a/s3api/middlewares/md5.go b/s3api/middlewares/md5.go index b1f97209..390bd887 100644 --- a/s3api/middlewares/md5.go +++ b/s3api/middlewares/md5.go @@ -16,7 +16,7 @@ package middlewares import ( "crypto/md5" - "encoding/base64" + "io" "github.com/gofiber/fiber/v2" "github.com/versity/versitygw/s3api/controllers" @@ -33,16 +33,18 @@ func VerifyMD5Body(logger s3log.AuditLogger) fiber.Handler { } if utils.IsBigDataAction(ctx) { - rdr := ctx.Locals("body-reader").(*utils.HashReader) - r := utils.NewHashReader(rdr, md5.New(), incomingSum, utils.HashTypeMd5) - ctx.Locals("body-reader", r) - } else { - sum := md5.Sum(ctx.Body()) - calculatedSum := base64.StdEncoding.EncodeToString(sum[:]) + wrapBodyReader(ctx, func(r io.Reader) io.Reader { + r, _ = utils.NewHashReader(r, incomingSum, utils.HashTypeMd5) + return r + }) + return ctx.Next() + } - if incomingSum != calculatedSum { - return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidDigest), &controllers.MetaOpts{Logger: logger}) - } + sum := md5.Sum(ctx.Body()) + calculatedSum := utils.Md5SumString(sum[:]) + + if incomingSum != calculatedSum { + return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidDigest), &controllers.MetaOpts{Logger: logger}) } return ctx.Next() diff --git a/s3api/utils/auth-reader.go b/s3api/utils/auth-reader.go new file mode 100644 index 00000000..03248352 --- /dev/null +++ b/s3api/utils/auth-reader.go @@ -0,0 +1,261 @@ +// Copyright 2023 Versity Software +// This file is licensed under the Apache License, Version 2.0 +// (the "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package utils + +import ( + "errors" + "fmt" + "io" + "os" + "strings" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" + "github.com/aws/smithy-go/logging" + "github.com/gofiber/fiber/v2" + "github.com/versity/versitygw/s3err" +) + +const ( + iso8601Format = "20060102T150405Z" + yyyymmdd = "20060102" +) + +// AuthReader is an io.Reader that validates the request authorization +// once the underlying reader returns io.EOF. This is needed for streaming +// data requests where the data size and checksum are not known until +// the data is completely read. +type AuthReader struct { + ctx *fiber.Ctx + auth AuthData + secret string + size int + r *HashReader + debug bool +} + +// NewAuthReader initializes an io.Reader that will verify the request +// v4 auth when the underlying reader returns io.EOF. This postpones the +// authorization check until the reader is consumed. So it is important that +// the consumer of this reader checks for the auth errors while reading. +func NewAuthReader(ctx *fiber.Ctx, r io.Reader, auth AuthData, secret string, debug bool) *AuthReader { + var hr *HashReader + hashPayload := ctx.Get("X-Amz-Content-Sha256") + if !IsSpecialPayload(hashPayload) { + hr, _ = NewHashReader(r, "", HashTypeSha256) + } else { + hr, _ = NewHashReader(r, "", HashTypeNone) + } + + return &AuthReader{ + ctx: ctx, + r: hr, + auth: auth, + secret: secret, + debug: debug, + } +} + +// Read allows *AuthReader to be used as an io.Reader +func (ar *AuthReader) Read(p []byte) (int, error) { + n, err := ar.r.Read(p) + ar.size += n + + if errors.Is(err, io.EOF) { + verr := ar.validateSignature() + if verr != nil { + return n, verr + } + } + + return n, err +} + +func (ar *AuthReader) validateSignature() error { + date := ar.ctx.Get("X-Amz-Date") + if date == "" { + return s3err.GetAPIError(s3err.ErrMissingDateHeader) + } + + hashPayload := ar.ctx.Get("X-Amz-Content-Sha256") + if !IsSpecialPayload(hashPayload) { + hexPayload := ar.r.Sum() + + // Compare the calculated hash with the hash provided + if hashPayload != hexPayload { + return s3err.GetAPIError(s3err.ErrContentSHA256Mismatch) + } + } + + // Parse the date and check the date validity + tdate, err := time.Parse(iso8601Format, date) + if err != nil { + return s3err.GetAPIError(s3err.ErrMalformedDate) + } + + return CheckValidSignature(ar.ctx, ar.auth, ar.secret, hashPayload, tdate, int64(ar.size), ar.debug) +} + +const ( + service = "s3" +) + +// CheckValidSignature validates the ctx v4 auth signature +func CheckValidSignature(ctx *fiber.Ctx, auth AuthData, secret, checksum string, tdate time.Time, contentLen int64, debug bool) error { + signedHdrs := strings.Split(auth.SignedHeaders, ";") + + // Create a new http request instance from fasthttp request + req, err := createHttpRequestFromCtx(ctx, signedHdrs, contentLen) + if err != nil { + return fmt.Errorf("create http request from context: %w", err) + } + + signer := v4.NewSigner() + + signErr := signer.SignHTTP(req.Context(), aws.Credentials{ + AccessKeyID: auth.Access, + SecretAccessKey: secret, + }, req, checksum, service, auth.Region, tdate, func(options *v4.SignerOptions) { + options.DisableURIPathEscaping = true + if debug { + options.LogSigning = true + options.Logger = logging.NewStandardLogger(os.Stderr) + } + }) + if signErr != nil { + return fmt.Errorf("sign generated http request: %w", err) + } + + genAuth, err := ParseAuthorization(req.Header.Get("Authorization")) + if err != nil { + return err + } + + if auth.Signature != genAuth.Signature { + return s3err.GetAPIError(s3err.ErrSignatureDoesNotMatch) + } + + return nil +} + +// AuthData is the parsed authorization data from the header +type AuthData struct { + Algorithm string + Access string + Region string + SignedHeaders string + Signature string + Date string +} + +// ParseAuthorization returns the parsed fields for the aws v4 auth header +// example authorization string from aws docs: +// Authorization: AWS4-HMAC-SHA256 +// Credential=AKIAIOSFODNN7EXAMPLE/20130524/us-east-1/s3/aws4_request, +// SignedHeaders=host;range;x-amz-date, +// Signature=fe5f80f77d5fa3beca038a248ff027d0445342fe2855ddc963176630326f1024 +func ParseAuthorization(authorization string) (AuthData, error) { + a := AuthData{} + + // authorization must start with: + // Authorization: + // followed by key=value pairs separated by "," + authParts := strings.Fields(authorization) + for i, el := range authParts { + authParts[i] = strings.TrimSpace(el) + } + + if len(authParts) < 3 { + return a, s3err.GetAPIError(s3err.ErrMissingFields) + } + + algo := authParts[0] + + kvData := strings.Join(authParts[1:], "") + kvPairs := strings.Split(kvData, ",") + // we are expecting at least Credential, SignedHeaders, and Signature + // key value pairs here + if len(kvPairs) < 3 { + return a, s3err.GetAPIError(s3err.ErrMissingFields) + } + + var access, region, signedHeaders, signature, date string + + for _, kv := range kvPairs { + keyValue := strings.Split(kv, "=") + if len(keyValue) != 2 { + switch { + case strings.HasPrefix(kv, "Credential"): + return a, s3err.GetAPIError(s3err.ErrCredMalformed) + case strings.HasPrefix(kv, "SignedHeaders"): + return a, s3err.GetAPIError(s3err.ErrInvalidQueryParams) + } + return a, s3err.GetAPIError(s3err.ErrMissingFields) + } + key := strings.TrimSpace(keyValue[0]) + value := strings.TrimSpace(keyValue[1]) + + switch key { + case "Credential": + creds := strings.Split(value, "/") + if len(creds) != 5 { + return a, s3err.GetAPIError(s3err.ErrCredMalformed) + } + if creds[3] != "s3" { + return a, s3err.GetAPIError(s3err.ErrSignatureIncorrService) + } + if creds[4] != "aws4_request" { + return a, s3err.GetAPIError(s3err.ErrSignatureTerminationStr) + } + _, err := time.Parse(yyyymmdd, creds[1]) + if err != nil { + return a, s3err.GetAPIError(s3err.ErrSignatureDateDoesNotMatch) + } + access = creds[0] + date = creds[1] + region = creds[2] + case "SignedHeaders": + signedHeaders = value + case "Signature": + signature = value + } + } + + return AuthData{ + Algorithm: algo, + Access: access, + Region: region, + SignedHeaders: signedHeaders, + Signature: signature, + Date: date, + }, nil +} + +var ( + specialValues = map[string]bool{ + "UNSIGNED-PAYLOAD": true, + "STREAMING-UNSIGNED-PAYLOAD-TRAILER": true, + "STREAMING-AWS4-HMAC-SHA256-PAYLOAD": true, + "STREAMING-AWS4-HMAC-SHA256-PAYLOAD-TRAILER": true, + "STREAMING-AWS4-ECDSA-P256-SHA256-PAYLOAD": true, + "STREAMING-AWS4-ECDSA-P256-SHA256-PAYLOAD-TRAILER": true, + } +) + +// IsSpecialPayload checks for streaming/unsigned authorization types +func IsSpecialPayload(str string) bool { + return specialValues[str] +} diff --git a/s3api/utils/csum-reader.go b/s3api/utils/csum-reader.go new file mode 100644 index 00000000..04ce90d7 --- /dev/null +++ b/s3api/utils/csum-reader.go @@ -0,0 +1,130 @@ +// Copyright 2023 Versity Software +// This file is licensed under the Apache License, Version 2.0 +// (the "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package utils + +import ( + "crypto/md5" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "errors" + "hash" + "io" + + "github.com/versity/versitygw/s3err" +) + +// HashType identifies the checksum algorithm to be used +type HashType string + +const ( + // HashTypeMd5 generates MD5 checksum for the data stream + HashTypeMd5 = "md5" + // HashTypeSha256 generates SHA256 checksum for the data stream + HashTypeSha256 = "sha256" + // HashTypeNone is a no-op checksum for the data stream + HashTypeNone = "none" +) + +// HashReader is an io.Reader that calculates the checksum +// as the data is read +type HashReader struct { + hashType HashType + hash hash.Hash + r io.Reader + sum string +} + +var ( + errInvalidHashType = errors.New("unsupported or invalid checksum type") +) + +// NewHashReader intializes an io.Reader from an underlying io.Reader that +// calculates the checksum while the reader is being read from. If the +// sum provided is not "", the reader will return an error when the underlying +// reader returns io.EOF if the checksum does not match the provided expected +// checksum. If the provided sum is "", then the Sum() method can still +// be used to get the current checksum for the data read so far. +func NewHashReader(r io.Reader, expectedSum string, ht HashType) (*HashReader, error) { + var hash hash.Hash + switch ht { + case HashTypeMd5: + hash = md5.New() + case HashTypeSha256: + hash = sha256.New() + case HashTypeNone: + hash = noop{} + default: + return nil, errInvalidHashType + } + + return &HashReader{ + hash: hash, + r: r, + sum: expectedSum, + hashType: ht, + }, nil +} + +// Read allows *HashReader to be used as an io.Reader +func (hr *HashReader) Read(p []byte) (int, error) { + n, readerr := hr.r.Read(p) + _, err := hr.hash.Write(p[:n]) + if err != nil { + return n, err + } + if errors.Is(readerr, io.EOF) && hr.sum != "" { + switch hr.hashType { + case HashTypeMd5: + sum := base64.StdEncoding.EncodeToString(hr.hash.Sum(nil)) + if sum != hr.sum { + return n, s3err.GetAPIError(s3err.ErrInvalidDigest) + } + case HashTypeSha256: + sum := hex.EncodeToString(hr.hash.Sum(nil)) + if sum != hr.sum { + return n, s3err.GetAPIError(s3err.ErrContentSHA256Mismatch) + } + default: + return n, errInvalidHashType + } + } + return n, readerr +} + +// Sum returns the checksum hash of the data read so far +func (hr *HashReader) Sum() string { + switch hr.hashType { + case HashTypeMd5: + return Md5SumString(hr.hash.Sum(nil)) + case HashTypeSha256: + return hex.EncodeToString(hr.hash.Sum(nil)) + default: + return "" + } +} + +// Md5SumString converts the hash bytes to the string checksum value +func Md5SumString(b []byte) string { + return base64.StdEncoding.EncodeToString(b) +} + +type noop struct{} + +func (n noop) Write(p []byte) (int, error) { return 0, nil } +func (n noop) Sum(b []byte) []byte { return []byte{} } +func (n noop) Reset() {} +func (n noop) Size() int { return 0 } +func (n noop) BlockSize() int { return 1 } diff --git a/s3api/utils/reader.go b/s3api/utils/reader.go deleted file mode 100644 index 6dad261d..00000000 --- a/s3api/utils/reader.go +++ /dev/null @@ -1,72 +0,0 @@ -// Copyright 2023 Versity Software -// This file is licensed under the Apache License, Version 2.0 -// (the "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -package utils - -import ( - "encoding/base64" - "encoding/hex" - "errors" - "hash" - "io" - - "github.com/versity/versitygw/s3err" -) - -type HashType string - -const ( - HashTypeMd5 = "md5" - HashTypeSha256 = "sha256" -) - -type HashReader struct { - hashType HashType - hash hash.Hash - r io.Reader - sum string - err error -} - -func NewHashReader(r io.Reader, hash hash.Hash, sum string, ht HashType) *HashReader { - return &HashReader{hash: hash, r: r, sum: sum, hashType: ht} -} - -func (hr *HashReader) Read(p []byte) (int, error) { - n, readerr := hr.r.Read(p) - _, err := hr.hash.Write(p[:n]) - if err != nil { - return n, err - } - if errors.Is(readerr, io.EOF) { - if hr.hashType == HashTypeMd5 { - sum := base64.StdEncoding.EncodeToString(hr.hash.Sum(nil)) - if sum != hr.sum { - hr.err = s3err.GetAPIError(s3err.ErrInvalidDigest) - return n, s3err.GetAPIError(s3err.ErrInvalidDigest) - } - } else if hr.hashType == HashTypeSha256 { - sum := hex.EncodeToString(hr.hash.Sum(nil)) - if sum != hr.sum { - hr.err = s3err.GetAPIError(s3err.ErrContentSHA256Mismatch) - return n, s3err.GetAPIError(s3err.ErrContentSHA256Mismatch) - } - } - } - return n, readerr -} - -func (hr *HashReader) Err() error { - return hr.err -} diff --git a/s3api/utils/utils.go b/s3api/utils/utils.go index 9358cf95..983437d7 100644 --- a/s3api/utils/utils.go +++ b/s3api/utils/utils.go @@ -50,12 +50,11 @@ func GetUserMetaData(headers *fasthttp.RequestHeader) (metadata map[string]strin return } -func CreateHttpRequestFromCtx(ctx *fiber.Ctx, signedHdrs []string) (*http.Request, error) { +func createHttpRequestFromCtx(ctx *fiber.Ctx, signedHdrs []string, contentLength int64) (*http.Request, error) { req := ctx.Request() var body io.Reader if IsBigDataAction(ctx) { body = req.BodyStream() - fmt.Println("create ctx body: ", body) } else { body = bytes.NewReader(req.Body()) } @@ -77,6 +76,8 @@ func CreateHttpRequestFromCtx(ctx *fiber.Ctx, signedHdrs []string) (*http.Reques // If content length is non 0, then the header will be included if !includeHeader("Content-Length", signedHdrs) { httpReq.ContentLength = 0 + } else { + httpReq.ContentLength = contentLength } // Set the Host header diff --git a/s3api/utils/utils_test.go b/s3api/utils/utils_test.go index 45e03311..a2c9ac01 100644 --- a/s3api/utils/utils_test.go +++ b/s3api/utils/utils_test.go @@ -55,7 +55,7 @@ func TestCreateHttpRequestFromCtx(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got, err := CreateHttpRequestFromCtx(tt.args.ctx, []string{"X-Amz-Mfa"}) + got, err := createHttpRequestFromCtx(tt.args.ctx, []string{"X-Amz-Mfa"}, 0) if (err != nil) != tt.wantErr { t.Errorf("CreateHttpRequestFromCtx() error = %v, wantErr %v", err, tt.wantErr) return