diff --git a/shared/relay/client/fallback_opener_test.go b/shared/relay/client/fallback_opener_test.go new file mode 100644 index 000000000..9123aca86 --- /dev/null +++ b/shared/relay/client/fallback_opener_test.go @@ -0,0 +1,175 @@ +package client + +import ( + "context" + "errors" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/relay/server" +) + +type fakeConn struct { + net.Conn + closed chan struct{} +} + +func newFakeConn() *fakeConn { + return &fakeConn{closed: make(chan struct{})} +} + +func (c *fakeConn) Close() error { + close(c.closed) + return nil +} + +func newTestConnRace(t *testing.T) *connRace { + t.Helper() + raceCtx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + _, cancelPreferred := context.WithCancel(raceCtx) + otherCtx, cancelOther := context.WithCancel(raceCtx) + timer := time.NewTimer(time.Hour) + timer.Stop() + return &connRace{ + opener: &FallbackOpener{}, + peerKey: "peerKey", + raceCtx: raceCtx, + otherCtx: otherCtx, + cancelPreferred: cancelPreferred, + cancelOther: cancelOther, + results: make(chan raceAttempt, 2), + fallbackTimer: timer, + } +} + +func TestHandleResult_PreferredSucceeds(t *testing.T) { + c := newTestConnRace(t) + + conn := newFakeConn() + o := c.handleResult(raceAttempt{conn: conn}) + + require.True(t, o.done) + require.NoError(t, o.err) + require.Same(t, net.Conn(conn), o.conn) + require.False(t, c.otherStarted, "fallback must not start once the preferred attempt wins") +} + +func TestHandleResult_ConnAlreadyExistsIsSuccess(t *testing.T) { + c := newTestConnRace(t) + + o := c.handleResult(raceAttempt{err: ErrConnAlreadyExists}) + + require.True(t, o.done) + require.ErrorIs(t, o.err, ErrConnAlreadyExists) + require.False(t, c.otherStarted) +} + +func TestHandleResult_PreferredFailsStartsOther(t *testing.T) { + c := newTestConnRace(t) + // The fallback attempt opens against a stalling listener so startOther's + // goroutine blocks on Connect until raceCtx is cancelled by t.Cleanup. + serverAddr, _ := stallingRelayListener(t) + c.opener.foreignStore = NewForeignRelaysStore(c.raceCtx, hmacTokenStore, "alice", 1280, newTransportFallback(), func(string) {}, keepUnusedServerTime) + c.remoteRelayServer = RelayServer{Addr: serverAddr} + c.preferForeign = false + + o := c.handleResult(raceAttempt{err: errors.New("boom")}) + + require.False(t, o.done, "a single failure must not settle the race") + require.True(t, c.otherStarted, "the fallback attempt must start after the preferred one fails") + require.EqualError(t, c.lastErr, "boom") +} + +func TestHandleResult_BothFailReturnsLastErr(t *testing.T) { + c := newTestConnRace(t) + c.otherStarted = true + c.settled = 1 + c.lastErr = errors.New("first") + + o := c.handleResult(raceAttempt{err: errors.New("second")}) + + require.True(t, o.done) + require.EqualError(t, o.err, "second") +} + +func TestOnTimeout_PrefersLastErr(t *testing.T) { + c := newTestConnRace(t) + c.lastErr = errors.New("dial failed") + + _, err := c.onTimeout() + require.EqualError(t, err, "dial failed") +} + +func TestOnTimeout_FallsBackToCtxErr(t *testing.T) { + c := newTestConnRace(t) + raceCtx, cancel := context.WithCancel(context.Background()) + cancel() + c.raceCtx = raceCtx + + _, err := c.onTimeout() + require.ErrorIs(t, err, context.Canceled) +} + +func TestDrainLoser_ClosesLateWinner(t *testing.T) { + r := &FallbackOpener{} + results := make(chan raceAttempt, 2) + + loser := newFakeConn() + results <- raceAttempt{conn: loser} + + done := make(chan struct{}) + go func() { + defer close(done) + r.drainLoser(results, 1, true) + }() + + select { + case <-loser.closed: + case <-time.After(2 * time.Second): + t.Fatal("losing connection was not closed") + } + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("drainLoser did not return") + } +} + +func TestDrainLoser_NoOtherAttempt(t *testing.T) { + r := &FallbackOpener{} + results := make(chan raceAttempt) + + done := make(chan struct{}) + go func() { + defer close(done) + r.drainLoser(results, 1, false) + }() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("drainLoser blocked with no outstanding attempt") + } +} + +func startTestRelayServer(t *testing.T, addr string) string { + t.Helper() + + srv, err := server.NewServer(newManagerTestServerConfig(addr)) + require.NoError(t, err) + + errChan := make(chan error, 1) + go func() { + if err := srv.Listen(server.ListenerConfig{Address: addr}); err != nil { + errChan <- err + } + }() + t.Cleanup(func() { _ = srv.Shutdown(context.Background()) }) + + require.NoError(t, waitForServerToStart(errChan)) + return addr +} diff --git a/shared/relay/client/foreign_relays_store_test.go b/shared/relay/client/foreign_relays_store_test.go new file mode 100644 index 000000000..19fe7bc3b --- /dev/null +++ b/shared/relay/client/foreign_relays_store_test.go @@ -0,0 +1,184 @@ +package client + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func newTestForeignStore(t *testing.T, ctx context.Context) *ForeignRelaysStore { + t.Helper() + return NewForeignRelaysStore(ctx, hmacTokenStore, "alice", 1280, newTransportFallback(), func(string) {}, keepUnusedServerTime) +} + +func TestForeignStore_AcquireDedupsConcurrentOpens(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52601") + server := RelayServer{Addr: "rel://" + addr} + + store := newTestForeignStore(t, ctx) + + const n = 8 + var wg sync.WaitGroup + results := make([]*foreignRelay, n) + for i := 0; i < n; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + fr, err := store.acquire(server) + require.NoError(t, err) + results[i] = fr + }(i) + } + wg.Wait() + + first := results[0] + require.NotNil(t, first) + for _, fr := range results { + require.Same(t, first, fr, "all acquires must share the same foreign relay") + } + + store.mu.RLock() + require.Len(t, store.clients, 1, "only one client entry must be stored") + require.Equal(t, n, first.inUse, "every acquire must be counted") + store.mu.RUnlock() +} + +func TestForeignStore_AcquireReleaseRefcount(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52602") + server := RelayServer{Addr: "rel://" + addr} + + store := newTestForeignStore(t, ctx) + + fr, err := store.acquire(server) + require.NoError(t, err) + _, err = store.acquire(server) + require.NoError(t, err) + + store.mu.RLock() + require.Equal(t, 2, fr.inUse) + store.mu.RUnlock() + + store.release(fr) + store.mu.RLock() + require.Equal(t, 1, fr.inUse) + require.Len(t, store.clients, 1, "release must not evict the client") + store.mu.RUnlock() +} + +func TestForeignStore_AcquireConnectFailure(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + + store := newTestForeignStore(t, ctx) + + // Nothing is listening on this port, so Connect fails. + _, err := store.acquire(RelayServer{Addr: "rel://127.0.0.1:1"}) + require.Error(t, err) + + store.mu.RLock() + require.Empty(t, store.clients, "a failed connect must not leave a client behind") + store.mu.RUnlock() +} + +func TestForeignStore_Evict(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + store := newTestForeignStore(t, ctx) + store.clients["rel://a"] = &foreignRelay{} + store.clients["rel://b"] = &foreignRelay{} + + store.evict("rel://a") + store.evict("rel://missing") + + require.NotContains(t, store.clients, "rel://a") + require.Contains(t, store.clients, "rel://b") +} + +func TestForeignStore_CleanupUnused_KeepsRecent(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52603") + store := newTestForeignStore(t, ctx) + + fr, err := store.acquire(RelayServer{Addr: "rel://" + addr}) + require.NoError(t, err) + store.release(fr) + + store.cleanupUnused() + + store.mu.RLock() + require.Len(t, store.clients, 1, "a freshly created client must be kept") + store.mu.RUnlock() +} + +func TestForeignStore_CleanupUnused_KeepsInUse(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52604") + store := newTestForeignStore(t, ctx) + + fr, err := store.acquire(RelayServer{Addr: "rel://" + addr}) + require.NoError(t, err) + + store.mu.Lock() + fr.created = time.Now().Add(-2 * keepUnusedServerTime) + store.mu.Unlock() + + store.cleanupUnused() + + store.mu.RLock() + require.Len(t, store.clients, 1, "an in-use client must be kept even when aged") + store.mu.RUnlock() +} + +func TestForeignStore_CleanupUnused_EvictsAgedIdle(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52605") + store := newTestForeignStore(t, ctx) + + fr, err := store.acquire(RelayServer{Addr: "rel://" + addr}) + require.NoError(t, err) + store.release(fr) + + store.mu.Lock() + fr.created = time.Now().Add(-2 * keepUnusedServerTime) + store.mu.Unlock() + + require.False(t, fr.client.HasConns(), "no peer connections were opened") + + store.cleanupUnused() + + store.mu.RLock() + require.Empty(t, store.clients, "an aged idle client must be evicted") + store.mu.RUnlock() +} + +func TestForeignStore_States(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + addr := startTestRelayServer(t, "127.0.0.1:52606") + store := newTestForeignStore(t, ctx) + + fr, err := store.acquire(RelayServer{Addr: "rel://" + addr}) + require.NoError(t, err) + store.release(fr) + + states := store.states() + require.Len(t, states, 1) + require.NotEmpty(t, states[0].URL) +}