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
4 changes: 3 additions & 1 deletion cmd/cleanup-lambda/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,9 @@ func deleteExpired(ctx context.Context, db *database.Connection, now time.Time,
}
defer func() {
if err != nil {
_ = tx.Rollback(ctx)
if rErr := tx.Rollback(ctx); rErr != nil {
log.Printf("rollback failed: %v", rErr)
}
}
}()

Expand Down
99 changes: 68 additions & 31 deletions cmd/configure_azure.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,9 @@ import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log"
"os"
"os/exec"
Expand All @@ -30,6 +32,18 @@ func validateAzureUUID(uuid, fieldName string) error {
return nil
}

// readTrimmedLine reads one line from reader and returns it with surrounding
// whitespace trimmed. io.EOF is tolerated when data was read — a final
// unterminated line from piped input (e.g. `printf "r" | cudly configure-azure`)
// is still valid input. io.EOF with no data, or any other error, is returned.
func readTrimmedLine(reader *bufio.Reader) (string, error) {
input, err := reader.ReadString('\n')
if err != nil && !(errors.Is(err, io.EOF) && input != "") {
return "", err
}
return strings.TrimSpace(input), nil
}

// AzureCredentials holds the Azure Service Principal credentials
type AzureCredentials struct {
TenantID string `json:"tenant_id"`
Expand Down Expand Up @@ -216,20 +230,20 @@ func collectAzureCredentials(reader *bufio.Reader) (AzureCredentials, error) {
func promptForAzureCredentialFields(reader *bufio.Reader, creds *AzureCredentials) error {
if creds.TenantID == "" {
fmt.Print("Azure Tenant ID: ")
input, err := reader.ReadString('\n')
input, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read tenant ID: %w", err)
}
creds.TenantID = strings.TrimSpace(input)
creds.TenantID = input
}

if creds.ClientID == "" {
fmt.Print("Client ID (appId): ")
input, err := reader.ReadString('\n')
input, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read client ID: %w", err)
}
creds.ClientID = strings.TrimSpace(input)
creds.ClientID = input
}

if creds.ClientSecret == "" {
Expand All @@ -244,11 +258,11 @@ func promptForAzureCredentialFields(reader *bufio.Reader, creds *AzureCredential

if creds.SubscriptionID == "" {
fmt.Print("Subscription ID: ")
input, err := reader.ReadString('\n')
input, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read subscription ID: %w", err)
}
creds.SubscriptionID = strings.TrimSpace(input)
creds.SubscriptionID = input
}

return nil
Expand Down Expand Up @@ -277,8 +291,10 @@ func runAzureSetupCommands(reader *bufio.Reader) error {

fmt.Println()
fmt.Print("Enter your Subscription ID from above: ")
subscriptionID, _ := reader.ReadString('\n')
subscriptionID = strings.TrimSpace(subscriptionID)
subscriptionID, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read subscription ID: %w", err)
}

if subscriptionID == "" {
return fmt.Errorf("subscription ID is required")
Expand All @@ -289,20 +305,38 @@ func runAzureSetupCommands(reader *bufio.Reader) error {
return err
}

if err := createAzureServicePrincipal(reader, subscriptionID); err != nil {
return err
}

fmt.Println()
fmt.Println("IMPORTANT: Copy the output above! You'll need:")
fmt.Println(" - appId -> Client ID")
fmt.Println(" - password -> Client Secret")
fmt.Println(" - tenant -> Tenant ID")
fmt.Printf(" - Subscription ID: %s\n", subscriptionID)
fmt.Println()

return nil
}

// createAzureServicePrincipal runs Step 3 of Azure setup: create service principal.
func createAzureServicePrincipal(reader *bufio.Reader, subscriptionID string) error {
fmt.Println()
fmt.Println("Step 3: Create Service Principal")
fmt.Println("---------------------------------")
fmt.Println("This creates an Azure Service Principal with Reservation Administrator role.")
fmt.Println()

// Build the create SP command - run directly without shell to avoid injection
// Using exec.Command directly with proper arguments
fmt.Printf("Command: az ad sp create-for-rbac --name CUDly --role \"Reservations Administrator\" --scopes /subscriptions/%s\n", subscriptionID)
fmt.Println()
fmt.Printf("[R]un, [S]kip? ")

choice, _ := reader.ReadString('\n')
choice = strings.ToLower(strings.TrimSpace(choice))
choice, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read choice: %w", err)
}
choice = strings.ToLower(choice)

if choice == "r" || choice == "run" || choice == "" {
fmt.Println()
Expand All @@ -317,24 +351,18 @@ func runAzureSetupCommands(reader *bufio.Reader) error {
if err := cmd.Run(); err != nil {
fmt.Printf("Command failed: %v\n", err)
fmt.Print("Continue anyway? [y/N]: ")
response, _ := reader.ReadString('\n')
if strings.ToLower(strings.TrimSpace(response)) != "y" {
response, readErr := readTrimmedLine(reader)
if readErr != nil {
return fmt.Errorf("failed to read response: %w", readErr)
}
if strings.ToLower(response) != "y" {
return fmt.Errorf("failed to create service principal: %w", err)
}
}
fmt.Println(strings.Repeat("-", 60))
} else {
fmt.Println("Skipping Create Service Principal")
}

fmt.Println()
fmt.Println("IMPORTANT: Copy the output above! You'll need:")
fmt.Println(" - appId -> Client ID")
fmt.Println(" - password -> Client Secret")
fmt.Println(" - tenant -> Tenant ID")
fmt.Printf(" - Subscription ID: %s\n", subscriptionID)
fmt.Println()

return nil
}

Expand All @@ -345,12 +373,15 @@ func promptAndRunExplicitCommand(reader *bufio.Reader, name, displayCmd string,
fmt.Println()
fmt.Printf("[R]un, [S]kip? ")

choice, _ := reader.ReadString('\n')
choice = strings.ToLower(strings.TrimSpace(choice))
choice, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read choice: %w", err)
}
choice = strings.ToLower(choice)

switch choice {
case "r", "run", "":
return executeExplicitCommand(displayCmd, program, args...)
return executeExplicitCommand(reader, displayCmd, program, args...)
case "s", "skip":
fmt.Printf("Skipping %s\n", name)
return nil
Expand All @@ -360,8 +391,12 @@ func promptAndRunExplicitCommand(reader *bufio.Reader, name, displayCmd string,
}
}

// executeExplicitCommand runs a command with explicit program and arguments
func executeExplicitCommand(displayCmd string, program string, args ...string) error {
// executeExplicitCommand runs a command with explicit program and arguments.
// The caller's reader is threaded through to the retry prompt so all input
// is consumed from one consistent buffered stream (a fresh
// bufio.NewReader(os.Stdin) here would drop input already buffered by the
// caller's reader, breaking piped input after earlier prompts).
func executeExplicitCommand(reader *bufio.Reader, displayCmd string, program string, args ...string) error {
fmt.Println()
fmt.Printf("Executing: %s\n", displayCmd)
fmt.Println(strings.Repeat("-", 60))
Expand All @@ -377,9 +412,11 @@ func executeExplicitCommand(displayCmd string, program string, args ...string) e
if err != nil {
fmt.Printf("Command failed: %v\n", err)
fmt.Print("Continue anyway? [y/N]: ")
reader := bufio.NewReader(os.Stdin)
response, _ := reader.ReadString('\n')
if strings.ToLower(strings.TrimSpace(response)) != "y" {
response, readErr := readTrimmedLine(reader)
if readErr != nil {
return fmt.Errorf("failed to read response: %w", readErr)
}
if strings.ToLower(response) != "y" {
return fmt.Errorf("command failed: %w", err)
}
}
Expand Down
55 changes: 39 additions & 16 deletions cmd/configure_gcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -177,8 +177,11 @@ func getGCPCredentialsFilePath(reader *bufio.Reader) (string, error) {

if credsFile == "" {
fmt.Print("Path to GCP service account JSON key file: ")
credsFile, _ = reader.ReadString('\n')
credsFile = strings.TrimSpace(credsFile)
var readErr error
credsFile, readErr = readTrimmedLine(reader)
if readErr != nil {
return "", fmt.Errorf("failed to read credentials file path: %w", readErr)
}
}

if credsFile == "" {
Expand Down Expand Up @@ -273,12 +276,9 @@ func runGCPSetupCommands(reader *bufio.Reader) (string, error) {
}

fmt.Println()
fmt.Print("Enter your Project ID from above: ")
projectID, _ := reader.ReadString('\n')
projectID = strings.TrimSpace(projectID)

if projectID == "" {
return "", fmt.Errorf("project ID is required")
projectID, err := readRequiredInputLine(reader, "Enter your Project ID from above: ", "project ID")
if err != nil {
return "", err
}

// Validate project ID to prevent command injection
Expand Down Expand Up @@ -358,19 +358,36 @@ func runGCPSetupCommands(reader *bufio.Reader) (string, error) {
return keyFile, nil
}

// readRequiredInputLine prints prompt, reads a line, trims whitespace, and
// returns an error if the result is empty.
func readRequiredInputLine(reader *bufio.Reader, prompt, fieldName string) (string, error) {
fmt.Print(prompt)
value, err := readTrimmedLine(reader)
if err != nil {
return "", fmt.Errorf("failed to read %s: %w", fieldName, err)
}
if value == "" {
return "", fmt.Errorf("%s is required", fieldName)
}
return value, nil
}

// promptAndRunGCPCommand shows a command and asks to run or skip.
// Takes explicit program and args to avoid command injection via string splitting.
func promptAndRunGCPCommand(reader *bufio.Reader, name, displayCmd string, program string, args ...string) error {
fmt.Printf("Command: %s\n", displayCmd)
fmt.Println()
fmt.Printf("[R]un, [S]kip? ")

choice, _ := reader.ReadString('\n')
choice = strings.ToLower(strings.TrimSpace(choice))
choice, err := readTrimmedLine(reader)
if err != nil {
return fmt.Errorf("failed to read choice: %w", err)
}
choice = strings.ToLower(choice)

switch choice {
case "r", "run", "":
return executeGCPCommand(displayCmd, program, args...)
return executeGCPCommand(reader, displayCmd, program, args...)
case "s", "skip":
fmt.Printf("Skipping %s\n", name)
return nil
Expand All @@ -380,8 +397,12 @@ func promptAndRunGCPCommand(reader *bufio.Reader, name, displayCmd string, progr
}
}

// executeGCPCommand runs a gcloud command with explicit program and arguments
func executeGCPCommand(displayCmd string, program string, args ...string) error {
// executeGCPCommand runs a gcloud command with explicit program and arguments.
// The caller's reader is threaded through to the retry prompt so all input
// is consumed from one consistent buffered stream (a fresh
// bufio.NewReader(os.Stdin) here would drop input already buffered by the
// caller's reader, breaking piped input after earlier prompts).
func executeGCPCommand(reader *bufio.Reader, displayCmd string, program string, args ...string) error {
fmt.Println()
fmt.Printf("Executing: %s\n", displayCmd)
fmt.Println(strings.Repeat("-", 60))
Expand All @@ -397,9 +418,11 @@ func executeGCPCommand(displayCmd string, program string, args ...string) error
if err != nil {
fmt.Printf("Command failed: %v\n", err)
fmt.Print("Continue anyway? [y/N]: ")
reader := bufio.NewReader(os.Stdin)
response, _ := reader.ReadString('\n')
if strings.ToLower(strings.TrimSpace(response)) != "y" {
response, readErr := readTrimmedLine(reader)
if readErr != nil {
return fmt.Errorf("failed to read response: %w", readErr)
}
if strings.ToLower(response) != "y" {
return fmt.Errorf("command failed: %w", err)
}
}
Expand Down
4 changes: 3 additions & 1 deletion cmd/rekey/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,9 @@ func rekeyOne(ctx context.Context, db *database.Connection, id, blob string, zer
return outcomeErrored
}
if _, err := tx.Exec(ctx, `UPDATE account_credentials SET encrypted_blob = $1 WHERE id = $2`, newBlob, id); err != nil {
_ = tx.Rollback(ctx)
if rErr := tx.Rollback(ctx); rErr != nil {
log.Printf("rekey: rollback id=%s: %v", id, rErr)
}
log.Printf("rekey: update id=%s: %v", id, err)
return outcomeErrored
}
Expand Down
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -109,7 +109,7 @@ require (
github.com/go-jose/go-jose/v4 v4.1.4
github.com/golang-migrate/migrate/v4 v4.19.1
github.com/google/uuid v1.6.0
github.com/jackc/pgx/v5 v5.8.0
github.com/jackc/pgx/v5 v5.9.2
github.com/pashagolub/pgxmock/v4 v4.9.0
github.com/testcontainers/testcontainers-go v0.42.0
github.com/testcontainers/testcontainers-go/modules/postgres v0.42.0
Expand Down
4 changes: 2 additions & 2 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -250,8 +250,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.8.0 h1:TYPDoleBBme0xGSAX3/+NujXXtpZn9HBONkQC7IEZSo=
github.com/jackc/pgx/v5 v5.8.0/go.mod h1:QVeDInX2m9VyzvNeiCJVjCkNFqzsNb43204HshNSZKw=
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/keybase/go-keychain v0.0.1 h1:way+bWYa6lDppZoZcgMbYsvC7GxljxrskdNInRtuthU=
Expand Down
Loading
Loading