fix(all): diagnose OAuth failures and center login providers
This commit is contained in:
@@ -11,6 +11,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
@@ -290,7 +291,22 @@ func (s *Service) callback(w http.ResponseWriter, r *http.Request) {
|
||||
clear(verifier)
|
||||
clear(nonce)
|
||||
if err != nil {
|
||||
_, _ = s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_rejected' WHERE id=? AND status='authorizing'`, id)
|
||||
_, _ = s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_rejected',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND status='authorizing'`, id)
|
||||
var failure *providerFailure
|
||||
if errors.As(err, &failure) {
|
||||
slog.Warn("OAuth provider authorization failed",
|
||||
"provider", failure.Provider,
|
||||
"stage", failure.Stage,
|
||||
"reason", failure.Reason,
|
||||
"http_status", failure.HTTPStatus,
|
||||
"oauth_error", failure.OAuthError)
|
||||
if failure.OAuthError == "invalid_client" {
|
||||
http.Error(w, "server OAuth configuration is invalid; contact the server administrator", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
slog.Warn("OAuth provider authorization failed", "provider", provider, "reason", "internal_error")
|
||||
}
|
||||
http.Error(w, "provider authorization failed", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
@@ -311,28 +327,81 @@ func (s *Service) callback(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
type providerIdentity struct{ issuer, subject string }
|
||||
|
||||
type providerFailure struct {
|
||||
Provider string
|
||||
Stage string
|
||||
Reason string
|
||||
HTTPStatus int
|
||||
OAuthError string
|
||||
}
|
||||
|
||||
func (e *providerFailure) Error() string {
|
||||
return fmt.Sprintf("provider=%s stage=%s reason=%s status=%d oauth_error=%s", e.Provider, e.Stage, e.Reason, e.HTTPStatus, e.OAuthError)
|
||||
}
|
||||
|
||||
func networkProviderFailure(ctx context.Context, provider, stage string, err error) error {
|
||||
reason := "network_error"
|
||||
if errors.Is(ctx.Err(), context.DeadlineExceeded) || errors.Is(err, context.DeadlineExceeded) {
|
||||
reason = "timeout"
|
||||
} else if errors.Is(ctx.Err(), context.Canceled) || errors.Is(err, context.Canceled) {
|
||||
reason = "cancelled"
|
||||
} else {
|
||||
var networkError net.Error
|
||||
if errors.As(err, &networkError) && networkError.Timeout() {
|
||||
reason = "timeout"
|
||||
}
|
||||
}
|
||||
return &providerFailure{Provider: provider, Stage: stage, Reason: reason}
|
||||
}
|
||||
|
||||
func rejectedProviderFailure(provider, stage string, response *http.Response) error {
|
||||
failure := &providerFailure{Provider: provider, Stage: stage, Reason: "http_rejected", HTTPStatus: response.StatusCode}
|
||||
var body struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
decoder := json.NewDecoder(io.LimitReader(response.Body, 8<<10))
|
||||
if decoder.Decode(&body) == nil {
|
||||
failure.OAuthError = safeOAuthError(body.Error)
|
||||
}
|
||||
return failure
|
||||
}
|
||||
|
||||
func invalidProviderResponse(provider, stage string) error {
|
||||
return &providerFailure{Provider: provider, Stage: stage, Reason: "invalid_response", HTTPStatus: http.StatusOK}
|
||||
}
|
||||
|
||||
func safeOAuthError(value string) string {
|
||||
switch value {
|
||||
case "invalid_request", "invalid_client", "invalid_grant", "unauthorized_client",
|
||||
"unsupported_grant_type", "invalid_scope", "access_denied", "server_error", "temporarily_unavailable":
|
||||
return value
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) exchangeIdentity(ctx context.Context, provider, code, verifier, nonce string) (providerIdentity, error) {
|
||||
values := url.Values{"client_id": {s.config.Providers[provider].ClientID}, "client_secret": {s.config.ProviderSecrets[provider]}, "grant_type": {"authorization_code"}, "code": {code}, "redirect_uri": {s.redirectURL(provider)}, "code_verifier": {verifier}}
|
||||
request, _ := http.NewRequestWithContext(ctx, http.MethodPost, providerTokenURL(provider), strings.NewReader(values.Encode()))
|
||||
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
response, err := s.client.Do(request)
|
||||
if err != nil {
|
||||
return providerIdentity{}, err
|
||||
return providerIdentity{}, networkProviderFailure(ctx, provider, "token_exchange", err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return providerIdentity{}, errors.New("token exchange rejected")
|
||||
return providerIdentity{}, rejectedProviderFailure(provider, "token_exchange", response)
|
||||
}
|
||||
var token struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
IDToken string `json:"id_token"`
|
||||
}
|
||||
if err := decodeProviderJSON(response.Body, &token); err != nil || token.AccessToken == "" {
|
||||
return providerIdentity{}, errors.New("invalid token response")
|
||||
return providerIdentity{}, invalidProviderResponse(provider, "token_exchange")
|
||||
}
|
||||
if provider == "google" {
|
||||
if token.IDToken == "" {
|
||||
return providerIdentity{}, errors.New("Google ID token missing")
|
||||
return providerIdentity{}, invalidProviderResponse(provider, "token_exchange")
|
||||
}
|
||||
identity, err := s.verifyGoogleIDToken(ctx, token.IDToken, nonce)
|
||||
token.AccessToken, token.IDToken = "", ""
|
||||
@@ -343,23 +412,23 @@ func (s *Service) exchangeIdentity(ctx context.Context, provider, code, verifier
|
||||
response, err = s.client.Do(userinfo)
|
||||
token.AccessToken = ""
|
||||
if err != nil {
|
||||
return providerIdentity{}, err
|
||||
return providerIdentity{}, networkProviderFailure(ctx, provider, "userinfo", err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return providerIdentity{}, errors.New("userinfo rejected")
|
||||
return providerIdentity{}, rejectedProviderFailure(provider, "userinfo", response)
|
||||
}
|
||||
var user struct {
|
||||
ID string `json:"id"`
|
||||
Sub string `json:"sub"`
|
||||
}
|
||||
if err := decodeProviderJSON(response.Body, &user); err != nil {
|
||||
return providerIdentity{}, err
|
||||
return providerIdentity{}, invalidProviderResponse(provider, "userinfo")
|
||||
}
|
||||
if provider == "discord" && user.ID != "" {
|
||||
return providerIdentity{issuer: "https://discord.com", subject: user.ID}, nil
|
||||
}
|
||||
return providerIdentity{}, errors.New("provider subject missing")
|
||||
return providerIdentity{}, invalidProviderResponse(provider, "userinfo")
|
||||
}
|
||||
|
||||
// verifyGoogleIDToken delegates signature and standard-claim verification to
|
||||
@@ -370,11 +439,11 @@ func (s *Service) verifyGoogleIDToken(ctx context.Context, idToken, nonce string
|
||||
request, _ := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
response, err := s.client.Do(request)
|
||||
if err != nil {
|
||||
return providerIdentity{}, err
|
||||
return providerIdentity{}, networkProviderFailure(ctx, "google", "id_token_verify", err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return providerIdentity{}, errors.New("Google ID token rejected")
|
||||
return providerIdentity{}, rejectedProviderFailure("google", "id_token_verify", response)
|
||||
}
|
||||
var claims struct {
|
||||
Issuer string `json:"iss"`
|
||||
@@ -384,12 +453,12 @@ func (s *Service) verifyGoogleIDToken(ctx context.Context, idToken, nonce string
|
||||
Expires string `json:"exp"`
|
||||
}
|
||||
if err := decodeProviderJSON(response.Body, &claims); err != nil {
|
||||
return providerIdentity{}, err
|
||||
return providerIdentity{}, invalidProviderResponse("google", "id_token_verify")
|
||||
}
|
||||
expires, err := strconv.ParseInt(claims.Expires, 10, 64)
|
||||
validIssuer := claims.Issuer == "https://accounts.google.com" || claims.Issuer == "accounts.google.com"
|
||||
if err != nil || !validIssuer || claims.Audience != s.config.Providers["google"].ClientID || claims.Subject == "" || claims.Nonce != nonce || s.store.now().Unix() >= expires {
|
||||
return providerIdentity{}, errors.New("Google ID token claims rejected")
|
||||
return providerIdentity{}, &providerFailure{Provider: "google", Stage: "id_token_verify", Reason: "invalid_claims"}
|
||||
}
|
||||
return providerIdentity{issuer: "https://accounts.google.com", subject: claims.Subject}, nil
|
||||
}
|
||||
|
||||
@@ -451,6 +451,94 @@ func TestProviderIdentityVerification(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderFailureIsStructuredAndSanitized(t *testing.T) {
|
||||
service, _ := testService(t)
|
||||
secretDescription := "provider leaked secret sentinel"
|
||||
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
return jsonResponse(http.StatusUnauthorized, `{"error":"invalid_client","error_description":"`+secretDescription+`"}`), nil
|
||||
})}
|
||||
_, err := service.exchangeIdentity(context.Background(), "discord", "code", "verifier", "nonce")
|
||||
var failure *providerFailure
|
||||
if !errors.As(err, &failure) {
|
||||
t.Fatalf("error %T does not expose provider failure", err)
|
||||
}
|
||||
if failure.Provider != "discord" || failure.Stage != "token_exchange" || failure.Reason != "http_rejected" ||
|
||||
failure.HTTPStatus != http.StatusUnauthorized || failure.OAuthError != "invalid_client" {
|
||||
t.Fatalf("failure=%+v", failure)
|
||||
}
|
||||
if strings.Contains(err.Error(), secretDescription) {
|
||||
t.Fatal("provider error description leaked through diagnostic error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderFailureRejectsUntrustedOAuthError(t *testing.T) {
|
||||
service, _ := testService(t)
|
||||
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
return jsonResponse(http.StatusBadRequest, `{"error":"access-token-sentinel"}`), nil
|
||||
})}
|
||||
_, err := service.exchangeIdentity(context.Background(), "discord", "code", "verifier", "nonce")
|
||||
var failure *providerFailure
|
||||
if !errors.As(err, &failure) || failure.OAuthError != "unknown" {
|
||||
t.Fatalf("failure=%+v err=%v", failure, err)
|
||||
}
|
||||
if strings.Contains(err.Error(), "access-token-sentinel") {
|
||||
t.Fatal("untrusted provider error leaked through diagnostic error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGoogleProviderNetworkFailureDoesNotLeakIDTokenURL(t *testing.T) {
|
||||
service, _ := testService(t)
|
||||
idToken := "signed-id-token-sentinel"
|
||||
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
if request.URL.Path == "/token" {
|
||||
return jsonResponse(http.StatusOK, `{"access_token":"provider-access","id_token":"`+idToken+`"}`), nil
|
||||
}
|
||||
return nil, &url.Error{Op: "Get", URL: "https://oauth2.googleapis.com/tokeninfo?id_token=" + idToken, Err: errors.New("transport sentinel")}
|
||||
})}
|
||||
_, err := service.exchangeIdentity(context.Background(), "google", "code", "verifier", "nonce")
|
||||
var failure *providerFailure
|
||||
if !errors.As(err, &failure) || failure.Provider != "google" || failure.Stage != "id_token_verify" || failure.Reason != "network_error" {
|
||||
t.Fatalf("failure=%+v err=%v", failure, err)
|
||||
}
|
||||
if strings.Contains(err.Error(), idToken) || strings.Contains(err.Error(), "transport sentinel") {
|
||||
t.Fatal("Google ID token URL or transport details leaked through diagnostic error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackFailureClearsShortLivedOAuthMaterial(t *testing.T) {
|
||||
service, store := testService(t)
|
||||
insertAuthorizingDevice(t, store, "failed-device", "discord")
|
||||
state := "failed-state"
|
||||
verifier, err := store.seal("failed-device", "pkce", []byte("verifier"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
nonce, err := store.seal("failed-device", "nonce", []byte("nonce"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := store.db.Exec(`UPDATE devices SET state_hash=?,verifier_cipher=?,nonce_cipher=? WHERE id='failed-device'`, store.digest("oauth-state", state), verifier, nonce); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
return jsonResponse(http.StatusUnauthorized, `{"error":"invalid_client"}`), nil
|
||||
})}
|
||||
request := httptest.NewRequest(http.MethodGet, "/auth/discord/callback?code=failed-code&state="+url.QueryEscape(state), nil)
|
||||
response := httptest.NewRecorder()
|
||||
service.Handler().ServeHTTP(response, request)
|
||||
if response.Code != http.StatusBadGateway || !strings.Contains(response.Body.String(), "server OAuth configuration is invalid") {
|
||||
t.Fatalf("status=%d body=%q", response.Code, response.Body.String())
|
||||
}
|
||||
var status string
|
||||
var stateHash, verifierCipher, nonceCipher []byte
|
||||
if err := store.db.QueryRow(`SELECT status,COALESCE(state_hash,X''),COALESCE(verifier_cipher,X''),COALESCE(nonce_cipher,X'') FROM devices WHERE id='failed-device'`).Scan(&status, &stateHash, &verifierCipher, &nonceCipher); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if status != "failed" || len(stateHash) != 0 || len(verifierCipher) != 0 || len(nonceCipher) != 0 {
|
||||
t.Fatalf("status=%q state=%d verifier=%d nonce=%d", status, len(stateHash), len(verifierCipher), len(nonceCipher))
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderScopesUseLeastPrivilege(t *testing.T) {
|
||||
if got := providerScope("discord"); got != "identify" {
|
||||
t.Fatalf("Discord scope=%q, want identify", got)
|
||||
|
||||
Reference in New Issue
Block a user