mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: implement oauth2 RFC 7009 token revocation endpoint (#20362)
Adds RFC 7009 token revocation endpoint
This commit is contained in:
@@ -26,12 +26,6 @@ import (
|
||||
"github.com/coder/coder/v2/cryptorand"
|
||||
)
|
||||
|
||||
// Constants for OAuth2 secret generation (RFC 7591)
|
||||
const (
|
||||
secretLength = 40 // Length of the actual secret part
|
||||
displaySecretLength = 6 // Length of visible part in UI (last 6 characters)
|
||||
)
|
||||
|
||||
// CreateDynamicClientRegistration returns an http.HandlerFunc that handles POST /oauth2/register
|
||||
func CreateDynamicClientRegistration(db database.Store, accessURL *url.URL, auditor *audit.Auditor, logger slog.Logger) http.HandlerFunc {
|
||||
return func(rw http.ResponseWriter, r *http.Request) {
|
||||
@@ -121,7 +115,7 @@ func CreateDynamicClientRegistration(db database.Store, accessURL *url.URL, audi
|
||||
}
|
||||
|
||||
// Create client secret - parse the formatted secret to get components
|
||||
parsedSecret, err := parseFormattedSecret(clientSecret)
|
||||
parsedSecret, err := ParseFormattedSecret(clientSecret)
|
||||
if err != nil {
|
||||
writeOAuth2RegistrationError(ctx, rw, http.StatusInternalServerError,
|
||||
"server_error", "Failed to parse generated secret")
|
||||
@@ -132,7 +126,7 @@ func CreateDynamicClientRegistration(db database.Store, accessURL *url.URL, audi
|
||||
_, err = db.InsertOAuth2ProviderAppSecret(dbauthz.AsSystemRestricted(ctx), database.InsertOAuth2ProviderAppSecretParams{
|
||||
ID: uuid.New(),
|
||||
CreatedAt: now,
|
||||
SecretPrefix: []byte(parsedSecret.prefix),
|
||||
SecretPrefix: []byte(parsedSecret.Prefix),
|
||||
HashedSecret: []byte(hashedSecret),
|
||||
DisplaySecret: createDisplaySecret(clientSecret),
|
||||
AppID: clientID,
|
||||
@@ -551,27 +545,6 @@ func writeOAuth2RegistrationError(_ context.Context, rw http.ResponseWriter, sta
|
||||
_ = json.NewEncoder(rw).Encode(errorResponse)
|
||||
}
|
||||
|
||||
// parsedSecret represents the components of a formatted OAuth2 secret
|
||||
type parsedSecret struct {
|
||||
prefix string
|
||||
secret string
|
||||
}
|
||||
|
||||
// parseFormattedSecret parses a formatted secret like "coder_prefix_secret"
|
||||
func parseFormattedSecret(secret string) (parsedSecret, error) {
|
||||
parts := strings.Split(secret, "_")
|
||||
if len(parts) != 3 {
|
||||
return parsedSecret{}, xerrors.Errorf("incorrect number of parts: %d", len(parts))
|
||||
}
|
||||
if parts[0] != "coder" {
|
||||
return parsedSecret{}, xerrors.Errorf("incorrect scheme: %s", parts[0])
|
||||
}
|
||||
return parsedSecret{
|
||||
prefix: parts[1],
|
||||
secret: parts[2],
|
||||
}, nil
|
||||
}
|
||||
|
||||
// createDisplaySecret creates a display version of the secret showing only the last few characters
|
||||
func createDisplaySecret(secret string) string {
|
||||
if len(secret) <= displaySecretLength {
|
||||
|
||||
@@ -1,15 +1,214 @@
|
||||
package oauth2provider
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/coderd/httpmw"
|
||||
"github.com/coder/coder/v2/coderd/userpassword"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrTokenNotBelongsToClient is returned when a token does not belong to the requesting client
|
||||
ErrTokenNotBelongsToClient = xerrors.New("token does not belong to requesting client")
|
||||
// ErrInvalidTokenFormat is returned when a token has an invalid format
|
||||
ErrInvalidTokenFormat = xerrors.New("invalid token format")
|
||||
)
|
||||
|
||||
// RevokeToken implements RFC 7009 OAuth2 Token Revocation
|
||||
// Authentication is unique for this endpoint in that it does not use the
|
||||
// standard token authentication middleware. Instead, it expects the token that
|
||||
// is being revoked to be valid.
|
||||
// TODO: Currently the token validation occurs in the revocation logic itself.
|
||||
// This code should be refactored to share token validation logic with other parts
|
||||
// of the OAuth2 provider/http middleware.
|
||||
func RevokeToken(db database.Store, logger slog.Logger) http.HandlerFunc {
|
||||
return func(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
app := httpmw.OAuth2ProviderApp(r)
|
||||
|
||||
// RFC 7009 requires POST method with application/x-www-form-urlencoded
|
||||
if r.Method != http.MethodPost {
|
||||
httpapi.WriteOAuth2Error(ctx, rw, http.StatusMethodNotAllowed, "invalid_request", "Method not allowed")
|
||||
return
|
||||
}
|
||||
|
||||
if err := r.ParseForm(); err != nil {
|
||||
httpapi.WriteOAuth2Error(ctx, rw, http.StatusBadRequest, "invalid_request", "Invalid form data")
|
||||
return
|
||||
}
|
||||
|
||||
// RFC 7009 requires 'token' parameter
|
||||
token := r.Form.Get("token")
|
||||
if token == "" {
|
||||
httpapi.WriteOAuth2Error(ctx, rw, http.StatusBadRequest, "invalid_request", "Missing token parameter")
|
||||
return
|
||||
}
|
||||
|
||||
// Determine if this is a refresh token (starts with "coder_") or API key
|
||||
// APIKeys do not have the SecretIdentifier prefix.
|
||||
const coderPrefix = SecretIdentifier + "_"
|
||||
isRefreshToken := strings.HasPrefix(token, coderPrefix)
|
||||
|
||||
// Revoke the token with ownership verification
|
||||
err := db.InTx(func(tx database.Store) error {
|
||||
if isRefreshToken {
|
||||
// Handle refresh token revocation
|
||||
return revokeRefreshTokenInTx(ctx, tx, token, app.ID)
|
||||
}
|
||||
// Handle API key revocation
|
||||
return revokeAPIKeyInTx(ctx, tx, token, app.ID)
|
||||
}, nil)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrTokenNotBelongsToClient) {
|
||||
// RFC 7009: Return success even if token doesn't belong to client (don't reveal token existence)
|
||||
logger.Debug(ctx, "token revocation failed: token does not belong to requesting client",
|
||||
slog.F("client_id", app.ID.String()),
|
||||
slog.F("app_name", app.Name))
|
||||
rw.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ErrInvalidTokenFormat) {
|
||||
// Invalid token format should return 400 bad request
|
||||
logger.Debug(ctx, "token revocation failed: invalid token format",
|
||||
slog.F("client_id", app.ID.String()),
|
||||
slog.F("app_name", app.Name))
|
||||
httpapi.WriteOAuth2Error(ctx, rw, http.StatusBadRequest, "invalid_request", "Invalid token format")
|
||||
return
|
||||
}
|
||||
logger.Error(ctx, "token revocation failed with internal server error",
|
||||
slog.Error(err),
|
||||
slog.F("client_id", app.ID.String()),
|
||||
slog.F("app_name", app.Name))
|
||||
httpapi.WriteOAuth2Error(ctx, rw, http.StatusInternalServerError, "server_error", "Internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
// RFC 7009: successful revocation returns HTTP 200
|
||||
rw.WriteHeader(http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
func revokeRefreshTokenInTx(ctx context.Context, db database.Store, token string, appID uuid.UUID) error {
|
||||
// Parse the refresh token using the existing function
|
||||
parsedToken, err := ParseFormattedSecret(token)
|
||||
if err != nil {
|
||||
return ErrInvalidTokenFormat
|
||||
}
|
||||
|
||||
// Try to find refresh token by prefix
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
dbToken, err := db.GetOAuth2ProviderAppTokenByPrefix(dbauthz.AsSystemOAuth2(ctx), []byte(parsedToken.Prefix))
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
// Token not found - return success per RFC 7009 (don't reveal token existence)
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("get oauth2 provider app token by prefix: %w", err)
|
||||
}
|
||||
|
||||
equal, err := userpassword.Compare(string(dbToken.RefreshHash), parsedToken.Secret)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("invalid refresh token: %w", err)
|
||||
}
|
||||
if !equal {
|
||||
return xerrors.Errorf("invalid refresh token")
|
||||
}
|
||||
|
||||
// Verify ownership
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
appSecret, err := db.GetOAuth2ProviderAppSecretByID(dbauthz.AsSystemOAuth2(ctx), dbToken.AppSecretID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get oauth2 provider app secret: %w", err)
|
||||
}
|
||||
if appSecret.AppID != appID {
|
||||
return ErrTokenNotBelongsToClient
|
||||
}
|
||||
|
||||
// Delete the associated API key, which should cascade to remove the refresh token
|
||||
// According to RFC 7009, when a refresh token is revoked, associated access tokens should be invalidated
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
err = db.DeleteAPIKeyByID(dbauthz.AsSystemOAuth2(ctx), dbToken.APIKeyID)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return xerrors.Errorf("delete api key: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func revokeAPIKeyInTx(ctx context.Context, db database.Store, token string, appID uuid.UUID) error {
|
||||
keyID, secret, err := httpmw.SplitAPIToken(token)
|
||||
if err != nil {
|
||||
return ErrInvalidTokenFormat
|
||||
}
|
||||
|
||||
// Get the API key
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
apiKey, err := db.GetAPIKeyByID(dbauthz.AsSystemOAuth2(ctx), keyID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
// API key not found - return success per RFC 7009 (don't reveal token existence)
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("get api key by id: %w", err)
|
||||
}
|
||||
|
||||
// Checking to see if the provided secret matches the stored hashed secret
|
||||
hashedSecret := sha256.Sum256([]byte(secret))
|
||||
if subtle.ConstantTimeCompare(apiKey.HashedSecret, hashedSecret[:]) != 1 {
|
||||
return xerrors.Errorf("invalid api key")
|
||||
}
|
||||
|
||||
// Verify the API key was created by OAuth2
|
||||
if apiKey.LoginType != database.LoginTypeOAuth2ProviderApp {
|
||||
return xerrors.New("api key is not an oauth2 token")
|
||||
}
|
||||
|
||||
// Find the associated OAuth2 token to verify ownership
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
dbToken, err := db.GetOAuth2ProviderAppTokenByAPIKeyID(dbauthz.AsSystemOAuth2(ctx), apiKey.ID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
// No associated OAuth2 token - return success per RFC 7009
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("get oauth2 provider app token by api key id: %w", err)
|
||||
}
|
||||
|
||||
// Verify the token belongs to the requesting app
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
appSecret, err := db.GetOAuth2ProviderAppSecretByID(dbauthz.AsSystemOAuth2(ctx), dbToken.AppSecretID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get oauth2 provider app secret for api key verification: %w", err)
|
||||
}
|
||||
|
||||
if appSecret.AppID != appID {
|
||||
return ErrTokenNotBelongsToClient
|
||||
}
|
||||
|
||||
// Delete the API key
|
||||
//nolint:gocritic // Using AsSystemOAuth2 for OAuth2 public token revocation endpoint
|
||||
err = db.DeleteAPIKeyByID(dbauthz.AsSystemOAuth2(ctx), apiKey.ID)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return xerrors.Errorf("delete api key for revocation: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func RevokeApp(db database.Store) http.HandlerFunc {
|
||||
return func(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
@@ -2,32 +2,68 @@ package oauth2provider
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/userpassword"
|
||||
"github.com/coder/coder/v2/cryptorand"
|
||||
)
|
||||
|
||||
type AppSecret struct {
|
||||
// Formatted contains the secret. This value is owned by the client, not the
|
||||
// server. It is formatted to include the prefix.
|
||||
Formatted string
|
||||
// Prefix is the ID of this secret owned by the server. When a client uses a
|
||||
// secret, this is the matching string to do a lookup on the hashed value. We
|
||||
// cannot use the hashed value directly because the server does not store the
|
||||
// salt.
|
||||
Prefix string
|
||||
const (
|
||||
// SecretIdentifier is the prefix added to all generated secrets.
|
||||
SecretIdentifier = "coder"
|
||||
)
|
||||
|
||||
// Constants for OAuth2 secret generation
|
||||
const (
|
||||
secretLength = 40 // Length of the actual secret part
|
||||
displaySecretLength = 6 // Length of visible part in UI (last 6 characters)
|
||||
)
|
||||
|
||||
type HashedAppSecret struct {
|
||||
AppSecret
|
||||
// Hashed is the server stored hash(secret,salt,...). Used for verifying a
|
||||
// secret.
|
||||
Hashed string
|
||||
}
|
||||
|
||||
type AppSecret struct {
|
||||
// Formatted contains the secret. This value is owned by the client, not the
|
||||
// server. It is formatted to include the prefix.
|
||||
Formatted string
|
||||
// Secret is the raw secret value. This value should only be known to the client.
|
||||
Secret string
|
||||
// Prefix is the ID of this secret owned by the server. When a client uses a
|
||||
// secret, this is the matching string to do a lookup on the hashed value. We
|
||||
// cannot use the hashed value directly because the server does not store the
|
||||
// salt.
|
||||
Prefix string
|
||||
}
|
||||
|
||||
// ParseFormattedSecret parses a formatted secret like "coder_<prefix>_<secret"
|
||||
func ParseFormattedSecret(formatted string) (AppSecret, error) {
|
||||
parts := strings.Split(formatted, "_")
|
||||
if len(parts) != 3 {
|
||||
return AppSecret{}, xerrors.Errorf("incorrect number of parts: %d", len(parts))
|
||||
}
|
||||
if parts[0] != SecretIdentifier {
|
||||
return AppSecret{}, xerrors.Errorf("incorrect scheme: %s", parts[0])
|
||||
}
|
||||
return AppSecret{
|
||||
Formatted: formatted,
|
||||
Prefix: parts[1],
|
||||
Secret: parts[2],
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GenerateSecret generates a secret to be used as a client secret, refresh
|
||||
// token, or authorization code.
|
||||
func GenerateSecret() (AppSecret, error) {
|
||||
func GenerateSecret() (HashedAppSecret, error) {
|
||||
// 40 characters matches the length of GitHub's client secrets.
|
||||
secret, err := cryptorand.String(40)
|
||||
secret, err := cryptorand.String(secretLength)
|
||||
if err != nil {
|
||||
return AppSecret{}, err
|
||||
return HashedAppSecret{}, err
|
||||
}
|
||||
|
||||
// This ID is prefixed to the secret so it can be used to look up the secret
|
||||
@@ -35,17 +71,20 @@ func GenerateSecret() (AppSecret, error) {
|
||||
// will not have the salt.
|
||||
prefix, err := cryptorand.String(10)
|
||||
if err != nil {
|
||||
return AppSecret{}, err
|
||||
return HashedAppSecret{}, err
|
||||
}
|
||||
|
||||
hashed, err := userpassword.Hash(secret)
|
||||
if err != nil {
|
||||
return AppSecret{}, err
|
||||
return HashedAppSecret{}, err
|
||||
}
|
||||
|
||||
return AppSecret{
|
||||
Formatted: fmt.Sprintf("coder_%s_%s", prefix, secret),
|
||||
Prefix: prefix,
|
||||
Hashed: hashed,
|
||||
return HashedAppSecret{
|
||||
AppSecret: AppSecret{
|
||||
Formatted: fmt.Sprintf("%s_%s_%s", SecretIdentifier, prefix, secret),
|
||||
Secret: secret,
|
||||
Prefix: prefix,
|
||||
},
|
||||
Hashed: hashed,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -185,19 +185,19 @@ func Tokens(db database.Store, lifetimes codersdk.SessionLifetime) http.HandlerF
|
||||
|
||||
func authorizationCodeGrant(ctx context.Context, db database.Store, app database.OAuth2ProviderApp, lifetimes codersdk.SessionLifetime, params tokenParams) (oauth2.Token, error) {
|
||||
// Validate the client secret.
|
||||
secret, err := parseFormattedSecret(params.clientSecret)
|
||||
secret, err := ParseFormattedSecret(params.clientSecret)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, errBadSecret
|
||||
}
|
||||
//nolint:gocritic // Users cannot read secrets so we must use the system.
|
||||
dbSecret, err := db.GetOAuth2ProviderAppSecretByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(secret.prefix))
|
||||
dbSecret, err := db.GetOAuth2ProviderAppSecretByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(secret.Prefix))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return oauth2.Token{}, errBadSecret
|
||||
}
|
||||
if err != nil {
|
||||
return oauth2.Token{}, err
|
||||
}
|
||||
equal, err := userpassword.Compare(string(dbSecret.HashedSecret), secret.secret)
|
||||
equal, err := userpassword.Compare(string(dbSecret.HashedSecret), secret.Secret)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, xerrors.Errorf("unable to compare secret: %w", err)
|
||||
}
|
||||
@@ -206,19 +206,19 @@ func authorizationCodeGrant(ctx context.Context, db database.Store, app database
|
||||
}
|
||||
|
||||
// Validate the authorization code.
|
||||
code, err := parseFormattedSecret(params.code)
|
||||
code, err := ParseFormattedSecret(params.code)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, errBadCode
|
||||
}
|
||||
//nolint:gocritic // There is no user yet so we must use the system.
|
||||
dbCode, err := db.GetOAuth2ProviderAppCodeByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(code.prefix))
|
||||
dbCode, err := db.GetOAuth2ProviderAppCodeByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(code.Prefix))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return oauth2.Token{}, errBadCode
|
||||
}
|
||||
if err != nil {
|
||||
return oauth2.Token{}, err
|
||||
}
|
||||
equal, err = userpassword.Compare(string(dbCode.HashedSecret), code.secret)
|
||||
equal, err = userpassword.Compare(string(dbCode.HashedSecret), code.Secret)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, xerrors.Errorf("unable to compare code: %w", err)
|
||||
}
|
||||
@@ -344,19 +344,19 @@ func authorizationCodeGrant(ctx context.Context, db database.Store, app database
|
||||
|
||||
func refreshTokenGrant(ctx context.Context, db database.Store, app database.OAuth2ProviderApp, lifetimes codersdk.SessionLifetime, params tokenParams) (oauth2.Token, error) {
|
||||
// Validate the token.
|
||||
token, err := parseFormattedSecret(params.refreshToken)
|
||||
token, err := ParseFormattedSecret(params.refreshToken)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, errBadToken
|
||||
}
|
||||
//nolint:gocritic // There is no user yet so we must use the system.
|
||||
dbToken, err := db.GetOAuth2ProviderAppTokenByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(token.prefix))
|
||||
dbToken, err := db.GetOAuth2ProviderAppTokenByPrefix(dbauthz.AsSystemRestricted(ctx), []byte(token.Prefix))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return oauth2.Token{}, errBadToken
|
||||
}
|
||||
if err != nil {
|
||||
return oauth2.Token{}, err
|
||||
}
|
||||
equal, err := userpassword.Compare(string(dbToken.RefreshHash), token.secret)
|
||||
equal, err := userpassword.Compare(string(dbToken.RefreshHash), token.Secret)
|
||||
if err != nil {
|
||||
return oauth2.Token{}, xerrors.Errorf("unable to compare token: %w", err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user