diff --git a/.gitignore b/.gitignore index 62c090bc..2b462d3f 100644 --- a/.gitignore +++ b/.gitignore @@ -62,4 +62,8 @@ tests/!s3cfg.local.default *.patch # grafana's local database (kept on filesystem for survival between instantiations) -metrics-exploration/grafana_data/** \ No newline at end of file +metrics-exploration/grafana_data/** + +# bats tools +/tests/bats-assert +/tests/bats-support \ No newline at end of file diff --git a/s3api/middlewares/url-decoder.go b/s3api/middlewares/url-decoder.go index 5918211f..e7e4c891 100644 --- a/s3api/middlewares/url-decoder.go +++ b/s3api/middlewares/url-decoder.go @@ -26,12 +26,11 @@ import ( func DecodeURL(logger s3log.AuditLogger, mm *metrics.Manager) fiber.Handler { return func(ctx *fiber.Ctx) error { - reqURL := ctx.Request().URI().String() - decoded, err := url.Parse(reqURL) + unescp, err := url.QueryUnescape(string(ctx.Request().URI().PathOriginal())) if err != nil { return controllers.SendResponse(ctx, s3err.GetAPIError(s3err.ErrInvalidURI), &controllers.MetaOpts{Logger: logger, MetricsMng: mm}) } - ctx.Path(decoded.Path) + ctx.Path(unescp) return ctx.Next() } } diff --git a/s3api/utils/utils.go b/s3api/utils/utils.go index 80553e96..c872f781 100644 --- a/s3api/utils/utils.go +++ b/s3api/utils/utils.go @@ -39,6 +39,10 @@ var ( bucketNameIpRegexp = regexp.MustCompile(`^(?:[0-9]{1,3}\.){3}[0-9]{1,3}$`) ) +const ( + upperhex = "0123456789ABCDEF" +) + func GetUserMetaData(headers *fasthttp.RequestHeader) (metadata map[string]string) { metadata = make(map[string]string) headers.DisableNormalizing() @@ -64,7 +68,9 @@ func createHttpRequestFromCtx(ctx *fiber.Ctx, signedHdrs []string, contentLength body = bytes.NewReader(req.Body()) } - httpReq, err := http.NewRequest(string(req.Header.Method()), string(ctx.Context().RequestURI()), body) + escapedURI := escapeOriginalURI(ctx) + + httpReq, err := http.NewRequest(string(req.Header.Method()), escapedURI, body) if err != nil { return nil, errors.New("error in creating an http request") } @@ -339,3 +345,74 @@ func IsValidOwnership(val types.ObjectOwnership) bool { return false } } + +func escapeOriginalURI(ctx *fiber.Ctx) string { + path := ctx.Path() + + // Escape the URI original path + escapedURI := escapePath(path) + + // Add the URI query params + query := string(ctx.Request().URI().QueryArgs().QueryString()) + if query != "" { + escapedURI = escapedURI + "?" + query + } + + return escapedURI +} + +// Escapes the path string +// Most of the parts copied from std url +func escapePath(s string) string { + hexCount := 0 + for i := 0; i < len(s); i++ { + c := s[i] + if shouldEscape(c) { + hexCount++ + } + } + + if hexCount == 0 { + return s + } + + var buf [64]byte + var t []byte + + required := len(s) + 2*hexCount + if required <= len(buf) { + t = buf[:required] + } else { + t = make([]byte, required) + } + + j := 0 + for i := 0; i < len(s); i++ { + switch c := s[i]; { + case shouldEscape(c): + t[j] = '%' + t[j+1] = upperhex[c>>4] + t[j+2] = upperhex[c&15] + j += 3 + default: + t[j] = s[i] + j++ + } + } + + return string(t) +} + +// Checks if the character needs to be escaped +func shouldEscape(c byte) bool { + if 'a' <= c && c <= 'z' || 'A' <= c && c <= 'Z' || '0' <= c && c <= '9' { + return false + } + + switch c { + case '-', '_', '.', '~', '/': + return false + } + + return true +} diff --git a/s3api/utils/utils_test.go b/s3api/utils/utils_test.go index a359e956..47c931ba 100644 --- a/s3api/utils/utils_test.go +++ b/s3api/utils/utils_test.go @@ -382,3 +382,125 @@ func TestIsValidOwnership(t *testing.T) { }) } } + +func Test_shouldEscape(t *testing.T) { + type args struct { + c byte + } + tests := []struct { + name string + args args + want bool + }{ + { + name: "shouldn't-escape-alphanum", + args: args{ + c: 'h', + }, + want: false, + }, + { + name: "shouldn't-escape-unreserved-char", + args: args{ + c: '_', + }, + want: false, + }, + { + name: "shouldn't-escape-unreserved-number", + args: args{ + c: '0', + }, + want: false, + }, + { + name: "shouldn't-escape-path-separator", + args: args{ + c: '/', + }, + want: false, + }, + { + name: "should-escape-special-char-1", + args: args{ + c: '&', + }, + want: true, + }, + { + name: "should-escape-special-char-2", + args: args{ + c: '*', + }, + want: true, + }, + { + name: "should-escape-special-char-3", + args: args{ + c: '(', + }, + want: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := shouldEscape(tt.args.c); got != tt.want { + t.Errorf("shouldEscape() = %v, want %v", got, tt.want) + } + }) + } +} + +func Test_escapePath(t *testing.T) { + type args struct { + s string + } + tests := []struct { + name string + args args + want string + }{ + { + name: "empty-string", + args: args{ + s: "", + }, + want: "", + }, + { + name: "alphanum-path", + args: args{ + s: "/test-bucket/test-key", + }, + want: "/test-bucket/test-key", + }, + { + name: "path-with-unescapable-chars", + args: args{ + s: "/test~bucket/test.key", + }, + want: "/test~bucket/test.key", + }, + { + name: "path-with-escapable-chars", + args: args{ + s: "/bucket-*(/test=key&", + }, + want: "/bucket-%2A%28/test%3Dkey%26", + }, + { + name: "path-with-space", + args: args{ + s: "/test-bucket/my key", + }, + want: "/test-bucket/my%20key", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := escapePath(tt.args.s); got != tt.want { + t.Errorf("escapePath() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/tests/integration/tests.go b/tests/integration/tests.go index d68d0c9a..bfe6e3b4 100644 --- a/tests/integration/tests.go +++ b/tests/integration/tests.go @@ -2692,10 +2692,29 @@ func PutObject_non_existing_bucket(s *S3Conf) error { func PutObject_special_chars(s *S3Conf) error { testName := "PutObject_special_chars" return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { - _, err := putObjects(s3client, []string{"my%key", "my^key", "my*key", "my.key", "my-key", "my_key", "my!key", "my'key", "my(key", "my)key", "my\\key", "my{}key", "my[]key", "my`key", "my+key", "my%25key", "my@key"}, bucket) + objs, err := putObjects(s3client, []string{ + "my!key", "my-key", "my_key", "my.key", "my'key", "my(key", "my)key", + "my&key", "my@key", "my=key", "my;key", "my:key", "my key", "my,key", + "my?key", "my\\key", "my^key", "my{}key", "my%key", "my`key", + "my[]key", "my~key", "my<>key", "my|key", "my#key", + }, bucket) if err != nil { return err } + + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + res, err := s3client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ + Bucket: &bucket, + }) + cancel() + if err != nil { + return err + } + + if !compareObjects(res.Contents, objs) { + return fmt.Errorf("expected the objects to be %v, instead got %v", objs, res.Contents) + } + return nil }) }