mirror of
https://github.com/zeromicro/go-zero.git
synced 2026-05-11 00:40:00 +08:00
feat: support rest.WithCorsHeaders to customize cors headers (#4284)
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user