diff --git a/internal/ucind/cluster.go b/internal/ucind/cluster.go index 6085b14e..9622b79f 100644 --- a/internal/ucind/cluster.go +++ b/internal/ucind/cluster.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "net/netip" + "sync" "time" "github.com/cenkalti/backoff/v4" @@ -16,6 +17,7 @@ import ( "github.com/psviderski/uncloud/api/pb" "github.com/psviderski/uncloud/internal/machine" "github.com/psviderski/uncloud/internal/machine/cluster" + "github.com/psviderski/uncloud/pkg/client" "google.golang.org/protobuf/types/known/emptypb" ) @@ -322,6 +324,78 @@ func (p *Provisioner) WaitClusterReady(ctx context.Context, c Cluster, timeout t return nil } +// WaitClusterMeshReady waits until every machine can reach every other machine through the WireGuard mesh network. +func (p *Provisioner) WaitClusterMeshReady(ctx context.Context, c Cluster, timeout time.Duration) error { + if len(c.Machines) < 2 { + return nil + } + + clients := make([]*client.Client, len(c.Machines)) + for i := range c.Machines { + cli, err := c.Machines[i].Connect(ctx) + if err != nil { + return fmt.Errorf("connect to machine '%s' over TCP '%s': %w", + c.Machines[i].Name, c.Machines[i].APIAddress, err) + } + clients[i] = cli + defer cli.Close() + } + + type machinePair struct { + source int + target int + } + pairs := make([]machinePair, 0, len(c.Machines)*(len(c.Machines)-1)) + for source := range c.Machines { + for target := range c.Machines { + if source != target { + pairs = append(pairs, machinePair{source: source, target: target}) + } + } + } + + boff := backoff.WithContext(backoff.NewExponentialBackOff( + backoff.WithInitialInterval(100*time.Millisecond), + backoff.WithMaxInterval(time.Second), + backoff.WithMaxElapsedTime(timeout), + ), ctx) + + checkMeshReady := func() error { + errs := make([]error, len(pairs)) + var wg sync.WaitGroup + for i, pair := range pairs { + wg.Go(func() { + source := c.Machines[pair.source] + target := c.Machines[pair.target] + + probeCtx, cancel := context.WithTimeout(ctx, 2*time.Second) + defer cancel() + resp, err := clients[pair.source].MachineClient.InspectMachine( + clients[pair.source].ProxySingleMachineContext(probeCtx, target.ID), + nil, + ) + if err != nil { + errs[i] = fmt.Errorf("reach machine '%s' from '%s': %w", target.Name, source.Name, err) + return + } + if len(resp.Machines) != 1 || resp.Machines[0].Machine == nil || + resp.Machines[0].Machine.Id != target.ID { + errs[i] = fmt.Errorf("machine '%s' returned an unexpected response when reached from '%s'", + target.Name, source.Name) + } + }) + } + wg.Wait() + + return errors.Join(errs...) + } + if err := backoff.Retry(checkMeshReady, boff); err != nil { + return fmt.Errorf("wait for cluster WireGuard mesh to be ready: %w", err) + } + + return nil +} + func (p *Provisioner) ListClusters(ctx context.Context) ([]Cluster, error) { nets, err := p.dockerCli.NetworkList(ctx, network.ListOptions{ Filters: filters.NewArgs( diff --git a/test/e2e/cluster_test.go b/test/e2e/cluster_test.go index 2422f3d1..dd5c366d 100644 --- a/test/e2e/cluster_test.go +++ b/test/e2e/cluster_test.go @@ -36,6 +36,10 @@ func createTestCluster( if envName != "" { c, err := p.InspectCluster(ctx, envName) if err == nil { + if waitReady { + require.NoError(t, p.WaitClusterReady(ctx, c, 90*time.Second)) + require.NoError(t, p.WaitClusterMeshReady(ctx, c, 90*time.Second)) + } return c, p } if !errors.Is(err, ucind.ErrNotFound) { @@ -60,6 +64,7 @@ func createTestCluster( if waitReady { require.NoError(t, p.WaitClusterReady(ctx, c, 90*time.Second)) + require.NoError(t, p.WaitClusterMeshReady(ctx, c, 90*time.Second)) } return c, p @@ -111,6 +116,7 @@ func TestClusterLifecycle(t *testing.T) { }, 30*time.Second, 50*time.Millisecond, "cluster store not reconciled on machine #%d", i+1) } }) + require.NoError(t, p.WaitClusterMeshReady(ctx, c, 90*time.Second)) t.Run("inspect", func(t *testing.T) { cluster, err := p.InspectCluster(ctx, name)