configurable cors settings

This commit is contained in:
pascal
2026-07-29 15:55:02 +02:00
parent f2c1070f95
commit 4b42adcbf3
5 changed files with 91 additions and 8 deletions

View File

@@ -13,7 +13,6 @@ import (
"github.com/gorilla/mux"
grpcMiddleware "github.com/grpc-ecosystem/go-grpc-middleware/v2"
"github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip"
"github.com/rs/cors"
"github.com/rs/xid"
log "github.com/sirupsen/logrus"
"google.golang.org/grpc"
@@ -24,13 +23,13 @@ import (
"github.com/netbirdio/netbird/encryption"
"github.com/netbirdio/netbird/formatter/hook"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/activity"
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbContext "github.com/netbirdio/netbird/management/server/context"
nbhttp "github.com/netbirdio/netbird/management/server/http"
@@ -122,7 +121,7 @@ func (s *BaseServer) EventStore() activity.Store {
func (s *BaseServer) APIHandler() http.Handler {
return Create(s, func() http.Handler {
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.Config.HttpConfig.CORSAllowedOrigins, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
if err != nil {
log.Fatalf("failed to create API handler: %v", err)
}
@@ -137,7 +136,7 @@ func (s *BaseServer) IDPHandler() http.Handler {
if !ok || embeddedIdP == nil {
return nil
}
return cors.AllowAll().Handler(embeddedIdP.Handler())
return nbhttp.CORSMiddleware(s.Config.HttpConfig.CORSAllowedOrigins).Handler(embeddedIdP.Handler())
}
func (s *BaseServer) Router() *mux.Router {

View File

@@ -125,6 +125,9 @@ type HttpServerConfig struct {
ExtraAuthAudience string
// AuthCallbackDomain contains the callback domain
AuthCallbackURL string
// CORSAllowedOrigins lists the browser origins allowed to read API responses,
// e.g. https://app.example.com. Any origin is allowed when left empty.
CORSAllowedOrigins []string
}
// Host represents a Netbird host (e.g. STUN, TURN, Signal)

View File

@@ -0,0 +1,55 @@
package http
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/stretchr/testify/assert"
)
func newCORSPreflight(origin string) *http.Request {
r := httptest.NewRequest(http.MethodOptions, "/api/peers", nil)
r.Header.Set("Origin", origin)
r.Header.Set("Access-Control-Request-Method", http.MethodPut)
r.Header.Set("Access-Control-Request-Headers", "authorization,content-type")
return r
}
func serveCORS(t *testing.T, allowedOrigins []string, r *http.Request) http.Header {
t.Helper()
w := httptest.NewRecorder()
CORSMiddleware(allowedOrigins).Handler(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})).ServeHTTP(w, r)
return w.Header()
}
func TestCORSMiddlewareAllowsConfiguredOrigin(t *testing.T) {
headers := serveCORS(t, []string{"https://app.example.com"}, newCORSPreflight("https://app.example.com"))
assert.Equal(t, "https://app.example.com", headers.Get("Access-Control-Allow-Origin"))
assert.Contains(t, headers.Get("Vary"), "Origin")
}
func TestCORSMiddlewareRejectsUnknownOrigin(t *testing.T) {
headers := serveCORS(t, []string{"https://app.example.com"}, newCORSPreflight("https://evil.example.com"))
assert.Empty(t, headers.Get("Access-Control-Allow-Origin"))
}
func TestCORSMiddlewareNeverAllowsCredentials(t *testing.T) {
for _, allowedOrigins := range [][]string{nil, {"https://app.example.com"}} {
headers := serveCORS(t, allowedOrigins, newCORSPreflight("https://app.example.com"))
assert.Empty(t, headers.Get("Access-Control-Allow-Credentials"))
assert.Empty(t, headers.Get("Access-Control-Expose-Headers"))
}
}
func TestCORSMiddlewareWithoutConfigAllowsAnyOrigin(t *testing.T) {
headers := serveCORS(t, nil, newCORSPreflight("https://evil.example.com"))
assert.Equal(t, "*", headers.Get("Access-Control-Allow-Origin"))
}

View File

@@ -60,8 +60,34 @@ import (
"github.com/netbirdio/netbird/management/server/telemetry"
)
// CORSMiddleware returns a CORS handler restricted to allowedOrigins. When none are
// configured it falls back to allowing any origin, preserving the behaviour of
// deployments that serve the dashboard and the API on different origins.
// AllowCredentials stays false: the API authenticates via the Authorization header
// only, so there are no ambient credentials for a foreign origin to abuse.
func CORSMiddleware(allowedOrigins []string) *cors.Cors {
if len(allowedOrigins) == 0 {
log.Warn("no CORS allowed origins configured, allowing any origin; set HttpConfig.CORSAllowedOrigins to the dashboard origin to restrict it")
return cors.AllowAll()
}
return cors.New(cors.Options{
AllowedOrigins: allowedOrigins,
AllowedMethods: []string{
http.MethodHead,
http.MethodGet,
http.MethodPost,
http.MethodPut,
http.MethodPatch,
http.MethodDelete,
},
AllowedHeaders: []string{"Authorization", "Content-Type"},
AllowCredentials: false,
})
}
// NewAPIHandler creates the Management service HTTP API handler registering all the available endpoints.
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, corsAllowedOrigins []string, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
// Register bypass paths for unauthenticated endpoints
if err := bypass.AddBypassPath("/api/instance"); err != nil {
@@ -98,7 +124,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
isValidChildAccount,
)
corsMiddleware := cors.AllowAll()
corsMiddleware := CORSMiddleware(corsAllowedOrigins)
metricsMiddleware := appMetrics.HTTPMiddleware()

View File

@@ -137,7 +137,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(context.Background(), apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(context.Background(), apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}
@@ -267,7 +267,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(context.Background(), apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(context.Background(), apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}