Skip to content
Open
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
2 changes: 1 addition & 1 deletion internal/command/mpg/v1/run_connect.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
118 changes: 95 additions & 23 deletions internal/command/mpg/v1/run_proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,84 +12,156 @@ 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
}

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
}
169 changes: 169 additions & 0 deletions internal/command/mpg/v1/run_proxy_test.go
Original file line number Diff line number Diff line change
@@ -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)
})
}
2 changes: 1 addition & 1 deletion internal/command/mpg/v2/run_connect.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
Loading
Loading