diff --git a/management/internals/network_map_db/pgsql/network.go b/management/internals/network_map_db/pgsql/network.go index 0b252b4e1..5d7d33bcc 100644 --- a/management/internals/network_map_db/pgsql/network.go +++ b/management/internals/network_map_db/pgsql/network.go @@ -2,8 +2,6 @@ package networkmap_pgsql import ( "context" - "database/sql" - "encoding/json" "reflect" "github.com/jackc/pgx/v5" @@ -25,7 +23,7 @@ func (pgc *PgStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdat return nmdata.Network{}, err } - n, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[accountnetwork]) + n, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[networkmapdb.AccountNetwork]) if err != nil { return nmdata.Network{}, err } @@ -39,11 +37,3 @@ func (pgc *PgStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdat return toret, nil } - -type accountnetwork struct { - Identifier sql.NullString - Net json.RawMessage - NetV6 json.RawMessage - Dns sql.NullString - Serial sql.NullInt64 -} diff --git a/management/internals/network_map_db/shared_types.go b/management/internals/network_map_db/shared_types.go index 2168cc143..2bb9f12cf 100644 --- a/management/internals/network_map_db/shared_types.go +++ b/management/internals/network_map_db/shared_types.go @@ -79,6 +79,14 @@ type Networkresource struct { Enabled sql.NullBool } +type AccountNetwork struct { + Identifier sql.NullString + Net []byte `nmap:"json"` + NetV6 []byte `nmap:"json"` + Dns sql.NullString + Serial sql.NullInt64 +} + func RecordTypeAndRdata(t, rdata string) (int, string, error) { switch t { case "A": diff --git a/management/internals/network_map_db/sqlite/account_setting.go b/management/internals/network_map_db/sqlite/account_setting.go index 134c26aff..db2c595ad 100644 --- a/management/internals/network_map_db/sqlite/account_setting.go +++ b/management/internals/network_map_db/sqlite/account_setting.go @@ -22,7 +22,7 @@ const ( settings_auto_update_always as auto_update_always, settings_metrics_push_enabled as metrics_push_enabled from accounts - where id=$1 + where id=? ` ) @@ -32,7 +32,7 @@ func (sc *SqliteStoreConn) GetAccountSettings(ctx context.Context, accountId str return nmdata.AccountSettingsInfo{}, err } - a, err := networkmapdb.CollectOneRowForSqlite[networkmapdb.Account](rows) + a, err := CollectOneRowForSqlite[networkmapdb.Account](rows) settingsInfo := nmdata.AccountSettingsInfo{} err = networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&a), reflect.ValueOf(&settingsInfo)) diff --git a/management/internals/network_map_db/sqlite/dns.go b/management/internals/network_map_db/sqlite/dns.go index 4dffd9f92..026467eae 100644 --- a/management/internals/network_map_db/sqlite/dns.go +++ b/management/internals/network_map_db/sqlite/dns.go @@ -13,7 +13,7 @@ const ( r.name as record_name, r.type as record_type, 'IN' record_class, r.ttl as record_ttl, r.content as record_rdata from zones left join records as r on r.zone_id = zones.id - where zones.account_id=$1 + where zones.account_id=? ` ) @@ -23,7 +23,7 @@ func (sc *SqliteStoreConn) GetAppliedZoneCandidates(ctx context.Context, account return nil, err } - zones, err := networkmapdb.CollectRowsForSqlite[networkmapdb.Zone](rows) + zones, err := CollectRowsForSqlite[networkmapdb.Zone](rows) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/sqlite/dns_setting.go b/management/internals/network_map_db/sqlite/dns_setting.go index 6cfc27d37..7c6e9e7eb 100644 --- a/management/internals/network_map_db/sqlite/dns_setting.go +++ b/management/internals/network_map_db/sqlite/dns_setting.go @@ -11,7 +11,7 @@ const ( GetDnsSettingsQuery = ` select dns_settings_disabled_management_groups from accounts - where id=$1 + where id=? ` ) diff --git a/management/internals/network_map_db/sqlite/domain.go b/management/internals/network_map_db/sqlite/domain.go index 241219a4a..572977c3b 100644 --- a/management/internals/network_map_db/sqlite/domain.go +++ b/management/internals/network_map_db/sqlite/domain.go @@ -10,7 +10,7 @@ const ( GetDomainsQuery = ` select domain, target_cluster from domains - where account_id=$1 and domain<>'' and target_cluster<>'' + where account_id=? and domain<>'' and target_cluster<>'' ` ) @@ -20,5 +20,5 @@ func (sc *SqliteStoreConn) GetDomains(ctx context.Context, accountId string) ([] return nil, err } - return networkmapdb.CollectRowsForSqlite[networkmapdb.Domain](rows) + return CollectRowsForSqlite[networkmapdb.Domain](rows) } diff --git a/management/internals/network_map_db/sqlite/group.go b/management/internals/network_map_db/sqlite/group.go index edd54f4db..c324d0d29 100644 --- a/management/internals/network_map_db/sqlite/group.go +++ b/management/internals/network_map_db/sqlite/group.go @@ -27,7 +27,7 @@ func (sc *SqliteStoreConn) GetGroups(ctx context.Context, accountId string) ([]n return nil, nil, err } - groups, err := networkmapdb.CollectRowsForSqlite[group](rows) + groups, err := CollectRowsForSqlite[group](rows) toret := make([]nmdata.Group, 0, len(groups)) resourceToGroupIdx := make(map[string]map[string]any) diff --git a/management/internals/network_map_db/sqlite/nameserver.go b/management/internals/network_map_db/sqlite/nameserver.go index a4a8ec890..618e1a1f3 100644 --- a/management/internals/network_map_db/sqlite/nameserver.go +++ b/management/internals/network_map_db/sqlite/nameserver.go @@ -11,7 +11,7 @@ const ( GetNameserversQuery = ` select id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled from name_server_groups - where account_id=$1 + where account_id=? ` ) @@ -21,7 +21,7 @@ func (sc *SqliteStoreConn) GetNameServerGroups(ctx context.Context, accountId st return nil, err } - nsgroups, err := networkmapdb.CollectRowsForSqlite[networkmapdb.NameserverGroup](rows) + nsgroups, err := CollectRowsForSqlite[networkmapdb.NameserverGroup](rows) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/sqlite/network.go b/management/internals/network_map_db/sqlite/network.go new file mode 100644 index 000000000..3fa85ecdf --- /dev/null +++ b/management/internals/network_map_db/sqlite/network.go @@ -0,0 +1,38 @@ +package networkmap_sqlite + +import ( + "context" + "reflect" + + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" +) + +const ( + GetNetworkQuery = ` + select network_identifier as identifier, network_net as net, network_net_v6 as net_v6, network_dns as dns, network_serial as serial + from accounts + where id=? + ` +) + +func (sc *SqliteStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) { + rows, err := sc.Conn.QueryContext(ctx, GetNetworkQuery, accountId) + if err != nil { + return nmdata.Network{}, err + } + + n, err := CollectOneRowForSqlite[networkmapdb.AccountNetwork](rows) + if err != nil { + return nmdata.Network{}, err + } + + toret := nmdata.Network{} + err = networkmapdb.FromSqlTypesToSharedTypes( + reflect.ValueOf(&n), reflect.ValueOf(&toret)) + if err != nil { + return nmdata.Network{}, err + } + + return toret, nil +} diff --git a/management/internals/network_map_db/sqlite/network_resource.go b/management/internals/network_map_db/sqlite/network_resource.go index 2b308ea0f..1d98a12e9 100644 --- a/management/internals/network_map_db/sqlite/network_resource.go +++ b/management/internals/network_map_db/sqlite/network_resource.go @@ -11,7 +11,7 @@ const ( GetNetworkResourcesQuery = ` select id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled from network_resources - where account_id=$1 + where account_id=? ` ) @@ -21,7 +21,7 @@ func (sc *SqliteStoreConn) GetNetworkResources(ctx context.Context, accountId st return nil, err } - netresorces, err := networkmapdb.CollectRowsForSqlite[networkmapdb.Networkresource](rows) + netresorces, err := CollectRowsForSqlite[networkmapdb.Networkresource](rows) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/sqlite/network_router.go b/management/internals/network_map_db/sqlite/network_router.go index d5887b33f..8c4c31cd6 100644 --- a/management/internals/network_map_db/sqlite/network_router.go +++ b/management/internals/network_map_db/sqlite/network_router.go @@ -25,7 +25,7 @@ func (sc *SqliteStoreConn) GetNetworkRouters(ctx context.Context, accountId stri return nil, err } - routers, err := networkmapdb.CollectRowsForSqlite[networkrouter](rows) + routers, err := CollectRowsForSqlite[networkrouter](rows) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/sqlite/sqlite_store.go b/management/internals/network_map_db/sqlite/sqlite_store.go index 980d0c493..f0618af3c 100644 --- a/management/internals/network_map_db/sqlite/sqlite_store.go +++ b/management/internals/network_map_db/sqlite/sqlite_store.go @@ -82,6 +82,44 @@ func (s *SqliteStore) UsingConn() *SqliteStoreConn { return &SqliteStoreConn{Conn: s.Db} } +func CollectOneRowForSqlite[T any](rows *sql.Rows) (T, error) { + defer rows.Close() + var r T + + if !rows.Next() { + if err := rows.Err(); err != nil { + return r, err + } + return r, ErrNoRows + } + err := rows.Scan(networkmapdb.StructFields(&r)...) + if err != nil { + return r, err + } + + return r, nil +} + +func CollectRowsForSqlite[T any](rows *sql.Rows) ([]T, error) { + defer rows.Close() + toret := make([]T, 0) + + for rows.Next() { + var r T + err := rows.Scan(networkmapdb.StructFields(&r)...) + if err != nil { + return nil, err + } + toret = append(toret, r) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return toret, nil +} + func (s *SqliteStoreConn) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) { return nil, nil, nil } @@ -91,9 +129,6 @@ func (s *SqliteStoreConn) GetPolicies(ctx context.Context, accountId string) ([] func (s *SqliteStoreConn) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) { return nil, nil } -func (s *SqliteStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) { - return nmdata.Network{}, nil -} func (s *SqliteStoreConn) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) { return nil, nil, nil } diff --git a/management/internals/network_map_db/struct_helpers.go b/management/internals/network_map_db/struct_helpers.go index 1bb0379d6..1719662fd 100644 --- a/management/internals/network_map_db/struct_helpers.go +++ b/management/internals/network_map_db/struct_helpers.go @@ -142,44 +142,6 @@ func StructFields(s any) []any { return toret } -func CollectOneRowForSqlite[T any](rows *sql.Rows) (T, error) { - defer rows.Close() - var r T - - if !rows.Next() { - if err := rows.Err(); err != nil { - return r, err - } - return r, ErrNoRows - } - err := rows.Scan(StructFields(&r)...) - if err != nil { - return r, err - } - - return r, nil -} - -func CollectRowsForSqlite[T any](rows *sql.Rows) ([]T, error) { - defer rows.Close() - toret := make([]T, 0) - - for rows.Next() { - var r T - err := rows.Scan(StructFields(&r)...) - if err != nil { - return nil, err - } - toret = append(toret, r) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return toret, nil -} - func ConvertAllToSharedTypes[T any, T1 any](allsrc []T) ([]T1, error) { toret := make([]T1, 0, len(allsrc)) for _, src := range allsrc {