diff --git a/management/internals/shared/grpc/components_encoder.go b/management/internals/shared/grpc/components_encoder.go index b13180fd5..2c948887d 100644 --- a/management/internals/shared/grpc/components_encoder.go +++ b/management/internals/shared/grpc/components_encoder.go @@ -6,7 +6,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/networkmap" - nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/proto" ) @@ -215,7 +215,9 @@ func (e *componentEncoder) encodeGroups() []*proto.GroupCompact { groupCompactResources := func() []*proto.ResourceCompact { var toret []*proto.ResourceCompact for _, r := range g.Resources { - toret = append(toret, e.resourceToProto(r)) + if pr := e.resourceToProto(r); pr != nil { + toret = append(toret, pr) + } } return toret } diff --git a/management/server/types/legacynmap/benchmark_test.go b/management/server/types/legacynmap/benchmark_test.go new file mode 100644 index 000000000..9206d2812 --- /dev/null +++ b/management/server/types/legacynmap/benchmark_test.go @@ -0,0 +1,350 @@ +//go:build nmapequiv + +// Account-load benchmark: the legacy store.GetAccount hydration (pgx fast +// path, as in production) vs the nmdata store's GetNetworkMapData, against the +// same Postgres copy as the equivalence test. +// +// NETBIRD_STORE_ENGINE_POSTGRES_DSN='...' go test -tags nmapequiv \ +// -run '^$' -bench . -benchtime 5x -timeout 60m \ +// ./management/server/types/legacynmap/ +// +// NETMAP_ACCOUNTS selects the accounts (comma-separated); by default the ten +// accounts with the most peers are used. Each account is a sub-benchmark, so +// the two paths can be compared per account. One warmup call runs untimed +// before each measurement so Postgres buffer-cache state is comparable. +// +// Reported metrics beyond ns/op and allocs: +// +// - queries/op round trips, counted client-side via a pgx tracer +// (GetNetworkMapData only — the legacy store's pool is internal) +// - xact/op committed transactions from pg_stat_database; the legacy +// pgx path runs autocommit statements, so this approximates its round +// trips, while GetNetworkMapData runs a single transaction +// - tup_returned/op, tup_fetched/op rows scanned/fetched server-side +// - blks_read/op, blks_hit/op buffer cache misses/hits +// +// The pg_stat_database numbers are database-global: run without concurrent +// load. The two stat snapshots per sub-benchmark add a small constant +// overhead to the server-side deltas. +package legacynmap_test + +import ( + "context" + "os" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache" + 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/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/networkmap" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" +) + +func BenchmarkGetAccount(b *testing.B) { + dsn := equivDSN() + if dsn == "" { + b.Skip("NETBIRD_STORE_ENGINE_POSTGRES_DSN not set") + } + ctx := context.Background() + + statsConn, err := pgx.Connect(ctx, dsn) + require.NoError(b, err, "connect stats connection") + b.Cleanup(func() { statsConn.Close(ctx) }) + + testStore, err := store.NewPostgresqlStore(ctx, dsn, nil, true) + require.NoError(b, err, "connect to postgres") + b.Cleanup(func() { testStore.Close(ctx) }) + + for _, accountID := range benchAccountIDs(b, ctx, statsConn) { + b.Run(accountID, func(b *testing.B) { + logAccountShape(b, ctx, statsConn, accountID) + benchDBLoad(b, ctx, statsConn, nil, func() error { + _, err := testStore.GetAccount(ctx, accountID) + return err + }) + }) + } +} + +func BenchmarkGetNetworkMapData(b *testing.B) { + dsn := equivDSN() + if dsn == "" { + b.Skip("NETBIRD_STORE_ENGINE_POSTGRES_DSN not set") + } + ctx := context.Background() + + statsConn, err := pgx.Connect(ctx, dsn) + require.NoError(b, err, "connect stats connection") + b.Cleanup(func() { statsConn.Close(ctx) }) + + tracer := &queryCountTracer{} + cfg, err := pgxpool.ParseConfig(dsn) + require.NoError(b, err, "parse dsn") + cfg.ConnConfig.Tracer = tracer + pool, err := pgxpool.NewWithConfig(ctx, cfg) + require.NoError(b, err, "connect nmdata store") + b.Cleanup(pool.Close) + nmStore := &networkmap_pgsql.PgStore{Pool: pool} + + for _, accountID := range benchAccountIDs(b, ctx, statsConn) { + b.Run(accountID, func(b *testing.B) { + logAccountShape(b, ctx, statsConn, accountID) + benchDBLoad(b, ctx, statsConn, tracer, func() error { + _, err := nmStore.GetNetworkMapData(ctx, accountID) + return err + }) + }) + } +} + +// BenchmarkAccountFullRound measures store load plus the full per-peer fan-out +// to *proto.SyncResponse for every peer of the account, the way the production +// account path runs it: index maps and per-peer twin building happen after +// GetAccount and are part of the measured op. BenchmarkNetworkMapDataFullRound +// is the equivalent for the nmdata path, whose index building happens inside +// GetNetworkMapData. Select both with -bench FullRound. +func BenchmarkAccountFullRound(b *testing.B) { + dsn := equivDSN() + if dsn == "" { + b.Skip("NETBIRD_STORE_ENGINE_POSTGRES_DSN not set") + } + ctx := context.Background() + + statsConn, err := pgx.Connect(ctx, dsn) + require.NoError(b, err, "connect stats connection") + b.Cleanup(func() { statsConn.Close(ctx) }) + + testStore, err := store.NewPostgresqlStore(ctx, dsn, nil, true) + require.NoError(b, err, "connect to postgres") + b.Cleanup(func() { testStore.Close(ctx) }) + + for _, accountID := range benchAccountIDs(b, ctx, statsConn) { + b.Run(accountID, func(b *testing.B) { + logAccountShape(b, ctx, statsConn, accountID) + benchDBLoad(b, ctx, statsConn, nil, func() error { + account, err := testStore.GetAccount(ctx, accountID) + if err != nil { + return err + } + buildAccountSyncResponses(ctx, account) + return nil + }) + }) + } +} + +func BenchmarkNetworkMapDataFullRound(b *testing.B) { + dsn := equivDSN() + if dsn == "" { + b.Skip("NETBIRD_STORE_ENGINE_POSTGRES_DSN not set") + } + ctx := context.Background() + + statsConn, err := pgx.Connect(ctx, dsn) + require.NoError(b, err, "connect stats connection") + b.Cleanup(func() { statsConn.Close(ctx) }) + + tracer := &queryCountTracer{} + cfg, err := pgxpool.ParseConfig(dsn) + require.NoError(b, err, "parse dsn") + cfg.ConnConfig.Tracer = tracer + pool, err := pgxpool.NewWithConfig(ctx, cfg) + require.NoError(b, err, "connect nmdata store") + b.Cleanup(pool.Close) + nmStore := &networkmap_pgsql.PgStore{Pool: pool} + + for _, accountID := range benchAccountIDs(b, ctx, statsConn) { + b.Run(accountID, func(b *testing.B) { + logAccountShape(b, ctx, statsConn, accountID) + benchDBLoad(b, ctx, statsConn, tracer, func() error { + nmData, err := nmStore.GetNetworkMapData(ctx, accountID) + if err != nil { + return err + } + buildDataSyncResponses(ctx, nmData) + return nil + }) + }) + } +} + +// buildAccountSyncResponses fans out to every peer like the controller's +// account path: index maps once, twin conversion and network-map computation +// per peer. +func buildAccountSyncResponses(ctx context.Context, account *types.Account) { + validated := make(map[string]struct{}, len(account.Peers)) + for peerID := range account.Peers { + validated[peerID] = struct{}{} + } + resourcePolicies := account.GetResourcePoliciesMap() + routers := account.GetResourceRoutersMap() + groupUsers := account.GetActiveGroupUsers() + settings := account.Settings + if settings == nil { + settings = &types.Settings{} + } + dnsCache := &cache.DNSConfigCache{} + + for peerID, peer := range account.Peers { + nm := account.GetPeerNetworkMapFromComponents( + ctx, peerID, nbdns.CustomZone{}, nil, validated, resourcePolicies, routers, nil, groupUsers, + ) + mgmtgrpc.ToSyncResponse( + ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, nm, equivDNSName, nil, + dnsCache, types.TwinAccountSettings(settings), settings.Extra, nil, 0, + ) + } +} + +// buildDataSyncResponses is the nmdata-path equivalent of +// buildAccountSyncResponses. +func buildDataSyncResponses(ctx context.Context, nmData *networkmap.NetworkMapData) { + validated := make(map[string]struct{}, len(nmData.Peers)) + for peerID := range nmData.Peers { + validated[peerID] = struct{}{} + } + nmData.ValidatedPeers = validated + dnsCache := &cache.DNSConfigCache{} + + for peerID, peer := range nmData.Peers { + components := nmData.GetPeerNetworkMapComponents(peerID, nmdata.CustomZone{}) + nm := &types.NetworkMap{Network: components.Network} + if !components.IsEmpty() { + nm = types.CalculateNetworkMapFromComponents(ctx, components) + } + mgmtgrpc.ToSyncResponse( + ctx, nil, nil, nil, peer, nil, nil, nm, equivDNSName, nil, + dnsCache, nmData.AccountSettings, nil, nil, 0, + ) + } +} + +// benchDBLoad runs op b.N times and reports server-side pg_stat_database +// deltas per op. A non-nil tracer additionally reports exact client round +// trips per op. +// +// Backends flush cumulative stats at most once per second and only while +// processing commands, so around each snapshot the load settles: sleep past +// the flush interval, then run one extra untimed op whose command end flushes +// everything pending. The trailing extra op lands inside the measured window, +// hence the b.N+1 denominator for the server-side metrics. +func benchDBLoad(b *testing.B, ctx context.Context, statsConn *pgx.Conn, tracer *queryCountTracer, op func() error) { + b.Helper() + + require.NoError(b, op(), "warmup") + settleDBStats(b, op) + + before, err := snapshotDBStats(ctx, statsConn) + require.NoError(b, err, "stats snapshot") + var queriesBefore int64 + if tracer != nil { + queriesBefore = tracer.queries.Load() + } + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := op(); err != nil { + b.Fatal(err) + } + } + b.StopTimer() + + settleDBStats(b, op) + after, err := snapshotDBStats(ctx, statsConn) + require.NoError(b, err, "stats snapshot") + + ops := float64(b.N + 1) + if tracer != nil { + b.ReportMetric(float64(tracer.queries.Load()-queriesBefore)/ops, "queries/op") + } + b.ReportMetric(float64(after.xactCommit-before.xactCommit)/ops, "xact/op") + b.ReportMetric(float64(after.tupReturned-before.tupReturned)/ops, "tup_returned/op") + b.ReportMetric(float64(after.tupFetched-before.tupFetched)/ops, "tup_fetched/op") + b.ReportMetric(float64(after.blksRead-before.blksRead)/ops, "blks_read/op") + b.ReportMetric(float64(after.blksHit-before.blksHit)/ops, "blks_hit/op") +} + +func settleDBStats(b *testing.B, op func() error) { + b.Helper() + time.Sleep(1100 * time.Millisecond) + require.NoError(b, op(), "stats flush op") + time.Sleep(100 * time.Millisecond) +} + +func benchAccountIDs(b *testing.B, ctx context.Context, conn *pgx.Conn) []string { + b.Helper() + + if ids := strings.TrimSpace(os.Getenv("NETMAP_ACCOUNTS")); ids != "" { + var out []string + for _, id := range strings.Split(ids, ",") { + if id = strings.TrimSpace(id); id != "" { + out = append(out, id) + } + } + return out + } + + rows, err := conn.Query(ctx, + "select account_id from peers group by account_id order by count(*) desc, account_id limit 10") + require.NoError(b, err, "list benchmark accounts") + ids, err := pgx.CollectRows(rows, pgx.RowTo[string]) + require.NoError(b, err, "collect benchmark accounts") + require.NotEmpty(b, ids, "no accounts found") + return ids +} + +func logAccountShape(b *testing.B, ctx context.Context, conn *pgx.Conn, accountID string) { + b.Helper() + + var peers, groups, users, policies, routes, resources, nsGroups int + err := conn.QueryRow(ctx, `select + (select count(*) from peers where account_id=$1), + (select count(*) from groups where account_id=$1), + (select count(*) from users where account_id=$1), + (select count(*) from policies where account_id=$1), + (select count(*) from routes where account_id=$1), + (select count(*) from network_resources where account_id=$1), + (select count(*) from name_server_groups where account_id=$1)`, accountID). + Scan(&peers, &groups, &users, &policies, &routes, &resources, &nsGroups) + require.NoError(b, err, "account shape") + b.Logf("account=%s peers=%d groups=%d users=%d policies=%d routes=%d resources=%d nsgroups=%d", + accountID, peers, groups, users, policies, routes, resources, nsGroups) +} + +type dbStats struct { + xactCommit int64 + tupReturned int64 + tupFetched int64 + blksRead int64 + blksHit int64 +} + +func snapshotDBStats(ctx context.Context, conn *pgx.Conn) (dbStats, error) { + var s dbStats + err := conn.QueryRow(ctx, `select xact_commit, tup_returned, tup_fetched, blks_read, blks_hit + from pg_stat_database where datname = current_database()`). + Scan(&s.xactCommit, &s.tupReturned, &s.tupFetched, &s.blksRead, &s.blksHit) + return s, err +} + +type queryCountTracer struct { + queries atomic.Int64 +} + +func (t *queryCountTracer) TraceQueryStart(ctx context.Context, _ *pgx.Conn, _ pgx.TraceQueryStartData) context.Context { + t.queries.Add(1) + return ctx +} + +func (t *queryCountTracer) TraceQueryEnd(context.Context, *pgx.Conn, pgx.TraceQueryEndData) {} diff --git a/shared/management/networkmap/decode.go b/shared/management/networkmap/decode.go index a7ae1665d..cb9465c5e 100644 --- a/shared/management/networkmap/decode.go +++ b/shared/management/networkmap/decode.go @@ -1,6 +1,7 @@ package networkmap import ( + "context" "encoding/base64" "fmt" "net" @@ -13,7 +14,7 @@ import ( nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/shared/management/domain" - nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/proto" "github.com/netbirdio/netbird/shared/management/types" ) @@ -25,7 +26,7 @@ import ( // ID scheme on the client side: // // Peers base64(wg_pub_key) // stable across snapshots -func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, error) { +func DecodeEnvelope(ctx context.Context, env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, error) { full := env.GetFull() if full == nil { return nil, fmt.Errorf("envelope has no Full payload") @@ -104,7 +105,12 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, var toret []nmdata.Resource for _, r := range gc.Resources { - toret = append(toret, resourceFromProto(r, peerIDByIndex)) + res := resourceFromProto(r, peerIDByIndex) + if res == (nmdata.Resource{}) { + log.WithContext(ctx).Warnf("skipping invalid resource in group compact: %s", r.ResourceId) + continue + } + toret = append(toret, res) } return toret @@ -264,7 +270,7 @@ func policiesForNetworkResource(resourceId string, allPolicies []*nmdata.Policy, networkResourceGroups := networkResourceGroups(resourceId, groups) for _, p := range allPolicies { - if p == nil || !p.Enabled { + if p == nil || !p.Enabled || len(p.Rules) == 0 { continue } diff --git a/shared/management/networkmap/decode_test.go b/shared/management/networkmap/decode_test.go index 634b5980e..3496a40d1 100644 --- a/shared/management/networkmap/decode_test.go +++ b/shared/management/networkmap/decode_test.go @@ -10,31 +10,31 @@ import ( func TestDecodePolicy(t *testing.T) { assert.Equal(t, + nmdata.Resource{Type: "peer", ID: "valid-id"}, resourceFromProto( &proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(1)}}, - []string{"invalid-id-0", "valid-id", "invalid-id-2"}), - nmdata.Resource{Type: "peer", ID: "valid-id"}) + []string{"invalid-id-0", "valid-id", "invalid-id-2"})) // check invalid peer index returns an empty resource assert.Equal(t, + nmdata.Resource{}, resourceFromProto( &proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(100)}}, - []string{"invalid-id-0", "valid-id", "invalid-id-2"}), - nmdata.Resource{}) + []string{"invalid-id-0", "valid-id", "invalid-id-2"})) assert.Equal(t, + nmdata.Resource{Type: "domain", ID: "domain"}, resourceFromProto( - &proto.ResourceCompact{Type: proto.ResourceCompactType_domain, ResourceId: &proto.ResourceCompact_Id{Id: "domain"}}, []string{}), - nmdata.Resource{Type: "domain", ID: "domain"}) + &proto.ResourceCompact{Type: proto.ResourceCompactType_domain, ResourceId: &proto.ResourceCompact_Id{Id: "domain"}}, []string{})) assert.Equal(t, + nmdata.Resource{Type: "host", ID: "host"}, resourceFromProto( - &proto.ResourceCompact{Type: proto.ResourceCompactType_host, ResourceId: &proto.ResourceCompact_Id{Id: "host"}}, []string{}), - nmdata.Resource{Type: "host", ID: "host"}) + &proto.ResourceCompact{Type: proto.ResourceCompactType_host, ResourceId: &proto.ResourceCompact_Id{Id: "host"}}, []string{})) assert.Equal(t, + nmdata.Resource{Type: "subnet", ID: "subnet"}, resourceFromProto( - &proto.ResourceCompact{Type: proto.ResourceCompactType_subnet, ResourceId: &proto.ResourceCompact_Id{Id: "subnet"}}, []string{}), - nmdata.Resource{Type: "subnet", ID: "subnet"}) + &proto.ResourceCompact{Type: proto.ResourceCompactType_subnet, ResourceId: &proto.ResourceCompact_Id{Id: "subnet"}}, []string{})) // an unknown resource type return an empty resource assert.Equal(t, + nmdata.Resource{}, resourceFromProto( - &proto.ResourceCompact{Type: proto.ResourceCompactType_unknown_type, ResourceId: &proto.ResourceCompact_Id{Id: "boom"}}, []string{}), - nmdata.Resource{}) + &proto.ResourceCompact{Type: proto.ResourceCompactType_unknown_type, ResourceId: &proto.ResourceCompact_Id{Id: "boom"}}, []string{})) } diff --git a/shared/management/networkmap/envelope.go b/shared/management/networkmap/envelope.go index 3f045a9eb..10bba5810 100644 --- a/shared/management/networkmap/envelope.go +++ b/shared/management/networkmap/envelope.go @@ -36,7 +36,7 @@ type EnvelopeResult struct { // dnsName is the account's DNS domain ("netbird.cloud" etc.); used when // rebuilding the per-peer FQDNs that proto.RemotePeerConfig carries. func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, localPeerKey, dnsName string) (*EnvelopeResult, error) { - components, err := DecodeEnvelope(env) + components, err := DecodeEnvelope(ctx, env) if err != nil { return nil, fmt.Errorf("decode envelope: %w", err) } diff --git a/shared/management/networkmap/networkmapcompute.go b/shared/management/networkmap/networkmapcompute.go index 0d4cf4b6f..f07ff0ba1 100644 --- a/shared/management/networkmap/networkmapcompute.go +++ b/shared/management/networkmap/networkmapcompute.go @@ -234,14 +234,20 @@ func (nmd *NetworkMapData) getPeersGroupsPoliciesRoutes( } for _, groupID := range r.PeerGroups { - relevantGroupIDs[groupID] = nmd.Groups[groupID] + if g := nmd.Groups[groupID]; g != nil { + relevantGroupIDs[groupID] = g + } } for _, groupID := range r.Groups { - relevantGroupIDs[groupID] = nmd.Groups[groupID] + if g := nmd.Groups[groupID]; g != nil { + relevantGroupIDs[groupID] = g + } } if r.Enabled { for _, groupID := range r.AccessControlGroups { - relevantGroupIDs[groupID] = nmd.Groups[groupID] + if g := nmd.Groups[groupID]; g != nil { + relevantGroupIDs[groupID] = g + } routeAccessControlGroups[groupID] = struct{}{} } } @@ -289,10 +295,14 @@ func (nmd *NetworkMapData) getPeersGroupsPoliciesRoutes( if _, needed := routeAccessControlGroups[destGroupID]; needed { policyRelevant = true for _, srcGroupID := range rule.Sources { - relevantGroupIDs[srcGroupID] = nmd.Groups[srcGroupID] + if g := nmd.Groups[srcGroupID]; g != nil { + relevantGroupIDs[srcGroupID] = g + } } for _, dstGroupID := range rule.Destinations { - relevantGroupIDs[dstGroupID] = nmd.Groups[dstGroupID] + if g := nmd.Groups[dstGroupID]; g != nil { + relevantGroupIDs[dstGroupID] = g + } } break } @@ -326,7 +336,9 @@ func (nmd *NetworkMapData) getPeersGroupsPoliciesRoutes( relevantPeerIDs[pid] = nmd.Peers[pid] } for _, dstGroupID := range rule.Destinations { - relevantGroupIDs[dstGroupID] = nmd.Groups[dstGroupID] + if g := nmd.Groups[dstGroupID]; g != nil { + relevantGroupIDs[dstGroupID] = g + } } } @@ -336,7 +348,9 @@ func (nmd *NetworkMapData) getPeersGroupsPoliciesRoutes( relevantPeerIDs[pid] = nmd.Peers[pid] } for _, srcGroupID := range rule.Sources { - relevantGroupIDs[srcGroupID] = nmd.Groups[srcGroupID] + if g := nmd.Groups[srcGroupID]; g != nil { + relevantGroupIDs[srcGroupID] = g + } } if rule.Protocol == string(types.PolicyRuleProtocolNetbirdSSH) { @@ -470,6 +484,9 @@ func (nmd *NetworkMapData) getPostureValidPeersSaveFailed(inputPeers []string, p dest = append(dest, peerID) continue } + if pname == "" { + continue + } if _, ok := (*postureFailedPeers)[pname]; !ok { (*postureFailedPeers)[pname] = make(map[string]struct{}) } diff --git a/shared/management/networkmap/nmdata/group.go b/shared/management/networkmap/nmdata/group.go index 9e795a06a..70fa6d3dc 100644 --- a/shared/management/networkmap/nmdata/group.go +++ b/shared/management/networkmap/nmdata/group.go @@ -19,9 +19,10 @@ func (g *Group) IsGroupAll() bool { func (g *Group) Copy() *Group { return &Group{ - ID: g.ID, - Name: g.Name, - PublicID: g.PublicID, - Peers: slices.Clone(g.Peers), + ID: g.ID, + Name: g.Name, + PublicID: g.PublicID, + Peers: slices.Clone(g.Peers), + Resources: slices.Clone(g.Resources), } } diff --git a/shared/management/networkmap/nmdata/group_test.go b/shared/management/networkmap/nmdata/group_test.go new file mode 100644 index 000000000..20aaa240f --- /dev/null +++ b/shared/management/networkmap/nmdata/group_test.go @@ -0,0 +1,84 @@ +package nmdata + +import ( + "reflect" + "testing" +) + +// TestGroupCopy_AllFieldsCopied fills every Group field with a unique non-zero +// value derived from its field path, so a field added to Group but forgotten +// in Copy fails here by name without the test needing an update. The unique +// per-path values also catch fields swapped inside Copy. +func TestGroupCopy_AllFieldsCopied(t *testing.T) { + src := &Group{} + seed := 0 + fillValue(t, reflect.ValueOf(src).Elem(), "Group", &seed) + + copied := src.Copy() + + srcV := reflect.ValueOf(src).Elem() + copiedV := reflect.ValueOf(copied).Elem() + for i := 0; i < srcV.NumField(); i++ { + name := srcV.Type().Field(i).Name + if !reflect.DeepEqual(srcV.Field(i).Interface(), copiedV.Field(i).Interface()) { + t.Errorf("field %s not copied: src=%#v copy=%#v", + name, srcV.Field(i).Interface(), copiedV.Field(i).Interface()) + } + } + + for i := 0; i < srcV.NumField(); i++ { + f := srcV.Field(i) + if f.Kind() != reflect.Slice || f.Len() == 0 { + continue + } + name := srcV.Type().Field(i).Name + fillValue(t, f.Index(0), name+"-mutated", &seed) + if reflect.DeepEqual(f.Interface(), copiedV.Field(i).Interface()) { + t.Errorf("field %s shares memory with the copy", name) + } + } +} + +// fillValue sets v to a deterministic non-zero value derived from its field +// path. Kinds it does not handle fail the test loudly, so the filler is +// extended together with the struct instead of silently under-testing new +// fields. +func fillValue(t *testing.T, v reflect.Value, path string, seed *int) { + t.Helper() + + switch v.Kind() { + case reflect.String: + v.SetString(path) + case reflect.Bool: + v.SetBool(true) + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + *seed++ + v.SetInt(int64(*seed)) + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + *seed++ + v.SetUint(uint64(*seed)) + case reflect.Float32, reflect.Float64: + *seed++ + v.SetFloat(float64(*seed)) + case reflect.Slice: + s := reflect.MakeSlice(v.Type(), 2, 2) + fillValue(t, s.Index(0), path+"[0]", seed) + fillValue(t, s.Index(1), path+"[1]", seed) + v.Set(s) + case reflect.Struct: + settable := 0 + for i := 0; i < v.NumField(); i++ { + f := v.Field(i) + if !f.CanSet() { + continue + } + settable++ + fillValue(t, f, path+"."+v.Type().Field(i).Name, seed) + } + if settable == 0 { + t.Fatalf("struct %s at %s has no settable fields — extend fillValue to construct it", v.Type(), path) + } + default: + t.Fatalf("unsupported kind %s at %s — extend fillValue", v.Kind(), path) + } +}