Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 24 additions & 17 deletions cmd/configure_gcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -153,20 +153,32 @@ func runConfigureGCP(cmd *cobra.Command, args []string) error {
}
store := NewAWSSecretsStore(secretsmanager.NewFromConfig(cfg))

credsFile, mintedKeyName, err := getGCPCredentialsFilePath(ctx, reader)
var stop context.CancelFunc
defer func() {
if stop != nil {
stop()
}
}()
beginUpload := func() context.Context {
if stop == nil {
ctx, stop = signal.NotifyContext(ctx, os.Interrupt, syscall.SIGTERM)
}
return ctx
}
credsFile, mintedKeyName, err := getGCPCredentialsFilePath(ctx, reader, beginUpload)
if err != nil {
return err
}

// Scoped to the upload (no stdin reads) so an interrupt cancels the
// Secrets Manager call and the minted-key cleanup still runs.
uploadCtx, stop := signal.NotifyContext(ctx, os.Interrupt, syscall.SIGTERM)
defer stop()
creds, err := uploadGCPCredentialsFile(uploadCtx, store, gcpOpts.StackName, credsFile, mintedKeyName, deleteGCPServiceAccountKey)
creds, err := uploadGCPCredentialsFile(beginUpload(), store, gcpOpts.StackName, credsFile, mintedKeyName, deleteGCPServiceAccountKey)
if err != nil {
return err
}

stop()
if mintedKeyName != "" {
fmt.Println("Removed the local copy of the minted key.")
}
printGCPConfigurationSuccess(creds)
return nil
}
Expand All @@ -183,8 +195,6 @@ func uploadGCPCredentialsFile(ctx context.Context, store SecretsStore, stackName
}
if rmErr := removeMintedGCPKey(credsFile); rmErr != nil {
err = errors.Join(err, rmErr)
} else if err == nil {
fmt.Println("Removed the local copy of the minted key.")
}
}()
}
Expand All @@ -209,7 +219,6 @@ func deleteMintedGCPKeyRemotely(cause error, keyName string, deleteKey func(cont
if err := deleteKey(ctx, keyName); err != nil {
return fmt.Errorf("%w; the minted key is still active in GCP and deleting it failed (%w); delete it with: %s", cause, err, gcloudDeleteKeyCommand(keyName))
}
fmt.Println("Deleted the minted key from GCP.")
return cause
}

Expand All @@ -226,11 +235,11 @@ func gcloudDeleteKeyCommand(keyName string) string {
// getGCPCredentialsFilePath determines the credentials file path from options
// or user input. mintedKeyName is the IAM key resource name when the setup
// wizard created the file, and empty otherwise.
func getGCPCredentialsFilePath(ctx context.Context, reader *bufio.Reader) (credsFile, mintedKeyName string, err error) {
func getGCPCredentialsFilePath(ctx context.Context, reader *bufio.Reader, beginUpload func() context.Context) (credsFile, mintedKeyName string, err error) {
if gcpOpts.CredentialsFile != "" {
credsFile = gcpOpts.CredentialsFile
} else if !gcpOpts.SkipSetup {
credsFile, mintedKeyName, err = runGCPSetupCommands(ctx, reader)
credsFile, mintedKeyName, err = runGCPSetupCommands(ctx, reader, beginUpload)
if err != nil {
return "", "", err
}
Expand Down Expand Up @@ -623,7 +632,7 @@ func writeServiceAccountKey(ctx context.Context, p gcpKeyProvisioner, saEmail, k
// Steps 4-6 (create SA, grant role, create key): performed via GCP IAM and
// Cloud Resource Manager SDK v1 APIs using ADC. Fail loud on any SDK error
// (no CLI fallback).
func runGCPSetupCommands(ctx context.Context, reader *bufio.Reader) (keyFile, keyName string, err error) {
func runGCPSetupCommands(ctx context.Context, reader *bufio.Reader, beginUpload func() context.Context) (keyFile, keyName string, err error) {
err = gcpStepLogin(reader)
if err != nil {
return "", "", err
Expand All @@ -644,7 +653,7 @@ func runGCPSetupCommands(ctx context.Context, reader *bufio.Reader) (keyFile, ke
return "", "", err
}

return gcpStepCreateKey(ctx, reader, saEmail)
return gcpStepCreateKey(reader, saEmail, beginUpload)
}

// gcpStepLogin runs the two interactive gcloud logins the wizard needs:
Expand Down Expand Up @@ -791,7 +800,7 @@ func gcpStepGrantRole(ctx context.Context, reader *bufio.Reader, projectID, saEm
// actually created; on skip or an unknown choice it returns empty strings so
// the caller knows to prompt for an existing credentials file instead of
// assuming one was written.
func gcpStepCreateKey(ctx context.Context, reader *bufio.Reader, saEmail string) (keyFile, keyName string, err error) {
func gcpStepCreateKey(reader *bufio.Reader, saEmail string, beginUpload func() context.Context) (keyFile, keyName string, err error) {
keyFile, err = newMintedGCPKeyPath()
if err != nil {
return "", "", err
Expand All @@ -810,12 +819,10 @@ func gcpStepCreateKey(ctx context.Context, reader *bufio.Reader, saEmail string)
}
switch strings.ToLower(strings.TrimSpace(choice)) {
case "r", "run", "":
keyName, err = createGCPServiceAccountKey(ctx, saEmail, keyFile)
keyName, err = createGCPServiceAccountKey(beginUpload(), saEmail, keyFile)
if err != nil {
return "", "", errors.Join(err, removeMintedGCPKey(keyFile))
}
fmt.Printf("Key written to temporary file: %s\n", keyFile)
fmt.Println()
return keyFile, keyName, nil
case "s", "skip":
fmt.Println("Skipping Create Key")
Expand Down
202 changes: 202 additions & 0 deletions cmd/configure_gcp_signal_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,202 @@
//go:build darwin || linux

package main

import (
"bytes"
"context"
"encoding/base64"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"os"
"os/exec"
"path/filepath"
"strings"
"syscall"
"testing"
"time"

"github.com/stretchr/testify/require"
)

func TestRunConfigureGCP_InterruptCleansMintedKey(t *testing.T) {
if dir := os.Getenv("CUDLY_GCP_SIGNAL_CHILD"); dir != "" {
if err := runGCPInterruptChild(dir); err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
os.Exit(0) // The test harness cannot write to the deliberately full stdout pipe.
}
for _, sig := range []syscall.Signal{syscall.SIGINT, syscall.SIGTERM} {
t.Run(sig.String(), func(t *testing.T) {
child, dir, stdout, stderr := startGCPInterruptChild(t, "upload")
defer stdout.Close()
var keyFile string
require.Eventually(t, func() bool {
matches, err := filepath.Glob(filepath.Join(dir, "cudly-gcp-key-*", "key.json"))
if err != nil || len(matches) != 1 {
return false
}
data, err := os.ReadFile(matches[0])
if err != nil || !bytes.Contains(data, []byte("fixture-secret")) {
return false
}
keyFile = matches[0]
return true
}, 10*time.Second, 10*time.Millisecond)
require.NoError(t, child.Process.Signal(sig))
err := child.Wait()
require.NoFileExists(t, keyFile)
require.NoDirExists(t, filepath.Dir(keyFile))
require.NoError(t, err, "child stderr: %s", stderr.String())
deleted, err := os.ReadFile(filepath.Join(dir, "deleted"))
require.NoError(t, err)
require.Equal(t, "/v1/projects/fixture-project/serviceAccounts/[email protected]/keys/fixture-key", string(deleted))
})
}
}

func TestRunConfigureGCP_InterruptDuringMint(t *testing.T) {
child, dir, stdout, stderr := startGCPInterruptChild(t, "mint")
defer stdout.Close()
require.Eventually(t, func() bool {
_, err := os.Stat(filepath.Join(dir, "creating"))
return err == nil
}, 10*time.Second, 10*time.Millisecond)
require.NoError(t, child.Process.Signal(os.Interrupt))
require.NoError(t, child.Wait(), "child stderr: %s", stderr.String())
matches, err := filepath.Glob(filepath.Join(dir, "cudly-gcp-key-*"))
require.NoError(t, err)
require.Empty(t, matches)
require.NoFileExists(t, filepath.Join(dir, "minted"))
}

func TestRunConfigureGCP_InterruptAtPromptExits(t *testing.T) {
child, dir, stdout, stderr := startGCPInterruptChild(t, "prompt")
defer stdout.Close()
var output strings.Builder
for !strings.Contains(output.String(), "the local copy is removed after upload) ") {
var b [1]byte
_, err := stdout.Read(b[:])
require.NoError(t, err)
output.WriteByte(b[0])
}
require.NoError(t, child.Process.Signal(os.Interrupt))
require.Error(t, child.Wait(), "child stderr: %s", stderr.String())
status, ok := child.ProcessState.Sys().(syscall.WaitStatus)
require.True(t, ok)
require.True(t, status.Signaled())
require.Equal(t, syscall.SIGINT, status.Signal())
require.NoFileExists(t, filepath.Join(dir, "minted"))
}

func startGCPInterruptChild(t *testing.T, phase string) (*exec.Cmd, string, io.ReadCloser, *bytes.Buffer) {
t.Helper()
dir := t.TempDir()
adc := filepath.Join(dir, "adc.json")
require.NoError(t, os.WriteFile(adc, []byte(`{"type":"authorized_user","client_id":"fixture","client_secret":"fixture","refresh_token":"fixture"}`), 0600))
require.NoError(t, os.WriteFile(filepath.Join(dir, "gcloud"), []byte("#!/bin/sh\n[ \"$*\" = \"config set project fixture-project\" ]\n"), 0700))
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
t.Cleanup(cancel)
child := exec.CommandContext(ctx, os.Args[0], "-test.run=^TestRunConfigureGCP_InterruptCleansMintedKey$")
child.Env = append(os.Environ(), "CUDLY_GCP_SIGNAL_CHILD="+dir, "CUDLY_GCP_SIGNAL_PHASE="+phase, "GOOGLE_APPLICATION_CREDENTIALS="+adc,
"TMPDIR="+dir, "PATH="+dir, "AWS_ACCESS_KEY_ID=fixture", "AWS_SECRET_ACCESS_KEY=fixture",
"AWS_SESSION_TOKEN=", "AWS_PROFILE=", "AWS_REGION=us-east-1", "AWS_EC2_METADATA_DISABLED=true")
stdin, err := child.StdinPipe()
require.NoError(t, err)
t.Cleanup(func() { _ = stdin.Close() })
stdout, err := child.StdoutPipe()
require.NoError(t, err)
stderr := new(bytes.Buffer)
child.Stderr = stderr
require.NoError(t, child.Start())
t.Cleanup(func() { _ = child.Process.Kill() })
input := "s\ns\ns\nfixture-project\ns\ns\n"
if phase != "prompt" {
input += "r\n"
}
_, err = io.WriteString(stdin, input)
require.NoError(t, err)
return child, dir, stdout, stderr
}

func runGCPInterruptChild(dir string) error {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if _, err := io.Copy(io.Discard, r.Body); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "application/json")
switch {
case r.URL.Path == "/token":
_, _ = io.WriteString(w, `{"access_token":"fixture","token_type":"Bearer","expires_in":3600}`)
case r.Method == http.MethodDelete:
if err := os.WriteFile(filepath.Join(dir, "deleted"), []byte(r.URL.Path), 0600); err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
_, _ = io.WriteString(w, `{}`)
case strings.HasSuffix(r.URL.Path, "/keys"):
if os.Getenv("CUDLY_GCP_SIGNAL_PHASE") == "mint" {
if err := os.WriteFile(filepath.Join(dir, "creating"), nil, 0600); err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
<-r.Context().Done()
return
}
if err := os.WriteFile(filepath.Join(dir, "minted"), nil, 0600); err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
if err := fillGCPChildStdout(); err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
material := base64.StdEncoding.EncodeToString([]byte(`{"type":"service_account","project_id":"fixture-project","client_email":"[email protected]","private_key":"fixture-secret"}`))
_, _ = fmt.Fprintf(w, `{"name":"projects/fixture-project/serviceAccounts/[email protected]/keys/fixture-key","privateKeyData":%q}`, material)
case r.Header.Get("X-Amz-Target") == "secretsmanager.ListSecrets":
<-r.Context().Done()
default:
http.Error(w, "unexpected fixture request", http.StatusBadRequest)
}
}))
defer server.Close()
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.Proxy = nil
transport.DialTLSContext = func(ctx context.Context, _, _ string) (net.Conn, error) {
return (&net.Dialer{}).DialContext(ctx, "tcp", strings.TrimPrefix(server.URL, "http://"))
}
http.DefaultTransport = transport
if err := os.Setenv("AWS_ENDPOINT_URL_SECRETS_MANAGER", server.URL); err != nil {
return err
}
gcpOpts = GCPConfigOptions{StackName: "fixture-stack"}
err := runConfigureGCP(nil, nil)
if !errors.Is(err, context.Canceled) {
return fmt.Errorf("expected cancellation from the real coordinator, got %w", err)
}
return nil
}

func fillGCPChildStdout() error {
if err := syscall.SetNonblock(1, true); err != nil {
return err
}
defer syscall.SetNonblock(1, false)
for _, chunk := range [][]byte{bytes.Repeat([]byte("x"), 4096), []byte("x")} {
for {
if _, err := syscall.Write(1, chunk); err != nil {
if errors.Is(err, syscall.EAGAIN) {
break
}
return err
}
}
}
return nil
}
Loading