diff --git a/internal/command/mpg/v1/run_connect.go b/internal/command/mpg/v1/run_connect.go index 4de567c915..8f10e81262 100644 --- a/internal/command/mpg/v1/run_connect.go +++ b/internal/command/mpg/v1/run_connect.go @@ -77,7 +77,7 @@ func RunConnect(ctx context.Context, clusterID string, resolvedOrgSlug string) ( } } - cluster, params, credentials, err := GetMpgProxyParams(ctx, localProxyPort, username, clusterID, resolvedOrgSlug) + cluster, params, credentials, err := GetMpgConnectParams(ctx, localProxyPort, username, clusterID, resolvedOrgSlug) if err != nil { return err } diff --git a/internal/command/mpg/v1/run_proxy.go b/internal/command/mpg/v1/run_proxy.go index eaadb2ceb8..a82b04a159 100644 --- a/internal/command/mpg/v1/run_proxy.go +++ b/internal/command/mpg/v1/run_proxy.go @@ -12,7 +12,7 @@ import ( ) func RunProxy(ctx context.Context, clusterID string, resolvedOrgSlug string, proxyPort string) error { - _, params, _, err := GetMpgProxyParams(ctx, proxyPort, "", clusterID, resolvedOrgSlug) + _, params, err := GetMpgProxyParams(ctx, proxyPort, clusterID, resolvedOrgSlug) if err != nil { return err } @@ -20,76 +20,148 @@ func RunProxy(ctx context.Context, clusterID string, resolvedOrgSlug string, pro return proxy.Connect(ctx, params) } -// GetMpgProxyParams builds proxy connection parameters for a given cluster. +// GetMpgProxyParams builds proxy connection parameters for a given cluster +// without requiring database credentials. // resolvedOrgSlug should already be the aliased slug suitable for wireguard tunnels. func GetMpgProxyParams( + ctx context.Context, + localProxyPort string, + clusterID string, + resolvedOrgSlug string, +) (*mpgv1.ManagedCluster, *proxy.ConnectParams, error) { + response, err := getCluster(ctx, clusterID) + if err != nil { + return nil, nil, err + } + + cluster, params, err := buildProxyParams(ctx, response, localProxyPort, resolvedOrgSlug) + if err != nil { + return nil, nil, err + } + + return cluster, params, nil +} + +// GetMpgConnectParams builds proxy connection parameters and resolves the +// database credentials needed by fly mpg connect. +func GetMpgConnectParams( ctx context.Context, localProxyPort string, username string, clusterID string, resolvedOrgSlug string, ) (*mpgv1.ManagedCluster, *proxy.ConnectParams, *mpgv1.GetManagedClusterCredentialsResponse, error) { - client := flyutil.ClientFromContext(ctx) - mpgClient := mpgv1.ClientFromContext(ctx) + response, err := getCluster(ctx, clusterID) + if err != nil { + return nil, nil, nil, err + } + + credentials, err := resolveConnectCredentials(ctx, response, username) + if err != nil { + return nil, nil, nil, err + } + + cluster, params, err := buildProxyParams(ctx, response, localProxyPort, resolvedOrgSlug) + if err != nil { + return nil, nil, nil, err + } + + return cluster, params, credentials, nil +} - // Get cluster details +func getCluster(ctx context.Context, clusterID string) (*mpgv1.GetManagedClusterResponse, error) { + mpgClient := mpgv1.ClientFromContext(ctx) response, err := mpgClient.GetManagedClusterById(ctx, clusterID) if err != nil { - return nil, nil, nil, fmt.Errorf("failed retrieving cluster %s: %w", clusterID, err) + return nil, fmt.Errorf("failed retrieving cluster %s: %w", clusterID, err) } - cluster := &response.Data + return &response, nil +} - // Get credentials - use user-specific endpoint if username provided, otherwise use default +func resolveConnectCredentials( + ctx context.Context, + response *mpgv1.GetManagedClusterResponse, + username string, +) (*mpgv1.GetManagedClusterCredentialsResponse, error) { var credentials mpgv1.GetManagedClusterCredentialsResponse if username != "" { - userCreds, err := mpgClient.GetUserCredentials(ctx, cluster.Id, username) + mpgClient := mpgv1.ClientFromContext(ctx) + userCreds, err := mpgClient.GetUserCredentials(ctx, response.Data.Id, username) if err != nil { - return nil, nil, nil, fmt.Errorf("failed retrieving credentials for user %s: %w", username, err) + return nil, fmt.Errorf("failed retrieving credentials for user %s: %w", username, err) } - // Convert user credentials to the standard format + credentials = mpgv1.GetManagedClusterCredentialsResponse{ User: userCreds.Data.User, Password: userCreds.Data.Password, - DBName: response.Credentials.DBName, // Use default DB name from cluster credentials + DBName: response.Credentials.DBName, } } else { credentials = response.Credentials } - // Validate cluster state (only for default credentials, user credentials don't have status) if username == "" { if credentials.Status == "initializing" { - return nil, nil, nil, fmt.Errorf("cluster is still initializing, wait a bit more") + return nil, fmt.Errorf("cluster is still initializing, wait a bit more") } if credentials.Status == "error" || credentials.Password == "" { - return nil, nil, nil, fmt.Errorf("error getting cluster password") + return nil, fmt.Errorf("error getting cluster password") } } else if credentials.Password == "" { - return nil, nil, nil, fmt.Errorf("error getting user password") + return nil, fmt.Errorf("error getting user password") } - if cluster.IpAssignments.Direct == "" { - return nil, nil, nil, fmt.Errorf("error getting cluster IP") + return &credentials, nil +} + +func buildProxyParams( + ctx context.Context, + response *mpgv1.GetManagedClusterResponse, + localProxyPort string, + resolvedOrgSlug string, +) (*mpgv1.ManagedCluster, *proxy.ConnectParams, error) { + cluster, params, err := proxyParams(response, localProxyPort, resolvedOrgSlug, flag.GetBindAddr(ctx), nil) + if err != nil { + return nil, nil, err } - // Establish wireguard tunnel + client := flyutil.ClientFromContext(ctx) + + // Establish wireguard tunnel after validating all prerequisites. agentclient, err := agent.Establish(ctx, client) if err != nil { - return nil, nil, nil, err + return nil, nil, err } dialer, err := agentclient.ConnectToTunnel(ctx, resolvedOrgSlug, "", false) if err != nil { - return nil, nil, nil, err + return nil, nil, err + } + + params.Dialer = dialer + + return cluster, params, nil +} + +func proxyParams( + response *mpgv1.GetManagedClusterResponse, + localProxyPort string, + resolvedOrgSlug string, + bindAddr string, + dialer agent.Dialer, +) (*mpgv1.ManagedCluster, *proxy.ConnectParams, error) { + cluster := &response.Data + if cluster.IpAssignments.Direct == "" { + return nil, nil, fmt.Errorf("error getting cluster IP") } return cluster, &proxy.ConnectParams{ Ports: []string{localProxyPort, "5432"}, OrganizationSlug: resolvedOrgSlug, Dialer: dialer, - BindAddr: flag.GetBindAddr(ctx), + BindAddr: bindAddr, RemoteHost: cluster.IpAssignments.Direct, - }, &credentials, nil + }, nil } diff --git a/internal/command/mpg/v1/run_proxy_test.go b/internal/command/mpg/v1/run_proxy_test.go new file mode 100644 index 0000000000..884d0daf9d --- /dev/null +++ b/internal/command/mpg/v1/run_proxy_test.go @@ -0,0 +1,169 @@ +package cmdv1 + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/superfly/flyctl/internal/mock" + "github.com/superfly/flyctl/internal/uiex/mpg" + mpgv1 "github.com/superfly/flyctl/internal/uiex/mpg/v1" + "github.com/superfly/flyctl/wg" +) + +type testDialer struct{} + +func (*testDialer) State() *wg.WireGuardState { return nil } + +func (*testDialer) Config() *wg.Config { return nil } + +func (*testDialer) DialContext(context.Context, string, string) (net.Conn, error) { + return nil, nil +} + +func TestProxyParamsIgnoreCredentials(t *testing.T) { + tests := []struct { + name string + credentials mpgv1.GetManagedClusterCredentialsResponse + }{ + { + name: "initializing credentials", + credentials: mpgv1.GetManagedClusterCredentialsResponse{ + Status: "initializing", + }, + }, + { + name: "empty password", + credentials: mpgv1.GetManagedClusterCredentialsResponse{ + Status: "ready", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + response := mpgv1.GetManagedClusterResponse{ + Data: mpgv1.ManagedCluster{ + IpAssignments: mpg.ManagedClusterIpAssignments{Direct: "fdaa:0:1234::2"}, + }, + Credentials: tt.credentials, + } + dialer := &testDialer{} + + cluster, params, err := proxyParams(&response, "15432", "test-org", "127.0.0.2", dialer) + require.NoError(t, err) + assert.Same(t, &response.Data, cluster) + assert.Equal(t, []string{"15432", "5432"}, params.Ports) + assert.Equal(t, "test-org", params.OrganizationSlug) + assert.Equal(t, "127.0.0.2", params.BindAddr) + assert.Equal(t, "fdaa:0:1234::2", params.RemoteHost) + assert.Same(t, dialer, params.Dialer) + }) + } +} + +func TestProxyParamsRequireDirectIP(t *testing.T) { + cluster, params, err := proxyParams(&mpgv1.GetManagedClusterResponse{}, "15432", "test-org", "127.0.0.1", nil) + + require.EqualError(t, err, "error getting cluster IP") + assert.Nil(t, cluster) + assert.Nil(t, params) +} + +func TestResolveDefaultConnectCredentials(t *testing.T) { + tests := []struct { + name string + credentials mpgv1.GetManagedClusterCredentialsResponse + err string + }{ + { + name: "initializing", + credentials: mpgv1.GetManagedClusterCredentialsResponse{Status: "initializing"}, + err: "cluster is still initializing, wait a bit more", + }, + { + name: "error status", + credentials: mpgv1.GetManagedClusterCredentialsResponse{Status: "error", Password: "password"}, + err: "error getting cluster password", + }, + { + name: "empty password", + credentials: mpgv1.GetManagedClusterCredentialsResponse{Status: "ready"}, + err: "error getting cluster password", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + response := &mpgv1.GetManagedClusterResponse{Credentials: tt.credentials} + + credentials, err := resolveConnectCredentials(context.Background(), response, "") + + require.EqualError(t, err, tt.err) + assert.Nil(t, credentials) + }) + } +} + +func TestResolveExplicitUserConnectCredentials(t *testing.T) { + response := &mpgv1.GetManagedClusterResponse{ + Data: mpgv1.ManagedCluster{Id: "cluster-id"}, + Credentials: mpgv1.GetManagedClusterCredentialsResponse{ + Status: "initializing", + DBName: "default-db", + }, + } + client := &mock.MpgV1Client{ + GetUserCredentialsFunc: func(_ context.Context, clusterID string, username string) (mpgv1.GetUserCredentialsResponse, error) { + assert.Equal(t, "cluster-id", clusterID) + assert.Equal(t, "app-user", username) + + result := mpgv1.GetUserCredentialsResponse{} + result.Data.User = "app-user" + result.Data.Password = "secret" + + return result, nil + }, + } + ctx := mpgv1.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, response, "app-user") + + require.NoError(t, err) + assert.Equal(t, "app-user", credentials.User) + assert.Equal(t, "secret", credentials.Password) + assert.Equal(t, "default-db", credentials.DBName) +} + +func TestResolveExplicitUserConnectCredentialsErrors(t *testing.T) { + t.Run("empty password", func(t *testing.T) { + client := &mock.MpgV1Client{ + GetUserCredentialsFunc: func(context.Context, string, string) (mpgv1.GetUserCredentialsResponse, error) { + return mpgv1.GetUserCredentialsResponse{}, nil + }, + } + ctx := mpgv1.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, &mpgv1.GetManagedClusterResponse{}, "app-user") + + require.EqualError(t, err, "error getting user password") + assert.Nil(t, credentials) + }) + + t.Run("request failure", func(t *testing.T) { + client := &mock.MpgV1Client{ + GetUserCredentialsFunc: func(context.Context, string, string) (mpgv1.GetUserCredentialsResponse, error) { + return mpgv1.GetUserCredentialsResponse{}, errors.New("request failed") + }, + } + ctx := mpgv1.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, &mpgv1.GetManagedClusterResponse{}, "app-user") + + require.EqualError(t, err, "failed retrieving credentials for user app-user: request failed") + assert.Nil(t, credentials) + }) +} diff --git a/internal/command/mpg/v2/run_connect.go b/internal/command/mpg/v2/run_connect.go index c1e3c65582..6113a141ea 100644 --- a/internal/command/mpg/v2/run_connect.go +++ b/internal/command/mpg/v2/run_connect.go @@ -77,7 +77,7 @@ func RunConnect(ctx context.Context, clusterID string, resolvedOrgSlug string, p } } - cluster, params, credentials, err := GetMpgProxyParams(ctx, localProxyPort, username, clusterID, resolvedOrgSlug) + cluster, params, credentials, err := GetMpgConnectParams(ctx, localProxyPort, username, clusterID, resolvedOrgSlug) if err != nil { return err } diff --git a/internal/command/mpg/v2/run_proxy.go b/internal/command/mpg/v2/run_proxy.go index 6a8597b27d..bcff1c67fa 100644 --- a/internal/command/mpg/v2/run_proxy.go +++ b/internal/command/mpg/v2/run_proxy.go @@ -12,7 +12,7 @@ import ( ) func RunProxy(ctx context.Context, clusterID string, resolvedOrgSlug string, proxyPort string) error { - _, params, _, err := GetMpgProxyParams(ctx, proxyPort, "", clusterID, resolvedOrgSlug) + _, params, err := GetMpgProxyParams(ctx, proxyPort, clusterID, resolvedOrgSlug) if err != nil { return err } @@ -20,76 +20,148 @@ func RunProxy(ctx context.Context, clusterID string, resolvedOrgSlug string, pro return proxy.Connect(ctx, params) } -// GetMpgProxyParams builds proxy connection parameters for a given cluster. +// GetMpgProxyParams builds proxy connection parameters for a given cluster +// without requiring database credentials. // resolvedOrgSlug should already be the aliased slug suitable for wireguard tunnels. func GetMpgProxyParams( + ctx context.Context, + localProxyPort string, + clusterID string, + resolvedOrgSlug string, +) (*mpgv2.ManagedCluster, *proxy.ConnectParams, error) { + response, err := getCluster(ctx, clusterID) + if err != nil { + return nil, nil, err + } + + cluster, params, err := buildProxyParams(ctx, response, localProxyPort, resolvedOrgSlug) + if err != nil { + return nil, nil, err + } + + return cluster, params, nil +} + +// GetMpgConnectParams builds proxy connection parameters and resolves the +// database credentials needed by fly mpg connect. +func GetMpgConnectParams( ctx context.Context, localProxyPort string, username string, clusterID string, resolvedOrgSlug string, ) (*mpgv2.ManagedCluster, *proxy.ConnectParams, *mpgv2.GetClusterCredentialsResponse, error) { - client := flyutil.ClientFromContext(ctx) - mpgClient := mpgv2.ClientFromContext(ctx) + response, err := getCluster(ctx, clusterID) + if err != nil { + return nil, nil, nil, err + } + + credentials, err := resolveConnectCredentials(ctx, response, username) + if err != nil { + return nil, nil, nil, err + } + + cluster, params, err := buildProxyParams(ctx, response, localProxyPort, resolvedOrgSlug) + if err != nil { + return nil, nil, nil, err + } + + return cluster, params, credentials, nil +} - // Get cluster details +func getCluster(ctx context.Context, clusterID string) (*mpgv2.GetClusterResponse, error) { + mpgClient := mpgv2.ClientFromContext(ctx) response, err := mpgClient.GetClusterById(ctx, clusterID) if err != nil { - return nil, nil, nil, fmt.Errorf("failed retrieving cluster %s: %w", clusterID, err) + return nil, fmt.Errorf("failed retrieving cluster %s: %w", clusterID, err) } - cluster := &response.Data + return &response, nil +} - // Get credentials - use user-specific endpoint if username provided, otherwise use default +func resolveConnectCredentials( + ctx context.Context, + response *mpgv2.GetClusterResponse, + username string, +) (*mpgv2.GetClusterCredentialsResponse, error) { var credentials mpgv2.GetClusterCredentialsResponse if username != "" { - userCreds, err := mpgClient.GetUserCredentials(ctx, cluster.Id, username) + mpgClient := mpgv2.ClientFromContext(ctx) + userCreds, err := mpgClient.GetUserCredentials(ctx, response.Data.Id, username) if err != nil { - return nil, nil, nil, fmt.Errorf("failed retrieving credentials for user %s: %w", username, err) + return nil, fmt.Errorf("failed retrieving credentials for user %s: %w", username, err) } - // Convert user credentials to the standard format + credentials = mpgv2.GetClusterCredentialsResponse{ User: userCreds.Data.User, Password: userCreds.Data.Password, - DBName: response.Credentials.DBName, // Use default DB name from cluster credentials + DBName: response.Credentials.DBName, } } else { credentials = response.Credentials } - // Validate cluster state (only for default credentials, user credentials don't have status) if username == "" { if credentials.Status == "initializing" { - return nil, nil, nil, fmt.Errorf("cluster is still initializing, wait a bit more") + return nil, fmt.Errorf("cluster is still initializing, wait a bit more") } if credentials.Status == "error" || credentials.Password == "" { - return nil, nil, nil, fmt.Errorf("error getting cluster password") + return nil, fmt.Errorf("error getting cluster password") } } else if credentials.Password == "" { - return nil, nil, nil, fmt.Errorf("error getting user password") + return nil, fmt.Errorf("error getting user password") } - if cluster.IpAssignments.Direct == "" { - return nil, nil, nil, fmt.Errorf("error getting cluster IP") + return &credentials, nil +} + +func buildProxyParams( + ctx context.Context, + response *mpgv2.GetClusterResponse, + localProxyPort string, + resolvedOrgSlug string, +) (*mpgv2.ManagedCluster, *proxy.ConnectParams, error) { + cluster, params, err := proxyParams(response, localProxyPort, resolvedOrgSlug, flag.GetBindAddr(ctx), nil) + if err != nil { + return nil, nil, err } - // Establish wireguard tunnel + client := flyutil.ClientFromContext(ctx) + + // Establish wireguard tunnel after validating all prerequisites. agentclient, err := agent.Establish(ctx, client) if err != nil { - return nil, nil, nil, err + return nil, nil, err } dialer, err := agentclient.ConnectToTunnel(ctx, resolvedOrgSlug, "", false) if err != nil { - return nil, nil, nil, err + return nil, nil, err + } + + params.Dialer = dialer + + return cluster, params, nil +} + +func proxyParams( + response *mpgv2.GetClusterResponse, + localProxyPort string, + resolvedOrgSlug string, + bindAddr string, + dialer agent.Dialer, +) (*mpgv2.ManagedCluster, *proxy.ConnectParams, error) { + cluster := &response.Data + if cluster.IpAssignments.Direct == "" { + return nil, nil, fmt.Errorf("error getting cluster IP") } return cluster, &proxy.ConnectParams{ Ports: []string{localProxyPort, "5432"}, OrganizationSlug: resolvedOrgSlug, Dialer: dialer, - BindAddr: flag.GetBindAddr(ctx), + BindAddr: bindAddr, RemoteHost: cluster.IpAssignments.Direct, - }, &credentials, nil + }, nil } diff --git a/internal/command/mpg/v2/run_proxy_test.go b/internal/command/mpg/v2/run_proxy_test.go new file mode 100644 index 0000000000..90134aef15 --- /dev/null +++ b/internal/command/mpg/v2/run_proxy_test.go @@ -0,0 +1,169 @@ +package cmdv2 + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/superfly/flyctl/internal/mock" + "github.com/superfly/flyctl/internal/uiex/mpg" + mpgv2 "github.com/superfly/flyctl/internal/uiex/mpg/v2" + "github.com/superfly/flyctl/wg" +) + +type testDialer struct{} + +func (*testDialer) State() *wg.WireGuardState { return nil } + +func (*testDialer) Config() *wg.Config { return nil } + +func (*testDialer) DialContext(context.Context, string, string) (net.Conn, error) { + return nil, nil +} + +func TestProxyParamsIgnoreCredentials(t *testing.T) { + tests := []struct { + name string + credentials mpgv2.GetClusterCredentialsResponse + }{ + { + name: "initializing credentials", + credentials: mpgv2.GetClusterCredentialsResponse{ + Status: "initializing", + }, + }, + { + name: "empty password", + credentials: mpgv2.GetClusterCredentialsResponse{ + Status: "ready", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + response := mpgv2.GetClusterResponse{ + Data: mpgv2.ManagedCluster{ + IpAssignments: mpg.ManagedClusterIpAssignments{Direct: "fdaa:0:1234::2"}, + }, + Credentials: tt.credentials, + } + dialer := &testDialer{} + + cluster, params, err := proxyParams(&response, "15432", "test-org", "127.0.0.2", dialer) + require.NoError(t, err) + assert.Same(t, &response.Data, cluster) + assert.Equal(t, []string{"15432", "5432"}, params.Ports) + assert.Equal(t, "test-org", params.OrganizationSlug) + assert.Equal(t, "127.0.0.2", params.BindAddr) + assert.Equal(t, "fdaa:0:1234::2", params.RemoteHost) + assert.Same(t, dialer, params.Dialer) + }) + } +} + +func TestProxyParamsRequireDirectIP(t *testing.T) { + cluster, params, err := proxyParams(&mpgv2.GetClusterResponse{}, "15432", "test-org", "127.0.0.1", nil) + + require.EqualError(t, err, "error getting cluster IP") + assert.Nil(t, cluster) + assert.Nil(t, params) +} + +func TestResolveDefaultConnectCredentials(t *testing.T) { + tests := []struct { + name string + credentials mpgv2.GetClusterCredentialsResponse + err string + }{ + { + name: "initializing", + credentials: mpgv2.GetClusterCredentialsResponse{Status: "initializing"}, + err: "cluster is still initializing, wait a bit more", + }, + { + name: "error status", + credentials: mpgv2.GetClusterCredentialsResponse{Status: "error", Password: "password"}, + err: "error getting cluster password", + }, + { + name: "empty password", + credentials: mpgv2.GetClusterCredentialsResponse{Status: "ready"}, + err: "error getting cluster password", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + response := &mpgv2.GetClusterResponse{Credentials: tt.credentials} + + credentials, err := resolveConnectCredentials(context.Background(), response, "") + + require.EqualError(t, err, tt.err) + assert.Nil(t, credentials) + }) + } +} + +func TestResolveExplicitUserConnectCredentials(t *testing.T) { + response := &mpgv2.GetClusterResponse{ + Data: mpgv2.ManagedCluster{Id: "cluster-id"}, + Credentials: mpgv2.GetClusterCredentialsResponse{ + Status: "initializing", + DBName: "default-db", + }, + } + client := &mock.MpgV2Client{ + GetUserCredentialsFunc: func(_ context.Context, clusterID string, username string) (mpgv2.GetUserCredentialsResponse, error) { + assert.Equal(t, "cluster-id", clusterID) + assert.Equal(t, "app-user", username) + + result := mpgv2.GetUserCredentialsResponse{} + result.Data.User = "app-user" + result.Data.Password = "secret" + + return result, nil + }, + } + ctx := mpgv2.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, response, "app-user") + + require.NoError(t, err) + assert.Equal(t, "app-user", credentials.User) + assert.Equal(t, "secret", credentials.Password) + assert.Equal(t, "default-db", credentials.DBName) +} + +func TestResolveExplicitUserConnectCredentialsErrors(t *testing.T) { + t.Run("empty password", func(t *testing.T) { + client := &mock.MpgV2Client{ + GetUserCredentialsFunc: func(context.Context, string, string) (mpgv2.GetUserCredentialsResponse, error) { + return mpgv2.GetUserCredentialsResponse{}, nil + }, + } + ctx := mpgv2.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, &mpgv2.GetClusterResponse{}, "app-user") + + require.EqualError(t, err, "error getting user password") + assert.Nil(t, credentials) + }) + + t.Run("request failure", func(t *testing.T) { + client := &mock.MpgV2Client{ + GetUserCredentialsFunc: func(context.Context, string, string) (mpgv2.GetUserCredentialsResponse, error) { + return mpgv2.GetUserCredentialsResponse{}, errors.New("request failed") + }, + } + ctx := mpgv2.NewContextWithClient(context.Background(), client) + + credentials, err := resolveConnectCredentials(ctx, &mpgv2.GetClusterResponse{}, "app-user") + + require.EqualError(t, err, "failed retrieving credentials for user app-user: request failed") + assert.Nil(t, credentials) + }) +}