diff --git a/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx index d12aa75c9..010e98d64 100644 --- a/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx +++ b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx @@ -70,10 +70,10 @@ function OIDCButtons() { onClick={() => { window.location.href = provider.loginURL - + "?continue=" - + encodeURIComponent( - safeContinueUrl.pathname + safeContinueUrl.search, - ); + + "?continue=" + + encodeURIComponent( + safeContinueUrl.pathname + safeContinueUrl.search, + ); }} > diff --git a/pkg/bootstrap/builder.go b/pkg/bootstrap/builder.go index 10bd9d225..805520d95 100644 --- a/pkg/bootstrap/builder.go +++ b/pkg/bootstrap/builder.go @@ -113,12 +113,12 @@ func (b *Builder) Build() (*probod.FullConfig, error) { Google: probod.OIDCProviderConfig{ ClientID: b.getEnv("AUTH_GOOGLE_CLIENT_ID"), ClientSecret: b.getEnv("AUTH_GOOGLE_CLIENT_SECRET"), - Enabled: b.getEnv("AUTH_GOOGLE_CLIENT_ID") != "", + Enabled: b.getEnv("AUTH_GOOGLE_CLIENT_ID") != "" && b.getEnv("AUTH_GOOGLE_CLIENT_SECRET") != "", }, Microsoft: probod.OIDCProviderConfig{ ClientID: b.getEnv("AUTH_MICROSOFT_CLIENT_ID"), ClientSecret: b.getEnv("AUTH_MICROSOFT_CLIENT_SECRET"), - Enabled: b.getEnv("AUTH_MICROSOFT_CLIENT_ID") != "", + Enabled: b.getEnv("AUTH_MICROSOFT_CLIENT_ID") != "" && b.getEnv("AUTH_MICROSOFT_CLIENT_SECRET") != "", }, }, TrustCenter: probod.TrustCenterConfig{ diff --git a/pkg/iam/oidc/service.go b/pkg/iam/oidc/service.go index 95885824e..96585ad2b 100644 --- a/pkg/iam/oidc/service.go +++ b/pkg/iam/oidc/service.go @@ -34,6 +34,7 @@ import ( "sync" "time" + "go.gearno.de/kit/httpclient" "go.gearno.de/kit/log" "go.gearno.de/kit/pg" "go.probo.inc/probo/pkg/coredata" @@ -64,10 +65,11 @@ type ( } Service struct { - pg *pg.Client - baseURL string - logger *log.Logger - providers map[coredata.OIDCProvider]*providerInfo + pg *pg.Client + baseURL string + logger *log.Logger + httpClient *http.Client + providers map[coredata.OIDCProvider]*providerInfo jwksMu sync.RWMutex jwksCache map[string]*jwksEntry @@ -157,11 +159,12 @@ func NewService( logger *log.Logger, ) *Service { s := &Service{ - pg: pgClient, - baseURL: baseURL, - logger: logger.Named("oidc"), - providers: make(map[coredata.OIDCProvider]*providerInfo), - jwksCache: make(map[string]*jwksEntry), + pg: pgClient, + baseURL: baseURL, + logger: logger.Named("oidc"), + httpClient: httpclient.DefaultPooledClient(httpclient.WithLogger(logger)), + providers: make(map[coredata.OIDCProvider]*providerInfo), + jwksCache: make(map[string]*jwksEntry), } if google.Enabled { @@ -488,7 +491,7 @@ func (s *Service) verifyAndParseIDToken(ctx context.Context, info *providerInfo, } if claims.Nonce != expectedNonce { - return nil, fmt.Errorf("cannot validate nonce: expected %q, got %q", expectedNonce, claims.Nonce) + return nil, fmt.Errorf("cannot validate nonce: mismatch") } if time.Now().After(time.Unix(int64(claims.ExpiresAt), 0)) { @@ -504,7 +507,7 @@ func (s *Service) getSigningKey(ctx context.Context, jwksURL string, kid string) s.jwksMu.RUnlock() if !ok || time.Since(entry.fetchedAt) > jwksCacheTTL { - keys, err := fetchJWKS(ctx, jwksURL) + keys, err := fetchJWKS(ctx, s.httpClient, jwksURL) if err != nil { return nil, err } @@ -522,7 +525,7 @@ func (s *Service) getSigningKey(ctx context.Context, jwksURL string, kid string) } // Key not found in cache, try refreshing - keys, err := fetchJWKS(ctx, jwksURL) + keys, err := fetchJWKS(ctx, s.httpClient, jwksURL) if err != nil { return nil, err } @@ -541,13 +544,13 @@ func (s *Service) getSigningKey(ctx context.Context, jwksURL string, kid string) return nil, fmt.Errorf("cannot find signing key %q in JWKS", kid) } -func fetchJWKS(ctx context.Context, jwksURL string) ([]jwk, error) { +func fetchJWKS(ctx context.Context, httpClient *http.Client, jwksURL string) ([]jwk, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, jwksURL, nil) if err != nil { return nil, fmt.Errorf("cannot create jwks request: %w", err) } - resp, err := http.DefaultClient.Do(req) + resp, err := httpClient.Do(req) if err != nil { return nil, fmt.Errorf("cannot fetch jwks: %w", err) } diff --git a/pkg/iam/saml/gc.go b/pkg/iam/saml/gc.go index d66a75226..c44826f59 100644 --- a/pkg/iam/saml/gc.go +++ b/pkg/iam/saml/gc.go @@ -71,6 +71,10 @@ func (gc *GarbageCollector) Run(ctx context.Context) error { gc.logger.ErrorCtx(ctx, "cannot run initial cleanup", log.Error(err)) } + if gc.interval <= 0 { + return fmt.Errorf("cannot run SAML garbage collector: interval must be greater than zero") + } + ticker := time.NewTicker(gc.interval) defer ticker.Stop() diff --git a/pkg/iam/saml_domain_verifier.go b/pkg/iam/saml_domain_verifier.go index 89ae5f0f7..a35588fc1 100644 --- a/pkg/iam/saml_domain_verifier.go +++ b/pkg/iam/saml_domain_verifier.go @@ -69,6 +69,10 @@ func (v *SAMLDomainVerifier) Run(ctx context.Context) error { v.runOnce(ctx) + if v.interval <= 0 { + return fmt.Errorf("cannot run SAML domain verifier: interval must be greater than zero") + } + ticker := time.NewTicker(v.interval) defer ticker.Stop()