diff --git a/api/holodeck/v1alpha1/types.go b/api/holodeck/v1alpha1/types.go index 61c8a8f90..7e9168585 100644 --- a/api/holodeck/v1alpha1/types.go +++ b/api/holodeck/v1alpha1/types.go @@ -744,7 +744,6 @@ type NVIDIAContainerToolkit struct { Version string `json:"version,omitempty"` } - // LoadBalancer defines load balancer configuration for HA clusters type LoadBalancer struct { // Enabled enables creation of a Network Load Balancer diff --git a/pkg/provider/aws/create.go b/pkg/provider/aws/create.go index 90590e733..77c7c8a4d 100644 --- a/pkg/provider/aws/create.go +++ b/pkg/provider/aws/create.go @@ -424,9 +424,12 @@ func (p *Provider) createEC2Instance(cache *AWS) error { return fmt.Errorf("error getting AMI: %w", err) } - // if the root volume size is not set, use the default size + // Use a local variable to avoid mutating the package-level storageSizeGB, + // which would cause a data race if multiple Create() calls run concurrently + // and pointer aliasing issues with the AWS SDK. + volumeSize := storageSizeGB if p.Spec.RootVolumeSizeGB != nil { - storageSizeGB = *p.Spec.RootVolumeSizeGB + volumeSize = *p.Spec.RootVolumeSizeGB } instanceIn := &ec2.RunInstancesInput{ @@ -439,7 +442,7 @@ func (p *Provider) createEC2Instance(cache *AWS) error { { DeviceName: aws.String("/dev/sda1"), Ebs: &types.EbsBlockDevice{ - VolumeSize: &storageSizeGB, + VolumeSize: &volumeSize, VolumeType: types.VolumeTypeGp2, }, }, diff --git a/pkg/provider/aws/create_test.go b/pkg/provider/aws/create_test.go index 73355ae3b..5b01a3bf1 100644 --- a/pkg/provider/aws/create_test.go +++ b/pkg/provider/aws/create_test.go @@ -24,10 +24,11 @@ import ( "sync" "testing" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "github.com/NVIDIA/holodeck/api/holodeck/v1alpha1" internalaws "github.com/NVIDIA/holodeck/internal/aws" "github.com/NVIDIA/holodeck/internal/logger" - metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/service/ec2" diff --git a/pkg/provider/aws/status.go b/pkg/provider/aws/status.go index 574487ed3..8065e7c20 100644 --- a/pkg/provider/aws/status.go +++ b/pkg/provider/aws/status.go @@ -162,30 +162,16 @@ func update(env *v1alpha1.Environment, cachePath string) error { return err } - // if the cache file does not exist, check if the directory exists - // if the directory does not exist, create it - // if the directory exists, create the cache file - if _, err := os.Stat(cachePath); os.IsNotExist(err) { - dir := filepath.Dir(cachePath) - if _, err := os.Stat(dir); os.IsNotExist(err) { - err := os.MkdirAll(dir, 0750) - if err != nil { - return err - } - } - _, err := os.Create(cachePath) // nolint:gosec - if err != nil { - return err - } + // Ensure the parent directory exists + if err := os.MkdirAll(filepath.Dir(cachePath), 0750); err != nil { + return err } - // write to file - err = os.WriteFile(cachePath, data, 0600) - if err != nil { + if err := os.WriteFile(cachePath, data, 0600); err != nil { return err } - - return nil + // Enforce permissions even if the file already existed with broader perms + return os.Chmod(cachePath, 0600) } // updateAvailableCondition is used to mark a given resource as "available". diff --git a/pkg/provisioner/provisioner.go b/pkg/provisioner/provisioner.go index c269c8757..fac577d4c 100644 --- a/pkg/provisioner/provisioner.go +++ b/pkg/provisioner/provisioner.go @@ -20,10 +20,10 @@ import ( "bytes" "fmt" "io" - "log" "os" "path/filepath" "strings" + "sync" "text/template" "time" @@ -214,30 +214,45 @@ func (p *Provisioner) provision() error { if err != nil { return fmt.Errorf("failed to create session: %w", err) } + defer func() { _ = session.Close() }() + reader, writer := io.Pipe() session.Stdout = writer session.Stderr = writer + var wg sync.WaitGroup + copyErrCh := make(chan error, 1) + wg.Add(1) go func() { - defer func() { _ = writer.Close() }() - _, err := io.Copy(os.Stdout, reader) - if err != nil { - log.Fatalf("Failed to copy from reader: %v", err) + defer wg.Done() + if _, err := io.Copy(os.Stdout, reader); err != nil { + copyErrCh <- fmt.Errorf("failed to copy from reader: %w", err) } }() - defer func() { _ = session.Close() }() script := p.tpl.String() // run the script - err = session.Start(script) - if err != nil { + if err = session.Start(script); err != nil { + _ = writer.Close() // unblock io.Copy goroutine + wg.Wait() return fmt.Errorf("failed to start session: %w", err) } - if err := session.Wait(); err != nil { + if err = session.Wait(); err != nil { + _ = writer.Close() + wg.Wait() return fmt.Errorf("failed to wait for session: %w", err) } + _ = writer.Close() + wg.Wait() + + select { + case copyErr := <-copyErrCh: + return copyErr + default: + } + return nil }