Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 38 additions & 2 deletions internal/httpserver/cors.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ package httpserver
import (
"errors"
"net/http"
"strings"

"github.com/gorilla/handlers"

Expand All @@ -15,19 +16,22 @@ const (
fieldCORS = "cors"
fieldCORSEnabled = "enabled"
fieldCORSAllowedOrigins = "allowed_origins"
fieldCORSAllowedHeaders = "allowed_headers"
)

// CORSConfig contains struct configuration for allowing CORS headers.
type CORSConfig struct {
Enabled bool `json:"enabled" yaml:"enabled"`
AllowedOrigins []string `json:"allowed_origins" yaml:"allowed_origins"`
AllowedHeaders []string `json:"allowed_headers" yaml:"allowed_headers"`
}

// NewServerCORSConfig returns a new server CORS config with default fields.
func NewServerCORSConfig() CORSConfig {
return CORSConfig{
Enabled: false,
AllowedOrigins: []string{},
AllowedHeaders: []string{},
}
}

Expand All @@ -40,17 +44,46 @@ func (conf CORSConfig) WrapHandler(handler http.Handler) (http.Handler, error) {
if len(conf.AllowedOrigins) == 0 {
return nil, errors.New("must specify at least one allowed origin")
}
return handlers.CORS(
corsHandler := handlers.CORS(
handlers.AllowedOrigins(conf.AllowedOrigins),
handlers.AllowedMethods([]string{"GET", "HEAD", "POST", "PUT", "PATCH", "DELETE"}),
)(handler), nil
handlers.AllowedHeaders(conf.AllowedHeaders),
)(handler)

if !conf.allowsAllHeaders() {
return corsHandler, nil
}

return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requestedHeaders := r.Header.Get("Access-Control-Request-Headers")
if r.Method != http.MethodOptions || requestedHeaders == "" {
corsHandler.ServeHTTP(w, r)
return
}

request := r.Clone(r.Context())
request.Header = r.Header.Clone()
request.Header.Del("Access-Control-Request-Headers")
w.Header().Set("Access-Control-Allow-Headers", requestedHeaders)
corsHandler.ServeHTTP(w, request)
}), nil
}

func (conf CORSConfig) allowsAllHeaders() bool {
for _, header := range conf.AllowedHeaders {
if strings.TrimSpace(header) == "*" {
return true
}
}
return false
}

// ServerCORSFieldSpec returns a field spec for an http server CORS component.
func ServerCORSFieldSpec() docs.FieldSpec {
return docs.FieldObject(fieldCORS, "Adds Cross-Origin Resource Sharing headers.").WithChildren(
docs.FieldBool(fieldCORSEnabled, "Whether to allow CORS requests.").HasDefault(false),
docs.FieldString(fieldCORSAllowedOrigins, "An explicit list of origins that are allowed for CORS requests.").Array().HasDefault([]any{}),
docs.FieldString(fieldCORSAllowedHeaders, "An explicit list of headers allowed in CORS requests. Specify `*` to allow any requested header.").Array().HasDefault([]any{}),
).AtVersion("3.63.0").Advanced()
}

Expand All @@ -63,5 +96,8 @@ func CORSConfigFromParsed(pConf *docs.ParsedConfig) (conf CORSConfig, err error)
if conf.AllowedOrigins, err = pConf.FieldStringList(fieldCORSAllowedOrigins); err != nil {
return
}
if conf.AllowedHeaders, err = pConf.FieldStringList(fieldCORSAllowedHeaders); err != nil {
return
}
return
}
42 changes: 42 additions & 0 deletions internal/httpserver/cors_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,48 @@ func TestAPIEnableCORSOrigins(t *testing.T) {
assert.Empty(t, response.Header().Get("Access-Control-Allow-Origin"))
}

func TestAPIEnableCORSAllowedHeaders(t *testing.T) {
conf := NewServerCORSConfig()
conf.Enabled = true
conf.AllowedOrigins = []string{"*"}
conf.AllowedHeaders = []string{"Content-Type"}

handler, err := conf.WrapHandler(http.NewServeMux())
require.NoError(t, err)

request, _ := http.NewRequest("OPTIONS", "/version", http.NoBody)
request.Header.Set("Origin", "meow")
request.Header.Set("Access-Control-Request-Method", "POST")
request.Header.Set("Access-Control-Request-Headers", "content-type")

response := httptest.NewRecorder()
handler.ServeHTTP(response, request)

assert.Equal(t, http.StatusOK, response.Code)
assert.Equal(t, "Content-Type", response.Header().Get("Access-Control-Allow-Headers"))
}

func TestAPIEnableCORSAllHeaders(t *testing.T) {
conf := NewServerCORSConfig()
conf.Enabled = true
conf.AllowedOrigins = []string{"*"}
conf.AllowedHeaders = []string{"*"}

handler, err := conf.WrapHandler(http.NewServeMux())
require.NoError(t, err)

request, _ := http.NewRequest("OPTIONS", "/version", http.NoBody)
request.Header.Set("Origin", "meow")
request.Header.Set("Access-Control-Request-Method", "POST")
request.Header.Set("Access-Control-Request-Headers", "content-type, x-request-id")

response := httptest.NewRecorder()
handler.ServeHTTP(response, request)

assert.Equal(t, http.StatusOK, response.Code)
assert.Equal(t, "content-type, x-request-id", response.Header().Get("Access-Control-Allow-Headers"))
}

func TestAPIEnableCORSNoHeaders(t *testing.T) {
conf := NewServerCORSConfig()
conf.Enabled = true
Expand Down
4 changes: 4 additions & 0 deletions internal/impl/io/input_http_server.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ const (
hsiFieldCORS = "cors"
hsiFieldCORSEnabled = "enabled"
hsiFieldCORSAllowedOrigins = "allowed_origins"
hsiFieldCORSAllowedHeaders = "allowed_headers"
hsiFieldResponse = "sync_response"
hsiFieldResponseStatus = "status"
hsiFieldResponseHeaders = "headers"
Expand Down Expand Up @@ -147,6 +148,9 @@ func corsConfigFromParsed(pConf *service.ParsedConfig) (conf httpserver.CORSConf
if conf.AllowedOrigins, err = pConf.FieldStringList(hsiFieldCORSAllowedOrigins); err != nil {
return
}
if conf.AllowedHeaders, err = pConf.FieldStringList(hsiFieldCORSAllowedHeaders); err != nil {
return
}
return
}

Expand Down
4 changes: 3 additions & 1 deletion internal/impl/io/input_http_server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1324,6 +1324,7 @@ http_server:
cors:
enabled: true
allowed_origins: [ foo, bar ]
allowed_headers: [ Content-Type ]
`, freePort)

server, err := mock.NewManager().NewInput(conf)
Expand All @@ -1343,7 +1344,7 @@ http_server:

req.Header.Add("Origin", "foo")
req.Header.Add("Access-Control-Request-Method", "POST")
req.Header.Set("Content-Type", "text/plain")
req.Header.Set("Access-Control-Request-Headers", "content-type")

if resp, cerr = http.DefaultClient.Do(req); cerr == nil {
succeeded = true
Expand All @@ -1354,6 +1355,7 @@ http_server:

assert.Equal(t, "200 OK", resp.Status)
assert.Equal(t, "foo", resp.Header.Get("Access-Control-Allow-Origin"))
assert.Equal(t, "Content-Type", resp.Header.Get("Access-Control-Allow-Headers"))
}

// TestHTTPServerReload tests that the server can be closed and recreated on the same port
Expand Down
4 changes: 4 additions & 0 deletions internal/impl/io/input_http_server_wasm.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ const (
hsiFieldCORS = "cors"
hsiFieldCORSEnabled = "enabled"
hsiFieldCORSAllowedOrigins = "allowed_origins"
hsiFieldCORSAllowedHeaders = "allowed_headers"
hsiFieldResponse = "sync_response"
hsiFieldResponseStatus = "status"
hsiFieldResponseHeaders = "headers"
Expand Down Expand Up @@ -142,6 +143,9 @@ func corsConfigFromParsed(pConf *service.ParsedConfig) (conf httpserver.CORSConf
if conf.AllowedOrigins, err = pConf.FieldStringList(hsiFieldCORSAllowedOrigins); err != nil {
return
}
if conf.AllowedHeaders, err = pConf.FieldStringList(hsiFieldCORSAllowedHeaders); err != nil {
return
}
return
}

Expand Down