diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index 1c78af9d0..16a8addb5 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -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 { diff --git a/management/internals/server/config/config.go b/management/internals/server/config/config.go index a77d5c19b..d55212b69 100644 --- a/management/internals/server/config/config.go +++ b/management/internals/server/config/config.go @@ -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) diff --git a/management/server/http/cors_test.go b/management/server/http/cors_test.go new file mode 100644 index 000000000..865afd990 --- /dev/null +++ b/management/server/http/cors_test.go @@ -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")) +} diff --git a/management/server/http/handler.go b/management/server/http/handler.go index a57f44b3c..22bba9f7c 100644 --- a/management/server/http/handler.go +++ b/management/server/http/handler.go @@ -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() diff --git a/management/server/http/testing/testing_tools/channel/channel.go b/management/server/http/testing/testing_tools/channel/channel.go index 8b05b2ddf..1f922075b 100644 --- a/management/server/http/testing/testing_tools/channel/channel.go +++ b/management/server/http/testing/testing_tools/channel/channel.go @@ -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) }