diff --git a/k3k-kubelet/kubelet.go b/k3k-kubelet/kubelet.go index db0c122a..14b19b81 100644 --- a/k3k-kubelet/kubelet.go +++ b/k3k-kubelet/kubelet.go @@ -247,8 +247,8 @@ func (k *kubelet) newProviderFunc(cfg config) nodeutil.NewProviderFunc { cfg.AgentHostname, k.port, k.agentIP, - k.hostMgr, - utilProvider.VirtualClient, + utilProvider.Host.Manager, + utilProvider.Virtual.Client, k.virtualCluster, cfg.Version, cfg.MirrorHostNodes, diff --git a/k3k-kubelet/provider/provider.go b/k3k-kubelet/provider/provider.go index ef8c27ff..893ab137 100644 --- a/k3k-kubelet/provider/provider.go +++ b/k3k-kubelet/provider/provider.go @@ -46,15 +46,20 @@ import ( // check at compile time if the Provider implements the nodeutil.Provider interface var _ nodeutil.Provider = (*Provider)(nil) +// ClusterContext includes the controller runtime manager and clients +type ClusterContext struct { + Config rest.Config + Client client.Client + CoreClient cv1.CoreV1Interface + Manager manager.Manager +} + // Provider implements nodetuil.Provider from virtual Kubelet. // TODO: Implement NotifyPods and the required usage so that this can be an async provider type Provider struct { + Host ClusterContext + Virtual ClusterContext Translator translate.ToHostTranslator - HostClient client.Client - VirtualClient client.Client - VirtualManager manager.Manager - ClientConfig rest.Config - CoreClient cv1.CoreV1Interface ClusterNamespace string ClusterName string serverIP string @@ -71,18 +76,29 @@ func New(hostConfig rest.Config, hostMgr, virtualMgr manager.Manager, logger log return nil, err } + virtualCoreClient, err := cv1.NewForConfig(virtualMgr.GetConfig()) + if err != nil { + return nil, err + } + translator := translate.ToHostTranslator{ ClusterName: name, ClusterNamespace: namespace, } p := Provider{ - HostClient: hostMgr.GetClient(), - VirtualClient: virtualMgr.GetClient(), - VirtualManager: virtualMgr, + Host: ClusterContext{ + Manager: hostMgr, + Client: hostMgr.GetClient(), + CoreClient: coreClient, + Config: hostConfig, + }, + Virtual: ClusterContext{ + Manager: virtualMgr, + Client: virtualMgr.GetClient(), + CoreClient: virtualCoreClient, + }, Translator: translator, - ClientConfig: hostConfig, - CoreClient: coreClient, ClusterNamespace: namespace, ClusterName: name, logger: logger.WithValues("cluster", name), @@ -128,7 +144,7 @@ func (p *Provider) GetContainerLogs(ctx context.Context, namespace, name, contai options.SinceTime = &sinceTime } - closer, err := p.CoreClient.Pods(p.ClusterNamespace).GetLogs(hostPodName, &options).Stream(ctx) + closer, err := p.Host.CoreClient.Pods(p.ClusterNamespace).GetLogs(hostPodName, &options).Stream(ctx) if err != nil { logger.Error(err, "Error getting logs from container") } @@ -144,7 +160,7 @@ func (p *Provider) RunInContainer(ctx context.Context, namespace, name, containe logger := p.logger.WithValues("namespace", namespace, "name", name, "pod", hostPodName, "container", containerName) logger.V(1).Info("RunInContainer") - req := p.CoreClient.RESTClient().Post(). + req := p.Host.CoreClient.RESTClient().Post(). Resource("pods"). Name(hostPodName). Namespace(p.ClusterNamespace). @@ -159,7 +175,7 @@ func (p *Provider) RunInContainer(ctx context.Context, namespace, name, containe Stderr: attach.Stderr() != nil, }, scheme.ParameterCodec) - exec, err := remotecommand.NewSPDYExecutor(&p.ClientConfig, http.MethodPost, req.URL()) + exec, err := remotecommand.NewSPDYExecutor(&p.Host.Config, http.MethodPost, req.URL()) if err != nil { logger.Error(err, "Error creating SPDY executor") return err @@ -189,7 +205,7 @@ func (p *Provider) AttachToContainer(ctx context.Context, namespace, name, conta logger := p.logger.WithValues("namespace", namespace, "name", name, "pod", hostPodName, "container", containerName) logger.V(1).Info("AttachToContainer") - req := p.CoreClient.RESTClient().Post(). + req := p.Host.CoreClient.RESTClient().Post(). Resource("pods"). Name(hostPodName). Namespace(p.ClusterNamespace). @@ -203,7 +219,7 @@ func (p *Provider) AttachToContainer(ctx context.Context, namespace, name, conta Stderr: attach.Stderr() != nil, }, scheme.ParameterCodec) - exec, err := remotecommand.NewSPDYExecutor(&p.ClientConfig, http.MethodPost, req.URL()) + exec, err := remotecommand.NewSPDYExecutor(&p.Host.Config, http.MethodPost, req.URL()) if err != nil { logger.Error(err, "Error creating SPDY executor") return err @@ -229,13 +245,13 @@ func (p *Provider) AttachToContainer(ctx context.Context, namespace, name, conta func (p *Provider) GetStatsSummary(ctx context.Context) (*v1alpha1stats.Summary, error) { p.logger.V(1).Info("GetStatsSummary") - node, err := p.CoreClient.Nodes().Get(ctx, p.agentHostname, metav1.GetOptions{}) + node, err := p.Host.CoreClient.Nodes().Get(ctx, p.agentHostname, metav1.GetOptions{}) if err != nil { p.logger.Error(err, "Unable to get nodes of cluster") return nil, err } - res, err := p.CoreClient.RESTClient(). + res, err := p.Host.CoreClient.RESTClient(). Get(). Resource("nodes"). Name(node.Name). @@ -322,13 +338,13 @@ func (p *Provider) PortForward(ctx context.Context, namespace, name string, port logger := p.logger.WithValues("namespace", namespace, "name", name, "pod", hostPodName, "port", port) logger.V(1).Info("PortForward") - req := p.CoreClient.RESTClient().Post(). + req := p.Host.CoreClient.RESTClient().Post(). Resource("pods"). Name(hostPodName). Namespace(p.ClusterNamespace). SubResource("portforward") - transport, upgrader, err := spdy.RoundTripperFor(&p.ClientConfig) + transport, upgrader, err := spdy.RoundTripperFor(&p.Host.Config) if err != nil { logger.Error(err, "Error creating RoundTripper for PortForward") return err @@ -372,7 +388,7 @@ func (p *Provider) createPod(ctx context.Context, pod *corev1.Pod) error { } var cluster v1beta1.Cluster - if err := p.HostClient.Get(ctx, clusterKey, &cluster); err != nil { + if err := p.Host.Client.Get(ctx, clusterKey, &cluster); err != nil { logger.Error(err, "Error getting Virtual Cluster definition") return err } @@ -384,7 +400,7 @@ func (p *Provider) createPod(ctx context.Context, pod *corev1.Pod) error { } var virtualPod corev1.Pod - if err := p.VirtualClient.Get(ctx, key, &virtualPod); err != nil { + if err := p.Virtual.Client.Get(ctx, key, &virtualPod); err != nil { logger.Error(err, "Error getting Pod from Virtual Cluster") return err } @@ -478,12 +494,12 @@ func (p *Provider) createPod(ctx context.Context, pod *corev1.Pod) error { configureNetworking(hostPod, virtualPod.Name, virtualPod.Namespace, p.serverIP, p.dnsIP) // set ownerReference to the cluster object - if err := controllerutil.SetControllerReference(&cluster, hostPod, p.HostClient.Scheme()); err != nil { + if err := controllerutil.SetControllerReference(&cluster, hostPod, p.Host.Client.Scheme()); err != nil { logger.Error(err, "Unable to set owner reference for pod") return err } - if err := p.HostClient.Create(ctx, hostPod); err != nil { + if err := p.Host.Client.Create(ctx, hostPod); err != nil { logger.Error(err, "Error creating pod on host cluster") return err } @@ -588,14 +604,14 @@ func (p *Provider) updatePod(ctx context.Context, pod *corev1.Pod) error { } var hostPod corev1.Pod - if err := p.HostClient.Get(ctx, hostKey, &hostPod); err != nil { + if err := p.Host.Client.Get(ctx, hostKey, &hostPod); err != nil { logger.Error(err, "Unable to get Pod to update from host cluster") return err } updatePod(&hostPod, pod) - if err := p.HostClient.Update(ctx, &hostPod); err != nil { + if err := p.Host.Client.Update(ctx, &hostPod); err != nil { logger.Error(err, "Unable to update Pod in host cluster") return err } @@ -606,7 +622,7 @@ func (p *Provider) updatePod(ctx context.Context, pod *corev1.Pod) error { hostPod.Spec.EphemeralContainers = pod.Spec.EphemeralContainers - if _, err := p.CoreClient.Pods(p.ClusterNamespace).UpdateEphemeralContainers(ctx, hostPod.Name, &hostPod, metav1.UpdateOptions{}); err != nil { + if _, err := p.Host.CoreClient.Pods(p.ClusterNamespace).UpdateEphemeralContainers(ctx, hostPod.Name, &hostPod, metav1.UpdateOptions{}); err != nil { logger.Error(err, "Error when updating ephemeral containers in host pod") return err } @@ -624,14 +640,14 @@ func (p *Provider) updatePod(ctx context.Context, pod *corev1.Pod) error { } var virtualPod corev1.Pod - if err := p.VirtualClient.Get(ctx, key, &virtualPod); err != nil { + if err := p.Virtual.Client.Get(ctx, key, &virtualPod); err != nil { logger.Error(err, "Unable to get pod to update from virtual cluster") return err } updatePod(&virtualPod, pod) - if err := p.VirtualClient.Update(ctx, &virtualPod); err != nil { + if err := p.Virtual.Client.Update(ctx, &virtualPod); err != nil { logger.Error(err, "Unable to update Pod in virtual cluster") return err } @@ -642,7 +658,7 @@ func (p *Provider) updatePod(ctx context.Context, pod *corev1.Pod) error { virtualPod.Spec.EphemeralContainers = pod.Spec.EphemeralContainers - if _, err := p.CoreClient.Pods(p.ClusterNamespace).UpdateEphemeralContainers(ctx, virtualPod.Name, &virtualPod, metav1.UpdateOptions{}); err != nil { + if _, err := p.Host.CoreClient.Pods(p.ClusterNamespace).UpdateEphemeralContainers(ctx, virtualPod.Name, &virtualPod, metav1.UpdateOptions{}); err != nil { logger.Error(err, "Error when updating ephemeral containers in virtual pod") return err } @@ -691,7 +707,7 @@ func (p *Provider) deletePod(ctx context.Context, pod *corev1.Pod) error { logger := p.logger.WithValues("namespace", pod.Namespace, "name", pod.Name, "pod", hostPodName) logger.V(1).Info("DeletePod") - err := p.CoreClient.Pods(p.ClusterNamespace).Delete(ctx, hostPodName, metav1.DeleteOptions{}) + err := p.Host.CoreClient.Pods(p.ClusterNamespace).Delete(ctx, hostPodName, metav1.DeleteOptions{}) if err != nil { if apierrors.IsNotFound(err) { logger.Info("Pod to delete not found in host cluster") @@ -753,7 +769,7 @@ func (p *Provider) getPodFromHostCluster(ctx context.Context, hostPodName string } var pod corev1.Pod - if err := p.HostClient.Get(ctx, key, &pod); err != nil { + if err := p.Host.Client.Get(ctx, key, &pod); err != nil { return nil, err } @@ -781,7 +797,7 @@ func (p *Provider) GetPods(ctx context.Context) ([]*corev1.Pod, error) { var podList corev1.PodList - err = p.HostClient.List(ctx, &podList, &client.ListOptions{LabelSelector: selector}) + err = p.Host.Client.List(ctx, &podList, &client.ListOptions{LabelSelector: selector}) if err != nil { p.logger.Error(err, "Error listing pods from host cluster") return nil, err diff --git a/k3k-kubelet/provider/token.go b/k3k-kubelet/provider/token.go index 2081e518..47f66758 100644 --- a/k3k-kubelet/provider/token.go +++ b/k3k-kubelet/provider/token.go @@ -3,11 +3,14 @@ package provider import ( "context" "fmt" + "strconv" "strings" "k8s.io/apimachinery/pkg/types" "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/controller/controllerutil" + authv1 "k8s.io/api/authentication/v1" corev1 "k8s.io/api/core/v1" apierrors "k8s.io/apimachinery/pkg/api/errors" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -20,22 +23,36 @@ const ( serviceAccountTokenMountPath = "/var/run/secrets/kubernetes.io/serviceaccount" ) -// transformTokens copies the serviceaccount tokens used by pod's serviceaccount to a secret on the host cluster and mount it +// transformTokens copies the serviceaccount tokens used by virtualPod's serviceaccount to a secret on the host cluster and mount it // to look like the serviceaccount token -func (p *Provider) transformTokens(ctx context.Context, pod, tPod *corev1.Pod) error { - logger := p.logger.WithValues("namespace", pod.Namespace, "name", pod.Name, "serviceAccountNameod", pod.Spec.ServiceAccountName) - logger.V(1).Info("Transforming token") +func (p *Provider) transformTokens(ctx context.Context, virtualPod, hostPod *corev1.Pod) error { + logger := p.logger.WithValues("namespace", virtualPod.Namespace, "name", virtualPod.Name, "serviceAccountName", virtualPod.Spec.ServiceAccountName) + logger.V(1).Info("Transforming service account tokens") + // transform projected service account token + if err := p.transformProjectedTokens(ctx, virtualPod, hostPod); err != nil { + return err + } + + // transform kube-api-access token for all containers in virtualPod + if err := p.transformKubeAccessToken(ctx, virtualPod, hostPod); err != nil { + return err + } + + return nil +} + +func (p *Provider) transformKubeAccessToken(ctx context.Context, virtualPod, hostPod *corev1.Pod) error { // skip this process if the kube-api-access is already removed from the pod // this is needed in case users already adds their own custom tokens like in rancher imported clusters - if !isKubeAccessVolumeFound(pod) { + if !hasKubeAccessVolume(virtualPod) { return nil } - virtualSecretName := k3kcontroller.SafeConcatNameWithPrefix(pod.Spec.ServiceAccountName, "token") + virtualSecretName := k3kcontroller.SafeConcatNameWithPrefix(virtualPod.Spec.ServiceAccountName, "token") - virtualSecret := virtualSecret(virtualSecretName, pod.Namespace, pod.Spec.ServiceAccountName) - if err := p.VirtualClient.Create(ctx, virtualSecret); err != nil { + virtualSecret := virtualSecret(virtualSecretName, virtualPod.Namespace, virtualPod.Spec.ServiceAccountName) + if err := p.Virtual.Client.Create(ctx, virtualSecret); err != nil { if !apierrors.IsAlreadyExists(err) { return err } @@ -46,7 +63,7 @@ func (p *Provider) transformTokens(ctx context.Context, pod, tPod *corev1.Pod) e Name: virtualSecret.Name, Namespace: virtualSecret.Namespace, } - if err := p.VirtualClient.Get(ctx, virtualSecretKey, virtualSecret); err != nil { + if err := p.Virtual.Client.Get(ctx, virtualSecretKey, virtualSecret); err != nil { return err } // To avoid race conditions we need to check if the secret's data has been populated @@ -55,23 +72,107 @@ func (p *Provider) transformTokens(ctx context.Context, pod, tPod *corev1.Pod) e return fmt.Errorf("token secret %s/%s data is empty", virtualSecret.Namespace, virtualSecret.Name) } - hostSecret := virtualSecret.DeepCopy() - hostSecret.Type = "" - hostSecret.Annotations = make(map[string]string) + hostSecret, err := p.translateAndCreateHostTokenSecret(ctx, virtualSecret) + if err != nil { + return err + } - p.Translator.TranslateTo(hostSecret) + hostPod.Spec.ServiceAccountName = "" + hostPod.Spec.DeprecatedServiceAccount = "" + hostPod.Spec.AutomountServiceAccountToken = ptr.To(false) - if err := p.HostClient.Create(ctx, hostSecret); err != nil { - if !apierrors.IsAlreadyExists(err) { - return err + removeKubeAccessVolume(hostPod) + addKubeAccessVolume(hostPod, hostSecret.Name) + + return nil +} + +// transformProjectedTokens will iterate over the host pod projected volume sources +// and transform projected tokens to use a requested token secret from the virtual cluster +// instead the automatically generated secret on the host cluster. +func (p *Provider) transformProjectedTokens(ctx context.Context, virtualPod, hostPod *corev1.Pod) error { + for i, volume := range hostPod.Spec.Volumes { + if strings.HasPrefix(volume.Name, kubeAPIAccessPrefix) { + continue + } + + if volume.Projected == nil { + continue + } + + for j, source := range volume.Projected.Sources { + if source.ServiceAccountToken == nil { + continue + } + + projectedSecret, err := p.requestTokenSecret(ctx, source.ServiceAccountToken, virtualPod) + if err != nil { + return err + } + + hostSecret, err := p.translateAndCreateHostTokenSecret(ctx, projectedSecret) + if err != nil { + return err + } + // replace the projected token volume with a projected secret + hostPod.Spec.Volumes[i].Projected.Sources[j].ServiceAccountToken = nil + hostPod.Spec.Volumes[i].Projected.Sources[j].Secret = &corev1.SecretProjection{ + LocalObjectReference: corev1.LocalObjectReference{ + Name: hostSecret.Name, + }, + } } } - p.translateToken(tPod, hostSecret.Name) - return nil } +func (p *Provider) requestTokenSecret(ctx context.Context, token *corev1.ServiceAccountTokenProjection, virtualPod *corev1.Pod) (*corev1.Secret, error) { + namespace := virtualPod.Namespace + serviceAccountName := virtualPod.Spec.ServiceAccountName + + var audiences []string + if token.Audience != "" { + audiences = []string{token.Audience} + } + + tokenRequest := &authv1.TokenRequest{ + ObjectMeta: metav1.ObjectMeta{ + Name: serviceAccountName, + Namespace: namespace, + }, + Spec: authv1.TokenRequestSpec{ + Audiences: audiences, + ExpirationSeconds: token.ExpirationSeconds, + BoundObjectRef: &authv1.BoundObjectReference{ + Name: virtualPod.Name, + UID: virtualPod.UID, + Kind: "Pod", + APIVersion: "v1", + }, + }, + } + + tokenResp, err := p.Virtual.CoreClient.ServiceAccounts(namespace).CreateToken(ctx, serviceAccountName, tokenRequest, metav1.CreateOptions{}) + if err != nil { + return nil, err + } + + // create a virtual secret with that token + virtualSecret := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{ + // creating unique name for the virtual secret based on the request attributes + Name: generateTokenSecretName(serviceAccountName, token.Path, tokenResp), + Namespace: namespace, + }, + Data: map[string][]byte{ + token.Path: []byte(tokenResp.Status.Token), + }, + } + + return virtualSecret, nil +} + func virtualSecret(name, namespace, serviceAccountName string) *corev1.Secret { return &corev1.Secret{ TypeMeta: metav1.TypeMeta{ @@ -89,17 +190,25 @@ func virtualSecret(name, namespace, serviceAccountName string) *corev1.Secret { } } -// translateToken will remove the serviceaccount from the pod and replace the kube-api-access volume -// with a custom token volume and mount it to all containers within the pod -func (p *Provider) translateToken(pod *corev1.Pod, hostSecretName string) { - pod.Spec.ServiceAccountName = "" - pod.Spec.DeprecatedServiceAccount = "" - pod.Spec.AutomountServiceAccountToken = ptr.To(false) - removeKubeAccessVolume(pod) - addKubeAccessVolume(pod, hostSecretName) +func (p *Provider) translateAndCreateHostTokenSecret(ctx context.Context, projectedToken *corev1.Secret) (*corev1.Secret, error) { + hostSecret := projectedToken.DeepCopy() + hostSecret.Type = "" + hostSecret.Annotations = make(map[string]string) + + p.Translator.TranslateTo(hostSecret) + + data := hostSecret.Data + if _, err := controllerutil.CreateOrUpdate(ctx, p.Host.Client, hostSecret, func() error { + hostSecret.Data = data + return nil + }); err != nil { + return nil, err + } + + return hostSecret, nil } -func isKubeAccessVolumeFound(pod *corev1.Pod) bool { +func hasKubeAccessVolume(pod *corev1.Pod) bool { for _, volume := range pod.Spec.Volumes { if strings.HasPrefix(volume.Name, kubeAPIAccessPrefix) { return true @@ -171,4 +280,29 @@ func addKubeAccessVolume(pod *corev1.Pod, hostSecretName string) { MountPath: serviceAccountTokenMountPath, }) } + + for i := range pod.Spec.EphemeralContainers { + pod.Spec.EphemeralContainers[i].VolumeMounts = append(pod.Spec.EphemeralContainers[i].VolumeMounts, corev1.VolumeMount{ + Name: tokenVolumeName, + MountPath: serviceAccountTokenMountPath, + }) + } +} + +func generateTokenSecretName(serviceAccountName, tokenPath string, tokenReq *authv1.TokenRequest) string { + nameComponents := []string{serviceAccountName} + + if tokenReq.Spec.Audiences != nil { + nameComponents = append(nameComponents, tokenReq.Spec.Audiences...) + } + + if exp := tokenReq.Spec.ExpirationSeconds; exp != nil { + nameComponents = append(nameComponents, strconv.FormatInt(*exp, 10)) + } + + if tokenPath != "" { + nameComponents = append(nameComponents, tokenPath) + } + + return k3kcontroller.SafeConcatNameWithPrefix(nameComponents...) } diff --git a/k3k-kubelet/provider/token_test.go b/k3k-kubelet/provider/token_test.go new file mode 100644 index 00000000..73eb834d --- /dev/null +++ b/k3k-kubelet/provider/token_test.go @@ -0,0 +1,327 @@ +package provider + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "k8s.io/utils/ptr" + + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + + k3kcontroller "github.com/rancher/k3k/pkg/controller" +) + +func Test_hasKubeAccessVolume(t *testing.T) { + tests := []struct { + name string + pod *corev1.Pod + want bool + }{ + { + name: "no volumes", + pod: &corev1.Pod{}, + want: false, + }, + { + name: "volume with kube-api-access prefix", + pod: &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "kube-api-access-abc123"}, + }, + }, + }, + want: true, + }, + { + name: "exact kube-api-access name", + pod: &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "kube-api-access"}, + }, + }, + }, + want: true, + }, + { + name: "volume without kube-api-access prefix", + pod: &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "my-volume"}, + }, + }, + }, + want: false, + }, + { + name: "multiple volumes with one kube-api-access", + pod: &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "config-volume"}, + {Name: "kube-api-access-xyz"}, + {Name: "data-volume"}, + }, + }, + }, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, hasKubeAccessVolume(tt.pod)) + }) + } +} + +func Test_removeKubeAccessVolume(t *testing.T) { + t.Run("removes volume and all volume mounts from containers", func(t *testing.T) { + pod := &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "config-volume"}, + {Name: "kube-api-access-abc"}, + {Name: "data-volume"}, + }, + InitContainers: []corev1.Container{ + { + Name: "init", + VolumeMounts: []corev1.VolumeMount{ + {Name: "config-volume", MountPath: "/config"}, + {Name: "kube-api-access-abc", MountPath: serviceAccountTokenMountPath}, + }, + }, + }, + Containers: []corev1.Container{ + { + Name: "main", + VolumeMounts: []corev1.VolumeMount{ + {Name: "kube-api-access-abc", MountPath: serviceAccountTokenMountPath}, + {Name: "data-volume", MountPath: "/data"}, + }, + }, + }, + EphemeralContainers: []corev1.EphemeralContainer{ + { + EphemeralContainerCommon: corev1.EphemeralContainerCommon{ + Name: "debug", + VolumeMounts: []corev1.VolumeMount{ + {Name: "kube-api-access-abc", MountPath: serviceAccountTokenMountPath}, + }, + }, + }, + }, + }, + } + + removeKubeAccessVolume(pod) + + // Verify volume was removed + assert.Equal(t, 2, len(pod.Spec.Volumes)) + assert.Equal(t, "config-volume", pod.Spec.Volumes[0].Name) + assert.Equal(t, "data-volume", pod.Spec.Volumes[1].Name) + + // Verify init container mount was removed + assert.Equal(t, 1, len(pod.Spec.InitContainers[0].VolumeMounts)) + assert.Equal(t, "config-volume", pod.Spec.InitContainers[0].VolumeMounts[0].Name) + + // Verify container mount was removed + assert.Equal(t, 1, len(pod.Spec.Containers[0].VolumeMounts)) + assert.Equal(t, "data-volume", pod.Spec.Containers[0].VolumeMounts[0].Name) + + // Verify ephemeral container mount was removed + assert.Equal(t, 0, len(pod.Spec.EphemeralContainers[0].VolumeMounts)) + }) + + t.Run("no kube-api-access volume present", func(t *testing.T) { + pod := &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "config-volume"}, + }, + Containers: []corev1.Container{ + { + Name: "main", + VolumeMounts: []corev1.VolumeMount{ + {Name: "config-volume", MountPath: "/config"}, + }, + }, + }, + }, + } + + removeKubeAccessVolume(pod) + + assert.Equal(t, 1, len(pod.Spec.Volumes)) + assert.Equal(t, "config-volume", pod.Spec.Volumes[0].Name) + assert.Equal(t, 1, len(pod.Spec.Containers[0].VolumeMounts)) + }) +} + +func Test_addKubeAccessVolume(t *testing.T) { + tokenVolumeName := k3kcontroller.SafeConcatNameWithPrefix(kubeAPIAccessPrefix) + hostSecretName := "host-secret-token" + + pod := &corev1.Pod{ + Spec: corev1.PodSpec{ + Volumes: []corev1.Volume{ + {Name: "existing-volume"}, + }, + InitContainers: []corev1.Container{ + {Name: "init"}, + }, + Containers: []corev1.Container{ + {Name: "main"}, + {Name: "sidecar"}, + }, + EphemeralContainers: []corev1.EphemeralContainer{ + { + EphemeralContainerCommon: corev1.EphemeralContainerCommon{ + Name: "debug", + }, + }, + }, + }, + } + + addKubeAccessVolume(pod, hostSecretName) + + // Verify volume was added + assert.Equal(t, 2, len(pod.Spec.Volumes)) + addedVol := pod.Spec.Volumes[1] + assert.Equal(t, tokenVolumeName, addedVol.Name) + assert.Equal(t, hostSecretName, addedVol.Secret.SecretName) + + // Verify init container mount was added + assert.Equal(t, 1, len(pod.Spec.InitContainers[0].VolumeMounts)) + assert.Equal(t, tokenVolumeName, pod.Spec.InitContainers[0].VolumeMounts[0].Name) + assert.Equal(t, serviceAccountTokenMountPath, pod.Spec.InitContainers[0].VolumeMounts[0].MountPath) + + // Verify all container mounts were added + for _, c := range pod.Spec.Containers { + assert.Equal(t, 1, len(c.VolumeMounts), "container %s should have mount", c.Name) + assert.Equal(t, tokenVolumeName, c.VolumeMounts[0].Name) + assert.Equal(t, serviceAccountTokenMountPath, c.VolumeMounts[0].MountPath) + } + + // Verify ephemeral container mounts were added + for _, c := range pod.Spec.EphemeralContainers { + assert.Equal(t, 1, len(c.VolumeMounts), "ephemeral container %s should have mount", c.Name) + assert.Equal(t, tokenVolumeName, c.VolumeMounts[0].Name) + assert.Equal(t, serviceAccountTokenMountPath, c.VolumeMounts[0].MountPath) + } +} + +func Test_virtualSecret(t *testing.T) { + s := virtualSecret("my-secret", "my-ns", "my-sa") + + assert.Equal(t, "my-secret", s.Name) + assert.Equal(t, "my-ns", s.Namespace) + assert.Equal(t, corev1.SecretTypeServiceAccountToken, s.Type) + assert.Equal(t, "my-sa", s.Annotations[corev1.ServiceAccountNameKey]) + assert.Equal(t, "Secret", s.Kind) + assert.Equal(t, "v1", s.APIVersion) +} + +func Test_generateTokenSecretName(t *testing.T) { + tests := []struct { + name string + serviceAccountName string + tokenPath string + tokenReq *authv1.TokenRequest + want string + }{ + { + name: "no audiences, no expiration, no path", + serviceAccountName: "default", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{}, + }, + want: "k3k-default", + }, + { + name: "no audiences, with expiration", + serviceAccountName: "default", + tokenPath: "token", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + ExpirationSeconds: ptr.To(int64(3600)), + }, + }, + want: "k3k-default-3600-token", + }, + { + name: "with single audience and expiration", + serviceAccountName: "my-sa", + tokenPath: "token", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + Audiences: []string{"api"}, + ExpirationSeconds: ptr.To(int64(3600)), + }, + }, + want: "k3k-my-sa-api-3600-token", + }, + { + name: "with multiple audiences and expiration", + serviceAccountName: "my-sa", + tokenPath: "token", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + Audiences: []string{"api", "vault"}, + ExpirationSeconds: ptr.To(int64(3600)), + }, + }, + want: "k3k-my-sa-api-vault-3600-token", + }, + { + name: "with audiences, no expiration", + serviceAccountName: "my-sa", + tokenPath: "vault-token", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + Audiences: []string{"api"}, + }, + }, + want: "k3k-my-sa-api-vault-token", + }, + { + name: "different paths produce different names", + serviceAccountName: "my-sa", + tokenPath: "other-path", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + Audiences: []string{"api"}, + ExpirationSeconds: ptr.To(int64(3600)), + }, + }, + want: "k3k-my-sa-api-3600-other-path", + }, + { + name: "long name gets truncated with hash", + serviceAccountName: "my-very-long-service-account-name", + tokenPath: "some-very-long-token-path-value", + tokenReq: &authv1.TokenRequest{ + Spec: authv1.TokenRequestSpec{ + Audiences: []string{"some-very-long-audience-string"}, + ExpirationSeconds: ptr.To(int64(3600)), + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := generateTokenSecretName(tt.serviceAccountName, tt.tokenPath, tt.tokenReq) + if tt.want != "" { + assert.Equal(t, tt.want, got) + } + + assert.Less(t, len(got), 64, "name should be under 64 characters") + }) + } +}