Skip to content
Draft
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
7 changes: 3 additions & 4 deletions management/internals/server/boot.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand Down Expand Up @@ -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)
}
Expand All @@ -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 {
Expand Down
3 changes: 3 additions & 0 deletions management/internals/server/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
55 changes: 55 additions & 0 deletions management/server/http/cors_test.go
Original file line number Diff line number Diff line change
@@ -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"))
}
30 changes: 28 additions & 2 deletions management/server/http/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -60,8 +60,34 @@
"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) {

Check warning on line 90 in management/server/http/handler.go

View check run for this annotation

SonarQubeCloud / SonarCloud Code Analysis

This function has 25 parameters, which is greater than the 10 authorized.

See more on https://sonarcloud.io/project/issues?id=netbirdio_netbird&issues=AZ-uN_YosYyBWXkvpg5e&open=AZ-uN_YosYyBWXkvpg5e&pullRequest=6964

// Register bypass paths for unauthenticated endpoints
if err := bypass.AddBypassPath("/api/instance"); err != nil {
Expand Down Expand Up @@ -98,7 +124,7 @@
isValidChildAccount,
)

corsMiddleware := cors.AllowAll()
corsMiddleware := CORSMiddleware(corsAllowedOrigins)

metricsMiddleware := appMetrics.HTTPMiddleware()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down
Loading