From 50f1ebbb5ee543a2c8c1ffa56ea11e6b23f7336a Mon Sep 17 00:00:00 2001 From: Flechazo <2558755403@qq.com> Date: Wed, 30 Sep 2026 18:29:38 +0800 Subject: [PATCH] fix(all): diagnose OAuth failures and center login providers --- go/internal/server/auth/service.go | 95 ++++++++++++++++++++---- go/internal/server/auth/service_test.go | 88 ++++++++++++++++++++++ plugins/LoginUI/Plugin.cs | 97 +++++++++++++++++++++++++ versions.json | 2 +- 4 files changed, 268 insertions(+), 14 deletions(-) diff --git a/go/internal/server/auth/service.go b/go/internal/server/auth/service.go index c17b4e3..baeac66 100644 --- a/go/internal/server/auth/service.go +++ b/go/internal/server/auth/service.go @@ -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 } diff --git a/go/internal/server/auth/service_test.go b/go/internal/server/auth/service_test.go index 2d7b124..f5515a0 100644 --- a/go/internal/server/auth/service_test.go +++ b/go/internal/server/auth/service_test.go @@ -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) diff --git a/plugins/LoginUI/Plugin.cs b/plugins/LoginUI/Plugin.cs index f893879..5a6f576 100644 --- a/plugins/LoginUI/Plugin.cs +++ b/plugins/LoginUI/Plugin.cs @@ -297,6 +297,8 @@ public sealed class Plugin : BaseUnityPlugin google.gameObject.SetActive(ProviderEnabled("google")); discord.gameObject.SetActive(ProviderEnabled("discord")); + int providerCount = (google.gameObject.activeSelf ? 1 : 0) + + (discord.gameObject.activeSelf ? 1 : 0); discord.transform.SetSiblingIndex(0); google.transform.SetSiblingIndex(1); discord.gameObject.name = "Button - Discord"; @@ -304,6 +306,7 @@ public sealed class Plugin : BaseUnityPlugin ReplaceClick(google, introUI, "google"); ReplaceClick(discord, introUI, "discord"); ApplyDiscordBrand(box, logo, title); + ConfigureProviderGrid(panel, providerCount); Canvas.ForceUpdateCanvases(); if (panel is RectTransform panelRect) @@ -317,6 +320,100 @@ public sealed class Plugin : BaseUnityPlugin } } + private static void ConfigureProviderGrid(Transform panel, int providerCount) + { + GridLayoutGroup grid = panel.GetComponentInChildren(true); + if (grid == null) + { + throw new MissingMemberException("Login provider grid was not found"); + } + int columns = Math.Max(1, providerCount); + grid.constraint = GridLayoutGroup.Constraint.FixedColumnCount; + grid.constraintCount = columns; + if (grid.transform is RectTransform gridRect) + { + float width = grid.padding.horizontal + grid.cellSize.x * columns + + grid.spacing.x * Math.Max(0, columns - 1); + UpdateBetterGridProfiles(grid, columns); + UpdateBetterLocatorProfiles(grid, width); + gridRect.SetSizeWithCurrentAnchors(RectTransform.Axis.Horizontal, width); + LayoutRebuilder.ForceRebuildLayoutImmediate(gridRect); + } + } + + private static void UpdateBetterGridProfiles(GridLayoutGroup grid, int columns) + { + Type type = grid.GetType(); + if (type.FullName != "TheraBytes.BetterUi.BetterGridLayoutGroup") + { + return; + } + const BindingFlags flags = BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic; + UpdateBetterGridSettings(type.GetField("settingsFallback", flags)?.GetValue(grid), columns); + object collection = type.GetField("customSettings", flags)?.GetValue(grid); + IEnumerable items = collection?.GetType().GetProperty("Items", flags)?.GetValue(collection, null) as IEnumerable; + if (items == null) + { + return; + } + foreach (object settings in items) + { + UpdateBetterGridSettings(settings, columns); + } + } + + private static void UpdateBetterGridSettings(object settings, int columns) + { + if (settings == null) + { + return; + } + Type type = settings.GetType(); + FieldInfo constraint = type.GetField("Constraint", BindingFlags.Instance | BindingFlags.Public); + FieldInfo count = type.GetField("ConstraintCount", BindingFlags.Instance | BindingFlags.Public); + if (constraint != null) + { + constraint.SetValue(settings, Enum.ToObject(constraint.FieldType, (int)GridLayoutGroup.Constraint.FixedColumnCount)); + } + count?.SetValue(settings, columns); + } + + private static void UpdateBetterLocatorProfiles(GridLayoutGroup grid, float width) + { + const BindingFlags flags = BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic; + foreach (Component component in grid.GetComponents()) + { + Type type = component?.GetType(); + if (type?.FullName != "TheraBytes.BetterUi.BetterLocator") + { + continue; + } + UpdateBetterRectTransformData(type.GetField("transformFallback", flags)?.GetValue(component), width); + object collection = type.GetField("transformConfigs", flags)?.GetValue(component); + IEnumerable items = collection?.GetType().GetProperty("Items", flags)?.GetValue(collection, null) as IEnumerable; + if (items == null) + { + continue; + } + foreach (object data in items) + { + UpdateBetterRectTransformData(data, width); + } + } + } + + private static void UpdateBetterRectTransformData(object data, float width) + { + FieldInfo sizeField = data?.GetType().GetField("SizeDelta", BindingFlags.Instance | BindingFlags.Public); + if (sizeField == null || sizeField.FieldType != typeof(Vector2)) + { + return; + } + Vector2 size = (Vector2)sizeField.GetValue(data); + size.x = width; + sizeField.SetValue(data, size); + } + private static void ReplaceClick(Button button, object introUI, string provider) { // Assigning a fresh event removes both serialized persistent calls and diff --git a/versions.json b/versions.json index 9f14ca4..67ff049 100644 --- a/versions.json +++ b/versions.json @@ -6,6 +6,6 @@ "plugins": { "local_identity": "0.6.0", "capture_environment": "0.2.0", - "login_ui": "0.1.0" + "login_ui": "0.1.1" } }