diff --git a/management/server/types/legacynmap/benchmark_test.go b/management/server/types/legacynmap/benchmark_test.go index 9206d2812..22e291e00 100644 --- a/management/server/types/legacynmap/benchmark_test.go +++ b/management/server/types/legacynmap/benchmark_test.go @@ -94,7 +94,7 @@ func BenchmarkGetNetworkMapData(b *testing.B) { pool, err := pgxpool.NewWithConfig(ctx, cfg) require.NoError(b, err, "connect nmdata store") b.Cleanup(pool.Close) - nmStore := &networkmap_pgsql.PgStore{Pool: pool} + nmStore := nmDataStore(b, &networkmap_pgsql.PgStore{Pool: pool}) for _, accountID := range benchAccountIDs(b, ctx, statsConn) { b.Run(accountID, func(b *testing.B) { @@ -161,7 +161,7 @@ func BenchmarkNetworkMapDataFullRound(b *testing.B) { pool, err := pgxpool.NewWithConfig(ctx, cfg) require.NoError(b, err, "connect nmdata store") b.Cleanup(pool.Close) - nmStore := &networkmap_pgsql.PgStore{Pool: pool} + nmStore := nmDataStore(b, &networkmap_pgsql.PgStore{Pool: pool}) for _, accountID := range benchAccountIDs(b, ctx, statsConn) { b.Run(accountID, func(b *testing.B) { diff --git a/management/server/types/legacynmap/equivalence_test.go b/management/server/types/legacynmap/equivalence_test.go index 79eafd784..3fd5514d3 100644 --- a/management/server/types/legacynmap/equivalence_test.go +++ b/management/server/types/legacynmap/equivalence_test.go @@ -42,6 +42,7 @@ import ( "strings" "testing" + "github.com/golang/mock/gomock" "github.com/stretchr/testify/require" "google.golang.org/protobuf/encoding/prototext" goproto "google.golang.org/protobuf/proto" @@ -51,8 +52,11 @@ import ( nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" + "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" + "github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/management/server/types/legacynmap" @@ -86,9 +90,10 @@ func TestNetworkMapProtoEquivalence(t *testing.T) { require.NoError(t, err, "connect to postgres") t.Cleanup(func() { testStore.Close(ctx) }) - nmStore, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) + pgStore, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) require.NoError(t, err, "connect nmdata store") - t.Cleanup(func() { nmStore.Pool.Close() }) + t.Cleanup(func() { pgStore.Pool.Close() }) + nmStore := nmDataStore(t, pgStore) accountIDs := equivAccountIDs(t, dsn) require.NotEmpty(t, accountIDs, "no accounts selected") @@ -121,7 +126,7 @@ func TestNetworkMapProtoEquivalence(t *testing.T) { // checkAccount compares both paths for every peer of one account. Nothing is // retained across peers, so memory stays flat within an account. -func checkAccount(ctx context.Context, t *testing.T, nmStore *networkmap_pgsql.PgStore, account *types.Account, maxPeers int, stats *equivStats) { +func checkAccount(ctx context.Context, t *testing.T, nmStore *networkmapdb.NetworkMapDBStoreImpl, account *types.Account, maxPeers int, stats *equivStats) { t.Helper() if len(account.Peers) == 0 { @@ -218,6 +223,23 @@ func checkAccount(ctx context.Context, t *testing.T, nmStore *networkmap_pgsql.P } } +// nmDataStore wraps a raw connection store the way production's factory does. +// The validator marks every peer validated and the extra settings are empty: +// checkAccount overwrites ValidatedPeers anyway, and neither reaches the +// compared network map. +func nmDataStore(tb testing.TB, s networkmapdb.NetworkMapDBStore) *networkmapdb.NetworkMapDBStoreImpl { + tb.Helper() + + extraSettings := settings.NewMockManager(gomock.NewController(tb)) + extraSettings.EXPECT().GetExtraSettings(gomock.Any(), gomock.Any()).Return(&types.ExtraSettings{}, nil).AnyTimes() + + return &networkmapdb.NetworkMapDBStoreImpl{ + Store: s, + IntegratedPeerValidator: &validator.IntegratedValidatorImpl{}, + ExtraSettingsManager: extraSettings, + } +} + func equivDSN() string { if dsn := os.Getenv("NETBIRD_STORE_ENGINE_POSTGRES_DSN"); dsn != "" { return dsn