diff --git a/internal/xds/server/kubejwt/jwtinterceptor.go b/internal/xds/server/kubejwt/jwtinterceptor.go index 905c34ef766..86917e61c13 100644 --- a/internal/xds/server/kubejwt/jwtinterceptor.go +++ b/internal/xds/server/kubejwt/jwtinterceptor.go @@ -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" @@ -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) } return nil @@ -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) } } diff --git a/internal/xds/server/kubejwt/jwtinterceptor_test.go b/internal/xds/server/kubejwt/jwtinterceptor_test.go index d7d639ebfdf..bf473981b75 100644 --- a/internal/xds/server/kubejwt/jwtinterceptor_test.go +++ b/internal/xds/server/kubejwt/jwtinterceptor_test.go @@ -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" @@ -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, }, } } @@ -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") }) } } @@ -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") }) } } diff --git a/internal/xds/server/kubejwt/tokenreview.go b/internal/xds/server/kubejwt/tokenreview.go index 4bef9066813..611ec6dd807 100644 --- a/internal/xds/server/kubejwt/tokenreview.go +++ b/internal/xds/server/kubejwt/tokenreview.go @@ -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" @@ -42,19 +44,19 @@ 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. @@ -62,11 +64,11 @@ func (i *JWTAuthInterceptor) validateKubeJWT(ctx context.Context, token, nodeID 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]) } }