feat: support rest.WithCorsHeaders to customize cors headers (#4284)

This commit is contained in:
Kevin Wan
2024-07-30 17:29:44 +08:00
committed by GitHub
parent c2421beb25
commit 2588a36555
4 changed files with 144 additions and 0 deletions

View File

@@ -26,6 +26,11 @@ const (
originHeader = "Origin"
)
// AddAllowHeaders sets the allowed headers.
func AddAllowHeaders(header http.Header, headers ...string) {
header.Add(allowHeaders, strings.Join(headers, ", "))
}
// NotAllowedHandler handles cross domain not allowed requests.
// At most one origin can be specified, other origins are ignored if given, default to be *.
func NotAllowedHandler(fn func(w http.ResponseWriter), origins ...string) http.Handler {

View File

@@ -3,11 +3,80 @@ package cors
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/stretchr/testify/assert"
)
func TestAddAllowHeaders(t *testing.T) {
tests := []struct {
name string
initial string
headers []string
expected string
}{
{
name: "single header",
initial: "",
headers: []string{"Content-Type"},
expected: "Content-Type",
},
{
name: "multiple headers",
initial: "",
headers: []string{"Content-Type", "Authorization", "X-Requested-With"},
expected: "Content-Type, Authorization, X-Requested-With",
},
{
name: "add to existing headers",
initial: "Origin, Accept",
headers: []string{"Content-Type"},
expected: "Origin, Accept, Content-Type",
},
{
name: "no headers",
initial: "",
headers: []string{},
expected: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
header := http.Header{}
headers := make(map[string]struct{})
if tt.initial != "" {
header.Set(allowHeaders, tt.initial)
vals := strings.Split(tt.initial, ", ")
for _, v := range vals {
headers[v] = struct{}{}
}
}
for _, h := range tt.headers {
headers[h] = struct{}{}
}
AddAllowHeaders(header, tt.headers...)
var actual []string
vals := header.Values(allowHeaders)
for _, v := range vals {
bunch := strings.Split(v, ", ")
for _, b := range bunch {
if len(b) > 0 {
actual = append(actual, b)
}
}
}
var expect []string
for k := range headers {
expect = append(expect, k)
}
assert.ElementsMatch(t, expect, actual)
})
}
}
func TestCorsHandlerWithOrigins(t *testing.T) {
tests := []struct {
name string