From 6c17e87cf61cd9c62ea4ef4e25d4adb5603ffc88 Mon Sep 17 00:00:00 2001 From: boojack Date: Sun, 12 Jul 2026 23:34:33 +0800 Subject: [PATCH] fix(auth): support OAuth client auth auto-detection - support providers requiring client_secret_basic while preserving POST fallback - stop logging user-info claims and mapped profile data - cover both client authentication styles with PKCE --- internal/idp/oauth2/oauth2.go | 5 +- internal/idp/oauth2/oauth2_test.go | 75 ++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 4 deletions(-) diff --git a/internal/idp/oauth2/oauth2.go b/internal/idp/oauth2/oauth2.go index a6915e41..c281b6d5 100644 --- a/internal/idp/oauth2/oauth2.go +++ b/internal/idp/oauth2/oauth2.go @@ -6,7 +6,6 @@ import ( "encoding/json" "fmt" "io" - "log/slog" "net/http" "time" @@ -54,7 +53,7 @@ func (p *IdentityProvider) ExchangeToken(ctx context.Context, redirectURL, code, Endpoint: oauth2.Endpoint{ AuthURL: p.config.AuthUrl, TokenURL: p.config.TokenUrl, - AuthStyle: oauth2.AuthStyleInParams, + AuthStyle: oauth2.AuthStyleAutoDetect, }, } @@ -112,7 +111,6 @@ func (p *IdentityProvider) UserInfo(ctx context.Context, token string) (*idp.Ide if err := json.Unmarshal(body, &claims); err != nil { return nil, errors.Wrap(err, "failed to unmarshal response body") } - slog.Info("user info claims", "claims", claims) userInfo := &idp.IdentityProviderUserInfo{} if v, ok := claims[p.config.FieldMapping.Identifier].(string); ok { userInfo.Identifier = v @@ -140,6 +138,5 @@ func (p *IdentityProvider) UserInfo(ctx context.Context, token string) (*idp.Ide userInfo.AvatarURL = v } } - slog.Info("user info", "userInfo", userInfo) return userInfo, nil } diff --git a/internal/idp/oauth2/oauth2_test.go b/internal/idp/oauth2/oauth2_test.go index 0bec3d69..293a3e31 100644 --- a/internal/idp/oauth2/oauth2_test.go +++ b/internal/idp/oauth2/oauth2_test.go @@ -163,6 +163,81 @@ func TestIdentityProvider(t *testing.T) { assert.Equal(t, wantUserInfo, userInfoResult) } +func TestIdentityProviderExchangeTokenClientAuthentication(t *testing.T) { + const ( + clientID = "test-client-id" + clientSecret = "test-client-secret" + code = "test-code" + accessToken = "test-access-token" + codeVerifier = "test-code-verifier" + ) + + tests := []struct { + name string + acceptBasicAuth bool + expectedRequests int + }{ + { + name: "client secret basic", + acceptBasicAuth: true, + expectedRequests: 1, + }, + { + name: "client secret post fallback", + acceptBasicAuth: false, + expectedRequests: 2, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + requestCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestCount++ + require.NoError(t, r.ParseForm()) + require.Equal(t, code, r.Form.Get("code")) + require.Equal(t, codeVerifier, r.Form.Get("code_verifier")) + + username, password, hasBasicAuth := r.BasicAuth() + if test.acceptBasicAuth { + require.True(t, hasBasicAuth) + require.Equal(t, clientID, username) + require.Equal(t, clientSecret, password) + require.Empty(t, r.Form.Get("client_id")) + require.Empty(t, r.Form.Get("client_secret")) + } else if hasBasicAuth { + http.Error(w, `{"error":"invalid_client"}`, http.StatusUnauthorized) + return + } else { + require.Equal(t, clientID, r.Form.Get("client_id")) + require.Equal(t, clientSecret, r.Form.Get("client_secret")) + } + + w.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(w).Encode(map[string]any{ + "access_token": accessToken, + "token_type": "Bearer", + })) + })) + defer server.Close() + + provider, err := NewIdentityProvider(&storepb.OAuth2Config{ + ClientId: clientID, + ClientSecret: clientSecret, + TokenUrl: server.URL, + UserInfoUrl: "https://example.com/oauth2/userinfo", + FieldMapping: &storepb.FieldMapping{Identifier: "sub"}, + }) + require.NoError(t, err) + + token, err := provider.ExchangeToken(context.Background(), "https://example.com/auth/callback", code, codeVerifier) + require.NoError(t, err) + assert.Equal(t, accessToken, token) + assert.Equal(t, test.expectedRequests, requestCount) + }) + } +} + func TestIdentityProviderUserInfoUsesContext(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel()