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
This commit is contained in:
boojack
2026-07-12 23:34:33 +08:00
parent 4bc3928029
commit 6c17e87cf6
2 changed files with 76 additions and 4 deletions
+1 -4
View File
@@ -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
}
+75
View File
@@ -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()