diff --git a/docs/usages/configuration.md b/docs/usages/configuration.md index 12f16240..4509c32f 100644 --- a/docs/usages/configuration.md +++ b/docs/usages/configuration.md @@ -82,7 +82,7 @@ At least one join or Azure authentication method must be configured. `azure.boot | `azure.servicePrincipal.tenantId` | string | Microsoft Entra tenant ID for the service principal. | `70a036f6-8e4d-4615-bad6-149c02e7720d` | | `azure.servicePrincipal.clientId` | string | Application client ID. | `00000000-0000-0000-0000-000000000000` | | `azure.servicePrincipal.clientSecret` | string | Application client secret. Mutually exclusive with `clientSecretFile`. | `` | -| `azure.servicePrincipal.clientSecretFile` | string | Path to a protected regular file (no group/other access; e.g., 0600) containing the application client secret. Mutually exclusive with `clientSecret`. | `/run/credentials/aks-flex-node-sp` | +| `azure.servicePrincipal.clientSecretFile` | string | Path to a protected regular file (no group/other access; e.g., 0600) containing either the application client secret or a PEM/unencrypted PFX certificate and private key. PFX certificate files must use a `.pfx` suffix. | `/run/credentials/aks-flex-node-sp` | ## Agent diff --git a/docs/usages/joining-nodes.md b/docs/usages/joining-nodes.md index e41bc1e6..f0781c5b 100644 --- a/docs/usages/joining-nodes.md +++ b/docs/usages/joining-nodes.md @@ -107,7 +107,7 @@ Minimal config shape: } ``` -The credential file contains only the client secret. The agent requires this to be a non-empty regular file with no group/world access (for example, mode 0600). Inline `clientSecret` remains supported, but cannot be configured together with `clientSecretFile`. +The credential file contains either the client secret or a PEM/unencrypted PFX application certificate and private key; the agent detects the credential type from its contents. PFX certificate files must use a `.pfx` suffix. The file must be a non-empty regular file with no group/world access (for example, mode 0600). Only one of `clientSecret` or `clientSecretFile` can be configured. ## Authentication Mode Selection diff --git a/pkg/aksmachine/client_armapi.go b/pkg/aksmachine/client_armapi.go index 31ede750..7114abde 100644 --- a/pkg/aksmachine/client_armapi.go +++ b/pkg/aksmachine/client_armapi.go @@ -165,6 +165,19 @@ func getCredential(cfg *config.Config, logger *slog.Logger, clientOpts azcore.Cl "tenantID", cfg.Azure.ServicePrincipal.TenantID, "clientID", cfg.Azure.ServicePrincipal.ClientID, ) + if cfg.Azure.ServicePrincipal.ClientSecretFile != "" { + certificates, privateKey, err := cfg.Azure.ServicePrincipal.LoadClientCertificate() + if err != nil { + return nil, fmt.Errorf("load service principal client certificate: %w", err) + } + return azidentity.NewClientCertificateCredential( + cfg.Azure.ServicePrincipal.TenantID, + cfg.Azure.ServicePrincipal.ClientID, + certificates, + privateKey, + &azidentity.ClientCertificateCredentialOptions{ClientOptions: clientOpts}, + ) + } return azidentity.NewClientSecretCredential( cfg.Azure.ServicePrincipal.TenantID, cfg.Azure.ServicePrincipal.ClientID, diff --git a/pkg/aksmachine/client_armapi_test.go b/pkg/aksmachine/client_armapi_test.go index 67b31297..63429b91 100644 --- a/pkg/aksmachine/client_armapi_test.go +++ b/pkg/aksmachine/client_armapi_test.go @@ -1,7 +1,10 @@ package aksmachine import ( + "io" + "log/slog" "math" + "path/filepath" "strings" "testing" @@ -150,6 +153,7 @@ func TestAzureClientOptionsFromConfig(t *testing.T) { if !ok { t.Fatal("ResourceManager cloud service is missing") } + if service.Endpoint != "https://management.example.test" { t.Fatalf("ResourceManager endpoint = %q, want https://management.example.test", service.Endpoint) } @@ -161,6 +165,28 @@ func TestAzureClientOptionsFromConfig(t *testing.T) { } } +func TestGetCredentialClientCertificateLoadError(t *testing.T) { + t.Parallel() + + cfg := testARMConfig(testClusterResourceID, "flex-node-1", "1.34.0") + cfg.Azure.ServicePrincipal = &config.ServicePrincipalConfig{ + TenantID: "tenant", + ClientID: "client", + ClientSecretFile: filepath.Join(t.TempDir(), "missing"), + } + credential, err := getCredential( + cfg, + slog.New(slog.NewTextHandler(io.Discard, nil)), + azureClientOptionsFromConfig(cfg), + ) + if err == nil || !strings.Contains(err.Error(), "load service principal client certificate") { + t.Fatalf("getCredential() error = %v, want client certificate load error", err) + } + if credential != nil { + t.Fatalf("getCredential() = %T, want nil", credential) + } +} + func TestBuildK8sProfile(t *testing.T) { t.Parallel() diff --git a/pkg/cmd/token/kubelogin/kubelogin.go b/pkg/cmd/token/kubelogin/kubelogin.go index 68bbb02d..9495a88b 100644 --- a/pkg/cmd/token/kubelogin/kubelogin.go +++ b/pkg/cmd/token/kubelogin/kubelogin.go @@ -7,6 +7,7 @@ import ( "io" "os" + "github.com/Azure/AKSFlexNode/pkg/config" "github.com/Azure/kubelogin/pkg/token" "github.com/spf13/cobra" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -24,6 +25,7 @@ const aksAADServerID = "6dae42f8-4368-4678-94ff-3960e28e3630" var flagServerID string var flagPopEnabled bool var flagPopClaims string +var flagClientCertificateFile string var Command = &cobra.Command{ Use: "kubelogin", @@ -47,6 +49,10 @@ func init() { &flagPopClaims, "pop-claims", "", "Comma-separated list of key=value claims to include in the PoP token (e.g., 'u=cluster-resource-id').", ) + Command.Flags().StringVar( + &flagClientCertificateFile, "client-certificate-file", "", + "Path to the service principal client certificate file.", + ) } func run(ctx context.Context, out io.Writer) error { @@ -63,6 +69,12 @@ func run(ctx context.Context, out io.Writer) error { tokOpts.ServerID = flagServerID tokOpts.IsPoPTokenEnabled = flagPopEnabled tokOpts.PoPTokenClaims = flagPopClaims + if flagClientCertificateFile != "" { + if err := validateClientCertificateFile(flagClientCertificateFile); err != nil { + return err + } + tokOpts.ClientCert = flagClientCertificateFile + } // TODO: logging to show login details provider, err := token.GetTokenProvider(tokOpts) if err != nil { @@ -76,6 +88,13 @@ func run(ctx context.Context, out io.Writer) error { return outputToken(out, ec, accessToken) } +func validateClientCertificateFile(certificateFile string) error { + if err := config.ValidateServicePrincipalCertificateFile(certificateFile); err != nil { + return fmt.Errorf("validate client certificate file: %w", err) + } + return nil +} + const execInfoEnv = "KUBERNETES_EXEC_INFO" var scheme = runtime.NewScheme() diff --git a/pkg/cmd/token/kubelogin/kubelogin_test.go b/pkg/cmd/token/kubelogin/kubelogin_test.go new file mode 100644 index 00000000..8d531e4a --- /dev/null +++ b/pkg/cmd/token/kubelogin/kubelogin_test.go @@ -0,0 +1,98 @@ +package kubelogin + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestValidateClientCertificateFile(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + validFile := filepath.Join(dir, "client-certificate.pem") + writeTestClientCertificate(t, validFile) + insecureFile := filepath.Join(dir, "client-certificate-insecure.pem") + writeTestClientCertificate(t, insecureFile) + if err := os.Chmod(insecureFile, 0o644); err != nil { + t.Fatalf("os.Chmod: %v", err) + } + invalidFile := filepath.Join(dir, "client-certificate-invalid.pem") + if err := os.WriteFile(invalidFile, []byte("not-a-certificate"), 0o600); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } + + tests := []struct { + name string + path string + wantErr string + }{ + {name: "valid protected certificate file", path: validFile}, + { + name: "rejects insecure permissions", + path: insecureFile, + wantErr: "must not be accessible by group or other users", + }, + { + name: "rejects malformed certificate file contents", + path: invalidFile, + wantErr: "parse service principal client certificate file", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + err := validateClientCertificateFile(tt.path) + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("validateClientCertificateFile() error = %v, want %q", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("validateClientCertificateFile() error = %v", err) + } + }) + } +} + +func writeTestClientCertificate(t *testing.T, path string) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("rsa.GenerateKey: %v", err) + } + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "test-client"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + } + certificate, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatalf("x509.CreateCertificate: %v", err) + } + privateKeyData, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + t.Fatalf("x509.MarshalPKCS8PrivateKey: %v", err) + } + data := append( + pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate}), + pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privateKeyData})..., + ) + if err := os.WriteFile(path, data, 0o600); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } +} diff --git a/pkg/config/adapter.go b/pkg/config/adapter.go index 99285104..dfd302b3 100644 --- a/pkg/config/adapter.go +++ b/pkg/config/adapter.go @@ -74,12 +74,21 @@ func ToAgentConfig(cfg *Config, machineName string) *agentconfig.AgentConfig { ac.Kubelet.Auth.BootstrapToken = cfg.Azure.BootstrapToken.Token case cfg.IsSPConfigured(): - ac.Kubelet.Auth.ExecCredential = buildExecCredential(map[string]string{ - "AAD_LOGIN_METHOD": "spn", - "AAD_SERVICE_PRINCIPAL_CLIENT_ID": cfg.Azure.ServicePrincipal.ClientID, - "AAD_SERVICE_PRINCIPAL_CLIENT_SECRET": cfg.Azure.ServicePrincipal.ClientSecret, - "AZURE_TENANT_ID": cfg.Azure.ServicePrincipal.TenantID, - }) + env := map[string]string{ + "AAD_LOGIN_METHOD": "spn", + "AAD_SERVICE_PRINCIPAL_CLIENT_ID": cfg.Azure.ServicePrincipal.ClientID, + "AZURE_TENANT_ID": cfg.Azure.ServicePrincipal.TenantID, + } + if cfg.Azure.ServicePrincipal.clientCertificateData() != "" { + ac.Kubelet.Auth.ExecCredential = buildExecCredential( + env, + "--client-certificate-file", + cfg.Azure.ServicePrincipal.ClientSecretFile, + ) + } else { + env["AAD_SERVICE_PRINCIPAL_CLIENT_SECRET"] = cfg.Azure.ServicePrincipal.ClientSecret + ac.Kubelet.Auth.ExecCredential = buildExecCredential(env) + } case cfg.IsMIConfigured(): env := map[string]string{ @@ -114,16 +123,17 @@ func ResolveMachineGoalState(ctx context.Context, log *slog.Logger, cfg *Config, // buildExecCredential creates an ExecConfig that invokes the aks-flex-node // binary as a credential plugin. The binary's `token kubelogin` subcommand // uses kubelogin to obtain an Azure AD token for the AKS API server. -func buildExecCredential(env map[string]string) *clientcmdapi.ExecConfig { +func buildExecCredential(env map[string]string, args ...string) *clientcmdapi.ExecConfig { execEnv := make([]clientcmdapi.ExecEnvVar, 0, len(env)) for k, v := range env { execEnv = append(execEnv, clientcmdapi.ExecEnvVar{Name: k, Value: v}) } + execArgs := append([]string{"token", "kubelogin", "--server-id", aksAADServerID}, args...) return &clientcmdapi.ExecConfig{ APIVersion: "client.authentication.k8s.io/v1", Command: flexNodeBinaryPath, - Args: []string{"token", "kubelogin", "--server-id", aksAADServerID}, + Args: execArgs, Env: execEnv, InteractiveMode: clientcmdapi.NeverExecInteractiveMode, ProvideClusterInfo: false, diff --git a/pkg/config/adapter_test.go b/pkg/config/adapter_test.go index 6fa6f763..510af5e4 100644 --- a/pkg/config/adapter_test.go +++ b/pkg/config/adapter_test.go @@ -3,6 +3,7 @@ package config import ( "os" "path/filepath" + "slices" "testing" "github.com/Azure/unbounded/pkg/agent/goalstates" @@ -141,6 +142,7 @@ func TestToAgentConfig_ServicePrincipalClientSecretFile(t *testing.T) { if err := os.WriteFile(clientSecretFile, []byte("file-secret\n"), 0o600); err != nil { t.Fatalf("os.WriteFile: %v", err) } + cfg := &Config{ Azure: AzureConfig{ ServicePrincipal: &ServicePrincipalConfig{ @@ -167,6 +169,40 @@ func TestToAgentConfig_ServicePrincipalClientSecretFile(t *testing.T) { } } +func TestToAgentConfig_ServicePrincipalCertificateFile(t *testing.T) { + t.Parallel() + + certificateFile := filepath.Join(t.TempDir(), "client-certificate.pem") + writeTestClientCertificate(t, certificateFile) + cfg := &Config{ + Azure: AzureConfig{ + ServicePrincipal: &ServicePrincipalConfig{ + TenantID: "tenant-123", + ClientID: "client-456", + ClientSecretFile: certificateFile, + }, + }, + } + if err := cfg.Azure.ServicePrincipal.validate(); err != nil { + t.Fatalf("ServicePrincipalConfig.validate() error = %v", err) + } + + exec := ToAgentConfig(cfg, "kube1").Kubelet.Auth.ExecCredential + if exec == nil { + t.Fatal("ExecCredential should be set for SP auth") + } + envMap := make(map[string]string) + for _, e := range exec.Env { + envMap[e.Name] = e.Value + } + if !slices.Equal(exec.Args, []string{"token", "kubelogin", "--server-id", aksAADServerID, "--client-certificate-file", certificateFile}) { + t.Fatalf("Args=%v, want kubelogin client-certificate-file args", exec.Args) + } + if _, ok := envMap["AAD_SERVICE_PRINCIPAL_CLIENT_SECRET"]; ok { + t.Fatal("AAD_SERVICE_PRINCIPAL_CLIENT_SECRET should not be set for certificate auth") + } +} + func TestToAgentConfig_ManagedIdentity(t *testing.T) { t.Parallel() diff --git a/pkg/config/config.go b/pkg/config/config.go index 99ff01da..24d3574b 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -1,7 +1,12 @@ package config import ( + "bytes" + "crypto" + "crypto/rsa" + "crypto/x509" "encoding/json" + "encoding/pem" "fmt" "net/url" "os" @@ -9,6 +14,9 @@ import ( "regexp" "strings" "time" + "unicode/utf8" + + "github.com/Azure/azure-sdk-for-go/sdk/azidentity" "github.com/Azure/AKSFlexNode/pkg/logger" agentconfig "github.com/Azure/unbounded/pkg/agent/config" @@ -75,10 +83,11 @@ type AzureConfig struct { // ServicePrincipalConfig holds Azure service principal authentication configuration. // When provided, service principal authentication will be used instead of Azure CLI. type ServicePrincipalConfig struct { - TenantID string `json:"tenantId"` // Azure AD tenant ID - ClientID string `json:"clientId"` // Azure AD application (client) ID - ClientSecret string `json:"clientSecret,omitempty"` // Azure AD application client secret - ClientSecretFile string `json:"clientSecretFile,omitempty"` // File containing the Azure AD application client secret + TenantID string `json:"tenantId"` // Azure AD tenant ID + ClientID string `json:"clientId"` // Azure AD application (client) ID + ClientSecret string `json:"clientSecret,omitempty"` // Azure AD application client secret + ClientSecretFile string `json:"clientSecretFile,omitempty"` // File containing an Azure AD application client secret or certificate + clientCertificatePEM string } // ManagedIdentityConfig holds managed identity authentication configuration. @@ -346,6 +355,9 @@ func (cfg *Config) DeepCopy() *Config { if err := json.Unmarshal(data, &out); err != nil { return nil } + if cfg.Azure.ServicePrincipal != nil && out.Azure.ServicePrincipal != nil { + out.Azure.ServicePrincipal.clientCertificatePEM = cfg.Azure.ServicePrincipal.clientCertificatePEM + } return &out } @@ -599,46 +611,141 @@ func (c *ServicePrincipalConfig) validate() error { return fmt.Errorf("only one of azure.servicePrincipal.clientSecret or azure.servicePrincipal.clientSecretFile can be configured") } if c.ClientSecret == "" && c.ClientSecretFile == "" { - return fmt.Errorf("azure.servicePrincipal.clientSecret or azure.servicePrincipal.clientSecretFile is required when service principal is configured") + return fmt.Errorf("one of azure.servicePrincipal.clientSecret or azure.servicePrincipal.clientSecretFile is required when service principal is configured") } if c.ClientSecretFile != "" { - clientSecret, err := loadServicePrincipalClientSecret(c.ClientSecretFile) + absolutePath, err := filepath.Abs(c.ClientSecretFile) + if err != nil { + return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: resolve absolute path: %w", err) + } + c.ClientSecretFile = absolutePath + data, err := LoadServicePrincipalCredentialFile(c.ClientSecretFile) if err != nil { return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: %w", err) } - c.ClientSecret = clientSecret - c.ClientSecretFile = "" + certificates, privateKey, certificateErr := azidentity.ParseCertificates(data, nil) + if certificateErr == nil { + if err := validatePKCS12FileSuffix(c.ClientSecretFile, data); err != nil { + return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: %w", err) + } + clientCertificatePEM, err := MarshalClientCertificatePEM(certificates, privateKey) + if err != nil { + return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: %w", err) + } + c.clientCertificatePEM = clientCertificatePEM + } else { + if credentialLooksLikeCertificate(data) { + return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: parse service principal client certificate file: %w", certificateErr) + } + c.ClientSecret = strings.TrimRight(string(data), "\r\n") + if c.ClientSecret == "" { + return fmt.Errorf("invalid azure.servicePrincipal.clientSecretFile: service principal credential file is empty") + } + c.ClientSecretFile = "" + } + } + return nil +} + +// ValidateServicePrincipalCertificateFile verifies that a certificate credential +// file is protected, well-formed, and compatible with service principal auth. +func ValidateServicePrincipalCertificateFile(path string) error { + data, err := LoadServicePrincipalCredentialFile(path) + if err != nil { + return err + } + if err := validatePKCS12FileSuffix(path, data); err != nil { + return err + } + certificates, privateKey, err := azidentity.ParseCertificates(data, nil) + if err != nil { + return fmt.Errorf("parse service principal client certificate file: %w", err) + } + if _, err := MarshalClientCertificatePEM(certificates, privateKey); err != nil { + return err } return nil } -// loadServicePrincipalClientSecret reads a service principal secret from a -// protected file. A trailing line ending is ignored to support standard secret -// file creation tools. -func loadServicePrincipalClientSecret(path string) (string, error) { +// LoadServicePrincipalCredentialFile reads a service principal credential file +// after verifying it is a protected regular file. +func LoadServicePrincipalCredentialFile(path string) ([]byte, error) { cleanPath := filepath.Clean(path) info, err := os.Lstat(cleanPath) if err != nil { - return "", fmt.Errorf("stat service principal client secret file: %w", err) + return nil, fmt.Errorf("stat service principal credential file: %w", err) } if info.Mode()&os.ModeSymlink != 0 { - return "", fmt.Errorf("service principal client secret file must not be a symlink") + return nil, fmt.Errorf("service principal credential file must not be a symlink") } if !info.Mode().IsRegular() { - return "", fmt.Errorf("service principal client secret file must be a regular file") + return nil, fmt.Errorf("service principal credential file must be a regular file") } if info.Mode().Perm()&0o077 != 0 { - return "", fmt.Errorf("service principal client secret file must not be accessible by group or other users") + return nil, fmt.Errorf("service principal credential file must not be accessible by group or other users") } data, err := os.ReadFile(cleanPath) if err != nil { - return "", fmt.Errorf("read service principal client secret file: %w", err) + return nil, fmt.Errorf("read service principal credential file: %w", err) } - secret := strings.TrimRight(string(data), "\r\n") - if secret == "" { - return "", fmt.Errorf("service principal client secret file is empty") + return data, nil +} + +func validatePKCS12FileSuffix(path string, data []byte) error { + if utf8.Valid(data) { + return nil } - return secret, nil + if !strings.EqualFold(filepath.Ext(path), ".pfx") { + return fmt.Errorf("PFX service principal credential files must use a .pfx suffix") + } + return nil +} + +func credentialLooksLikeCertificate(data []byte) bool { + return bytes.Contains(data, []byte("-----BEGIN CERTIFICATE-----")) || + bytes.Contains(data, []byte("-----BEGIN PRIVATE KEY-----")) || + bytes.Contains(data, []byte("-----BEGIN RSA PRIVATE KEY-----")) || + !utf8.Valid(data) +} + +// LoadClientCertificate loads the service principal certificate and private key +// from ClientSecretFile. +func (c *ServicePrincipalConfig) LoadClientCertificate() ([]*x509.Certificate, crypto.PrivateKey, error) { + data := []byte(c.clientCertificatePEM) + if len(data) == 0 { + var err error + data, err = LoadServicePrincipalCredentialFile(c.ClientSecretFile) + if err != nil { + return nil, nil, err + } + } + certificates, privateKey, err := azidentity.ParseCertificates(data, nil) + if err != nil { + return nil, nil, fmt.Errorf("parse service principal client certificate file: %w", err) + } + return certificates, privateKey, nil +} + +// MarshalClientCertificatePEM converts a client certificate chain and private +// key into PKCS#8 PEM form for consumers that require PEM files. +func MarshalClientCertificatePEM(certificates []*x509.Certificate, privateKey crypto.PrivateKey) (string, error) { + if _, ok := privateKey.(*rsa.PrivateKey); !ok { + return "", fmt.Errorf("service principal client certificate private key must be RSA") + } + privateKeyData, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + return "", fmt.Errorf("marshal service principal client certificate private key: %w", err) + } + var data []byte + for _, certificate := range certificates { + data = append(data, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate.Raw})...) + } + data = append(data, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privateKeyData})...) + return string(data), nil +} + +func (c *ServicePrincipalConfig) clientCertificateData() string { + return c.clientCertificatePEM } func (c *ManagedIdentityConfig) validate() error { diff --git a/pkg/config/config_test.go b/pkg/config/config_test.go index 7f5659a2..7d92dd8f 100644 --- a/pkg/config/config_test.go +++ b/pkg/config/config_test.go @@ -1,8 +1,14 @@ package config import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" "encoding/json" + "encoding/pem" "fmt" + "math/big" "os" "path/filepath" "sort" @@ -1812,7 +1818,7 @@ func TestAuthenticationMethodValidation(t *testing.T) { }, }, wantErr: true, - errMsg: "azure.servicePrincipal.clientSecret or azure.servicePrincipal.clientSecretFile is required when service principal is configured", + errMsg: "one of azure.servicePrincipal.clientSecret or azure.servicePrincipal.clientSecretFile is required when service principal is configured", }, { name: "managed identity authentication enabled", @@ -2127,6 +2133,7 @@ func TestServicePrincipalClientSecretFile(t *testing.T) { if err := os.WriteFile(validFile, []byte("file-secret\r\n"), 0o600); err != nil { t.Fatalf("os.WriteFile: %v", err) } + emptyFile := filepath.Join(dir, "empty") if err := os.WriteFile(emptyFile, nil, 0o600); err != nil { t.Fatalf("os.WriteFile: %v", err) @@ -2155,12 +2162,12 @@ func TestServicePrincipalClientSecretFile(t *testing.T) { { name: "rejects missing file", config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecretFile: filepath.Join(dir, "missing")}, - wantErr: "stat service principal client secret file", + wantErr: "stat service principal credential file", }, { name: "rejects empty file", config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecretFile: emptyFile}, - wantErr: "service principal client secret file is empty", + wantErr: "service principal credential file is empty", }, { name: "rejects insecure permissions", @@ -2193,6 +2200,153 @@ func TestServicePrincipalClientSecretFile(t *testing.T) { } } +func TestServicePrincipalCertificateFile(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + validFile := filepath.Join(dir, "client-certificate.pem") + writeTestClientCertificate(t, validFile) + invalidFile := filepath.Join(dir, "invalid-certificate.pem") + if err := os.WriteFile(invalidFile, []byte("-----BEGIN CERTIFICATE-----\ninvalid\n-----END CERTIFICATE-----\n"), 0o600); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } + insecureFile := filepath.Join(dir, "insecure") + writeTestClientCertificate(t, insecureFile) + if err := os.Chmod(insecureFile, 0o644); err != nil { + t.Fatalf("os.Chmod: %v", err) + } + + tests := []struct { + name string + config *ServicePrincipalConfig + wantErr string + }{ + { + name: "loads certificate", + config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecretFile: validFile}, + }, + { + name: "rejects secret and certificate", + config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecret: "secret", ClientSecretFile: validFile}, + wantErr: "only one of", + }, + { + name: "rejects malformed certificate", + config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecretFile: invalidFile}, + wantErr: "parse service principal client certificate file", + }, + { + name: "rejects insecure permissions", + config: &ServicePrincipalConfig{TenantID: "tenant", ClientID: "client", ClientSecretFile: insecureFile}, + wantErr: "must not be accessible by group or other users", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + err := tt.config.validate() + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("ServicePrincipalConfig.validate() error = %v, want %q", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("ServicePrincipalConfig.validate() error = %v", err) + } + certificates, privateKey, err := tt.config.LoadClientCertificate() + if err != nil { + t.Fatalf("ServicePrincipalConfig.LoadClientCertificate() error = %v", err) + } + if len(certificates) != 1 || privateKey == nil { + t.Fatalf("LoadClientCertificate() returned %d certificates and key %T", len(certificates), privateKey) + } + copied := (&Config{Azure: AzureConfig{ServicePrincipal: tt.config}}).DeepCopy() + if copied == nil || copied.Azure.ServicePrincipal.clientCertificateData() == "" { + t.Fatal("DeepCopy() did not preserve normalized client certificate data") + } + }) + } +} + +func TestValidatePKCS12FileSuffix(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + path string + data []byte + wantErr string + }{ + { + name: "allows utf8 data with non-pfx suffix", + path: "client-secret", + data: []byte("client-secret"), + }, + { + name: "allows binary data with pfx suffix", + path: "client-certificate.pfx", + data: []byte{0xff, 0x00, 0x01}, + }, + { + name: "rejects binary data without pfx suffix", + path: "client-certificate", + data: []byte{0xff, 0x00, 0x01}, + wantErr: "must use a .pfx suffix", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + err := validatePKCS12FileSuffix(tt.path, tt.data) + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("validatePKCS12FileSuffix() error = %v, want %q", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("validatePKCS12FileSuffix() error = %v", err) + } + }) + } +} + +func writeTestClientCertificate(t *testing.T, path string) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("rsa.GenerateKey: %v", err) + } + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "test-client"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + } + certificate, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatalf("x509.CreateCertificate: %v", err) + } + privateKeyData, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + t.Fatalf("x509.MarshalPKCS8PrivateKey: %v", err) + } + data := append( + pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate}), + pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privateKeyData})..., + ) + if err := os.WriteFile(path, data, 0o600); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } +} + func TestValidateBootstrapTokenAPIServerURL(t *testing.T) { t.Parallel()