fix: 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)
|
||||
|
||||
@@ -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<GridLayoutGroup>(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<Component>())
|
||||
{
|
||||
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
|
||||
|
||||
+1
-1
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user