diff --git a/custom/conf/app.example.ini b/custom/conf/app.example.ini index 194868f..64667e4 100644 --- a/custom/conf/app.example.ini +++ b/custom/conf/app.example.ini @@ -1725,6 +1725,9 @@ LEVEL = Info ;; For more information about the possible values see https://openid.net/specs/openid-connect-core-1_0.html#ScopeClaims ;OPENID_CONNECT_SCOPES = ;; +;; Use Proof Key for Code Exchange (PKCE) with the S256 method for OpenID Connect login sources. +;ENABLE_OPENID_CONNECT_PKCE = false +;; ;; Automatically create user accounts for new oauth2 users. ;ENABLE_AUTO_REGISTRATION = false ;; diff --git a/modules/setting/oauth2.go b/modules/setting/oauth2.go index 0c2db10..f43a374 100644 --- a/modules/setting/oauth2.go +++ b/modules/setting/oauth2.go @@ -52,18 +52,20 @@ func (accountLinking OAuth2AccountLinkingType) isValid() bool { // OAuth2Client settings var OAuth2Client struct { - RegisterEmailConfirm bool - OpenIDConnectScopes []string - EnableAutoRegistration bool - Username OAuth2UsernameType - UpdateAvatar bool - AccountLinking OAuth2AccountLinkingType + RegisterEmailConfirm bool + OpenIDConnectScopes []string + EnableOpenIDConnectPKCE bool + EnableAutoRegistration bool + Username OAuth2UsernameType + UpdateAvatar bool + AccountLinking OAuth2AccountLinkingType } func loadOAuth2ClientFrom(rootCfg ConfigProvider) { sec := rootCfg.Section("oauth2_client") OAuth2Client.RegisterEmailConfirm = sec.Key("REGISTER_EMAIL_CONFIRM").MustBool(Service.RegisterEmailConfirm) OAuth2Client.OpenIDConnectScopes = parseScopes(sec, "OPENID_CONNECT_SCOPES") + OAuth2Client.EnableOpenIDConnectPKCE = sec.Key("ENABLE_OPENID_CONNECT_PKCE").MustBool() OAuth2Client.EnableAutoRegistration = sec.Key("ENABLE_AUTO_REGISTRATION").MustBool() OAuth2Client.Username = OAuth2UsernameType(sec.Key("USERNAME").MustString(string(OAuth2UsernameNickname))) if !OAuth2Client.Username.isValid() { diff --git a/modules/setting/oauth2_test.go b/modules/setting/oauth2_test.go index a2235c6..fccff21 100644 --- a/modules/setting/oauth2_test.go +++ b/modules/setting/oauth2_test.go @@ -76,3 +76,15 @@ DEFAULT_APPLICATIONS = loadOAuth2From(cfg) assert.Nil(t, OAuth2.DefaultApplications) } + +func TestOAuth2ClientOpenIDConnectPKCE(t *testing.T) { + cfg, _ := NewConfigProviderFromData(``) + loadOAuth2ClientFrom(cfg) + assert.False(t, OAuth2Client.EnableOpenIDConnectPKCE) + + cfg, _ = NewConfigProviderFromData(`[oauth2_client] +ENABLE_OPENID_CONNECT_PKCE = true +`) + loadOAuth2ClientFrom(cfg) + assert.True(t, OAuth2Client.EnableOpenIDConnectPKCE) +} diff --git a/services/auth/source/oauth2/providers.go b/services/auth/source/oauth2/providers.go index a0cb886..3fc3a3b 100644 --- a/services/auth/source/oauth2/providers.go +++ b/services/auth/source/oauth2/providers.go @@ -20,7 +20,6 @@ import ( "gitea.dev/modules/setting" "github.com/markbates/goth" - "github.com/markbates/goth/providers/openidConnect" ) // Provider is an interface for describing a single OAuth2 provider @@ -210,7 +209,7 @@ func GetOIDCEndSessionEndpoint(providerName string) string { return "" } - oidcProvider, ok := provider.(*openidConnect.Provider) + oidcProvider, ok := asOpenIDConnectProvider(provider) if !ok || oidcProvider.OpenIDConfig == nil { return "" } diff --git a/services/auth/source/oauth2/providers_openid.go b/services/auth/source/oauth2/providers_openid.go index 74bfa1a..d1f1ef1 100644 --- a/services/auth/source/oauth2/providers_openid.go +++ b/services/auth/source/oauth2/providers_openid.go @@ -53,6 +53,9 @@ func (o *OpenIDProvider) CreateGothProvider(providerName, callbackURL string, so // A single entry is sufficient because the admin explicitly chooses one claim (e.g. "oid" for Azure AD). provider.UserIdClaims = []string{source.ExternalIDClaim} } + if setting.OAuth2Client.EnableOpenIDConnectPKCE { + return newOpenIDConnectPKCEProvider(provider), nil + } return provider, nil } diff --git a/services/auth/source/oauth2/providers_openid_pkce.go b/services/auth/source/oauth2/providers_openid_pkce.go new file mode 100644 index 0000000..e14dfc1 --- /dev/null +++ b/services/auth/source/oauth2/providers_openid_pkce.go @@ -0,0 +1,139 @@ +// Copyright 2026 The Gitea Authors. All rights reserved. +// SPDX-License-Identifier: MIT + +package oauth2 + +import ( + "encoding/json" + "errors" + "fmt" + "net/url" + "strings" + + "github.com/markbates/goth" + "github.com/markbates/goth/providers/openidConnect" + go_oauth2 "golang.org/x/oauth2" +) + +type openIDConnectPKCEProvider struct { + *openidConnect.Provider +} + +type openIDConnectPKCESession struct { + Session openidConnect.Session `json:"session"` + CodeVerifier string `json:"code_verifier"` +} + +type openIDConnectPKCEParams struct { + goth.Params + codeVerifier string +} + +func newOpenIDConnectPKCEProvider(provider *openidConnect.Provider) goth.Provider { + return &openIDConnectPKCEProvider{Provider: provider} +} + +func asOpenIDConnectProvider(provider goth.Provider) (*openidConnect.Provider, bool) { + switch provider := provider.(type) { + case *openidConnect.Provider: + return provider, true + case *openIDConnectPKCEProvider: + return provider.Provider, true + default: + return nil, false + } +} + +func (p *openIDConnectPKCEProvider) BeginAuth(state string) (goth.Session, error) { + session, err := p.Provider.BeginAuth(state) + if err != nil { + return nil, err + } + + openIDSession, ok := session.(*openidConnect.Session) + if !ok { + return nil, fmt.Errorf("unexpected OpenID Connect session type %T", session) + } + + codeVerifier := go_oauth2.GenerateVerifier() + authURL, err := addOpenIDConnectPKCEChallenge(openIDSession.AuthURL, codeVerifier) + if err != nil { + return nil, err + } + openIDSession.AuthURL = authURL + + return &openIDConnectPKCESession{ + Session: *openIDSession, + CodeVerifier: codeVerifier, + }, nil +} + +func (p *openIDConnectPKCEProvider) UnmarshalSession(data string) (goth.Session, error) { + session := &openIDConnectPKCESession{} + if err := json.NewDecoder(strings.NewReader(data)).Decode(session); err != nil { + return nil, err + } + if session.CodeVerifier == "" { + return nil, errors.New("OpenID Connect PKCE session is missing the code verifier") + } + return session, nil +} + +func (p *openIDConnectPKCEProvider) FetchUser(session goth.Session) (goth.User, error) { + pkceSession, ok := session.(*openIDConnectPKCESession) + if !ok { + return goth.User{}, fmt.Errorf("unexpected OpenID Connect PKCE session type %T", session) + } + return p.Provider.FetchUser(&pkceSession.Session) +} + +func (s *openIDConnectPKCESession) GetAuthURL() (string, error) { + return s.Session.GetAuthURL() +} + +func (s *openIDConnectPKCESession) Authorize(provider goth.Provider, params goth.Params) (string, error) { + pkceProvider, ok := provider.(*openIDConnectPKCEProvider) + if !ok { + return "", fmt.Errorf("unexpected OpenID Connect PKCE provider type %T", provider) + } + if s.CodeVerifier == "" { + return "", errors.New("OpenID Connect PKCE session is missing the code verifier") + } + return s.Session.Authorize(pkceProvider.Provider, openIDConnectPKCEParams{ + Params: params, + codeVerifier: s.CodeVerifier, + }) +} + +func (s *openIDConnectPKCESession) Marshal() string { + data, _ := json.Marshal(s) + return string(data) +} + +func (s *openIDConnectPKCESession) String() string { + return s.Marshal() +} + +func (p openIDConnectPKCEParams) Get(key string) string { + if key == "code_verifier" { + return p.codeVerifier + } + return p.Params.Get(key) +} + +func addOpenIDConnectPKCEChallenge(rawURL, codeVerifier string) (string, error) { + authURL, err := url.Parse(rawURL) + if err != nil { + return "", fmt.Errorf("parse OpenID Connect authorization URL: %w", err) + } + query := authURL.Query() + query.Set("code_challenge", go_oauth2.S256ChallengeFromVerifier(codeVerifier)) + query.Set("code_challenge_method", "S256") + authURL.RawQuery = query.Encode() + return authURL.String(), nil +} + +var ( + _ goth.Provider = &openIDConnectPKCEProvider{} + _ goth.Session = &openIDConnectPKCESession{} +) diff --git a/services/auth/source/oauth2/providers_openid_pkce_test.go b/services/auth/source/oauth2/providers_openid_pkce_test.go new file mode 100644 index 0000000..e887f93 --- /dev/null +++ b/services/auth/source/oauth2/providers_openid_pkce_test.go @@ -0,0 +1,126 @@ +// Copyright 2026 The Gitea Authors. All rights reserved. +// SPDX-License-Identifier: MIT + +package oauth2 + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "gitea.dev/modules/setting" + "gitea.dev/modules/test" + + "github.com/markbates/goth/providers/openidConnect" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + go_oauth2 "golang.org/x/oauth2" +) + +func TestOpenIDConnectPKCEProvider(t *testing.T) { + var codeChallenge string + var codeVerifier string + var serverURL string + + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/.well-known/openid-configuration": + response.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(response).Encode(map[string]string{ + "issuer": serverURL, + "authorization_endpoint": serverURL + "/authorize", + "token_endpoint": serverURL + "/token", + })) + case "/token": + require.NoError(t, request.ParseForm()) + codeVerifier = request.Form.Get("code_verifier") + if codeVerifier == "" || go_oauth2.S256ChallengeFromVerifier(codeVerifier) != codeChallenge { + http.Error(response, "invalid PKCE verifier", http.StatusBadRequest) + return + } + claims, err := json.Marshal(map[string]any{ + "aud": "client-id", + "exp": time.Now().Add(time.Hour).Unix(), + "iss": serverURL, + "sub": "person-1", + }) + require.NoError(t, err) + response.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(response).Encode(map[string]any{ + "access_token": "access-token", + "expires_in": 3600, + "id_token": "e30." + base64.RawURLEncoding.EncodeToString(claims) + ".signature", + "token_type": "Bearer", + })) + default: + http.NotFound(response, request) + } + })) + defer server.Close() + serverURL = server.URL + + defer test.MockVariableValue(&setting.OAuth2Client.EnableOpenIDConnectPKCE, true)() + createdProvider, err := (&OpenIDProvider{}).CreateGothProvider( + "example", + "https://git.example.com/user/oauth2/example/callback", + &Source{ + ClientID: "client-id", + ClientSecret: "client-secret", + OpenIDConnectAutoDiscoveryURL: server.URL + "/.well-known/openid-configuration", + Scopes: []string{"openid"}, + }, + ) + require.NoError(t, err) + pkceProvider, ok := createdProvider.(*openIDConnectPKCEProvider) + require.True(t, ok) + pkceProvider.Provider.SkipUserInfoRequest = true + pkceProvider.SetName("example") + session, err := pkceProvider.BeginAuth("test-state") + require.NoError(t, err) + authURL, err := session.GetAuthURL() + require.NoError(t, err) + parsedAuthURL, err := url.Parse(authURL) + require.NoError(t, err) + assert.Equal(t, "S256", parsedAuthURL.Query().Get("code_challenge_method")) + codeChallenge = parsedAuthURL.Query().Get("code_challenge") + assert.NotEmpty(t, codeChallenge) + + restoredSession, err := pkceProvider.UnmarshalSession(session.Marshal()) + require.NoError(t, err) + accessToken, err := restoredSession.Authorize(pkceProvider, url.Values{"code": {"authorization-code"}}) + require.NoError(t, err) + assert.Equal(t, "access-token", accessToken) + assert.NotEmpty(t, codeVerifier) + + user, err := pkceProvider.FetchUser(restoredSession) + require.NoError(t, err) + assert.Equal(t, "person-1", user.UserID) + assert.Equal(t, "example", user.Provider) +} + +func TestOpenIDConnectPKCESessionRejectsMissingVerifier(t *testing.T) { + provider := &openIDConnectPKCEProvider{Provider: &openidConnect.Provider{}} + session := &openIDConnectPKCESession{} + + _, err := session.Authorize(provider, url.Values{"code": {"authorization-code"}}) + assert.EqualError(t, err, "OpenID Connect PKCE session is missing the code verifier") + _, err = provider.UnmarshalSession(fmt.Sprintf(`{"session":{"AuthURL":%q}}`, "https://id.example.com/authorize")) + assert.EqualError(t, err, "OpenID Connect PKCE session is missing the code verifier") +} + +func TestAsOpenIDConnectProvider(t *testing.T) { + provider := &openidConnect.Provider{} + + actual, ok := asOpenIDConnectProvider(provider) + require.True(t, ok) + assert.Same(t, provider, actual) + + actual, ok = asOpenIDConnectProvider(newOpenIDConnectPKCEProvider(provider)) + require.True(t, ok) + assert.Same(t, provider, actual) +}