diff --git a/pkg/openid/client/login.go b/pkg/openid/client/login.go index 3af5c1e..65adcbf 100644 --- a/pkg/openid/client/login.go +++ b/pkg/openid/client/login.go @@ -113,13 +113,13 @@ func (c *Client) authCodeURL(ctx context.Context, authCodeParams openid.Authoriz ctx, span := otel.StartSpan(ctx, "Client.PushedAuthorizationRequest") defer span.End() - clientAuth, err := c.ClientAuthenticationParams() - if err != nil { - return "", fmt.Errorf("generating client authentication parameters: %w", err) - } - endpoint := c.cfg.Provider().PushedAuthorizationRequestEndpoint() resp, err := retry.DoValue(ctx, func(ctx context.Context) (*openid.PushedAuthorizationResponse, error) { + clientAuth, err := c.ClientAuthenticationParams() + if err != nil { + return nil, fmt.Errorf("generating client authentication parameters: %w", err) + } + body, err := c.oauthPostRequest(ctx, endpoint, authCodeParams.RequestParams().With(clientAuth)) if err != nil { if errors.Is(err, ErrOpenIDServer) { diff --git a/pkg/openid/client/login_test.go b/pkg/openid/client/login_test.go index 37176ad..e0a893a 100644 --- a/pkg/openid/client/login_test.go +++ b/pkg/openid/client/login_test.go @@ -1,9 +1,13 @@ package client_test import ( + "net/http" + "net/http/httptest" "net/url" + "sync" "testing" + "github.com/lestrrat-go/jwx/v3/jwt" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/oauth2" @@ -36,6 +40,68 @@ func TestLogin_PushedAuthorizationRequest(t *testing.T) { assert.ElementsMatch(t, query["client_id"], []string{idp.OpenIDConfig.Client().ClientID()}) } +func TestLogin_PushedAuthorizationRequest_RetryMintsNewAssertion(t *testing.T) { + var mu sync.Mutex + var assertions []string + + attempts := 0 + parServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + + mu.Lock() + assertions = append(assertions, r.PostForm.Get("client_assertion")) + attempts++ + attempt := attempts + mu.Unlock() + + if attempt == 1 { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"request_uri":"urn:ietf:params:oauth:request_uri:some-uri","expires_in":60}`)) + })) + defer parServer.Close() + + cfg := mock.Config() + idp := mock.NewIdentityProvider(cfg) + defer idp.Close() + idp.OpenIDConfig.TestProvider.SetPushedAuthorizationRequestEndpoint(parServer.URL) + + req := idp.GetRequest(mock.Ingress + "/oauth2/login") + result, err := idp.RelyingPartyHandler.Client.Login(req) + require.NoError(t, err) + + parsed, err := url.Parse(result.AuthCodeURL) + require.NoError(t, err) + assert.Equal(t, "urn:ietf:params:oauth:request_uri:some-uri", parsed.Query().Get("request_uri")) + + mu.Lock() + defer mu.Unlock() + require.Len(t, assertions, 2, "expected the 503 to be retried exactly once") + + parseJti := func(t *testing.T, assertion string) string { + t.Helper() + + tok, err := jwt.ParseString( + assertion, + jwt.WithVerify(false), // the assertion is signed for the IdP, not for us + jwt.WithValidate(false), + ) + require.NoError(t, err) + + jti, ok := tok.JwtID() + require.True(t, ok) + return jti + } + + first := parseJti(t, assertions[0]) + second := parseJti(t, assertions[1]) + assert.NotEmpty(t, first) + assert.NotEqual(t, first, second, "retry reused the client assertion; each attempt must mint a new jti") +} + func TestLogin_URL(t *testing.T) { type loginURLTest struct { name string