Skip to content
19 changes: 11 additions & 8 deletions internal/xds/server/kubejwt/jwtinterceptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,13 @@ package kubejwt

import (
"context"
"fmt"
"strings"

discoveryv3 "github.com/envoyproxy/go-control-plane/envoy/service/discovery/v3"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
"k8s.io/client-go/kubernetes"

"github.com/envoyproxy/gateway/internal/logging"
Expand Down Expand Up @@ -51,18 +52,20 @@ func (i *JWTAuthInterceptor) authenticate(ctx context.Context, msg any) error {

md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return fmt.Errorf("missing metadata")
return status.Errorf(codes.Unauthenticated, "missing metadata for node %s", nodeID)
}

authHeader := md.Get("authorization")
if len(authHeader) == 0 {
return fmt.Errorf("missing authorization token in metadata: %s", md)
return status.Errorf(codes.Unauthenticated, "missing authorization header for node %s", nodeID)
}
token := strings.TrimPrefix(authHeader[0], "Bearer ")

if err := i.validateKubeJWT(ctx, token, nodeID); err != nil {
i.logger.Error(err, "failed to validate token")
return fmt.Errorf("failed to validate token: %w", err)
if s, ok := status.FromError(err); ok {
return status.Errorf(s.Code(), "failed to validate the token for node %s: %s", nodeID, s.Message())
}
return status.Errorf(codes.Unauthenticated, "failed to validate the token for node %s: %v", nodeID, err)
Comment thread
cnvergence marked this conversation as resolved.
}

return nil
Expand Down Expand Up @@ -109,15 +112,15 @@ func extractNodeID(m any) (string, error) {
switch req := m.(type) {
case *discoveryv3.DeltaDiscoveryRequest:
if req.Node == nil || req.Node.Id == "" {
return "", fmt.Errorf("missing node ID")
return "", status.Error(codes.InvalidArgument, "missing node ID")
}
return req.Node.Id, nil
case *discoveryv3.DiscoveryRequest:
if req.Node == nil || req.Node.Id == "" {
return "", fmt.Errorf("missing node ID")
return "", status.Error(codes.InvalidArgument, "missing node ID")
}
return req.Node.Id, nil
default:
return "", fmt.Errorf("unexpected message type: %T", m)
return "", status.Errorf(codes.InvalidArgument, "unexpected message type: %T", m)
}
}
79 changes: 47 additions & 32 deletions internal/xds/server/kubejwt/jwtinterceptor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,9 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"

egv1a1 "github.com/envoyproxy/gateway/api/v1alpha1"
"github.com/envoyproxy/gateway/internal/logging"
Expand All @@ -42,85 +44,96 @@ func (m *mockServerStream) RecvMsg(_ any) error { return nil }
// authTestCase defines a single test case for the authenticate method.
// The same cases are run against both the stream (RecvMsg) and unary interceptor paths.
type authTestCase struct {
name string
msg any
ctx context.Context
wantErr string
name string
msg any
ctx context.Context
wantErr string
wantCode codes.Code
}

func authTestCases() []authTestCase {
return []authTestCase{
{
name: "unknown message type is rejected (fail-closed)",
msg: &discoveryv3.DiscoveryResponse{},
ctx: context.Background(),
wantErr: "unexpected message type",
name: "unknown message type is rejected (fail-closed)",
msg: &discoveryv3.DiscoveryResponse{},
ctx: context.Background(),
wantErr: "unexpected message type",
wantCode: codes.InvalidArgument,
},
{
name: "nil message is rejected",
msg: nil,
ctx: context.Background(),
wantErr: "unexpected message type",
name: "nil message is rejected",
msg: nil,
ctx: context.Background(),
wantErr: "unexpected message type",
wantCode: codes.InvalidArgument,
},
{
name: "DeltaDiscoveryRequest with nil node is rejected",
msg: &discoveryv3.DeltaDiscoveryRequest{},
ctx: context.Background(),
wantErr: "missing node ID",
name: "DeltaDiscoveryRequest with nil node is rejected",
msg: &discoveryv3.DeltaDiscoveryRequest{},
ctx: context.Background(),
wantErr: "missing node ID",
wantCode: codes.InvalidArgument,
},
{
name: "DiscoveryRequest with nil node is rejected",
msg: &discoveryv3.DiscoveryRequest{},
ctx: context.Background(),
wantErr: "missing node ID",
name: "DiscoveryRequest with nil node is rejected",
msg: &discoveryv3.DiscoveryRequest{},
ctx: context.Background(),
wantErr: "missing node ID",
wantCode: codes.InvalidArgument,
},
{
name: "DeltaDiscoveryRequest with empty node ID is rejected",
msg: &discoveryv3.DeltaDiscoveryRequest{
Node: &corev3.Node{Id: ""},
},
ctx: context.Background(),
wantErr: "missing node ID",
ctx: context.Background(),
wantErr: "missing node ID",
wantCode: codes.InvalidArgument,
},
{
name: "DiscoveryRequest with empty node ID is rejected",
msg: &discoveryv3.DiscoveryRequest{
Node: &corev3.Node{Id: ""},
},
ctx: context.Background(),
wantErr: "missing node ID",
ctx: context.Background(),
wantErr: "missing node ID",
wantCode: codes.InvalidArgument,
},
{
name: "DeltaDiscoveryRequest without metadata is rejected",
msg: &discoveryv3.DeltaDiscoveryRequest{
Node: &corev3.Node{Id: "pod-1"},
},
ctx: context.Background(),
wantErr: "missing metadata",
ctx: context.Background(),
wantErr: "missing metadata",
wantCode: codes.Unauthenticated,
},
{
name: "DiscoveryRequest without metadata is rejected",
msg: &discoveryv3.DiscoveryRequest{
Node: &corev3.Node{Id: "pod-1"},
},
ctx: context.Background(),
wantErr: "missing metadata",
ctx: context.Background(),
wantErr: "missing metadata",
wantCode: codes.Unauthenticated,
},
{
name: "DeltaDiscoveryRequest without auth header is rejected",
msg: &discoveryv3.DeltaDiscoveryRequest{
Node: &corev3.Node{Id: "pod-1"},
},
ctx: metadata.NewIncomingContext(context.Background(), metadata.MD{}),
wantErr: "missing authorization token",
ctx: metadata.NewIncomingContext(context.Background(), metadata.MD{}),
wantErr: "missing authorization header",
wantCode: codes.Unauthenticated,
},
{
name: "DiscoveryRequest without auth header is rejected",
msg: &discoveryv3.DiscoveryRequest{
Node: &corev3.Node{Id: "pod-1"},
},
ctx: metadata.NewIncomingContext(context.Background(), metadata.MD{}),
wantErr: "missing authorization token",
ctx: metadata.NewIncomingContext(context.Background(), metadata.MD{}),
wantErr: "missing authorization header",
wantCode: codes.Unauthenticated,
},
}
}
Expand All @@ -141,6 +154,7 @@ func TestAuthenticate_Stream(t *testing.T) {
err := ws.RecvMsg(tt.msg)
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErr)
assert.Equal(t, tt.wantCode, status.Code(err), "unexpected gRPC status code")
})
}
}
Expand All @@ -157,6 +171,7 @@ func TestAuthenticate_Unary(t *testing.T) {
_, err := unary(tt.ctx, tt.msg, &grpc.UnaryServerInfo{}, handler)
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErr)
assert.Equal(t, tt.wantCode, status.Code(err), "unexpected gRPC status code")
})
}
}
Expand Down
14 changes: 8 additions & 6 deletions internal/xds/server/kubejwt/tokenreview.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ import (
"fmt"
"slices"

"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
authenticationv1 "k8s.io/api/authentication/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apiserver/pkg/authentication/serviceaccount"
Expand Down Expand Up @@ -42,31 +44,31 @@ func (i *JWTAuthInterceptor) validateKubeJWT(ctx context.Context, token, nodeID

tokenReview, err := i.clientset.AuthenticationV1().TokenReviews().Create(ctx, tokenReview, metav1.CreateOptions{})
if err != nil {
return fmt.Errorf("failed to call TokenReview API to verify service account JWT: %w", err)
return status.Errorf(codes.Internal, "failed to call TokenReview API to verify service account JWT: %v", err)
}

if tokenReview.Status.Error != "" {
return fmt.Errorf("token review found error: %s", tokenReview.Status.Error)
return status.Errorf(codes.Unauthenticated, "token review found error: %s", tokenReview.Status.Error)
}

if !slices.Contains(tokenReview.Status.User.Groups, "system:serviceaccounts") {
return fmt.Errorf("the token is not a service account")
return status.Error(codes.Unauthenticated, "the token is not a service account")
}

if !tokenReview.Status.Authenticated {
return fmt.Errorf("token is not authenticated")
return status.Error(codes.Unauthenticated, "token is not authenticated")
}

// Check if the node ID in the request matches the pod name in the token review response.
// This is used to prevent a client from accessing the xDS resource of another one.
if tokenReview.Status.User.Extra != nil {
podName := tokenReview.Status.User.Extra[serviceaccount.PodNameKey]
if podName[0] == "" {
return fmt.Errorf("pod name not found in token review response")
return status.Error(codes.Unauthenticated, "pod name not found in token review response")
}

if podName[0] != nodeID {
return fmt.Errorf("pod name mismatch: expected %s, got %s", nodeID, podName[0])
return status.Errorf(codes.Unauthenticated, "pod name mismatch: expected %s, got %s", nodeID, podName[0])
}
}

Expand Down