// 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{} )