mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-23 23:29:08 +02:00
configurable cors settings
This commit is contained in:
@@ -13,7 +13,6 @@ import (
|
|||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
grpcMiddleware "github.com/grpc-ecosystem/go-grpc-middleware/v2"
|
grpcMiddleware "github.com/grpc-ecosystem/go-grpc-middleware/v2"
|
||||||
"github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip"
|
"github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip"
|
||||||
"github.com/rs/cors"
|
|
||||||
"github.com/rs/xid"
|
"github.com/rs/xid"
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
@@ -24,13 +23,13 @@ import (
|
|||||||
|
|
||||||
"github.com/netbirdio/netbird/encryption"
|
"github.com/netbirdio/netbird/encryption"
|
||||||
"github.com/netbirdio/netbird/formatter/hook"
|
"github.com/netbirdio/netbird/formatter/hook"
|
||||||
|
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||||
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
|
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
|
||||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||||
"github.com/netbirdio/netbird/management/server/activity"
|
"github.com/netbirdio/netbird/management/server/activity"
|
||||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
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"
|
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||||
nbContext "github.com/netbirdio/netbird/management/server/context"
|
nbContext "github.com/netbirdio/netbird/management/server/context"
|
||||||
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
||||||
@@ -122,7 +121,7 @@ func (s *BaseServer) EventStore() activity.Store {
|
|||||||
|
|
||||||
func (s *BaseServer) APIHandler() http.Handler {
|
func (s *BaseServer) APIHandler() http.Handler {
|
||||||
return Create(s, func() 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 {
|
if err != nil {
|
||||||
log.Fatalf("failed to create API handler: %v", err)
|
log.Fatalf("failed to create API handler: %v", err)
|
||||||
}
|
}
|
||||||
@@ -137,7 +136,7 @@ func (s *BaseServer) IDPHandler() http.Handler {
|
|||||||
if !ok || embeddedIdP == nil {
|
if !ok || embeddedIdP == nil {
|
||||||
return 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 {
|
func (s *BaseServer) Router() *mux.Router {
|
||||||
|
|||||||
@@ -125,6 +125,9 @@ type HttpServerConfig struct {
|
|||||||
ExtraAuthAudience string
|
ExtraAuthAudience string
|
||||||
// AuthCallbackDomain contains the callback domain
|
// AuthCallbackDomain contains the callback domain
|
||||||
AuthCallbackURL string
|
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)
|
// Host represents a Netbird host (e.g. STUN, TURN, Signal)
|
||||||
|
|||||||
@@ -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"))
|
||||||
|
}
|
||||||
@@ -60,8 +60,34 @@ import (
|
|||||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
"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.
|
// 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
|
// Register bypass paths for unauthenticated endpoints
|
||||||
if err := bypass.AddBypassPath("/api/instance"); err != nil {
|
if err := bypass.AddBypassPath("/api/instance"); err != nil {
|
||||||
@@ -98,7 +124,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
|
|||||||
isValidChildAccount,
|
isValidChildAccount,
|
||||||
)
|
)
|
||||||
|
|
||||||
corsMiddleware := cors.AllowAll()
|
corsMiddleware := CORSMiddleware(corsAllowedOrigins)
|
||||||
|
|
||||||
metricsMiddleware := appMetrics.HTTPMiddleware()
|
metricsMiddleware := appMetrics.HTTPMiddleware()
|
||||||
|
|
||||||
|
|||||||
@@ -137,7 +137,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
|
|||||||
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
|
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
|
||||||
|
|
||||||
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
|
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 {
|
if err != nil {
|
||||||
t.Fatalf("Failed to create API handler: %v", err)
|
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)
|
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
|
||||||
|
|
||||||
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
|
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 {
|
if err != nil {
|
||||||
t.Fatalf("Failed to create API handler: %v", err)
|
t.Fatalf("Failed to create API handler: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user