diff --git a/cmd/wonderwall/main.go b/cmd/wonderwall/main.go index 16fbf80..cc91cb0 100644 --- a/cmd/wonderwall/main.go +++ b/cmd/wonderwall/main.go @@ -85,7 +85,7 @@ func run() error { } func standalone(ctx context.Context, cfg *config.Config, crypt crypto.Crypter) (*handler.Standalone, error) { - openidConfig, err := openidconfig.NewConfig(cfg) + openidConfig, err := openidconfig.NewConfig(ctx, cfg) if err != nil { return nil, err } diff --git a/pkg/openid/config/config.go b/pkg/openid/config/config.go index 220b1e8..8b5e9a2 100644 --- a/pkg/openid/config/config.go +++ b/pkg/openid/config/config.go @@ -1,6 +1,8 @@ package config import ( + "context" + wonderwallconfig "github.com/nais/wonderwall/pkg/config" ) @@ -22,13 +24,13 @@ func (c *openidconfig) Provider() Provider { return c.providerConfig } -func NewConfig(cfg *wonderwallconfig.Config) (Config, error) { +func NewConfig(ctx context.Context, cfg *wonderwallconfig.Config) (Config, error) { clientCfg, err := NewClientConfig(cfg) if err != nil { return nil, err } - providerCfg, err := NewProviderConfig(cfg) + providerCfg, err := NewProviderConfig(ctx, cfg) if err != nil { return nil, err } diff --git a/pkg/openid/config/provider.go b/pkg/openid/config/provider.go index 031759d..de5f77e 100644 --- a/pkg/openid/config/provider.go +++ b/pkg/openid/config/provider.go @@ -1,15 +1,18 @@ package config import ( + "context" "encoding/json" "fmt" "net/http" "net/url" "slices" + "time" "github.com/lestrrat-go/jwx/v3/jwa" log "github.com/sirupsen/logrus" + httpinternal "github.com/nais/wonderwall/internal/http" "github.com/nais/wonderwall/pkg/config" "github.com/nais/wonderwall/pkg/openid/acr" ) @@ -83,13 +86,30 @@ func (p *provider) SidClaimRequired() bool { return p.metadata.FrontchannelLogoutSupported && p.metadata.FrontchannelLogoutSessionSupported } -func NewProviderConfig(cfg *config.Config) (Provider, error) { - response, err := http.Get(cfg.OpenID.WellKnownURL) +// wellKnownTimeout bounds the fetch of the provider's metadata document at startup. +const wellKnownTimeout = 10 * time.Second + +func NewProviderConfig(ctx context.Context, cfg *config.Config) (Provider, error) { + ctx, cancel := context.WithTimeout(ctx, wellKnownTimeout) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, cfg.OpenID.WellKnownURL, nil) + if err != nil { + return nil, fmt.Errorf("creating request for well known configuration: %w", err) + } + + client := &http.Client{Transport: httpinternal.Transport()} + + response, err := client.Do(req) if err != nil { return nil, fmt.Errorf("fetching well known configuration: %w", err) } defer func() { _ = response.Body.Close() }() + if response.StatusCode != http.StatusOK { + return nil, fmt.Errorf("fetching well known configuration: %s responded with HTTP %d", cfg.OpenID.WellKnownURL, response.StatusCode) + } + providerCfg := new(ProviderMetadata) if err := json.NewDecoder(response.Body).Decode(providerCfg); err != nil { return nil, fmt.Errorf("decoding well known configuration: %w", err)