From 82b1c7da222382a5f1013a3fabc0965a703109f4 Mon Sep 17 00:00:00 2001 From: Bethuel Mmbaga Date: Fri, 11 Sep 2026 19:09:13 +0300 Subject: [PATCH] [management] Harden OIDC issuer validation and discovery (#7435) --- management/server/identity_provider.go | 37 ++++++++------ management/server/identity_provider_test.go | 47 ++++++++++++++++- management/server/types/identity_provider.go | 12 ++++- .../server/types/identity_provider_test.go | 51 +++++++++++++++++++ 4 files changed, 129 insertions(+), 18 deletions(-) diff --git a/management/server/identity_provider.go b/management/server/identity_provider.go index 86bbcd893..764d8598b 100644 --- a/management/server/identity_provider.go +++ b/management/server/identity_provider.go @@ -23,6 +23,9 @@ import ( "github.com/netbirdio/netbird/shared/management/status" ) +// maxDiscoveryDocumentSize caps the discovery document read at 1 MiB. Providers serve a few kilobytes. +const maxDiscoveryDocumentSize = 1 << 20 + // oidcProviderJSON represents the OpenID Connect discovery document type oidcProviderJSON struct { Issuer string `json:"issuer"` @@ -35,6 +38,10 @@ func validateOIDCIssuer(ctx context.Context, issuer string) error { httpClient := &http.Client{ Timeout: 10 * time.Second, + // An issuer that redirects its own discovery document is misconfigured. + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, } req, err := http.NewRequestWithContext(ctx, http.MethodGet, wellKnown, nil) @@ -48,22 +55,22 @@ func validateOIDCIssuer(ctx context.Context, issuer string) error { } defer resp.Body.Close() - body, err := io.ReadAll(resp.Body) - if err != nil { - return fmt.Errorf("%w: unable to read response body: %v", types.ErrIdentityProviderIssuerUnreachable, err) + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("%w: %s", types.ErrIdentityProviderIssuerUnreachable, resp.Status) } - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("%w: %s: %s", types.ErrIdentityProviderIssuerUnreachable, resp.Status, body) + body, err := io.ReadAll(io.LimitReader(resp.Body, maxDiscoveryDocumentSize+1)) + if err != nil || len(body) > maxDiscoveryDocumentSize { + return fmt.Errorf("%w: failed to decode provider discovery object", types.ErrIdentityProviderIssuerUnreachable) } var p oidcProviderJSON if err := json.Unmarshal(body, &p); err != nil { - return fmt.Errorf("%w: failed to decode provider discovery object: %v", types.ErrIdentityProviderIssuerUnreachable, err) + return fmt.Errorf("%w: failed to decode provider discovery object", types.ErrIdentityProviderIssuerUnreachable) } if p.Issuer != issuer { - return fmt.Errorf("%w: expected %q got %q", types.ErrIdentityProviderIssuerMismatch, issuer, p.Issuer) + return fmt.Errorf("%w: %q", types.ErrIdentityProviderIssuerMismatch, issuer) } return nil @@ -151,15 +158,15 @@ func (am *DefaultAccountManager) CreateIdentityProvider(ctx context.Context, acc return nil, status.NewPermissionDeniedError() } - if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil { - return nil, err - } - embeddedManager, ok := am.idpManager.(*idp.EmbeddedIdPManager) if !ok { return nil, status.Errorf(status.Internal, "identity provider management requires embedded IdP") } + if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil { + return nil, err + } + // Generate ID if not provided if idpConfig.ID == "" { idpConfig.ID = generateIdentityProviderID(idpConfig.Type) @@ -188,15 +195,15 @@ func (am *DefaultAccountManager) UpdateIdentityProvider(ctx context.Context, acc return nil, status.NewPermissionDeniedError() } - if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil { - return nil, err - } - embeddedManager, ok := am.idpManager.(*idp.EmbeddedIdPManager) if !ok { return nil, status.Errorf(status.Internal, "identity provider management requires embedded IdP") } + if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil { + return nil, err + } + idpConfig.ID = idpID idpConfig.AccountID = accountID diff --git a/management/server/identity_provider_test.go b/management/server/identity_provider_test.go index ecc47337c..c7a8af1d2 100644 --- a/management/server/identity_provider_test.go +++ b/management/server/identity_provider_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "path/filepath" + "strings" "testing" "time" @@ -121,7 +122,7 @@ func createManagerWithEmbeddedIdPModeAndSetup( } func TestDefaultAccountManager_CreateIdentityProvider_Validation(t *testing.T) { - manager, _, err := createManager(t) + manager, _, err := createManagerWithEmbeddedIdP(t) require.NoError(t, err) userID := "testingUser" @@ -233,7 +234,7 @@ func TestUpdateUserAuthWithSingleModeKeepsConfiguredDomain(t *testing.T) { } func TestDefaultAccountManager_UpdateIdentityProvider_Validation(t *testing.T) { - manager, _, err := createManager(t) + manager, _, err := createManagerWithEmbeddedIdP(t) require.NoError(t, err) userID := "testingUser" @@ -355,3 +356,45 @@ func TestValidateOIDCIssuer_TrailingSlash(t *testing.T) { require.Error(t, err) assert.True(t, errors.Is(err, types.ErrIdentityProviderIssuerMismatch)) } + +func TestValidateOIDCIssuer_DoesNotFollowRedirects(t *testing.T) { + var reached bool + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + reached = true + w.WriteHeader(http.StatusForbidden) + })) + t.Cleanup(target.Close) + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, target.URL+"/redirect-target", http.StatusFound) + })) + t.Cleanup(srv.Close) + + err := validateOIDCIssuer(context.Background(), srv.URL) + require.Error(t, err) + assert.False(t, reached, "Redirects are not followed") +} + +func TestValidateOIDCIssuer_BoundsResponseSize(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"issuer":"` + strings.Repeat("a", maxDiscoveryDocumentSize) + `"}`)) + })) + t.Cleanup(srv.Close) + + err := validateOIDCIssuer(context.Background(), srv.URL) + require.ErrorIs(t, err, types.ErrIdentityProviderIssuerUnreachable) + assert.NotErrorIs(t, err, types.ErrIdentityProviderIssuerMismatch) +} + +func TestValidateOIDCIssuer_RejectsTrailingContent(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"issuer":"http://` + r.Host + `"} {"issuer":"second"}`)) + })) + t.Cleanup(srv.Close) + + err := validateOIDCIssuer(context.Background(), srv.URL) + require.ErrorIs(t, err, types.ErrIdentityProviderIssuerUnreachable, + "Content after the first object is not a valid discovery document") +} diff --git a/management/server/types/identity_provider.go b/management/server/types/identity_provider.go index 0c1f9509c..f75b5319d 100644 --- a/management/server/types/identity_provider.go +++ b/management/server/types/identity_provider.go @@ -3,6 +3,7 @@ package types import ( "errors" "net/url" + "strings" ) // Identity provider validation errors @@ -99,7 +100,16 @@ func (idp *IdentityProvider) Validate() error { } if idp.Issuer != "" { parsedURL, err := url.Parse(idp.Issuer) - if err != nil || parsedURL.Scheme == "" || parsedURL.Host == "" { + if err != nil || parsedURL.Host == "" { + return ErrIdentityProviderIssuerInvalid + } + if parsedURL.Scheme != "https" { + return ErrIdentityProviderIssuerInvalid + } + if parsedURL.User != nil { + return ErrIdentityProviderIssuerInvalid + } + if strings.ContainsAny(idp.Issuer, "?#") { return ErrIdentityProviderIssuerInvalid } } diff --git a/management/server/types/identity_provider_test.go b/management/server/types/identity_provider_test.go index 6ddc563f2..53385c5ac 100644 --- a/management/server/types/identity_provider_test.go +++ b/management/server/types/identity_provider_test.go @@ -135,3 +135,54 @@ func TestIdentityProvider_Validate(t *testing.T) { }) } } + +func TestIdentityProvider_ValidateRejectsNonOriginIssuers(t *testing.T) { + issuers := []string{ + "https://idp.example.com/realms/nb?foo=bar", + "https://idp.example.com/realms/nb#section", + "https://user:pass@idp.example.com", + "ftp://idp.example.com", + "ldap://idp.example.com", + "http://idp.example.com", + } + + for _, issuer := range issuers { + t.Run(issuer, func(t *testing.T) { + idp := &IdentityProvider{ + Name: "test", + Type: IdentityProviderTypeOIDC, + Issuer: issuer, + ClientID: "client-id", + } + assert.ErrorIs(t, idp.Validate(), ErrIdentityProviderIssuerInvalid) + }) + } +} + +func TestIdentityProvider_ValidateAcceptsOriginAndPath(t *testing.T) { + for _, issuer := range []string{"https://idp.example.com", "https://idp.example.com/realms/nb", "https://127.0.0.1:5556/dex"} { + t.Run(issuer, func(t *testing.T) { + idp := &IdentityProvider{ + Name: "test", + Type: IdentityProviderTypeOIDC, + Issuer: issuer, + ClientID: "client-id", + } + assert.NoError(t, idp.Validate()) + }) + } +} + +func TestIdentityProviderValidateRejectsBareDelimiters(t *testing.T) { + for _, issuer := range []string{"https://idp.example.com/realms/nb?", "https://idp.example.com/realms/nb#"} { + t.Run(issuer, func(t *testing.T) { + idp := &IdentityProvider{ + Name: "test", + Type: IdentityProviderTypeOIDC, + Issuer: issuer, + ClientID: "client-id", + } + assert.ErrorIs(t, idp.Validate(), ErrIdentityProviderIssuerInvalid) + }) + } +}