diff --git a/coderd/oauth2provider/tokens.go b/coderd/oauth2provider/tokens.go index 461dc7c7825..7df380a8ffc 100644 --- a/coderd/oauth2provider/tokens.go +++ b/coderd/oauth2provider/tokens.go @@ -188,17 +188,9 @@ func extractTokenRequest(r *http.Request, logger slog.Logger, primary *url.URL, Scope: p.String(vals, "", "scope"), } - // RFC 6749 §2.3.1: confidential clients may authenticate via HTTP Basic. - if user, pass, ok := r.BasicAuth(); ok && user != "" { - if req.ClientID != "" && req.ClientID != user { - return codersdk.OAuth2TokenRequest{}, nil, errConflictingClientAuth - } - if req.ClientSecret != "" && req.ClientSecret != pass { - return codersdk.OAuth2TokenRequest{}, nil, errConflictingClientAuth - } - - req.ClientID = user - req.ClientSecret = pass + req.ClientID, req.ClientSecret, err = mergeBasicClientAuth(r, req.ClientID, req.ClientSecret) + if err != nil { + return codersdk.OAuth2TokenRequest{}, nil, err } // Grant-specific required checks that can be satisfied via HTTP Basic. @@ -255,6 +247,57 @@ func extractTokenRequest(r *http.Request, logger slog.Logger, primary *url.URL, return req, nil, nil } +// mergeBasicClientAuth combines a confidential client's HTTP Basic +// credentials (RFC 6749 §2.3.1) with the form client_id and client_secret. +// Without a Basic header, or with an empty Basic username, the form values +// pass through unchanged. Otherwise each form field must be empty or equal +// to its header counterpart, or the result is errConflictingClientAuth. An +// empty Basic password still counts as a presented password, so a form +// secret beside it is a conflict. +func mergeBasicClientAuth(r *http.Request, clientID, clientSecret string) (mergedID, mergedSecret string, err error) { + user, pass, ok := r.BasicAuth() + if !ok || user == "" { + return clientID, clientSecret, nil + } + if clientID != "" && clientID != user { + return "", "", errConflictingClientAuth + } + if clientSecret != "" && clientSecret != pass { + return "", "", errConflictingClientAuth + } + return user, pass, nil +} + +// authenticateClient checks a client secret and confirms it belongs to the +// app named by client_id. That id arrives unverified, so without the app +// check a valid secret for one app could issue a token for another. It +// returns the matched secret row. Every authentication failure returns +// errBadSecret so the response does not reveal which step failed; a +// datastore failure returns the underlying error. Callers skip it for public +// clients, which have no secret and are bound by PKCE and the token's app id +// instead. +func authenticateClient(ctx context.Context, db database.Store, app database.OAuth2ProviderApp, clientSecret string) (database.OAuth2ProviderAppSecret, error) { + secret, err := ParseFormattedSecret(clientSecret) + if err != nil { + return database.OAuth2ProviderAppSecret{}, errBadSecret + } + //nolint:gocritic // OAuth2 system context, users cannot read secrets + dbSecret, err := db.GetOAuth2ProviderAppSecretByPrefix(dbauthz.AsSystemOAuth2(ctx), []byte(secret.Prefix)) + if errors.Is(err, sql.ErrNoRows) { + return database.OAuth2ProviderAppSecret{}, errBadSecret + } + if err != nil { + return database.OAuth2ProviderAppSecret{}, err + } + if !apikey.ValidateHash(dbSecret.HashedSecret, secret.Secret) { + return database.OAuth2ProviderAppSecret{}, errBadSecret + } + if dbSecret.AppID != app.ID { + return database.OAuth2ProviderAppSecret{}, errBadSecret + } + return dbSecret, nil +} + // writeTokenError renders an RFC 6749 §5.2 error body. Descriptions can quote // what the client sent, so they are confined and capped here rather than at each // call site, leaving the guarantee with the endpoint. @@ -427,31 +470,10 @@ func authorizationCodeGrant(ctx context.Context, db database.Store, logger slog. // client instead. var appSecretID uuid.NullUUID if !app.IsPublic() { - secret, err := ParseFormattedSecret(req.ClientSecret) - if err != nil { - return codersdk.OAuth2TokenResponse{}, errBadSecret - } - //nolint:gocritic // OAuth2 system context, users cannot read secrets - dbSecret, err := db.GetOAuth2ProviderAppSecretByPrefix(dbauthz.AsSystemOAuth2(ctx), []byte(secret.Prefix)) - if errors.Is(err, sql.ErrNoRows) { - return codersdk.OAuth2TokenResponse{}, errBadSecret - } + dbSecret, err := authenticateClient(ctx, db, app, req.ClientSecret) if err != nil { return codersdk.OAuth2TokenResponse{}, err } - - equalSecret := apikey.ValidateHash(dbSecret.HashedSecret, secret.Secret) - if !equalSecret { - return codersdk.OAuth2TokenResponse{}, errBadSecret - } - - // The secret must belong to the app named by client_id, which arrives - // unverified in the request. Otherwise a valid secret for one app - // could issue a token for another. - if dbSecret.AppID != app.ID { - return codersdk.OAuth2TokenResponse{}, errBadSecret - } - appSecretID = uuid.NullUUID{UUID: dbSecret.ID, Valid: true} } diff --git a/coderd/oauth2provider/tokens_internal_test.go b/coderd/oauth2provider/tokens_internal_test.go index 3c787240c62..645f751006f 100644 --- a/coderd/oauth2provider/tokens_internal_test.go +++ b/coderd/oauth2provider/tokens_internal_test.go @@ -18,8 +18,11 @@ import ( "cdr.dev/slog/v3/sloggers/slogjson" "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" + "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/testutil" ) // parseScopes parses a space-delimited scope string into a slice of scopes @@ -914,6 +917,106 @@ func TestExtractTokenRequest_UnrecognizedParametersLogged(t *testing.T) { } } +// Every failure is errBadSecret so the caller cannot tell a malformed secret +// from a valid one for the wrong app (RFC 6749 §5.2). The OtherAppsSecret row +// covers the belongs-to-app check, the step a retyped copy of the check would +// be most likely to lose. +func TestAuthenticateClient(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + + seed := func() (database.OAuth2ProviderApp, database.OAuth2ProviderAppSecret, string) { + app := dbgen.OAuth2ProviderApp(t, db, database.OAuth2ProviderApp{}) + secret, err := GenerateSecret() + require.NoError(t, err) + dbSecret := dbgen.OAuth2ProviderAppSecret(t, db, database.OAuth2ProviderAppSecret{ + AppID: app.ID, + SecretPrefix: []byte(secret.Prefix), + HashedSecret: secret.Hashed, + }) + return app, dbSecret, secret.Formatted + } + app, dbSecret, formatted := seed() + _, _, otherFormatted := seed() + + unknown, err := GenerateSecret() + require.NoError(t, err) + parsed, err := ParseFormattedSecret(formatted) + require.NoError(t, err) + + tests := []struct { + name string + secret string + want error + }{ + {name: "Empty", secret: "", want: errBadSecret}, + {name: "Malformed", secret: "not-a-secret", want: errBadSecret}, + {name: "UnknownPrefix", secret: unknown.Formatted, want: errBadSecret}, + {name: "WrongHash", secret: SecretIdentifier + "_" + parsed.Prefix + "_" + unknown.Secret, want: errBadSecret}, + {name: "OtherAppsSecret", secret: otherFormatted, want: errBadSecret}, + {name: "OwnSecret", secret: formatted}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + + got, err := authenticateClient(ctx, db, app, test.secret) + if test.want != nil { + require.ErrorIs(t, err, test.want) + return + } + require.NoError(t, err) + require.Equal(t, dbSecret.ID, got.ID) + }) + } +} + +func TestMergeBasicClientAuth(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + basicUser string + basicPass string + bodyID string + bodySecret string + wantID string + wantSecret string + wantErr error + }{ + {name: "NoHeader", bodyID: "id", bodySecret: "s", wantID: "id", wantSecret: "s"}, + {name: "HeaderOnly", basicUser: "id", basicPass: "s", wantID: "id", wantSecret: "s"}, + {name: "HeaderAndMatchingBody", basicUser: "id", basicPass: "s", bodyID: "id", bodySecret: "s", wantID: "id", wantSecret: "s"}, + // A header with an empty username is treated as no header at all. + {name: "HeaderEmptyUser", basicPass: "pass", bodyID: "id", bodySecret: "s", wantID: "id", wantSecret: "s"}, + // An empty Basic password is still a presented password, so a body + // secret beside it is a conflict rather than a fallback. + {name: "HeaderEmptyPasswordBodySecret", basicUser: "id", bodySecret: "s", wantErr: errConflictingClientAuth}, + {name: "ConflictingID", basicUser: "id", basicPass: "s", bodyID: "other", wantErr: errConflictingClientAuth}, + {name: "ConflictingSecret", basicUser: "id", basicPass: "s", bodySecret: "other", wantErr: errConflictingClientAuth}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + r := &http.Request{Header: http.Header{}} + if test.basicUser != "" || test.basicPass != "" { + r.SetBasicAuth(test.basicUser, test.basicPass) + } + id, secret, err := mergeBasicClientAuth(r, test.bodyID, test.bodySecret) + if test.wantErr != nil { + require.ErrorIs(t, err, test.wantErr) + return + } + require.NoError(t, err) + require.Equal(t, test.wantID, id) + require.Equal(t, test.wantSecret, secret) + }) + } +} + // TestRefreshTokenGrant_Scopes tests that scopes can be requested during refresh func TestRefreshTokenGrant_Scopes(t *testing.T) { t.Parallel()