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