diff --git a/management/internals/network_map_db/db_store.go b/management/internals/network_map_db/db_store.go index 55a6798c4..e32a2008c 100644 --- a/management/internals/network_map_db/db_store.go +++ b/management/internals/network_map_db/db_store.go @@ -19,7 +19,7 @@ const ( ) type NetworkMapDBStore interface { - GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, error) + GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string][]*nmdata.Group, error) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, error) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, error) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) diff --git a/management/internals/network_map_db/pgsql/group.go b/management/internals/network_map_db/pgsql/group.go index 5fa42b22f..f33de66b9 100644 --- a/management/internals/network_map_db/pgsql/group.go +++ b/management/internals/network_map_db/pgsql/group.go @@ -23,33 +23,38 @@ const ( ` ) -func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, error) { +func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string][]*nmdata.Group, error) { c, err := pg.Pool.Acquire(ctx) if err != nil { - return nil, err + return nil, nil, err } return GetGroupsViaPgxConnection(ctx, c.Conn(), accountId) } -func GetGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Group, error) { +func GetGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Group, map[string][]*nmdata.Group, error) { rows, err := con.Query(ctx, GetGroupsQuery, accountId) if err != nil { - return nil, err + return nil, nil, err } groups, err := pgx.CollectRows(rows, pgx.RowToStructByName[group]) toret := make([]nmdata.Group, 0, len(groups)) + toretidx := make(map[string][]*nmdata.Group) + for _, g := range groups { dg := nmdata.Group{} err := networkmapdb.FromSqlTypesToSharedTypes( reflect.ValueOf(&g), reflect.ValueOf(&dg)) if err != nil { - return nil, err + return nil, nil, err } toret = append(toret, dg) + for _, resource := range dg.Resources { + toretidx[resource.ID] = append(toretidx[resource.ID], &dg) + } } - return toret, err + return toret, toretidx, err } type group struct { diff --git a/management/internals/network_map_db/pgsql/network_map_data.go b/management/internals/network_map_db/pgsql/network_map_data.go index 34a23d339..896937323 100644 --- a/management/internals/network_map_db/pgsql/network_map_data.go +++ b/management/internals/network_map_db/pgsql/network_map_data.go @@ -22,7 +22,7 @@ func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*ne // if err != nil { // return rollbackAndReturnError(ctx, tx, err) // } - groups, err := GetGroupsViaPgxConnection(ctx, tx.Conn(), accountId) + groups, netResourceToGroups, err := GetGroupsViaPgxConnection(ctx, tx.Conn(), accountId) if err != nil { return rollbackAndReturnError(ctx, tx, err) }