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)
+}