From 09031a30e5addc6c648f66e384b0ff95d142f047 Mon Sep 17 00:00:00 2001 From: niksis02 Date: Fri, 15 Aug 2025 03:41:01 +0400 Subject: [PATCH] feat: bucket cors implementation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes #1003 **Changes Introduced:** 1. **S3 Bucket CORS Actions** * Implemented the following S3 bucket CORS APIs: * `PutBucketCors` – Configure CORS rules for a bucket. * `GetBucketCors` – Retrieve the current CORS configuration for a bucket. * `DeleteBucketCors` – Remove CORS configuration from a bucket. 2. **CORS Preflight Handling** * Added an `OPTIONS` endpoint to handle browser preflight requests. * The endpoint evaluates incoming requests against bucket CORS rules and returns the appropriate `Access-Control-*` headers. 3. **CORS Middleware** * Implemented middleware that: * Checks if a bucket has CORS configured. * Detects the `Origin` header in the request. * Adds the necessary `Access-Control-*` headers to the response when the request matches the bucket CORS configuration. --- auth/bucket_cors.go | 306 ++++++++++ auth/bucket_cors_test.go | 610 +++++++++++++++++++ backend/azure/azure.go | 24 + backend/backend.go | 4 +- backend/posix/posix.go | 51 ++ backend/s3proxy/s3.go | 126 ++++ metrics/actions.go | 1 + s3api/controllers/backend_moq_test.go | 26 +- s3api/controllers/bucket-get.go | 11 +- s3api/controllers/bucket-get_test.go | 40 +- s3api/controllers/bucket-put.go | 57 +- s3api/controllers/bucket-put_test.go | 70 ++- s3api/controllers/options.go | 112 ++++ s3api/middlewares/apply-bucket-cors.go | 104 ++++ s3api/router.go | 55 ++ s3err/s3err.go | 56 ++ tests/integration/group-tests.go | 65 ++ tests/integration/tests.go | 796 ++++++++++++++++++++++++- tests/integration/utils.go | 166 +++++- 19 files changed, 2654 insertions(+), 26 deletions(-) create mode 100644 auth/bucket_cors.go create mode 100644 auth/bucket_cors_test.go create mode 100644 s3api/controllers/options.go create mode 100644 s3api/middlewares/apply-bucket-cors.go diff --git a/auth/bucket_cors.go b/auth/bucket_cors.go new file mode 100644 index 00000000..309b93ac --- /dev/null +++ b/auth/bucket_cors.go @@ -0,0 +1,306 @@ +// 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 auth + +import ( + "encoding/xml" + "fmt" + "net/http" + "regexp" + "strings" + + "github.com/versity/versitygw/s3api/debuglogger" + "github.com/versity/versitygw/s3err" +) + +// headerRegex is the regexp to validate http header names +var headerRegex = regexp.MustCompile(`^[!#$%&'*+\-.^_` + "`" + `|~0-9A-Za-z]+$`) + +type CORSHeader string +type CORSHTTPMethod string + +// IsValid validates the CORS http header +// the rules are based on http RFC +// https://datatracker.ietf.org/doc/html/rfc7230#section-3.2 +// +// Empty values are considered as valid +func (ch CORSHeader) IsValid() bool { + return ch == "" || headerRegex.MatchString(ch.String()) +} + +// String converts the header value to 'string' +func (ch CORSHeader) String() string { + return string(ch) +} + +// IsValid validates the cors http request method: +// the methods are case sensitive +func (cm CORSHTTPMethod) IsValid() bool { + return cm == http.MethodGet || cm == http.MethodHead || cm == http.MethodPut || + cm == http.MethodPost || cm == http.MethodDelete +} + +// String converts the method value to 'string' +func (cm CORSHTTPMethod) String() string { + return string(cm) +} + +type CORSConfiguration struct { + Rules []CORSRule `xml:"CORSRule"` +} + +// Validate validates the cors configuration rules +func (cc *CORSConfiguration) Validate() error { + if cc == nil || cc.Rules == nil { + debuglogger.Logf("invalid CORS configuration") + return s3err.GetAPIError(s3err.ErrMalformedXML) + } + + if len(cc.Rules) == 0 { + debuglogger.Logf("empty CORS config rules") + return s3err.GetAPIError(s3err.ErrMalformedXML) + } + + // validate each CORS rule + for _, rule := range cc.Rules { + if err := rule.Validate(); err != nil { + return err + } + } + + return nil +} + +type CORSAllowanceConfig struct { + Origin string + Methods string + ExposedHeaders string + AllowCredentials string + MaxAge *int32 +} + +// IsAllowed walks through the CORS rules and finds the first one allowing access. +// If no rule grants access, returns 'AccessForbidden' +func (cc *CORSConfiguration) IsAllowed(origin string, method CORSHTTPMethod, headers []CORSHeader) (*CORSAllowanceConfig, error) { + for _, rule := range cc.Rules { + // find the first rule granting access + if isAllowed, wilcardOrigin := rule.Match(origin, method, headers); isAllowed { + o := origin + allowCredentials := "true" + if wilcardOrigin { + o = "*" + allowCredentials = "false" + } + + return &CORSAllowanceConfig{ + Origin: o, + AllowCredentials: allowCredentials, + Methods: rule.GetAllowedMethods(), + ExposedHeaders: rule.GetExposeHeaders(), + MaxAge: rule.MaxAgeSeconds, + }, nil + } + } + + // if no matching rule is found, return AccessForbidden + return nil, s3err.GetAPIError(s3err.ErrCORSForbidden) +} + +type CORSRule struct { + AllowedMethods []CORSHTTPMethod `xml:"AllowedMethod"` + AllowedHeaders []CORSHeader `xml:"AllowedHeader"` + ExposeHeaders []CORSHeader `xml:"ExposeHeader"` + AllowedOrigins []string `xml:"AllowedOrigin"` + ID *string + MaxAgeSeconds *int32 +} + +// Validate validates and returns error if CORS configuration has invalid rule +func (cr *CORSRule) Validate() error { + // validate CORS allowed headers + for _, header := range cr.AllowedHeaders { + if !header.IsValid() { + debuglogger.Logf("invalid CORS allowed header: %s", header) + return s3err.GetInvalidCORSHeaderErr(header.String()) + } + } + // validate CORS allowed methods + for _, method := range cr.AllowedMethods { + if !method.IsValid() { + debuglogger.Logf("invalid CORS allowed method: %s", method) + return s3err.GetUnsopportedCORSMethodErr(method.String()) + } + } + // validate CORS expose headers + for _, header := range cr.ExposeHeaders { + if !header.IsValid() { + debuglogger.Logf("invalid CORS exposed header: %s", header) + return s3err.GetInvalidCORSHeaderErr(header.String()) + } + } + + return nil +} + +// Match matches the provided origin, method and headers with the +// CORS configuration rule +// if the matching origin is "*", it returns true as the first argument +func (cr *CORSRule) Match(origin string, method CORSHTTPMethod, headers []CORSHeader) (bool, bool) { + wildcardOrigin := false + originFound := false + + // check if the provided origin exists in CORS AllowedOrigins + for _, or := range cr.AllowedOrigins { + if wildcardMatch(or, origin) { + originFound = true + if or == "*" { + // mark wildcardOrigin as true, if "*" is found in AllowedOrigins + wildcardOrigin = true + } + break + } + } + + if !originFound { + return false, false + } + + // cache the CORS AllowedMethods in a map + allowedMethods := cacheCORSMethods(cr.AllowedMethods) + // check if the provided method exists in CORS AllowedMethods + if _, ok := allowedMethods[method]; !ok { + return false, false + } + + // check is CORS rule allowed headers match + // with the requested allowed headers + for _, reqHeader := range headers { + match := false + for _, header := range cr.AllowedHeaders { + if wildcardMatch(header.String(), reqHeader.String()) { + match = true + break + } + } + + if !match { + return false, false + } + } + + return true, wildcardOrigin +} + +// GetExposeHeaders returns comma separated CORS expose headers +func (cr *CORSRule) GetExposeHeaders() string { + var result strings.Builder + + for i, h := range cr.ExposeHeaders { + if i > 0 { + result.WriteString(", ") + } + result.WriteString(h.String()) + } + + return result.String() +} + +// GetAllowedMethods returns comma separated CORS allowed methods +func (cr *CORSRule) GetAllowedMethods() string { + var result strings.Builder + + for i, m := range cr.AllowedMethods { + if i > 0 { + result.WriteString(", ") + } + result.WriteString(m.String()) + } + + return result.String() +} + +// ParseCORSOutput parses raw bytes to 'CORSConfiguration' +func ParseCORSOutput(data []byte) (*CORSConfiguration, error) { + var config CORSConfiguration + err := xml.Unmarshal(data, &config) + if err != nil { + debuglogger.Logf("unmarshal cors output: %v", err) + return nil, fmt.Errorf("failed to parse cors config: %w", err) + } + + return &config, nil +} + +func cacheCORSMethods(input []CORSHTTPMethod) map[CORSHTTPMethod]struct{} { + result := make(map[CORSHTTPMethod]struct{}, len(input)) + for _, el := range input { + result[el] = struct{}{} + } + + return result +} + +// ParseCORSHeaders parses/validates Access-Control-Request-Headers +// and returns []CORSHeaders +func ParseCORSHeaders(headers string) ([]CORSHeader, error) { + result := []CORSHeader{} + if headers == "" { + return result, nil + } + + headersSplitted := strings.Split(headers, ",") + for _, h := range headersSplitted { + corsHeader := CORSHeader(strings.TrimSpace(h)) + if corsHeader == "" || !corsHeader.IsValid() { + debuglogger.Logf("invalid access control header: %s", h) + return nil, s3err.GetInvalidCORSRequestHeaderErr(h) + } + result = append(result, corsHeader) + } + + return result, nil +} + +func wildcardMatch(pattern, input string) bool { + pIdx, sIdx := 0, 0 + starIdx, matchIdx := -1, 0 + + for sIdx < len(input) { + if pIdx < len(pattern) && pattern[pIdx] == input[sIdx] { + // exact match of current char + sIdx++ + pIdx++ + } else if pIdx < len(pattern) && pattern[pIdx] == '*' { + // remember star position + starIdx = pIdx + matchIdx = sIdx + pIdx++ + } else if starIdx != -1 { + // backtrack: try to match more characters with '*' + pIdx = starIdx + 1 + matchIdx++ + sIdx = matchIdx + } else { + return false + } + } + + // skip trailing stars + for pIdx < len(pattern) && pattern[pIdx] == '*' { + pIdx++ + } + + return pIdx == len(pattern) +} diff --git a/auth/bucket_cors_test.go b/auth/bucket_cors_test.go new file mode 100644 index 00000000..5ff7c07c --- /dev/null +++ b/auth/bucket_cors_test.go @@ -0,0 +1,610 @@ +// 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 auth + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/versity/versitygw/s3err" +) + +func TestCORSHeader_IsValid(t *testing.T) { + tests := []struct { + name string + header CORSHeader + want bool + }{ + {"empty", "", true}, + {"valid", "X-Custom-Header", true}, + {"invalid_1", "Invalid Header", false}, + {"invalid_2", "invalid/header", false}, + {"invalid_3", "Invalid\tHeader", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.header.IsValid(); got != tt.want { + t.Errorf("IsValid() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestCORSHTTPMethod_IsValid(t *testing.T) { + tests := []struct { + name string + method CORSHTTPMethod + want bool + }{ + {"GET valid", http.MethodGet, true}, + {"HEAD valid", http.MethodHead, true}, + {"PUT valid", http.MethodPut, true}, + {"POST valid", http.MethodPost, true}, + {"DELETE valid", http.MethodDelete, true}, + {"get valid", "get", false}, + {"put valid", "put", false}, + {"post valid", "post", false}, + {"head valid", "head", false}, + {"invalid", "FOO", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.method.IsValid(); got != tt.want { + t.Errorf("IsValid() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestCORSConfiguration_Validate(t *testing.T) { + tests := []struct { + name string + cfg *CORSConfiguration + want error + }{ + {"nil config", nil, s3err.GetAPIError(s3err.ErrMalformedXML)}, + {"nil rules", &CORSConfiguration{}, s3err.GetAPIError(s3err.ErrMalformedXML)}, + {"empty rules", &CORSConfiguration{Rules: []CORSRule{}}, s3err.GetAPIError(s3err.ErrMalformedXML)}, + {"invalid rule", &CORSConfiguration{Rules: []CORSRule{{AllowedHeaders: []CORSHeader{"Invalid Header"}}}}, s3err.GetInvalidCORSHeaderErr("Invalid Header")}, + {"valid rule", &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"origin"}, + AllowedHeaders: []CORSHeader{"X-Test"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + ExposeHeaders: []CORSHeader{"X-Expose"}, + }}}, nil}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.cfg.Validate() + assert.EqualValues(t, tt.want, err) + }) + } +} + +func TestCORSConfiguration_IsAllowed(t *testing.T) { + type input struct { + cfg *CORSConfiguration + origin string + method CORSHTTPMethod + headers []CORSHeader + } + type output struct { + result *CORSAllowanceConfig + err error + } + tests := []struct { + name string + input input + output output + }{ + { + name: "allowed exact origin", + input: input{ + cfg: &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }}}, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{ + result: &CORSAllowanceConfig{ + Origin: "http://allowed.com", + AllowCredentials: "true", + Methods: http.MethodGet, + ExposedHeaders: "", + MaxAge: nil, + }, + err: nil, + }, + }, + { + name: "allowed wildcard origin", + input: input{ + cfg: &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"*"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }}}, + origin: "anything", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{ + result: &CORSAllowanceConfig{ + Origin: "*", + AllowCredentials: "false", + Methods: http.MethodGet, + ExposedHeaders: "", + MaxAge: nil, + }, + err: nil, + }, + }, + { + name: "forbidden no matching origin", + input: input{ + cfg: &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"http://nope.com"}, + }}}, + origin: "http://not-allowed.com", + method: http.MethodGet, + }, + output: output{ + result: nil, + err: s3err.GetAPIError(s3err.ErrCORSForbidden), + }, + }, + { + name: "forbidden method not allowed", + input: input{ + cfg: &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodPost}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }}}, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{ + result: nil, + err: s3err.GetAPIError(s3err.ErrCORSForbidden), + }, + }, + { + name: "forbidden header not allowed", + input: input{ + cfg: &CORSConfiguration{Rules: []CORSRule{{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }}}, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Nope"}, + }, + output: output{ + result: nil, + err: s3err.GetAPIError(s3err.ErrCORSForbidden), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.input.cfg.IsAllowed(tt.input.origin, tt.input.method, tt.input.headers) + assert.EqualValues(t, tt.output.err, err) + assert.EqualValues(t, tt.output.result, got) + }) + } +} + +func TestCORSRule_Validate(t *testing.T) { + tests := []struct { + name string + rule CORSRule + want error + }{ + { + name: "valid rule", + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + want: nil, + }, + { + name: "invalid allowed methods", + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{"invalid_method"}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + want: s3err.GetUnsopportedCORSMethodErr("invalid_method"), + }, + { + name: "invalid allowed header", + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"Invalid Header"}, + }, + want: s3err.GetInvalidCORSHeaderErr("Invalid Header"), + }, + { + name: "invalid allowed header", + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"Content-Length"}, + ExposeHeaders: []CORSHeader{"Content-Encoding", "invalid header"}, + }, + want: s3err.GetInvalidCORSHeaderErr("invalid header"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.rule.Validate() + assert.EqualValues(t, tt.want, err) + }) + } +} + +func TestCORSRule_Match(t *testing.T) { + type input struct { + rule CORSRule + origin string + method CORSHTTPMethod + headers []CORSHeader + } + type output struct { + isAllowed bool + isWildcard bool + } + tests := []struct { + name string + input input + output output + }{ + { + name: "exact origin and method match", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{isAllowed: true, isWildcard: false}, + }, + { + name: "wildcard origin match", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"*"}, + AllowedMethods: []CORSHTTPMethod{http.MethodPost}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://random.com", + method: http.MethodPost, + headers: []CORSHeader{"X-Test"}, + }, + output: output{isAllowed: true, isWildcard: true}, + }, + { + name: "wildcard containing origin match", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://random*"}, + AllowedMethods: []CORSHTTPMethod{http.MethodPost}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://random.com", + method: http.MethodPost, + headers: []CORSHeader{"X-Test"}, + }, + output: output{isAllowed: true, isWildcard: false}, + }, + { + name: "wildcard allowed headers match", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://something.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodPost}, + AllowedHeaders: []CORSHeader{"X-*"}, + }, + origin: "http://something.com", + method: http.MethodPost, + headers: []CORSHeader{"X-Test", "X-Something", "X-Anyting"}, + }, + output: output{isAllowed: true, isWildcard: false}, + }, + { + name: "origin mismatch", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://notallowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{isAllowed: false, isWildcard: false}, + }, + { + name: "method mismatch", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodPost}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Test"}, + }, + output: output{isAllowed: false, isWildcard: false}, + }, + { + name: "header mismatch", + input: input{ + rule: CORSRule{ + AllowedOrigins: []string{"http://allowed.com"}, + AllowedMethods: []CORSHTTPMethod{http.MethodGet}, + AllowedHeaders: []CORSHeader{"X-Test"}, + }, + origin: "http://allowed.com", + method: http.MethodGet, + headers: []CORSHeader{"X-Other"}, + }, + output: output{isAllowed: false, isWildcard: false}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + isAllowed, wild := tt.input.rule.Match(tt.input.origin, tt.input.method, tt.input.headers) + assert.Equal(t, tt.output.isAllowed, isAllowed) + assert.Equal(t, tt.output.isWildcard, wild) + }) + } +} + +func TestGetExposeHeaders(t *testing.T) { + tests := []struct { + name string + rule CORSRule + want string + }{ + {"multiple headers", CORSRule{ExposeHeaders: []CORSHeader{"Content-Length", "Content-Type", "Content-Encoding"}}, "Content-Length, Content-Type, Content-Encoding"}, + {"single header", CORSRule{ExposeHeaders: []CORSHeader{"Authorization"}}, "Authorization"}, + {"no headers", CORSRule{}, ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.rule.GetExposeHeaders() + assert.Equal(t, tt.want, got) + }) + } +} + +func TestGetAllowedMethods(t *testing.T) { + tests := []struct { + name string + rule CORSRule + want string + }{ + {"multiple methods", CORSRule{AllowedMethods: []CORSHTTPMethod{http.MethodGet, http.MethodPost, http.MethodPut}}, "GET, POST, PUT"}, + {"single method", CORSRule{AllowedMethods: []CORSHTTPMethod{http.MethodGet}}, "GET"}, + {"no methods", CORSRule{}, ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.rule.GetAllowedMethods() + assert.Equal(t, tt.want, got) + }) + } +} + +func TestParseCORSOutput(t *testing.T) { + tests := []struct { + name string + data string + want bool + }{ + {"valid", ``, true}, + {"invalid xml", ``, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg, err := ParseCORSOutput([]byte(tt.data)) + if (err == nil) != tt.want { + t.Errorf("ParseCORSOutput() err = %v, want success=%v", err, tt.want) + } + if tt.want && cfg == nil { + t.Errorf("Expected non-nil config") + } + }) + } +} + +func TestCacheCORSProps(t *testing.T) { + tests := []struct { + name string + in []CORSHTTPMethod + want map[string]struct{} + }{ + { + name: "empty CORSHTTPMethod slice", + in: []CORSHTTPMethod{}, + want: map[string]struct{}{}, + }, + { + name: "single CORSHTTPMethod", + in: []CORSHTTPMethod{http.MethodGet}, + want: map[string]struct{}{http.MethodGet: {}}, + }, + { + name: "multiple CORSHTTPMethods", + in: []CORSHTTPMethod{http.MethodGet, http.MethodPost, http.MethodPut}, + want: map[string]struct{}{ + http.MethodGet: {}, + http.MethodPost: {}, + http.MethodPut: {}, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := cacheCORSMethods(tt.in) + assert.Equal(t, len(tt.want), len(got)) + for key := range tt.want { + _, ok := got[CORSHTTPMethod(key)] + assert.True(t, ok) + } + }) + } +} + +func TestParseCORSHeaders(t *testing.T) { + tests := []struct { + name string + in string + want []CORSHeader + err error + }{ + { + name: "empty string", + in: "", + want: []CORSHeader{}, + err: nil, + }, + { + name: "single valid header", + in: "X-Test", + want: []CORSHeader{"X-Test"}, + err: nil, + }, + { + name: "multiple valid headers with spaces", + in: "X-Test, Content-Type, Authorization", + want: []CORSHeader{"X-Test", "Content-Type", "Authorization"}, + err: nil, + }, + { + name: "header with leading/trailing spaces", + in: " X-Test ", + want: []CORSHeader{"X-Test"}, + err: nil, + }, + { + name: "contains invalid header", + in: "X-Test, Invalid Header, Content-Type", + want: nil, + err: s3err.GetInvalidCORSRequestHeaderErr(" Invalid Header"), + }, + { + name: "only invalid header", + in: "Invalid Header", + want: nil, + err: s3err.GetInvalidCORSRequestHeaderErr("Invalid Header"), + }, + { + name: "multiple commas in a row", + in: "X-Test,,Content-Type", + want: nil, + err: s3err.GetInvalidCORSRequestHeaderErr(""), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := ParseCORSHeaders(tt.in) + assert.EqualValues(t, tt.err, err) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestWildcardMatch(t *testing.T) { + tests := []struct { + name string + pattern string + input string + want bool + }{ + // Exact match, no wildcards + {"exact match", "hello", "hello", true}, + {"exact mismatch", "hello", "hell", false}, + // Single '*' matching zero chars + {"star matches zero chars", "he*lo", "helo", true}, + // Single '*' matching multiple chars + {"star matches multiple chars", "he*o", "heyyyyyo", true}, + // '*' at start + {"star at start", "*world", "hello world", true}, + // '*' at end + {"star at end", "hello*", "hello there", true}, + // '*' matches whole string + {"only star", "*", "anything", true}, + {"only star empty", "*", "", true}, + // Multiple '*'s + {"multiple stars", "a*b*c", "axxxbzzzzyc", true}, + {"multiple stars no match", "a*b*c", "axxxbzzzzy", false}, + // Backtracking needed + {"backtracking required", "a*b*c", "ab123c", true}, + // No match with star present + {"star but mismatch", "he*world", "hey there", false}, + // Trailing stars in pattern + {"trailing stars match", "abc**", "abc", true}, + {"trailing stars match longer", "abc**", "abccc", true}, + // Empty pattern cases + {"empty pattern and empty input", "", "", true}, + {"empty pattern non-empty input", "", "a", false}, + {"only stars pattern with empty input", "***", "", true}, + // Pattern longer than input + {"pattern longer no star", "abcd", "abc", false}, + // Input longer but no star + {"input longer no star", "abc", "abcd", false}, + // Complex interleaved match + {"complex interleaved", "*a*b*cd*", "xxaYYbZZcd123", true}, + // Star match at the end after mismatch + {"mismatch then star match", "ab*xyz", "abzzzxyz", true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := wildcardMatch(tt.pattern, tt.input) + assert.Equal(t, tt.want, got) + }) + } +} diff --git a/backend/azure/azure.go b/backend/azure/azure.go index 7af49210..f153533f 100644 --- a/backend/azure/azure.go +++ b/backend/azure/azure.go @@ -60,6 +60,7 @@ const ( keyOwnership key = "Ownership" keyTags key = "Tags" keyPolicy key = "Policy" + keyCors key = "Cors" keyBucketLock key = "Bucketlock" keyObjRetention key = "Objectretention" keyObjLegalHold key = "Objectlegalhold" @@ -1506,6 +1507,29 @@ func (az *Azure) DeleteBucketPolicy(ctx context.Context, bucket string) error { return az.PutBucketPolicy(ctx, bucket, nil) } +func (az *Azure) PutBucketCors(ctx context.Context, bucket string, cors []byte) error { + if cors == nil { + return az.deleteContainerMetaData(ctx, bucket, string(keyCors)) + } + + return az.setContainerMetaData(ctx, bucket, string(keyCors), cors) +} + +func (az *Azure) GetBucketCors(ctx context.Context, bucket string) ([]byte, error) { + p, err := az.getContainerMetaData(ctx, bucket, string(keyCors)) + if err != nil { + return nil, err + } + if len(p) == 0 { + return nil, s3err.GetAPIError(s3err.ErrNoSuchCORSConfiguration) + } + return p, nil +} + +func (az *Azure) DeleteBucketCors(ctx context.Context, bucket string) error { + return az.PutBucketCors(ctx, bucket, nil) +} + func (az *Azure) PutObjectLockConfiguration(ctx context.Context, bucket string, config []byte) error { cfg, err := az.getContainerMetaData(ctx, bucket, string(keyBucketLock)) if err != nil { diff --git a/backend/backend.go b/backend/backend.go index 946ee735..b0ce7ff0 100644 --- a/backend/backend.go +++ b/backend/backend.go @@ -46,7 +46,7 @@ type Backend interface { PutBucketOwnershipControls(_ context.Context, bucket string, ownership types.ObjectOwnership) error GetBucketOwnershipControls(_ context.Context, bucket string) (types.ObjectOwnership, error) DeleteBucketOwnershipControls(_ context.Context, bucket string) error - PutBucketCors(context.Context, []byte) error + PutBucketCors(_ context.Context, bucket string, cors []byte) error GetBucketCors(_ context.Context, bucket string) ([]byte, error) DeleteBucketCors(_ context.Context, bucket string) error @@ -153,7 +153,7 @@ func (BackendUnsupported) GetBucketOwnershipControls(_ context.Context, bucket s func (BackendUnsupported) DeleteBucketOwnershipControls(_ context.Context, bucket string) error { return s3err.GetAPIError(s3err.ErrNotImplemented) } -func (BackendUnsupported) PutBucketCors(context.Context, []byte) error { +func (BackendUnsupported) PutBucketCors(context.Context, string, []byte) error { return s3err.GetAPIError(s3err.ErrNotImplemented) } func (BackendUnsupported) GetBucketCors(_ context.Context, bucket string) ([]byte, error) { diff --git a/backend/posix/posix.go b/backend/posix/posix.go index bfe73337..0659c6d5 100644 --- a/backend/posix/posix.go +++ b/backend/posix/posix.go @@ -103,6 +103,7 @@ const ( bucketLockKey = "bucket-lock" objectRetentionKey = "object-retention" objectLegalHoldKey = "object-legal-hold" + corskey = "cors" versioningKey = "versioning" deleteMarkerKey = "delete-marker" versionIdKey = "version-id" @@ -4659,6 +4660,56 @@ func (p *Posix) DeleteBucketPolicy(ctx context.Context, bucket string) error { return p.PutBucketPolicy(ctx, bucket, nil) } +func (p *Posix) PutBucketCors(_ context.Context, bucket string, cors []byte) error { + _, err := os.Stat(bucket) + if errors.Is(err, fs.ErrNotExist) { + return s3err.GetAPIError(s3err.ErrNoSuchBucket) + } + if err != nil { + return fmt.Errorf("stat bucket: %w", err) + } + + if cors == nil { + err = p.meta.DeleteAttribute(bucket, "", corskey) + if err != nil && !errors.Is(err, meta.ErrNoSuchKey) { + return fmt.Errorf("remove cors: %w", err) + } + + return nil + } + + err = p.meta.StoreAttribute(nil, bucket, "", corskey, cors) + if err != nil { + return fmt.Errorf("set cors: %w", err) + } + + return nil +} + +func (p *Posix) GetBucketCors(_ context.Context, bucket string) ([]byte, error) { + _, err := os.Stat(bucket) + if errors.Is(err, fs.ErrNotExist) { + return nil, s3err.GetAPIError(s3err.ErrNoSuchBucket) + } + if err != nil { + return nil, fmt.Errorf("stat bucket: %w", err) + } + + cors, err := p.meta.RetrieveAttribute(nil, bucket, "", corskey) + if errors.Is(err, meta.ErrNoSuchKey) { + return nil, s3err.GetAPIError(s3err.ErrNoSuchCORSConfiguration) + } + if err != nil { + return nil, err + } + + return cors, nil +} + +func (p *Posix) DeleteBucketCors(ctx context.Context, bucket string) error { + return p.PutBucketCors(ctx, bucket, nil) +} + func (p *Posix) isBucketObjectLockEnabled(bucket string) error { cfg, err := p.meta.RetrieveAttribute(nil, bucket, "", bucketLockKey) if errors.Is(err, fs.ErrNotExist) { diff --git a/backend/s3proxy/s3.go b/backend/s3proxy/s3.go index d29db3df..150b4a8e 100644 --- a/backend/s3proxy/s3.go +++ b/backend/s3proxy/s3.go @@ -17,6 +17,7 @@ package s3proxy import ( "bytes" "context" + "encoding/xml" "errors" "fmt" "io" @@ -1496,6 +1497,45 @@ func (s *S3Proxy) DeleteObjectTagging(ctx context.Context, bucket, object string return handleError(err) } +func (s *S3Proxy) PutBucketCors(ctx context.Context, bucket string, cors []byte) error { + cfg, err := auth.ParseCORSOutput(cors) + if err != nil { + return handleError(err) + } + + _, err = s.client.PutBucketCors(ctx, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: parseGatewayCORSToSDKConfig(cfg), + }) + + return handleError(err) +} + +func (s *S3Proxy) GetBucketCors(ctx context.Context, bucket string) ([]byte, error) { + resp, err := s.client.GetBucketCors(ctx, &s3.GetBucketCorsInput{ + Bucket: &bucket, + }) + if err != nil { + return nil, handleError(err) + } + + config := parseSdkCORSToGatewayConfig(resp.CORSRules) + data, err := xml.Marshal(config) + if err != nil { + return nil, handleError(err) + } + + return data, nil +} + +func (s *S3Proxy) DeleteBucketCors(ctx context.Context, bucket string) error { + _, err := s.client.DeleteBucketCors(ctx, &s3.DeleteBucketCorsInput{ + Bucket: &bucket, + }) + + return handleError(err) +} + func (s *S3Proxy) PutBucketPolicy(ctx context.Context, bucket string, policy []byte) error { return handleError(s.putMetaBucketObj(ctx, bucket, policy, metaPrefixPolicy)) } @@ -1741,3 +1781,89 @@ func convertObjectVersions(versions []types.ObjectVersion) []s3response.ObjectVe return result } + +func parseGatewayCORSToSDKConfig(config *auth.CORSConfiguration) *types.CORSConfiguration { + if config == nil { + return nil + } + + result := &types.CORSConfiguration{ + CORSRules: make([]types.CORSRule, 0, len(config.Rules)), + } + + for _, cfg := range config.Rules { + result.CORSRules = append(result.CORSRules, types.CORSRule{ + AllowedMethods: convertCORSMethodsToString(cfg.AllowedMethods), + AllowedHeaders: convertCORSHeadersToString(cfg.AllowedHeaders), + ExposeHeaders: convertCORSHeadersToString(cfg.ExposeHeaders), + AllowedOrigins: cfg.AllowedOrigins, + ID: cfg.ID, + MaxAgeSeconds: cfg.MaxAgeSeconds, + }) + } + + return result +} + +// convertCORSHeadersToString []auth.CORSHeader to []string +func convertCORSHeadersToString(headers []auth.CORSHeader) []string { + result := make([]string, 0, len(headers)) + for _, h := range headers { + result = append(result, h.String()) + } + + return result +} + +// convertCORSMethodsToString converts []auth.CORSHTTPMethod to []string +func convertCORSMethodsToString(methods []auth.CORSHTTPMethod) []string { + result := make([]string, 0, len(methods)) + for _, m := range methods { + result = append(result, m.String()) + } + + return result +} + +// convertCORSHeaders converts []string to []auth.CORSHeader +func convertCORSHeaders(headers []string) []auth.CORSHeader { + result := make([]auth.CORSHeader, 0, len(headers)) + for _, h := range headers { + result = append(result, auth.CORSHeader(h)) + } + + return result +} + +// convertCORSMethods converts []string to []auth.CORSHTTPMethod +func convertCORSMethods(methods []string) []auth.CORSHTTPMethod { + result := make([]auth.CORSHTTPMethod, 0, len(methods)) + for _, m := range methods { + result = append(result, auth.CORSHTTPMethod(m)) + } + + return result +} + +func parseSdkCORSToGatewayConfig(rules []types.CORSRule) *auth.CORSConfiguration { + if rules == nil { + return nil + } + + result := &auth.CORSConfiguration{ + Rules: make([]auth.CORSRule, 0, len(rules)), + } + + for _, cfg := range rules { + result.Rules = append(result.Rules, auth.CORSRule{ + AllowedMethods: convertCORSMethods(cfg.AllowedMethods), + AllowedHeaders: convertCORSHeaders(cfg.AllowedHeaders), + ExposeHeaders: convertCORSHeaders(cfg.ExposeHeaders), + AllowedOrigins: cfg.AllowedOrigins, + ID: cfg.ID, + MaxAgeSeconds: cfg.MaxAgeSeconds, + }) + } + + return result +} diff --git a/metrics/actions.go b/metrics/actions.go index a207e107..4e2f79d0 100644 --- a/metrics/actions.go +++ b/metrics/actions.go @@ -75,6 +75,7 @@ var ( ActionPutBucketCors = "s3_PutBucketCors" ActionGetBucketCors = "s3_GetBucketCors" ActionDeleteBucketCors = "s3_DeleteBucketCors" + ActionOptions = "s3_Options" ActionPutBucketAnalyticsConfiguration = "s3_PutBucketAnalyticsConfiguration" ActionGetBucketAnalyticsConfiguration = "s3_GetBucketAnalyticsConfiguration" ActionListBucketAnalyticsConfigurations = "s3_ListBucketAnalyticsConfigurations" diff --git a/s3api/controllers/backend_moq_test.go b/s3api/controllers/backend_moq_test.go index d0ef8d70..b30ded2e 100644 --- a/s3api/controllers/backend_moq_test.go +++ b/s3api/controllers/backend_moq_test.go @@ -134,7 +134,7 @@ var _ backend.Backend = &BackendMock{} // PutBucketAclFunc: func(contextMoqParam context.Context, bucket string, data []byte) error { // panic("mock out the PutBucketAcl method") // }, -// PutBucketCorsFunc: func(contextMoqParam context.Context, bytes []byte) error { +// PutBucketCorsFunc: func(contextMoqParam context.Context, bucket string, cors []byte) error { // panic("mock out the PutBucketCors method") // }, // PutBucketOwnershipControlsFunc: func(contextMoqParam context.Context, bucket string, ownership types.ObjectOwnership) error { @@ -304,7 +304,7 @@ type BackendMock struct { PutBucketAclFunc func(contextMoqParam context.Context, bucket string, data []byte) error // PutBucketCorsFunc mocks the PutBucketCors method. - PutBucketCorsFunc func(contextMoqParam context.Context, bytes []byte) error + PutBucketCorsFunc func(contextMoqParam context.Context, bucket string, cors []byte) error // PutBucketOwnershipControlsFunc mocks the PutBucketOwnershipControls method. PutBucketOwnershipControlsFunc func(contextMoqParam context.Context, bucket string, ownership types.ObjectOwnership) error @@ -635,8 +635,10 @@ type BackendMock struct { PutBucketCors []struct { // ContextMoqParam is the contextMoqParam argument value. ContextMoqParam context.Context - // Bytes is the bytes argument value. - Bytes []byte + // Bucket is the bucket argument value. + Bucket string + // Cors is the cors argument value. + Cors []byte } // PutBucketOwnershipControls holds details about calls to the PutBucketOwnershipControls method. PutBucketOwnershipControls []struct { @@ -2192,21 +2194,23 @@ func (mock *BackendMock) PutBucketAclCalls() []struct { } // PutBucketCors calls PutBucketCorsFunc. -func (mock *BackendMock) PutBucketCors(contextMoqParam context.Context, bytes []byte) error { +func (mock *BackendMock) PutBucketCors(contextMoqParam context.Context, bucket string, cors []byte) error { if mock.PutBucketCorsFunc == nil { panic("BackendMock.PutBucketCorsFunc: method is nil but Backend.PutBucketCors was just called") } callInfo := struct { ContextMoqParam context.Context - Bytes []byte + Bucket string + Cors []byte }{ ContextMoqParam: contextMoqParam, - Bytes: bytes, + Bucket: bucket, + Cors: cors, } mock.lockPutBucketCors.Lock() mock.calls.PutBucketCors = append(mock.calls.PutBucketCors, callInfo) mock.lockPutBucketCors.Unlock() - return mock.PutBucketCorsFunc(contextMoqParam, bytes) + return mock.PutBucketCorsFunc(contextMoqParam, bucket, cors) } // PutBucketCorsCalls gets all the calls that were made to PutBucketCors. @@ -2215,11 +2219,13 @@ func (mock *BackendMock) PutBucketCors(contextMoqParam context.Context, bytes [] // len(mockedBackend.PutBucketCorsCalls()) func (mock *BackendMock) PutBucketCorsCalls() []struct { ContextMoqParam context.Context - Bytes []byte + Bucket string + Cors []byte } { var calls []struct { ContextMoqParam context.Context - Bytes []byte + Bucket string + Cors []byte } mock.lockPutBucketCors.RLock() calls = mock.calls.PutBucketCors diff --git a/s3api/controllers/bucket-get.go b/s3api/controllers/bucket-get.go index f4da5c34..141b267e 100644 --- a/s3api/controllers/bucket-get.go +++ b/s3api/controllers/bucket-get.go @@ -187,8 +187,17 @@ func (c S3ApiController) GetBucketCors(ctx *fiber.Ctx) (*Response, error) { } data, err := c.be.GetBucketCors(ctx.Context(), bucket) + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + output, err := auth.ParseCORSOutput(data) return &Response{ - Data: data, + Data: output, MetaOpts: &MetaOptions{ BucketOwner: parsedAcl.Owner, }, diff --git a/s3api/controllers/bucket-get_test.go b/s3api/controllers/bucket-get_test.go index 1787863d..73354b9d 100644 --- a/s3api/controllers/bucket-get_test.go +++ b/s3api/controllers/bucket-get_test.go @@ -17,7 +17,10 @@ package controllers import ( "context" "encoding/json" + "encoding/xml" + "errors" "fmt" + "net/http" "testing" "github.com/aws/aws-sdk-go-v2/service/s3" @@ -313,6 +316,20 @@ func TestS3ApiController_GetBucketVersioning(t *testing.T) { } func TestS3ApiController_GetBucketCors(t *testing.T) { + cors := &auth.CORSConfiguration{ + Rules: []auth.CORSRule{ + { + AllowedOrigins: []string{"origin"}, + AllowedMethods: []auth.CORSHTTPMethod{http.MethodPut}, + AllowedHeaders: []auth.CORSHeader{"X-Amz-Date"}, + }, + }, + } + beRes, err := xml.Marshal(cors) + assert.NoError(t, err) + + var nilResp *auth.CORSConfiguration + tests := []struct { name string input testInput @@ -341,7 +358,6 @@ func TestS3ApiController_GetBucketCors(t *testing.T) { }, output: testOutput{ response: &Response{ - Data: []byte{}, MetaOpts: &MetaOptions{ BucketOwner: "root", }, @@ -350,14 +366,30 @@ func TestS3ApiController_GetBucketCors(t *testing.T) { }, }, { - name: "successful response", + name: "invalid data from backend", input: testInput{ locals: defaultLocals, - beRes: []byte("mock_cors_resp"), + beRes: []byte("invalid_data"), }, output: testOutput{ response: &Response{ - Data: []byte("mock_cors_resp"), + Data: nilResp, + MetaOpts: &MetaOptions{ + BucketOwner: "root", + }, + }, + err: errors.New("failed to parse cors config:"), + }, + }, + { + name: "successful response", + input: testInput{ + locals: defaultLocals, + beRes: beRes, + }, + output: testOutput{ + response: &Response{ + Data: cors, MetaOpts: &MetaOptions{ BucketOwner: "root", }, diff --git a/s3api/controllers/bucket-put.go b/s3api/controllers/bucket-put.go index 7fbd0ea3..22909937 100644 --- a/s3api/controllers/bucket-put.go +++ b/s3api/controllers/bucket-put.go @@ -15,6 +15,7 @@ package controllers import ( + "bytes" "encoding/xml" "errors" "fmt" @@ -247,7 +248,61 @@ func (c S3ApiController) PutBucketCors(ctx *fiber.Ctx) (*Response, error) { }, err } - err = c.be.PutBucketCors(ctx.Context(), []byte{}) + body := ctx.Body() + + var corsConfig auth.CORSConfiguration + err = xml.Unmarshal(body, &corsConfig) + if err != nil { + debuglogger.Logf("invalid CORS request body: %v", err) + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, s3err.GetAPIError(s3err.ErrMalformedXML) + } + + // validate the CORS configuration rules + err = corsConfig.Validate() + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + algo, checksusms, err := utils.ParseChecksumHeadersAndSdkAlgo(ctx) + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + if algo != "" { + rdr, err := utils.NewHashReader(bytes.NewReader(body), checksusms[algo], utils.HashType(strings.ToLower(string(algo)))) + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + // Pass the same body to avoid data duplication + _, err = rdr.Read(body) + if err != nil { + debuglogger.Logf("failed to read hash calculation data: %v", err) + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + } + + err = c.be.PutBucketCors(ctx.Context(), bucket, body) return &Response{ MetaOpts: &MetaOptions{ BucketOwner: parsedAcl.Owner, diff --git a/s3api/controllers/bucket-put_test.go b/s3api/controllers/bucket-put_test.go index 239addc7..9555bd2d 100644 --- a/s3api/controllers/bucket-put_test.go +++ b/s3api/controllers/bucket-put_test.go @@ -463,6 +463,26 @@ func TestS3ApiController_PutObjectLockConfiguration(t *testing.T) { } func TestS3ApiController_PutBucketCors(t *testing.T) { + validBody, err := xml.Marshal(auth.CORSConfiguration{ + Rules: []auth.CORSRule{ + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []auth.CORSHTTPMethod{http.MethodPost}, + }, + }, + }) + assert.NoError(t, err) + + invalidCors, err := xml.Marshal(auth.CORSConfiguration{ + Rules: []auth.CORSRule{ + { + AllowedOrigins: []string{"origin"}, + AllowedMethods: []auth.CORSHTTPMethod{"invalid_method"}, + }, + }, + }) + assert.NoError(t, err) + tests := []struct { name string input testInput @@ -482,11 +502,54 @@ func TestS3ApiController_PutBucketCors(t *testing.T) { err: s3err.GetAPIError(s3err.ErrAccessDenied), }, }, + { + name: "invalid request body", + input: testInput{ + locals: defaultLocals, + body: []byte("invalid_body"), + }, + output: testOutput{ + response: &Response{ + MetaOpts: &MetaOptions{BucketOwner: "root"}, + }, + err: s3err.GetAPIError(s3err.ErrMalformedXML), + }, + }, + { + name: "invalid cors config", + input: testInput{ + locals: defaultLocals, + body: invalidCors, + }, + output: testOutput{ + response: &Response{ + MetaOpts: &MetaOptions{BucketOwner: "root"}, + }, + err: s3err.GetUnsopportedCORSMethodErr("invalid_method"), + }, + }, + { + name: "invalid checksum algo", + input: testInput{ + locals: defaultLocals, + body: validBody, + headers: map[string]string{ + "X-Amz-Sdk-Checksum-Algorithm": "invalid_algo", + }, + }, + output: testOutput{ + response: &Response{ + MetaOpts: &MetaOptions{BucketOwner: "root"}, + }, + err: s3err.GetAPIError(s3err.ErrInvalidChecksumAlgorithm), + }, + }, { name: "backend error", input: testInput{ locals: defaultLocals, beErr: s3err.GetAPIError(s3err.ErrNotImplemented), + body: validBody, }, output: testOutput{ response: &Response{ @@ -499,6 +562,7 @@ func TestS3ApiController_PutBucketCors(t *testing.T) { name: "success", input: testInput{ locals: defaultLocals, + body: validBody, }, output: testOutput{ response: &Response{ @@ -513,7 +577,7 @@ func TestS3ApiController_PutBucketCors(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { be := &BackendMock{ - PutBucketCorsFunc: func(contextMoqParam context.Context, bytes []byte) error { + PutBucketCorsFunc: func(contextMoqParam context.Context, bucket string, cors []byte) error { return tt.input.beErr }, GetBucketPolicyFunc: func(contextMoqParam context.Context, bucket string) ([]byte, error) { @@ -526,7 +590,9 @@ func TestS3ApiController_PutBucketCors(t *testing.T) { } testController(t, ctrl.PutBucketCors, tt.output.response, tt.output.err, ctxInputs{ - locals: tt.input.locals, + locals: tt.input.locals, + body: tt.input.body, + headers: tt.input.headers, }) }) } diff --git a/s3api/controllers/options.go b/s3api/controllers/options.go new file mode 100644 index 00000000..c15c1954 --- /dev/null +++ b/s3api/controllers/options.go @@ -0,0 +1,112 @@ +// 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 controllers + +import ( + "errors" + + "github.com/gofiber/fiber/v2" + "github.com/versity/versitygw/auth" + "github.com/versity/versitygw/s3api/debuglogger" + "github.com/versity/versitygw/s3api/middlewares" + "github.com/versity/versitygw/s3api/utils" + "github.com/versity/versitygw/s3err" +) + +func (s S3ApiController) CORSOptions(ctx *fiber.Ctx) (*Response, error) { + bucket := ctx.Params("bucket") + parsedAcl := utils.ContextKeyParsedAcl.Get(ctx).(auth.ACL) + // get headers + origin := ctx.Get("Origin") + method := auth.CORSHTTPMethod(ctx.Get("Access-Control-Request-Method")) + headers := ctx.Get("Access-Control-Request-Headers") + + // Origin is required + if origin == "" { + debuglogger.Logf("origin is missing: %v", origin) + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, s3err.GetAPIError(s3err.ErrMissingCORSOrigin) + } + + // check if allowed method is valid + if !method.IsValid() { + debuglogger.Logf("invalid cors method: %s", method) + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, s3err.GetInvalidCORSMethodErr(method.String()) + } + + // parse and validate headers + parsedHeaders, err := auth.ParseCORSHeaders(headers) + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + cors, err := s.be.GetBucketCors(ctx.Context(), bucket) + if err != nil { + debuglogger.Logf("failed to get bucket cors: %v", err) + if errors.Is(err, s3err.GetAPIError(s3err.ErrNoSuchCORSConfiguration)) { + err = s3err.GetAPIError(s3err.ErrCORSIsNotEnabled) + debuglogger.Logf("bucket cors is not set: %v", err) + } + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + corsConfig, err := auth.ParseCORSOutput(cors) + if err != nil { + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + allowConfig, err := corsConfig.IsAllowed(origin, method, parsedHeaders) + if err != nil { + debuglogger.Logf("cors access forbidden: %v", err) + return &Response{ + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, err + } + + return &Response{ + Headers: map[string]*string{ + "Access-Control-Allow-Origin": &allowConfig.Origin, + "Access-Control-Allow-Methods": &allowConfig.Methods, + "Access-Control-Expose-Headers": &allowConfig.ExposedHeaders, + "Access-Control-Allow-Credentials": &allowConfig.AllowCredentials, + "Access-Control-Max-Age": utils.ConvertPtrToStringPtr(allowConfig.MaxAge), + "Vary": &middlewares.VaryHdr, + }, + MetaOpts: &MetaOptions{ + BucketOwner: parsedAcl.Owner, + }, + }, nil +} diff --git a/s3api/middlewares/apply-bucket-cors.go b/s3api/middlewares/apply-bucket-cors.go new file mode 100644 index 00000000..71eb46fc --- /dev/null +++ b/s3api/middlewares/apply-bucket-cors.go @@ -0,0 +1,104 @@ +// 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 ( + "fmt" + + "github.com/gofiber/fiber/v2" + "github.com/versity/versitygw/auth" + "github.com/versity/versitygw/backend" + "github.com/versity/versitygw/s3api/debuglogger" + "github.com/versity/versitygw/s3err" +) + +// Vary http response header is always the same below +var VaryHdr = "Origin, Access-Control-Request-Headers, Access-Control-Request-Method" + +// ApplyBucketCORS retreives the bucket CORS configuration, +// checks if origin and method meets the cors rules and +// adds the necessary response headers. +// CORS check is applied only when 'Origin' request header is present +func ApplyBucketCORS(be backend.Backend) fiber.Handler { + return func(ctx *fiber.Ctx) error { + bucket := ctx.Params("bucket") + origin := ctx.Get("Origin") + // if the origin request header is empty, skip cors validation + if origin == "" { + return nil + } + + // if bucket cors is not set, skip the check + data, err := be.GetBucketCors(ctx.Context(), bucket) + if err != nil { + // If CORS is not configured, S3Error will have code NoSuchCORSConfiguration. + // In this case, we can safely continue. For any other error, we should log it. + s3Err, ok := err.(s3err.APIError) + if !ok || s3Err.Code != "NoSuchCORSConfiguration" { + debuglogger.Logf("failed to get bucket cors for bucket %q: %v", bucket, err) + } + return nil + } + + cors, err := auth.ParseCORSOutput(data) + if err != nil { + return nil + } + + method := auth.CORSHTTPMethod(ctx.Get("Access-Control-Request-Method")) + headers := ctx.Get("Access-Control-Request-Headers") + + // if request method is not specified with Access-Control-Request-Method + // override it with the actual request method + if method == "" { + method = auth.CORSHTTPMethod(ctx.Request().Header.Method()) + } else if !method.IsValid() { + // check if allowed method is valid + debuglogger.Logf("invalid cors method: %s", method) + return s3err.GetInvalidCORSMethodErr(method.String()) + } + + // parse and validate headers + parsedHeaders, err := auth.ParseCORSHeaders(headers) + if err != nil { + return err + } + + allowConfig, err := cors.IsAllowed(origin, method, parsedHeaders) + if err != nil { + // if bucket cors rules doesn't grant access, skip + // and don't add any response headers + return nil + } + + if allowConfig.MaxAge != nil { + ctx.Response().Header.Add("Access-Control-Max-Age", fmt.Sprint(*allowConfig.MaxAge)) + } + + for key, val := range map[string]string{ + "Access-Control-Allow-Origin": allowConfig.Origin, + "Access-Control-Allow-Methods": allowConfig.Methods, + "Access-Control-Expose-Headers": allowConfig.ExposedHeaders, + "Access-Control-Allow-Credentials": allowConfig.AllowCredentials, + "Vary": VaryHdr, + } { + if val != "" { + ctx.Response().Header.Add(key, val) + } + } + + return nil + } +} diff --git a/s3api/router.go b/s3api/router.go index 28be6af9..060ff2d4 100644 --- a/s3api/router.go +++ b/s3api/router.go @@ -116,6 +116,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), middlewares.ParseAcl(be), + middlewares.ApplyBucketCORS(be), )) bucketRouter.Put("", middlewares.MatchQueryArgs("ownershipControls"), @@ -128,6 +129,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -141,6 +143,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -154,6 +157,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -167,6 +171,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -180,6 +185,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -193,6 +199,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Put("", @@ -261,6 +268,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), )) // HeadBucket action @@ -269,11 +277,13 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ ctrl.HeadBucket, metrics.ActionHeadBucket, services, + middlewares.ApplyBucketCORS(be), middlewares.BucketObjectNameValidator(), middlewares.AuthorizePublicBucketAccess(be, metrics.ActionHeadBucket, auth.ListBucketAction, auth.PermissionRead), middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -289,6 +299,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Delete("", @@ -302,6 +313,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Delete("", @@ -315,6 +327,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Delete("", @@ -328,6 +341,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Delete("", @@ -396,6 +410,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -411,6 +426,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -424,6 +440,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -437,6 +454,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -450,6 +468,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -463,6 +482,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -476,6 +496,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -489,6 +510,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -502,6 +524,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -515,6 +538,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -626,6 +650,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) bucketRouter.Get("", @@ -638,6 +663,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -653,6 +679,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -667,6 +694,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -682,6 +710,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -695,6 +724,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -708,6 +738,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -721,6 +752,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -734,6 +766,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -747,6 +780,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Get("", @@ -759,6 +793,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -774,6 +809,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Delete("", @@ -787,6 +823,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Delete("", @@ -799,6 +836,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -813,6 +851,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Post("", @@ -827,6 +866,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Post("", @@ -840,6 +880,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Post("", @@ -853,6 +894,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -868,6 +910,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -881,6 +924,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -894,6 +938,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -907,6 +952,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -921,6 +967,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -934,6 +981,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) @@ -959,6 +1007,7 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) objectRouter.Put("", @@ -971,9 +1020,15 @@ func (sa *S3ApiRouter) Init(app *fiber.App, be backend.Backend, iam auth.IAMServ middlewares.VerifyPresignedV4Signature(root, iam, region, debug), middlewares.VerifyV4Signature(root, iam, region, debug), middlewares.VerifyMD5Body(), + middlewares.ApplyBucketCORS(be), middlewares.ParseAcl(be), )) + app.Options("/:bucket/*", controllers.ProcessHandlers(ctrl.CORSOptions, metrics.ActionOptions, services, + middlewares.BucketObjectNameValidator(), + middlewares.ParseAcl(be), + )) + // Return MethodNotAllowed for all the unmatched routes app.All("*", controllers.ProcessHandlers(ctrl.HandleErrorRoute(s3err.GetAPIError(s3err.ErrMethodNotAllowed)), metrics.ActionUndetected, services)) } diff --git a/s3err/s3err.go b/s3err/s3err.go index 8c7463af..ce249fa4 100644 --- a/s3err/s3err.go +++ b/s3err/s3err.go @@ -172,6 +172,10 @@ const ( ErrTrailerHeaderNotSupported ErrBadRequest ErrMissingUploadId + ErrNoSuchCORSConfiguration + ErrCORSForbidden + ErrMissingCORSOrigin + ErrCORSIsNotEnabled // Non-AWS errors ErrExistingObjectIsDirectory @@ -756,6 +760,26 @@ var errorCodeResponse = map[ErrorCode]APIError{ Description: "This operation does not accept partNumber without uploadId", HTTPStatusCode: http.StatusBadRequest, }, + ErrNoSuchCORSConfiguration: { + Code: "NoSuchCORSConfiguration", + Description: "The CORS configuration does not exist", + HTTPStatusCode: http.StatusNotFound, + }, + ErrCORSForbidden: { + Code: "AccessForbidden", + Description: "CORSResponse: This CORS request is not allowed. This is usually because the evalution of Origin, request method / Access-Control-Request-Method or Access-Control-Request-Headers are not whitelisted by the resource's CORS spec.", + HTTPStatusCode: http.StatusForbidden, + }, + ErrMissingCORSOrigin: { + Code: "BadRequest", + Description: "Insufficient information. Origin request header needed.", + HTTPStatusCode: http.StatusBadRequest, + }, + ErrCORSIsNotEnabled: { + Code: "AccessForbidden", + Description: "CORSResponse: CORS is not enabled for this bucket.", + HTTPStatusCode: http.StatusForbidden, + }, // non aws errors ErrExistingObjectIsDirectory: { @@ -935,3 +959,35 @@ func CreateExceedingRangeErr(objSize int64) APIError { HTTPStatusCode: http.StatusBadRequest, } } + +func GetInvalidCORSHeaderErr(header string) APIError { + return APIError{ + Code: "InvalidRequest", + Description: fmt.Sprintf(`AllowedHeader "%s" contains invalid character.`, header), + HTTPStatusCode: http.StatusBadRequest, + } +} + +func GetInvalidCORSRequestHeaderErr(header string) APIError { + return APIError{ + Code: "BadRequest", + Description: fmt.Sprintf(`Access-Control-Request-Headers "%s" contains invalid character.`, header), + HTTPStatusCode: http.StatusBadRequest, + } +} + +func GetUnsopportedCORSMethodErr(method string) APIError { + return APIError{ + Code: "InvalidRequest", + Description: fmt.Sprintf("Found unsupported HTTP method in CORS config. Unsupported method is %s", method), + HTTPStatusCode: http.StatusBadRequest, + } +} + +func GetInvalidCORSMethodErr(method string) APIError { + return APIError{ + Code: "BadRequest", + Description: fmt.Sprintf("Invalid Access-Control-Request-Method: %s", method), + HTTPStatusCode: http.StatusBadRequest, + } +} diff --git a/tests/integration/group-tests.go b/tests/integration/group-tests.go index 19fbd85f..ca34b148 100644 --- a/tests/integration/group-tests.go +++ b/tests/integration/group-tests.go @@ -527,6 +527,44 @@ func TestDeleteBucketPolicy(s *S3Conf) { DeleteBucketPolicy_success(s) } +func TestPutBucketCors(s *S3Conf) { + PutBucketCors_non_existing_bucket(s) + PutBucketCors_empty_cors_rules(s) + PutBucketCors_invalid_method(s) + PutBucketCors_invalid_header(s) + PutBucketCors_invalid_content_md5(s) + PutBucketCors_incorrect_content_md5(s) + PutBucketCors_success(s) +} + +func TestGetBucketCors(s *S3Conf) { + GetBucketCors_non_existing_bucket(s) + GetBucketCors_no_such_bucket_cors(s) + GetBucketCors_success(s) +} + +func TestDeleteBucketCors(s *S3Conf) { + DeleteBucketCors_non_existing_bucket(s) + DeleteBucketCors_success(s) +} + +func TestPreflightOPTIONSEndpoint(s *S3Conf) { + PreflightOPTIONS_non_existing_bucket(s) + PreflightOPTIONS_missing_origin(s) + PreflightOPTIONS_invalid_request_method(s) + PreflightOPTIONS_invalid_request_headers(s) + PreflightOPTIONS_unset_bucket_cors(s) + PreflightOPTIONS_access_forbidden(s) + PreflightOPTIONS_access_granted(s) +} + +func TestCORSMiddleware(s *S3Conf) { + CORSMiddleware_invalid_method(s) + CORSMiddleware_invalid_headers(s) + CORSMiddleware_access_forbidden(s) + CORSMiddleware_access_granted(s) +} + func TestPutObjectLockConfiguration(s *S3Conf) { PutObjectLockConfiguration_non_existing_bucket(s) PutObjectLockConfiguration_empty_config(s) @@ -660,6 +698,10 @@ func TestFullFlow(s *S3Conf) { TestPutBucketPolicy(s) TestGetBucketPolicy(s) TestDeleteBucketPolicy(s) + TestPutBucketCors(s) + TestGetBucketCors(s) + TestDeleteBucketCors(s) + TestPreflightOPTIONSEndpoint(s) TestPutObjectLockConfiguration(s) TestGetObjectLockConfiguration(s) TestPutObjectRetention(s) @@ -1272,6 +1314,29 @@ func GetIntTests() IntTests { "DeleteBucketPolicy_non_existing_bucket": DeleteBucketPolicy_non_existing_bucket, "DeleteBucketPolicy_remove_before_setting": DeleteBucketPolicy_remove_before_setting, "DeleteBucketPolicy_success": DeleteBucketPolicy_success, + "PutBucketCors_non_existing_bucket": PutBucketCors_non_existing_bucket, + "PutBucketCors_empty_cors_rules": PutBucketCors_empty_cors_rules, + "PutBucketCors_invalid_method": PutBucketCors_invalid_method, + "PutBucketCors_invalid_header": PutBucketCors_invalid_header, + "PutBucketCors_invalid_content_md5": PutBucketCors_invalid_content_md5, + "PutBucketCors_incorrect_content_md5": PutBucketCors_incorrect_content_md5, + "GetBucketCors_non_existing_bucket": GetBucketCors_non_existing_bucket, + "GetBucketCors_no_such_bucket_cors": GetBucketCors_no_such_bucket_cors, + "GetBucketCors_success": GetBucketCors_success, + "DeleteBucketCors_non_existing_bucket": DeleteBucketCors_non_existing_bucket, + "DeleteBucketCors_success": DeleteBucketCors_success, + "PutBucketCors_success": PutBucketCors_success, + "PreflightOPTIONS_non_existing_bucket": PreflightOPTIONS_non_existing_bucket, + "PreflightOPTIONS_missing_origin": PreflightOPTIONS_missing_origin, + "PreflightOPTIONS_invalid_request_method": PreflightOPTIONS_invalid_request_method, + "PreflightOPTIONS_invalid_request_headers": PreflightOPTIONS_invalid_request_headers, + "PreflightOPTIONS_unset_bucket_cors": PreflightOPTIONS_unset_bucket_cors, + "PreflightOPTIONS_access_forbidden": PreflightOPTIONS_access_forbidden, + "PreflightOPTIONS_access_granted": PreflightOPTIONS_access_granted, + "CORSMiddleware_invalid_method": CORSMiddleware_invalid_method, + "CORSMiddleware_invalid_headers": CORSMiddleware_invalid_headers, + "CORSMiddleware_access_forbidden": CORSMiddleware_access_forbidden, + "CORSMiddleware_access_granted": CORSMiddleware_access_granted, "PutObjectLockConfiguration_non_existing_bucket": PutObjectLockConfiguration_non_existing_bucket, "PutObjectLockConfiguration_empty_config": PutObjectLockConfiguration_empty_config, "PutObjectLockConfiguration_not_enabled_on_bucket_creation": PutObjectLockConfiguration_not_enabled_on_bucket_creation, diff --git a/tests/integration/tests.go b/tests/integration/tests.go index c70b3b13..40c420d1 100644 --- a/tests/integration/tests.go +++ b/tests/integration/tests.go @@ -13512,6 +13512,794 @@ func DeleteBucketPolicy_success(s *S3Conf) error { }) } +// Bucket CORS tests +func PutBucketCors_non_existing_bucket(s *S3Conf) error { + testName := "PutBucketCors_non_existing_bucket" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: getPtr("non-existing-bucket"), + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com"}, + AllowedMethods: []string{http.MethodGet}, + }, + }, + }, + }) + return checkApiErr(err, s3err.GetAPIError(s3err.ErrNoSuchBucket)) + }) +} + +func PutBucketCors_empty_cors_rules(s *S3Conf) error { + testName := "PutBucketCors_empty_cors_rules" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{}, + }, + }) + return checkApiErr(err, s3err.GetAPIError(s3err.ErrMalformedXML)) + }) +} + +func PutBucketCors_invalid_method(s *S3Conf) error { + testName := "PutBucketCors_invalid_method" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + for _, test := range []struct { + invalidMethod string + allowedMethods []string + }{ + {"get", []string{"get"}}, + {"put", []string{"put"}}, + {"post", []string{"post"}}, + {"head", []string{"head"}}, + {"delete", []string{"delete"}}, + {http.MethodPatch, []string{http.MethodGet, http.MethodPatch}}, + {http.MethodOptions, []string{http.MethodPost, http.MethodOptions}}, + {"invalid_method", []string{http.MethodGet, http.MethodHead, http.MethodPost, http.MethodPut, http.MethodDelete, "invalid_method"}}, + } { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com"}, + AllowedMethods: test.allowedMethods, + AllowedHeaders: []string{"X-Amz-Date"}, + ExposeHeaders: []string{"Authorization"}, + }, + }, + }, + }) + + if err := checkApiErr(err, s3err.GetUnsopportedCORSMethodErr(test.invalidMethod)); err != nil { + return err + } + } + + return nil + }) +} + +func PutBucketCors_invalid_header(s *S3Conf) error { + testName := "PutBucketCors_invalid_header" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + for _, test := range []struct { + invalidHeader string + headers []string + }{ + {"invalid header", []string{"X-Amz-Date", "X-Amz-Content-Sha256", "invalid header"}}, + {"X-Custom:Header", []string{"Authorization", "X-Custom:Header"}}, + {"X(Custom)", []string{"Content-Length", "X(Custom)"}}, + {"Bad/Header", []string{"Content-Encoding", "Bad/Header"}}, + {"X[Key]", []string{"Date", "X[Key]"}}, + {"Bad=Name", []string{"X-Amz-Custome-Header", "Bad=Name"}}, + {`X"Quote"`, []string{`X"Quote"`}}, + } { + // first check for allowed headers + cfg := &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com"}, + AllowedMethods: []string{http.MethodPost}, + AllowedHeaders: test.headers, + ExposeHeaders: []string{"Authorization"}, + }, + }, + } + + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: cfg, + }) + if err := checkApiErr(err, s3err.GetInvalidCORSHeaderErr(test.invalidHeader)); err != nil { + return err + } + + // second check for expose headers + cfg.CORSRules[0].AllowedHeaders = []string{"X-Amz-Date"} // set to any valid header + cfg.CORSRules[0].ExposeHeaders = test.headers + + err = putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: cfg, + }) + if err := checkApiErr(err, s3err.GetInvalidCORSHeaderErr(test.invalidHeader)); err != nil { + return err + } + } + + return nil + }) +} + +// TODO: report a bug for this case +func PutBucketCors_invalid_content_md5(s *S3Conf) error { + testName := "PutBucketCors_invalid_content_md5" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + return nil + }) +} + +func PutBucketCors_incorrect_content_md5(s *S3Conf) error { + testName := "PutBucketCors_incorrect_content_md5" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + return nil + }) +} + +func PutBucketCors_success(s *S3Conf) error { + testName := "PutBucketCors_success" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + maxAgePositive, maxAgeNegative := int32(3000), int32(-100) + return putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com"}, + AllowedMethods: []string{http.MethodPost, http.MethodPut}, + AllowedHeaders: []string{"X-Amz-Date"}, + ExposeHeaders: []string{"Authorization"}, + // weirdely negative max age seconds are also considered valid + MaxAgeSeconds: &maxAgeNegative, + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodDelete, http.MethodGet, http.MethodHead}, + AllowedHeaders: []string{"Content-Type", "Content-Encoding", "Content-MD5"}, + ExposeHeaders: []string{"Authorization", "X-Amz-Date", "X-Amz-Conten-Sha256"}, + ID: getPtr("id"), + MaxAgeSeconds: &maxAgePositive, + }, + { + AllowedOrigins: []string{"http://example.com", "https://something.net", "http://*origin.com"}, + AllowedMethods: []string{http.MethodGet}, + }, + }, + }, + }) + }) +} + +func GetBucketCors_non_existing_bucket(s *S3Conf) error { + testName := "GetBucketCors_non_existing_bucket" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err := s3client.GetBucketCors(ctx, &s3.GetBucketCorsInput{ + Bucket: getPtr("non-existing-bucket"), + }) + cancel() + + return checkApiErr(err, s3err.GetAPIError(s3err.ErrNoSuchBucket)) + }) +} + +func GetBucketCors_no_such_bucket_cors(s *S3Conf) error { + testName := "GetBucketCors_no_such_bucket_cors" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err := s3client.GetBucketCors(ctx, &s3.GetBucketCorsInput{ + Bucket: &bucket, + }) + cancel() + return checkApiErr(err, s3err.GetAPIError(s3err.ErrNoSuchCORSConfiguration)) + }) +} + +func GetBucketCors_success(s *S3Conf) error { + testName := "GetBucketCors_success" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + cfg := &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com", "helloworld.net"}, + AllowedMethods: []string{http.MethodPost, http.MethodPut, http.MethodHead}, + AllowedHeaders: []string{"X-Amz-Date", "X-Amz-Meta-Something"}, + ExposeHeaders: []string{"Authorization", "Content-Disposition"}, + MaxAgeSeconds: getPtr(int32(125)), + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodDelete, http.MethodGet, http.MethodHead}, + AllowedHeaders: []string{"Content-*"}, + ExposeHeaders: []string{"Authorization", "X-Amz-Date", "X-Amz-Conten-Sha256"}, + ID: getPtr("my_extra_unique_id"), + MaxAgeSeconds: getPtr(int32(-200)), + }, + }, + } + + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: cfg, + }) + if err != nil { + return err + } + + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + res, err := s3client.GetBucketCors(ctx, &s3.GetBucketCorsInput{ + Bucket: &bucket, + }) + cancel() + if err != nil { + return err + } + + return compareCorsConfig(cfg.CORSRules, res.CORSRules) + }) +} + +func DeleteBucketCors_non_existing_bucket(s *S3Conf) error { + testName := "DeleteBucketCors_non_existing_bucket" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err := s3client.DeleteBucketCors(ctx, &s3.DeleteBucketCorsInput{ + Bucket: getPtr("non-existing-bucket"), + }) + cancel() + return checkApiErr(err, s3err.GetAPIError(s3err.ErrNoSuchBucket)) + }) +} + +func DeleteBucketCors_success(s *S3Conf) error { + testName := "DeleteBucketCors_success" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + deletebucketcors := func() error { + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err := s3client.DeleteBucketCors(ctx, &s3.DeleteBucketCorsInput{ + Bucket: &bucket, + }) + cancel() + return err + } + + // should not return error when deleting unset bucket CORS + err := deletebucketcors() + if err != nil { + return err + } + + err = putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://origin.com"}, + AllowedMethods: []string{http.MethodPost}, + AllowedHeaders: []string{"X-Amz-Meta-Header"}, + ExposeHeaders: []string{"Content-Disposition"}, + MaxAgeSeconds: getPtr(int32(5000)), + }, + }, + }, + }) + if err != nil { + return err + } + + err = deletebucketcors() + if err != nil { + return err + } + + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err = s3client.GetBucketCors(ctx, &s3.GetBucketCorsInput{ + Bucket: &bucket, + }) + cancel() + + return checkApiErr(err, s3err.GetAPIError(s3err.ErrNoSuchCORSConfiguration)) + }) +} + +func PreflightOPTIONS_non_existing_bucket(s *S3Conf) error { + testName := "PreflightOPTIONS_non_existing_bucket" + return actionHandlerNoSetup(s, testName, func(s3client *s3.Client, bucket string) error { + res, err := makeOPTIONSRequest(s, "non-existing-bucket", "http://localhost:7070", http.MethodPost, "X-Amz-Date") + if err != nil { + return err + } + return checkApiErr(res.err, s3err.GetAPIError(s3err.ErrNoSuchBucket)) + }) +} + +func PreflightOPTIONS_missing_origin(s *S3Conf) error { + testName := "PreflightOPTIONS_missing_origin" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + res, err := makeOPTIONSRequest(s, bucket, "", http.MethodGet, "X-Custom-Header") + if err != nil { + return err + } + return checkApiErr(res.err, s3err.GetAPIError(s3err.ErrMissingCORSOrigin)) + }) +} + +func PreflightOPTIONS_invalid_request_method(s *S3Conf) error { + testName := "PreflightOPTIONS_invalid_request_method" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + for _, method := range []string{ + // should be case sensitive, all with capital letters + "get", "Get", "GEt", "geT", + "post", "Post", "POSt", "posT", + "put", "Put", "pUt", "puT", + "head", "Head", "HEAd", "heAD", + // unsupported methods + "PATCH", "CONNECT", "OPTIONS", + // nonsense strings + "something", "invalid_method", "method", + } { + res, err := makeOPTIONSRequest(s, bucket, "www.my-origin.com", method, "X-Custom-Header") + if err != nil { + return err + } + if err := checkApiErr(res.err, s3err.GetInvalidCORSMethodErr(method)); err != nil { + return err + } + } + + return nil + }) +} + +func PreflightOPTIONS_invalid_request_headers(s *S3Conf) error { + testName := "PreflightOPTIONS_invalid_request_headers" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + for _, test := range []struct { + invalidHeader string + headers string + }{ + {"invalid header", "X-Amz-Date,X-Amz-Content-Sha256,invalid header"}, // invalid 'space' in header name + {"X-Custom:Header", "Authorization,X-Custom:Header"}, // invalid char : + {"X(Custom)", "Content-Length,X(Custom)"}, // invalid char () + {" Bad/Header", "Content-Encoding, Bad/Header"}, // extra 'space', invalid char / + {"X[Key]", "Date,X[Key]"}, // invalid char '[]' + {"Bad=Name", "X-Amz-Custome-Header,Bad=Name"}, // invalid char = + {`X"Quote"`, `X"Quote"`}, // invalid quote " + {"NonAsciiŁ", "Content-Length,NonAsciiŁ"}, // non-ASCII character + {"Emoji😀", "X-Emoji,Emoji😀"}, // emoji invalid + {"bad@char", "Accept-Encoding,bad@char"}, // @ is invalid + {"tab\tchar", "tab\tchar,X-Something-Valid"}, // invalid encodign \t + } { + res, err := makeOPTIONSRequest(s, bucket, "www.my-origin.com", http.MethodGet, test.headers) + if err != nil { + return err + } + if err := checkApiErr(res.err, s3err.GetInvalidCORSRequestHeaderErr(test.invalidHeader)); err != nil { + return err + } + } + return nil + }) +} + +func PreflightOPTIONS_unset_bucket_cors(s *S3Conf) error { + testName := "PreflightOPTIONS_unset_bucket_cors" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + res, err := makeOPTIONSRequest(s, bucket, "http://example.com", http.MethodPost, "X-Amz-Date,Date") + if err != nil { + return err + } + return checkApiErr(res.err, s3err.GetAPIError(s3err.ErrCORSIsNotEnabled)) + }) +} + +func PreflightOPTIONS_access_forbidden(s *S3Conf) error { + testName := "PreflightOPTIONS_access_forbidden" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://example.com", "https://example.com"}, + AllowedMethods: []string{http.MethodGet}, + AllowedHeaders: []string{"X-Amz-Date", "X-Amz-Content-Sha256"}, + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodHead}, + }, + { + AllowedOrigins: []string{"http://origin*"}, + AllowedMethods: []string{http.MethodPost}, + AllowedHeaders: []string{"Authorization"}, + }, + { + AllowedOrigins: []string{"http://something.com"}, + AllowedMethods: []string{http.MethodPut}, + AllowedHeaders: []string{"X-Amz-*"}, + }, + }, + }, + }) + if err != nil { + return err + } + + for _, test := range []struct { + origin string + method string + headers string + }{ + // origin deson't match + {"http://non-matching-origin.net", http.MethodGet, "X-Amz-Date"}, + // method doesn't match + {"http://example.com", http.MethodPut, "X-Amz-Content-Sha256"}, + // header doesn't match + {"http://example.com", http.MethodGet, "X-Amz-Expected-Bucket-Owner"}, + // extra header + {"http://example.com", http.MethodGet, "X-Amz-Date,X-Amz-Content-Sha256,Extra-Header"}, + // extra header (2nd rule) + {"https://any-origin.com", http.MethodHead, "X-Amz-Extra-Header"}, + // origin match, method not (2nd rule) + {"https://any-origin.com", http.MethodPost, ""}, + // third rule: headers doesn't match + {"https://origin.com", http.MethodPost, "Content-Length"}, + // third rule: extra header + {"https://origin.com", http.MethodPost, "Authorization,Content-Disposition"}, + // third rule: origin doesn't match + {"https://www.origin.com", http.MethodPost, "Authorization"}, + // forth rule: header doesn't match the wildcard + {"https://something.com", http.MethodPut, "Authorization"}, + {"https://something.com", http.MethodPut, "X-Amz"}, + {"https://something.com", http.MethodPut, "X-Amz-Date,Content-Length"}, + } { + res, err := makeOPTIONSRequest(s, bucket, test.origin, test.method, test.headers) + if err != nil { + return err + } + + if err := checkApiErr(res.err, s3err.GetAPIError(s3err.ErrCORSForbidden)); err != nil { + return err + } + } + + return nil + }) +} + +func PreflightOPTIONS_access_granted(s *S3Conf) error { + testName := "PreflightOPTIONS_access_granted" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://example.com", "https://example.com"}, + AllowedMethods: []string{http.MethodGet, http.MethodHead}, + AllowedHeaders: []string{"X-Amz-Date", "X-Amz-Content-Sha256"}, + ExposeHeaders: []string{"Content-Type", "Content-Length"}, + MaxAgeSeconds: getPtr(int32(100)), + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodHead}, + AllowedHeaders: []string{"X-Amz-Meta-Something"}, + }, + { + AllowedOrigins: []string{"something.net"}, + AllowedMethods: []string{http.MethodPost, http.MethodPut}, + AllowedHeaders: []string{"Authorization"}, + ExposeHeaders: []string{"Content-Disposition", "Content-Encoding"}, + MaxAgeSeconds: getPtr(int32(3000)), + ID: getPtr("unique_id"), + }, + { + AllowedOrigins: []string{"http://www*"}, + AllowedMethods: []string{http.MethodGet}, + AllowedHeaders: []string{"x-amz-server-side-encryption"}, + ExposeHeaders: []string{"X-Amz-Expected-Bucket-Owner"}, + MaxAgeSeconds: getPtr(int32(5000)), + }, + { + AllowedOrigins: []string{"http://uniquie-origin.net"}, + AllowedMethods: []string{http.MethodPost, http.MethodPut}, + AllowedHeaders: []string{"X-Amz-*-Suffix"}, + ExposeHeaders: []string{"Authorization", "Content-Type"}, + MaxAgeSeconds: getPtr(int32(2000)), + }, + }, + }, + }) + if err != nil { + return err + } + + varyHdr := "Origin, Access-Control-Request-Headers, Access-Control-Request-Method" + + for _, test := range []struct { + origin string + method string + headers string + result PreflightResult + }{ + // first rule matches + {"http://example.com", http.MethodGet, "X-Amz-Date", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"http://example.com", http.MethodGet, "X-Amz-Content-Sha256", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"http://example.com", http.MethodHead, "", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"https://example.com", http.MethodGet, "X-Amz-Date,X-Amz-Content-Sha256", PreflightResult{"https://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + // second rule matches: origin is a wildcard + {"http://anything.com", http.MethodHead, "X-Amz-Meta-Something", PreflightResult{"*", "HEAD", "", "", "false", varyHdr, nil}}, + {"hello.com", http.MethodHead, "", PreflightResult{"*", "HEAD", "", "", "false", varyHdr, nil}}, + // third rule matches + {"something.net", http.MethodPut, "Authorization", PreflightResult{"something.net", "POST, PUT", "Content-Disposition, Content-Encoding", "3000", "true", varyHdr, nil}}, + {"something.net", http.MethodPost, "", PreflightResult{"something.net", "POST, PUT", "Content-Disposition, Content-Encoding", "3000", "true", varyHdr, nil}}, + // forth rule matches: origin contains wildcard + {"http://www.hello.world.com", http.MethodGet, "", PreflightResult{"http://www.hello.world.com", "GET", "X-Amz-Expected-Bucket-Owner", "5000", "true", varyHdr, nil}}, + {"http://www.example.com", http.MethodGet, "x-amz-server-side-encryption", PreflightResult{"http://www.example.com", "GET", "X-Amz-Expected-Bucket-Owner", "5000", "true", varyHdr, nil}}, + // fifth rule matches: allowed headers contains wildcard + {"http://uniquie-origin.net", http.MethodPost, "X-Amz-anything-Suffix", PreflightResult{"http://uniquie-origin.net", "POST, PUT", "Authorization, Content-Type", "2000", "true", varyHdr, nil}}, + {"http://uniquie-origin.net", http.MethodPut, "X-Amz-yyy-xxx-Suffix", PreflightResult{"http://uniquie-origin.net", "POST, PUT", "Authorization, Content-Type", "2000", "true", varyHdr, nil}}, + } { + err := testOPTIONSEdnpoint(s, bucket, test.origin, test.method, test.headers, &test.result) + if err != nil { + return err + } + } + + return nil + }) +} + +func CORSMiddleware_invalid_method(s *S3Conf) error { + testName := "CORSMiddleware_invalid_method" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://www.example.com"}, + AllowedMethods: []string{http.MethodPut}, + }, + }, + }, + }) + if err != nil { + return err + } + + // create a PutObject signed request + req, err := createSignedReq(http.MethodPut, s.endpoint, bucket+"/my-obj", s.awsID, s.awsSecret, "s3", s.awsRegion, nil, time.Now(), map[string]string{ + "Origin": "http://www.example.com", + "Access-Control-Request-Method": "invalid_method", + }) + if err != nil { + return err + } + + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + + result, err := extractCORSHeaders(resp) + if err != nil { + return err + } + + return checkApiErr(result.err, s3err.GetInvalidCORSMethodErr("invalid_method")) + }) +} + +func CORSMiddleware_invalid_headers(s *S3Conf) error { + testName := "CORSMiddleware_invalid_headers" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://www.example.com"}, + AllowedMethods: []string{http.MethodPut}, + }, + }, + }, + }) + if err != nil { + return err + } + + // create a PutObject signed request + req, err := createSignedReq(http.MethodPut, s.endpoint, bucket+"/my-obj", s.awsID, s.awsSecret, "s3", s.awsRegion, nil, time.Now(), map[string]string{ + "Origin": "http://www.example.com", + "Access-Control-Request-Headers": "invalid header", + }) + if err != nil { + return err + } + + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + + result, err := extractCORSHeaders(resp) + if err != nil { + return err + } + + return checkApiErr(result.err, s3err.GetInvalidCORSRequestHeaderErr("invalid header")) + }) +} + +func CORSMiddleware_access_forbidden(s *S3Conf) error { + testName := "CORSMiddleware_access_forbidden" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://example.com", "https://example.com"}, + AllowedMethods: []string{http.MethodGet}, + AllowedHeaders: []string{"X-Amz-Date", "X-Amz-Content-Sha256"}, + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodHead}, + }, + }, + }, + }) + if err != nil { + return err + } + + for _, test := range []struct { + origin string + method string + headers string + }{ + // origin deson't match + {"http://non-matching-origin.net", http.MethodGet, "X-Amz-Date"}, + // method doesn't match + {"http://example.com", http.MethodPut, "X-Amz-Content-Sha256"}, + // header doesn't match + {"http://example.com", http.MethodGet, "X-Amz-Expected-Bucket-Owner"}, + // extra header + {"http://example.com", http.MethodGet, "X-Amz-Date,X-Amz-Content-Sha256,Extra-Header"}, + // extra header (2nd rule) + {"https://any-origin.com", http.MethodHead, "X-Amz-Extra-Header"}, + // origin match, method not (2nd rule) + {"https://any-origin.com", http.MethodPost, ""}, + } { + req, err := createSignedReq(http.MethodPut, s.endpoint, bucket+"/my-obj", s.awsID, s.awsSecret, "s3", s.awsRegion, nil, time.Now(), map[string]string{ + "Origin": test.origin, + "Access-Control-Request-Headers": test.headers, + "Access-Control-Request-Method": test.method, + }) + if err != nil { + return err + } + + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + + res, err := extractCORSHeaders(resp) + if err != nil { + return err + } + + // no error expected, all the headers should be empty + if err := comparePreflightResult(&PreflightResult{}, res); err != nil { + return err + } + } + + return nil + }) +} + +func CORSMiddleware_access_granted(s *S3Conf) error { + testName := "CORSMiddleware_access_granted" + return actionHandler(s, testName, func(s3client *s3.Client, bucket string) error { + err := putBucketCors(s3client, &s3.PutBucketCorsInput{ + Bucket: &bucket, + CORSConfiguration: &types.CORSConfiguration{ + CORSRules: []types.CORSRule{ + { + AllowedOrigins: []string{"http://example.com", "https://example.com"}, + AllowedMethods: []string{http.MethodGet, http.MethodHead}, + AllowedHeaders: []string{"X-Amz-Date", "X-Amz-Content-Sha256"}, + ExposeHeaders: []string{"Content-Type", "Content-Length"}, + MaxAgeSeconds: getPtr(int32(100)), + }, + { + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{http.MethodHead}, + AllowedHeaders: []string{"X-Amz-Meta-Something"}, + }, + { + AllowedOrigins: []string{"something.net"}, + AllowedMethods: []string{http.MethodPost, http.MethodPut}, + AllowedHeaders: []string{"Authorization"}, + ExposeHeaders: []string{"Content-Disposition", "Content-Encoding"}, + MaxAgeSeconds: getPtr(int32(3000)), + ID: getPtr("unique_id"), + }, + }, + }, + }) + if err != nil { + return err + } + + varyHdr := "Origin, Access-Control-Request-Headers, Access-Control-Request-Method" + + for _, test := range []struct { + origin string + method string + headers string + result PreflightResult + }{ + // first rule matches + {"http://example.com", http.MethodGet, "X-Amz-Date", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"http://example.com", http.MethodGet, "X-Amz-Content-Sha256", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"http://example.com", http.MethodHead, "", PreflightResult{"http://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + {"https://example.com", http.MethodGet, "X-Amz-Date,X-Amz-Content-Sha256", PreflightResult{"https://example.com", "GET, HEAD", "Content-Type, Content-Length", "100", "true", varyHdr, nil}}, + // second rule matches + {"http://anything.com", http.MethodHead, "X-Amz-Meta-Something", PreflightResult{"*", "HEAD", "", "", "false", varyHdr, nil}}, + {"hello.com", http.MethodHead, "", PreflightResult{"*", "HEAD", "", "", "false", varyHdr, nil}}, + // third rule matches + {"something.net", http.MethodPut, "Authorization", PreflightResult{"something.net", "POST, PUT", "Content-Disposition, Content-Encoding", "3000", "true", varyHdr, nil}}, + {"something.net", http.MethodPost, "", PreflightResult{"something.net", "POST, PUT", "Content-Disposition, Content-Encoding", "3000", "true", varyHdr, nil}}, + } { + req, err := createSignedReq(http.MethodPut, s.endpoint, bucket+"/my-obj", s.awsID, s.awsSecret, "s3", s.awsRegion, nil, time.Now(), map[string]string{ + "Origin": test.origin, + "Access-Control-Request-Headers": test.headers, + "Access-Control-Request-Method": test.method, + }) + if err != nil { + return err + } + + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + + res, err := extractCORSHeaders(resp) + if err != nil { + return err + } + + if err := comparePreflightResult(&test.result, res); err != nil { + return err + } + } + + return nil + }) +} + // Object lock tests func PutObjectLockConfiguration_non_existing_bucket(s *S3Conf) error { testName := "PutObjectLockConfiguration_non_existing_bucket" @@ -16571,7 +17359,7 @@ func PublicBucket_public_bucket_policy(s *S3Conf) error { }) return err }, - ExpectedErr: s3err.GetAPIError(s3err.ErrNotImplemented), + ExpectedErr: nil, }, { Action: "GetBucketCors", @@ -16579,7 +17367,7 @@ func PublicBucket_public_bucket_policy(s *S3Conf) error { _, err := s3client.GetBucketCors(ctx, &s3.GetBucketCorsInput{Bucket: &bucket}) return err }, - ExpectedErr: s3err.GetAPIError(s3err.ErrNotImplemented), + ExpectedErr: nil, }, { Action: "DeleteBucketCors", @@ -16587,7 +17375,7 @@ func PublicBucket_public_bucket_policy(s *S3Conf) error { _, err := s3client.DeleteBucketCors(ctx, &s3.DeleteBucketCorsInput{Bucket: &bucket}) return err }, - ExpectedErr: s3err.GetAPIError(s3err.ErrNotImplemented), + ExpectedErr: nil, }, { Action: "CreateMultipartUpload", @@ -21321,7 +22109,6 @@ func Versioning_AccessControl_GetObjectVersion(s *S3Conf) error { } doc := genPolicyDoc("Allow", fmt.Sprintf(`"%s"`, testuser1.access), `"s3:GetObject"`, fmt.Sprintf(`"arn:aws:s3:::%s/*"`, bucket)) - fmt.Println(doc) ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) _, err = s3client.PutBucketPolicy(ctx, &s3.PutBucketPolicyInput{ Bucket: &bucket, @@ -21391,7 +22178,6 @@ func Versioning_AccessControl_HeadObjectVersion(s *S3Conf) error { } doc := genPolicyDoc("Allow", fmt.Sprintf(`"%s"`, testuser1.access), `"s3:GetObject"`, fmt.Sprintf(`"arn:aws:s3:::%s/*"`, bucket)) - fmt.Println(doc) ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) _, err = s3client.PutBucketPolicy(ctx, &s3.PutBucketPolicyInput{ Bucket: &bucket, diff --git a/tests/integration/utils.go b/tests/integration/utils.go index d7e820a7..23d34a71 100644 --- a/tests/integration/utils.go +++ b/tests/integration/utils.go @@ -703,7 +703,7 @@ func getString(str *string) string { return *str } -func getPtr(str string) *string { +func getPtr[T any](str T) *T { return &str } @@ -1629,3 +1629,167 @@ func getObject_zero_len_with_range_helper(testName, obj string, s *S3Conf) error return nil }) } + +func getInt32(ptr *int32) int32 { + if ptr == nil { + return 0 + } + + return *ptr +} + +func putBucketCors(client *s3.Client, input *s3.PutBucketCorsInput) error { + ctx, cancel := context.WithTimeout(context.Background(), shortTimeout) + _, err := client.PutBucketCors(ctx, input) + cancel() + return err +} + +func compareCorsConfig(expected, got []types.CORSRule) error { + if expected == nil && got == nil { + return nil + } + if got == nil { + return errors.New("nil CORS config") + } + + if len(expected) != len(got) { + return fmt.Errorf("expected CORS rules length to be %v, instead got %v", len(expected), len(got)) + } + + for i, r := range expected { + rule := got[i] + if !slices.Equal(r.AllowedOrigins, rule.AllowedOrigins) { + return fmt.Errorf("expected the allowed origins to be %v, instead got %v", r.AllowedOrigins, rule.AllowedOrigins) + } + if !slices.Equal(r.AllowedMethods, rule.AllowedMethods) { + return fmt.Errorf("expected the allowed methods to be %v, instead got %v", r.AllowedMethods, rule.AllowedMethods) + } + if !slices.Equal(r.AllowedHeaders, rule.AllowedHeaders) { + return fmt.Errorf("expected the allowed headers to be %v, instead got %v", r.AllowedHeaders, rule.AllowedHeaders) + } + if !slices.Equal(r.ExposeHeaders, rule.ExposeHeaders) { + return fmt.Errorf("expected the allowed origins to be %v, instead got %v", r.ExposeHeaders, rule.ExposeHeaders) + } + if getInt32(r.MaxAgeSeconds) != getInt32(rule.MaxAgeSeconds) { + return fmt.Errorf("expected the max age seconds to be %v, instead got %v", getInt32(r.MaxAgeSeconds), getInt32(rule.MaxAgeSeconds)) + } + if getString(r.ID) != getString(rule.ID) { + return fmt.Errorf("expected ID to be %v, instead got %v", getString(r.ID), getString(rule.ID)) + } + } + + return nil +} + +type PreflightResult struct { + Origin string + Methods string + ExposeHeaders string + MaxAge string + AllowCredentials string + Vary string + err error +} + +func extractCORSHeaders(resp *http.Response) (*PreflightResult, error) { + if resp.StatusCode >= 400 { + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read response body: %w", err) + } + + var errResp smithy.GenericAPIError + err = xml.Unmarshal(body, &errResp) + if err != nil { + return nil, fmt.Errorf("unmarshal respone body: %w", err) + } + + return &PreflightResult{ + err: &errResp, + }, nil + } + + return &PreflightResult{ + Origin: resp.Header.Get("Access-Control-Allow-Origin"), + Methods: resp.Header.Get("Access-Control-Allow-Methods"), + ExposeHeaders: resp.Header.Get("Access-Control-Expose-Headers"), + MaxAge: resp.Header.Get("Access-Control-Max-Age"), + AllowCredentials: resp.Header.Get("Access-Control-Allow-Credentials"), + Vary: resp.Header.Get("Vary"), + }, nil +} + +func makeOPTIONSRequest(s *S3Conf, bucket, origin, method string, headers string) (*PreflightResult, error) { + req, err := http.NewRequest(http.MethodOptions, fmt.Sprintf("%s/%s/object", s.endpoint, bucket), nil) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + + req.Header.Add("Origin", origin) + req.Header.Add("Access-Control-Request-Method", method) + req.Header.Add("Access-Control-Request-Headers", headers) + + resp, err := s.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("send request: %w", err) + } + + return extractCORSHeaders(resp) +} + +func comparePreflightResult(expected, got *PreflightResult) error { + if expected == nil { + return fmt.Errorf("nil expected preflight request result") + } + if got == nil { + return fmt.Errorf("expected the preflights result to be %v, instead got nil", *expected) + } + + if expected.err != nil { + if got.err == nil { + return fmt.Errorf("expected %w error, instaed got nil", expected.err) + } + + apiErr, ok := expected.err.(s3err.APIError) + if !ok { + return fmt.Errorf("expected s3err.APIError, instead got %w", expected.err) + } + + return checkApiErr(got.err, apiErr) + } + + if got.err != nil { + return fmt.Errorf("expected no error, instaed got %w", got.err) + } + + if expected.Origin != got.Origin { + return fmt.Errorf("expected the origin to be %v, instead got %v", expected.Origin, got.Origin) + } + if expected.Methods != got.Methods { + return fmt.Errorf("expected the allowed methods to be %v, instead got %v", expected.Methods, got.Methods) + } + if expected.ExposeHeaders != got.ExposeHeaders { + return fmt.Errorf("expected the expose headers to be %v, instead got %v", expected.ExposeHeaders, got.ExposeHeaders) + } + if expected.MaxAge != got.MaxAge { + return fmt.Errorf("expected the max age to be %v, instead got %v", expected.MaxAge, got.MaxAge) + } + if expected.AllowCredentials != got.AllowCredentials { + return fmt.Errorf("expected the allow credentials to be %v, instead got %v", expected.AllowCredentials, got.AllowCredentials) + } + if expected.Vary != got.Vary { + return fmt.Errorf("expected the Vary header to be %v, instead got %v", expected.Vary, got.Vary) + } + + return nil +} + +func testOPTIONSEdnpoint(s *S3Conf, bucket, origin, method string, headers string, expected *PreflightResult) error { + result, err := makeOPTIONSRequest(s, bucket, origin, method, headers) + if err != nil { + return err + } + + return comparePreflightResult(expected, result) +}