Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
264 changes: 264 additions & 0 deletions flyteadmin/auth/bearer_id_token_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,264 @@
package auth

import (
"context"
"crypto/rand"
"crypto/rsa"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"time"

"github.com/coreos/go-oidc/v3/oidc"
jwtgo "github.com/golang-jwt/jwt/v4"
"github.com/lestrrat-go/jwx/jwk"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"

"github.com/flyteorg/flyte/flyteadmin/auth/config"
"github.com/flyteorg/flyte/flyteadmin/auth/interfaces/mocks"
stdconfig "github.com/flyteorg/flyte/flytestdlib/config"
)

const (
testIDTokenClientID = "flyteadmin"
testIDTokenSubject = "user-subject"
testIDTokenKeyID = "test-key"
)

// fakeOIDCProvider serves an OpenID Connect discovery document and a JWKS for a generated RSA key, and mints ID tokens
// signed with it.
type fakeOIDCProvider struct {
server *httptest.Server
key *rsa.PrivateKey
}

func newFakeOIDCProvider(t *testing.T) *fakeOIDCProvider {
key, err := rsa.GenerateKey(rand.Reader, 2048)
assert.NoError(t, err)

p := &fakeOIDCProvider{key: key}
p.server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/.well-known/openid-configuration":
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"issuer": p.server.URL,
"authorization_endpoint": p.server.URL + "/auth",
"token_endpoint": p.server.URL + "/token",
"jwks_uri": p.server.URL + "/keys",
"id_token_signing_alg_values_supported": []string{"RS256"},
})
case "/keys":
pub, err := jwk.New(&key.PublicKey)
assert.NoError(t, err)
assert.NoError(t, pub.Set(jwk.KeyIDKey, testIDTokenKeyID))
assert.NoError(t, pub.Set(jwk.AlgorithmKey, "RS256"))
set := jwk.NewSet()
set.Add(pub)
_ = json.NewEncoder(w).Encode(set)
default:
w.WriteHeader(http.StatusNotFound)
}
}))

return p
}

func (p *fakeOIDCProvider) provider(t *testing.T) *oidc.Provider {
provider, err := oidc.NewProvider(oidc.ClientContext(context.Background(), p.server.Client()), p.server.URL)
assert.NoError(t, err)
return provider
}

func (p *fakeOIDCProvider) idToken(t *testing.T, audience string) string {
token := jwtgo.NewWithClaims(jwtgo.SigningMethodRS256, jwtgo.MapClaims{
"iss": p.server.URL,
"aud": audience,
"sub": testIDTokenSubject,
"email": "user@example.com",
"iat": time.Now().Unix(),
"exp": time.Now().Add(time.Hour).Unix(),
})
token.Header["kid"] = testIDTokenKeyID
signed, err := token.SignedString(p.key)
assert.NoError(t, err)
return signed
}

func newBearerIDTokenAuthContext(t *testing.T, provider *oidc.Provider) *mocks.AuthenticationContext {
return newBearerIDTokenAuthContextWithClientID(t, provider, testIDTokenClientID)
}

func newBearerIDTokenAuthContextWithClientID(t *testing.T, provider *oidc.Provider, clientID string) *mocks.AuthenticationContext {
resourceServer := &mocks.OAuth2ResourceServer{}
resourceServer.EXPECT().ValidateAccessToken(mock.Anything, mock.Anything, mock.Anything).
Return(nil, fmt.Errorf("not an access token issued by this server"))

authCtx := &mocks.AuthenticationContext{}
authCtx.EXPECT().Options().Return(&config.Config{
AuthorizedURIs: []stdconfig.URL{{URL: url.URL{Scheme: "https", Host: "flyte.example.com"}}},
UserAuth: config.UserAuthConfig{OpenID: config.OpenIDOptions{ClientID: clientID}},
})
authCtx.EXPECT().OAuth2ResourceServer().Return(resourceServer)
authCtx.EXPECT().OidcProvider().Return(provider)
return authCtx
}

func TestGetAuthenticationInterceptor_BearerIDToken(t *testing.T) {
idp := newFakeOIDCProvider(t)
defer idp.server.Close()
provider := idp.provider(t)

t.Run("id token sent with the Bearer scheme is accepted", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID)))

newCtx, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.NoError(t, err)
assert.Equal(t, testIDTokenSubject, IdentityContextFromContext(newCtx).UserID())
assert.True(t, IdentityContextFromContext(newCtx).Scopes().Has(ScopeAll))
})

t.Run("id token sent with the IDToken scheme still works", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, IDTokenScheme+" "+idp.idToken(t, testIDTokenClientID)))

newCtx, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.NoError(t, err)
assert.Equal(t, testIDTokenSubject, IdentityContextFromContext(newCtx).UserID())
})

t.Run("bearer id token for another audience is rejected", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, "some-other-client")))

_, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.Error(t, err)
assert.Equal(t, codes.Unauthenticated, status.Code(err))
})

t.Run("garbage bearer token is rejected", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" not.a.jwt"))

_, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.Error(t, err)
assert.Equal(t, codes.Unauthenticated, status.Code(err))
})

t.Run("no provider configured falls through to the usual rejection", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, nil)
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID)))

_, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.Error(t, err)
assert.Equal(t, codes.Unauthenticated, status.Code(err))
})

// An empty client id makes ParseIDTokenAndValidate skip the audience, issuer and expiry checks, so the fallback
// must not run at all in that configuration, even for a token the provider signed.
t.Run("empty client id never validates a bearer token as an id token", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContextWithClientID(t, provider, "")
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID)))

_, err := GetAuthenticationInterceptor(authCtx)(ctx)
assert.Error(t, err)
assert.Equal(t, codes.Unauthenticated, status.Code(err))
assert.Contains(t, err.Error(), "no OIDC client id configured")
})
}

func TestGRPCGetIdentityFromBearerIDToken(t *testing.T) {
idp := newFakeOIDCProvider(t)
defer idp.server.Close()
provider := idp.provider(t)

t.Run("no authorization metadata", func(t *testing.T) {
_, err := GRPCGetIdentityFromBearerIDToken(context.Background(), testIDTokenClientID, provider)
assert.Error(t, err)
})

t.Run("blank bearer token", func(t *testing.T) {
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "))
_, err := GRPCGetIdentityFromBearerIDToken(ctx, testIDTokenClientID, provider)
assert.Error(t, err)
})

t.Run("nil provider", func(t *testing.T) {
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID)))
_, err := GRPCGetIdentityFromBearerIDToken(ctx, testIDTokenClientID, nil)
assert.Error(t, err)
})

t.Run("valid token with user info metadata", func(t *testing.T) {
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(
DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID),
UserInfoMDKey, `{"email":"from-metadata@example.com"}`))
identityCtx, err := GRPCGetIdentityFromBearerIDToken(ctx, testIDTokenClientID, provider)
assert.NoError(t, err)
assert.Equal(t, testIDTokenSubject, identityCtx.UserID())
assert.Equal(t, "from-metadata@example.com", identityCtx.UserInfo().GetEmail())
})
}

func TestIdentityContextFromRequest_BearerIDToken(t *testing.T) {
idp := newFakeOIDCProvider(t)
defer idp.server.Close()
provider := idp.provider(t)
ctx := context.Background()

t.Run("id token sent with the Bearer scheme is accepted", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
req := httptest.NewRequest(http.MethodGet, "/api/v1/projects", nil)
req.Header.Set(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID))

identityCtx, err := IdentityContextFromRequest(ctx, req, authCtx)
assert.NoError(t, err)
assert.Equal(t, testIDTokenSubject, identityCtx.UserID())
})

t.Run("bearer id token for another audience is rejected", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, provider)
req := httptest.NewRequest(http.MethodGet, "/api/v1/projects", nil)
req.Header.Set(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, "some-other-client"))

identityCtx, err := IdentityContextFromRequest(ctx, req, authCtx)
assert.Error(t, err)
assert.Nil(t, identityCtx)
})

t.Run("empty client id never validates a bearer token as an id token", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContextWithClientID(t, provider, "")
req := httptest.NewRequest(http.MethodGet, "/api/v1/projects", nil)
req.Header.Set(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID))

identityCtx, err := IdentityContextFromRequest(ctx, req, authCtx)
assert.Error(t, err)
assert.Nil(t, identityCtx)
})

t.Run("no provider configured returns the access token error", func(t *testing.T) {
authCtx := newBearerIDTokenAuthContext(t, nil)
req := httptest.NewRequest(http.MethodGet, "/api/v1/projects", nil)
req.Header.Set(DefaultAuthorizationHeader, BearerScheme+" "+idp.idToken(t, testIDTokenClientID))

identityCtx, err := IdentityContextFromRequest(ctx, req, authCtx)
assert.Error(t, err)
assert.Nil(t, identityCtx)
assert.Contains(t, err.Error(), "not an access token issued by this server")
})
}
32 changes: 30 additions & 2 deletions flyteadmin/auth/handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -325,10 +325,19 @@ func GetAuthenticationInterceptor(authCtx interfaces.AuthenticationContext) func
}
logger.Debugf(ctx, "Failed to parse ID Token from context. Error: %v", idTokenErr)

// The bearer token may be an ID token from the userAuth provider sent with the wrong scheme.
identityContext, bearerIDTokenErr := GRPCGetIdentityFromBearerIDToken(ctx, authCtx.Options().UserAuth.OpenID.ClientID,
authCtx.OidcProvider())

if bearerIDTokenErr == nil {
return SetContextForIdentity(ctx, identityContext), nil
}
logger.Debugf(ctx, "Bearer token is not an ID Token either. Error: %v", bearerIDTokenErr)

// Only enforcement logic is present. The default case is to let things through.
if (isFromHTTP && !authCtx.Options().DisableForHTTP) ||
(!isFromHTTP && !authCtx.Options().DisableForGrpc) {
err := fmt.Errorf("id token err: %w, access token err: %w", fmt.Errorf("access token err: %w", accessTokenErr), idTokenErr)
err := fmt.Errorf("access token err: %w, id token err: %w, bearer id token err: %w", accessTokenErr, idTokenErr, bearerIDTokenErr)
return ctx, status.Errorf(codes.Unauthenticated, "token parse error %s", err)
}

Expand Down Expand Up @@ -430,8 +439,27 @@ func IdentityContextFromRequest(ctx context.Context, req *http.Request, authCtx
if len(headerValue) > 0 {
logger.Debugf(ctx, "Found authorization header at [%v] header. Validating.", authHeader)
if strings.HasPrefix(headerValue, BearerScheme+" ") {
tokenStr := strings.TrimPrefix(headerValue, BearerScheme+" ")
expectedAudience := GetPublicURL(ctx, req, authCtx.Options()).String()
return authCtx.OAuth2ResourceServer().ValidateAccessToken(ctx, expectedAudience, strings.TrimPrefix(headerValue, BearerScheme+" "))
identityCtx, accessTokenErr := authCtx.OAuth2ResourceServer().ValidateAccessToken(ctx, expectedAudience, tokenStr)
if accessTokenErr == nil {
return identityCtx, nil
}

// The bearer token may be an ID token from the userAuth provider sent with the wrong scheme. An empty
// client id would make ParseIDTokenAndValidate skip the audience, issuer and expiry checks, so the
// fallback requires one.
clientID := authCtx.Options().UserAuth.OpenID.ClientID
if provider := authCtx.OidcProvider(); provider != nil && clientID != "" {
identityCtx, idTokenErr := IdentityContextFromIDTokenToken(ctx, tokenStr, clientID, provider, nil)
if idTokenErr == nil {
return identityCtx, nil
}

return nil, fmt.Errorf("access token err: %w, bearer id token err: %v", accessTokenErr, idTokenErr)
}

return nil, accessTokenErr
}
}

Expand Down
37 changes: 36 additions & 1 deletion flyteadmin/auth/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -106,11 +106,46 @@ func GRPCGetIdentityFromIDToken(ctx context.Context, clientID string, provider *
return nil, errors.Errorf(ErrJwtValidation, "%v token is blank", IDTokenScheme)
}

return grpcIdentityFromIDTokenString(ctx, tokenStr, clientID, provider)
}

// GRPCGetIdentityFromBearerIDToken handles clients that present an OIDC ID token from the configured userAuth
// provider with the Bearer scheme instead of the IDToken scheme, for example flytectl's ExternalCommand auth type,
// which always sends its token as Bearer. It is tried only after the bearer token failed validation as an access
// token, and the token is verified exactly as an IDToken-scheme token would be.
func GRPCGetIdentityFromBearerIDToken(ctx context.Context, clientID string, provider *oidc.Provider) (
interfaces.IdentityContext, error) {

if provider == nil {
return nil, errors.Errorf(ErrJwtValidation, "no OIDC provider configured to validate a bearer token as an ID token")
}

// With an empty client id ParseIDTokenAndValidate skips the audience, issuer and expiry checks; never fall back
// to ID token validation in that mode.
if clientID == "" {
return nil, errors.Errorf(ErrJwtValidation, "no OIDC client id configured; a bearer token is not validated as an ID token")
}

tokenStr, err := grpcauth.AuthFromMD(ctx, BearerScheme)
if err != nil {
return nil, errors.Wrapf(ErrJwtValidation, err, "Could not retrieve bearer token from metadata")
}

if tokenStr == "" {
return nil, errors.Errorf(ErrJwtValidation, "%v token is blank", BearerScheme)
}

return grpcIdentityFromIDTokenString(ctx, tokenStr, clientID, provider)
}

func grpcIdentityFromIDTokenString(ctx context.Context, tokenStr, clientID string, provider *oidc.Provider) (
interfaces.IdentityContext, error) {

meta := metautils.ExtractIncoming(ctx)
userInfoDecoded := meta.Get(UserInfoMDKey)
userInfo := &service.UserInfoResponse{}
if len(userInfoDecoded) > 0 {
err = json.Unmarshal([]byte(userInfoDecoded), userInfo)
err := json.Unmarshal([]byte(userInfoDecoded), userInfo)
if err != nil {
logger.Infof(ctx, "Could not unmarshal user info from metadata %v", err)
}
Expand Down
Loading