mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-12 17:59:06 +02:00
Merge branch 'main' into fix/custom-domain-deletion-guard
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user