From 13feb4d639f32a89eaafdaeb947daa56694d851e Mon Sep 17 00:00:00 2001 From: GabrielFeliciano Date: Mon, 10 Aug 2026 19:09:08 -0300 Subject: [PATCH] refactor: allow request headers cors --- internal/httpserver/cors.go | 40 +++++++++++++++++++-- internal/httpserver/cors_test.go | 42 ++++++++++++++++++++++ internal/impl/io/input_http_server.go | 4 +++ internal/impl/io/input_http_server_test.go | 4 ++- internal/impl/io/input_http_server_wasm.go | 4 +++ 5 files changed, 91 insertions(+), 3 deletions(-) diff --git a/internal/httpserver/cors.go b/internal/httpserver/cors.go index 0c1ee1c24..055e60646 100644 --- a/internal/httpserver/cors.go +++ b/internal/httpserver/cors.go @@ -5,6 +5,7 @@ package httpserver import ( "errors" "net/http" + "strings" "github.com/gorilla/handlers" @@ -15,12 +16,14 @@ 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. @@ -28,6 +31,7 @@ func NewServerCORSConfig() CORSConfig { return CORSConfig{ Enabled: false, AllowedOrigins: []string{}, + AllowedHeaders: []string{}, } } @@ -40,10 +44,38 @@ 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. @@ -51,6 +83,7 @@ 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() } @@ -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 } diff --git a/internal/httpserver/cors_test.go b/internal/httpserver/cors_test.go index d7bfe9190..ba023f586 100644 --- a/internal/httpserver/cors_test.go +++ b/internal/httpserver/cors_test.go @@ -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 diff --git a/internal/impl/io/input_http_server.go b/internal/impl/io/input_http_server.go index 360056217..73d314879 100644 --- a/internal/impl/io/input_http_server.go +++ b/internal/impl/io/input_http_server.go @@ -58,6 +58,7 @@ const ( hsiFieldCORS = "cors" hsiFieldCORSEnabled = "enabled" hsiFieldCORSAllowedOrigins = "allowed_origins" + hsiFieldCORSAllowedHeaders = "allowed_headers" hsiFieldResponse = "sync_response" hsiFieldResponseStatus = "status" hsiFieldResponseHeaders = "headers" @@ -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 } diff --git a/internal/impl/io/input_http_server_test.go b/internal/impl/io/input_http_server_test.go index 85885c16c..5ef3751e7 100644 --- a/internal/impl/io/input_http_server_test.go +++ b/internal/impl/io/input_http_server_test.go @@ -1324,6 +1324,7 @@ http_server: cors: enabled: true allowed_origins: [ foo, bar ] + allowed_headers: [ Content-Type ] `, freePort) server, err := mock.NewManager().NewInput(conf) @@ -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 @@ -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 diff --git a/internal/impl/io/input_http_server_wasm.go b/internal/impl/io/input_http_server_wasm.go index 812c9ef7b..d3af3ab49 100644 --- a/internal/impl/io/input_http_server_wasm.go +++ b/internal/impl/io/input_http_server_wasm.go @@ -57,6 +57,7 @@ const ( hsiFieldCORS = "cors" hsiFieldCORSEnabled = "enabled" hsiFieldCORSAllowedOrigins = "allowed_origins" + hsiFieldCORSAllowedHeaders = "allowed_headers" hsiFieldResponse = "sync_response" hsiFieldResponseStatus = "status" hsiFieldResponseHeaders = "headers" @@ -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 }