diff --git a/db/channels.go b/db/channels.go index 9b6b7bb..094940d 100644 --- a/db/channels.go +++ b/db/channels.go @@ -7,7 +7,6 @@ import ( "context" "encoding/hex" "errors" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -132,11 +131,10 @@ func (s *Store) ListChannelMessages(ctx context.Context, channelID *int32, since ts := pgtype.Timestamptz{Time: since, Valid: !since.IsZero()} var messages []api.ChannelMessage var hasMore bool - iataFilter := strings.Join(iatas, ",") if channelID == nil { rows, err := s.q.ListAllChannelMessages(ctx, sqlc.ListAllChannelMessagesParams{ Column1: ts, - Column2: iataFilter, + Column2: iatas, Column3: scope, Column4: cursor, Limit: limit + 1, @@ -156,7 +154,7 @@ func (s *Store) ListChannelMessages(ctx context.Context, channelID *int32, since rows, err := s.q.ListChannelMessages(ctx, sqlc.ListChannelMessagesParams{ ChannelID: *channelID, Column2: ts, - Column3: iataFilter, + Column3: iatas, Column4: scope, Column5: cursor, Limit: limit + 1, @@ -187,11 +185,10 @@ func (s *Store) ListChannelMessages(ctx context.Context, channelID *int32, since } func (s *Store) ListChannelMessagesByHash(ctx context.Context, hash []byte, since time.Time, limit int32, iatas []string, scope string, cursor int64) (api.Page[api.ChannelMessage], error) { - iataFilter := strings.Join(iatas, ",") rows, err := s.q.ListChannelMessagesByHash(ctx, sqlc.ListChannelMessagesByHashParams{ ChannelHash: hash, Column2: pgtype.Timestamptz{Time: since, Valid: !since.IsZero()}, - Column3: iataFilter, + Column3: iatas, Column4: scope, Column5: cursor, Limit: limit + 1, @@ -220,10 +217,9 @@ func (s *Store) ListChannelMessagesByHash(ctx context.Context, hash []byte, sinc } func (s *Store) ListMessagesAfterID(ctx context.Context, afterID int64, iatas []string, scope string, limit int32) ([]api.ChannelMessage, error) { - iataFilter := strings.Join(iatas, ",") rows, err := s.q.ListMessagesAfterID(ctx, sqlc.ListMessagesAfterIDParams{ ID: afterID, - Column2: iataFilter, + Column2: iatas, Column3: scope, Limit: limit, }) diff --git a/db/channels_test.go b/db/channels_test.go index a09453b..66ee741 100644 --- a/db/channels_test.go +++ b/db/channels_test.go @@ -199,7 +199,7 @@ func TestListChannelMessages_AllChannels(t *testing.T) { mock.EXPECT(). ListAllChannelMessages(gomock.Any(), sqlc.ListAllChannelMessagesParams{ Column1: pgtype.Timestamptz{}, - Column2: "YVR", + Column2: []string{"YVR"}, Column3: "", Column4: int64(0), Limit: 3, @@ -263,7 +263,7 @@ func TestListChannelMessages_ByChannelID(t *testing.T) { ListChannelMessages(gomock.Any(), sqlc.ListChannelMessagesParams{ ChannelID: channelID, Column2: pgtype.Timestamptz{}, - Column3: "YVR", + Column3: []string{"YVR"}, Column4: "", Column5: int64(0), Limit: 3, diff --git a/db/nodes.go b/db/nodes.go index 2dcd25a..e416c78 100644 --- a/db/nodes.go +++ b/db/nodes.go @@ -10,7 +10,6 @@ import ( "errors" "fmt" "log" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -85,10 +84,9 @@ func (s *Store) ListNodes(ctx context.Context, nodeType int16, iatas []string, s if cursor > 0 { cursorTS = pgtype.Timestamptz{Time: time.UnixMilli(cursor), Valid: true} } - iataFilter := strings.Join(iatas, ",") rows, err := s.q.ListNodes(ctx, sqlc.ListNodesParams{ Column1: nodeType, - Column2: iataFilter, + Column2: iatas, Column3: tristate(supportsMultibytePaths), Column4: tristate(supportsMultibyteTraces), Column5: pubkey, diff --git a/db/nodes_test.go b/db/nodes_test.go index 997cadd..03e0221 100644 --- a/db/nodes_test.go +++ b/db/nodes_test.go @@ -379,7 +379,7 @@ func TestListNodes_IncludeNeighbors_PassesFlagAndMapsIDs(t *testing.T) { mock.EXPECT(). ListNodes(gomock.Any(), gomock.Eq(sqlc.ListNodesParams{ - Column1: int16(0), Column2: "", Column3: "any", Column4: "any", + Column1: int16(0), Column2: nil, Column3: "any", Column4: "any", Column5: nil, Column6: "", Column7: pgtype.Timestamptz{}, Limit: 11, Column9: "", Column10: true, })). @@ -409,7 +409,7 @@ func TestListNodes_ExcludeNeighbors_LeavesIDsNil(t *testing.T) { mock.EXPECT(). ListNodes(gomock.Any(), gomock.Eq(sqlc.ListNodesParams{ - Column1: int16(0), Column2: "", Column3: "any", Column4: "any", + Column1: int16(0), Column2: nil, Column3: "any", Column4: "any", Column5: nil, Column6: "", Column7: pgtype.Timestamptz{}, Limit: 11, Column9: "", Column10: false, })). diff --git a/db/observers.go b/db/observers.go index 0eb0ce9..744fa68 100644 --- a/db/observers.go +++ b/db/observers.go @@ -8,7 +8,6 @@ import ( "encoding/hex" "fmt" "log" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -35,9 +34,8 @@ func (s *Store) ListObservers(ctx context.Context, iatas []string, observerType, if cursor > 0 { cursorTS = pgtype.Timestamptz{Time: time.UnixMilli(cursor), Valid: true} } - iataFilter := strings.Join(iatas, ",") params := sqlc.ListObserversParams{ - Column1: iataFilter, + Column1: iatas, Column2: observerType, Column3: broker, Column4: status, diff --git a/db/packets.go b/db/packets.go index 3cedca7..09ec142 100644 --- a/db/packets.go +++ b/db/packets.go @@ -11,7 +11,6 @@ import ( "errors" "fmt" "log" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -72,11 +71,10 @@ func (s *Store) ListPackets(ctx context.Context, payloadType, routeType int16, i if !until.IsZero() { untilTS = pgtype.Timestamptz{Time: until, Valid: true} } - iataFilter := strings.Join(iatas, ",") rows, err := s.q.ListPackets(ctx, sqlc.ListPacketsParams{ Column1: payloadType, Column2: routeType, - Column3: iataFilter, + Column3: iatas, Column4: sinceTS, Column5: untilTS, Column6: cursorTS, @@ -125,12 +123,11 @@ func (s *Store) ListPackets(ctx context.Context, payloadType, routeType int16, i } func (s *Store) ListPacketsAfterID(ctx context.Context, afterObservationID int64, payloadType, routeType int16, iatas []string, scope string, limit int32) ([]api.PacketSummary, error) { - iataFilter := strings.Join(iatas, ",") rows, err := s.q.ListPacketsAfterID(ctx, sqlc.ListPacketsAfterIDParams{ ID: afterObservationID, Column2: payloadType, Column3: routeType, - Column4: iataFilter, + Column4: iatas, Column5: scope, Limit: limit, }) diff --git a/db/packets_test.go b/db/packets_test.go index 59ccb0b..5c20f33 100644 --- a/db/packets_test.go +++ b/db/packets_test.go @@ -330,6 +330,28 @@ func TestGetPacket_FirstToLastMs(t *testing.T) { } } +func TestListPacketsAfterID_PassesIATAsAsArray(t *testing.T) { + ctrl := gomock.NewController(t) + mock := mockdb.NewMockQuerier(ctrl) + + mock.EXPECT(). + ListPacketsAfterID(gomock.Any(), sqlc.ListPacketsAfterIDParams{ + ID: 0, + Column2: int16(-1), + Column3: int16(-1), + Column4: []string{"ALF", "YYZ"}, + Column5: "", + Limit: 50, + }). + Return([]sqlc.ListPacketsAfterIDRow{}, nil) + + store := &Store{q: mock} + _, err := store.ListPacketsAfterID(context.Background(), 0, -1, -1, []string{"ALF", "YYZ"}, "", 50) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } +} + func TestListNodeObservations_Pagination(t *testing.T) { ctrl := gomock.NewController(t) mock := mockdb.NewMockQuerier(ctrl) diff --git a/db/queries/queries.sql b/db/queries/queries.sql index 3a93fa4..62b287c 100644 --- a/db/queries/queries.sql +++ b/db/queries/queries.sql @@ -56,7 +56,7 @@ LEFT JOIN observer_scopes os ON os.scope_id = ts.id LEFT JOIN observers o ON o.id = os.observer_id LEFT JOIN packet_observations po ON po.observer_id = o.id LEFT JOIN nodes n ON n.default_scope_id = ts.id -WHERE ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])) GROUP BY ts.name ORDER BY ts.name; @@ -174,11 +174,11 @@ LEFT JOIN observer_brokers ob ON ob.observer_id = o.id LEFT JOIN observer_scopes os ON os.observer_id = o.id LEFT JOIN transport_scopes ts ON ts.id = os.scope_id WHERE - ($1::text = '' OR ( + (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR ( SELECT po.iata FROM packet_observations po WHERE po.observer_id = o.id ORDER BY po.heard_at DESC LIMIT 1 - ) = ANY(string_to_array($1::text, ','))) + ) = ANY($1::bpchar[])) AND ($2 = '' OR o.observer_type = $2) AND ($3 = '' OR ob.broker_name = $3) AND ($4 = '' OR CASE @@ -391,10 +391,10 @@ LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE ($1::smallint = -1 OR p.payload_type = $1::smallint) AND ($2::smallint = -1 OR p.route_type = $2::smallint) - AND ($3::text = '' OR EXISTS ( + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR EXISTS ( SELECT 1 FROM packet_observations po3 WHERE po3.packet_hash = p.packet_hash - AND po3.iata = ANY(string_to_array($3::text, ',')) + AND po3.iata = ANY($3::bpchar[]) )) AND ($4::timestamptz IS NULL OR p.first_heard_at >= $4) AND ($5::timestamptz IS NULL OR p.first_heard_at <= $5) @@ -424,7 +424,7 @@ LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE po.id > $1 AND ($2::smallint = -1 OR p.payload_type = $2::smallint) AND ($3::smallint = -1 OR p.route_type = $3::smallint) - AND ($4::text = '' OR po.iata = ANY(string_to_array($4::text, ','))) + AND (COALESCE(cardinality($4::bpchar[]), 0) = 0 OR po.iata = ANY($4::bpchar[])) AND ($5::text = '' OR ts.name = $5::text) ORDER BY po.id ASC LIMIT $6; @@ -538,7 +538,7 @@ LEFT JOIN node_iatas ni ON ni.node_id = n.id LEFT JOIN transport_scopes ts ON ts.id = n.default_scope_id WHERE ($1 = 0 OR n.node_type = $1) - AND ($2::text = '' OR n.id IN (SELECT node_id FROM node_iatas WHERE iata = ANY(string_to_array($2::text, ',')))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR n.id IN (SELECT node_id FROM node_iatas WHERE iata = ANY($2::bpchar[]))) AND ( $3::text = 'any' OR ($3::text = 'true' AND n.supports_multibyte_paths = TRUE) @@ -651,7 +651,7 @@ JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE cm.channel_id = $1 AND ($2::timestamptz IS NULL OR cm.sent_at >= $2) - AND ($3::text = '' OR po.iata = ANY(string_to_array($3::text, ','))) + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR po.iata = ANY($3::bpchar[])) AND ($4::text = '' OR ts.name = $4::text) AND ($5::bigint = 0 OR cm.id < $5::bigint) ORDER BY cm.id DESC @@ -669,7 +669,7 @@ JOIN packet_observations po ON po.packet_hash = cm.packet_hash JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE ($1::timestamptz IS NULL OR cm.sent_at >= $1) - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) AND ($3::text = '' OR ts.name = $3::text) AND ($4 = 0 OR cm.id < $4) ORDER BY cm.id DESC @@ -689,7 +689,7 @@ JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE c.channel_hash = $1 AND ($2::timestamptz IS NULL OR cm.sent_at >= $2) - AND ($3::text = '' OR po.iata = ANY(string_to_array($3::text, ','))) + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR po.iata = ANY($3::bpchar[])) AND ($4::text = '' OR ts.name = $4::text) AND ($5::bigint = 0 OR cm.id < $5::bigint) ORDER BY cm.id DESC @@ -706,7 +706,7 @@ JOIN packet_observations po ON po.packet_hash = cm.packet_hash JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE cm.id > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) AND ($3::text = '' OR ts.name = $3::text) ORDER BY cm.id ASC LIMIT $4; @@ -723,18 +723,18 @@ SELECT COUNT(DISTINCT po.iata) AS active_iatas FROM packet_observations po WHERE po.heard_at > NOW() - INTERVAL '24 hours' - AND ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))); + AND (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])); -- name: GetHourlyStats :many SELECT iata, hour, observation_count, unique_packets, active_observers FROM mv_hourly_iata_stats -WHERE ($1::text = '' OR iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR iata = ANY($1::bpchar[])) AND hour >= NOW() - $2::interval ORDER BY iata, hour; -- name: GetTopNodes :many SELECT * FROM mv_top_nodes_by_iata -WHERE ($1::text = '' OR iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR iata = ANY($1::bpchar[])) ORDER BY observation_count DESC LIMIT $2; @@ -746,7 +746,7 @@ SELECT FROM packet_observations po JOIN packets p ON p.packet_hash = po.packet_hash WHERE po.heard_at > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) GROUP BY p.payload_type ORDER BY count DESC; @@ -757,7 +757,7 @@ SELECT COUNT(DISTINCT n.id)::bigint AS count FROM nodes n LEFT JOIN node_iatas ni ON ni.node_id = n.id -WHERE ($1::text = '' OR ni.iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR ni.iata = ANY($1::bpchar[])) GROUP BY n.node_type ORDER BY count DESC; @@ -776,7 +776,7 @@ SELECT FROM packet_observations po JOIN observers o ON o.id = po.observer_id WHERE po.heard_at > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) GROUP BY o.id ORDER BY observation_count DESC LIMIT $3; @@ -785,7 +785,7 @@ LIMIT $3; SELECT preset, iata, source_type, count FROM mv_radio_presets WHERE ($1::text = '' OR preset = $1::text) - AND ($2::text = '' OR iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR iata = ANY($2::bpchar[])) ORDER BY preset, iata, source_type; -- name: GetScopeStats :many @@ -867,7 +867,7 @@ LEFT JOIN LATERAL ( LIMIT 1 ) best ON true WHERE p.trace_tag IS NOT NULL - AND ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))) + AND (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])) AND ($2::text = '' OR p.scope_id = (SELECT id FROM transport_scopes WHERE name = $2)) AND ($3::timestamptz IS NULL OR p.first_heard_at >= $3) AND ($4::timestamptz IS NULL OR p.first_heard_at <= $4) diff --git a/db/scopes.go b/db/scopes.go index 5a7bb27..af5cf9a 100644 --- a/db/scopes.go +++ b/db/scopes.go @@ -5,7 +5,6 @@ package db import ( "context" - "strings" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" "github.com/MeshCore-Beacon/beacon-server/internal/api" @@ -52,8 +51,7 @@ func (s *Store) GetScopeNames(ctx context.Context) ([]string, error) { // GetScopesByIATAs returns scope summaries filtered by the given IATA codes. func (s *Store) GetScopesByIATAs(ctx context.Context, iatas []string) ([]api.ScopeSummary, error) { - iataFilter := strings.Join(iatas, ",") - rows, err := s.q.GetScopesByIATAs(ctx, iataFilter) + rows, err := s.q.GetScopesByIATAs(ctx, iatas) if err != nil { return nil, err } diff --git a/db/scopes_test.go b/db/scopes_test.go index fc669c2..890e93a 100644 --- a/db/scopes_test.go +++ b/db/scopes_test.go @@ -48,7 +48,7 @@ func TestGetScopesByIATAs(t *testing.T) { mock := mockdb.NewMockQuerier(ctrl) mock.EXPECT(). - GetScopesByIATAs(gomock.Any(), "YVR,YYJ"). + GetScopesByIATAs(gomock.Any(), []string{"YVR", "YYJ"}). Return([]sqlc.GetScopesByIATAsRow{ {Name: "default", ObserverCount: 3, NodeCount: 10, IataCount: 2}, }, nil) diff --git a/db/sqlc/mock/querier.go b/db/sqlc/mock/querier.go index e889749..d89af04 100644 --- a/db/sqlc/mock/querier.go +++ b/db/sqlc/mock/querier.go @@ -477,7 +477,7 @@ func (mr *MockQuerierMockRecorder) GetScopeStats(ctx any) *gomock.Call { } // GetScopesByIATAs mocks base method. -func (m *MockQuerier) GetScopesByIATAs(ctx context.Context, dollar_1 string) ([]db.GetScopesByIATAsRow, error) { +func (m *MockQuerier) GetScopesByIATAs(ctx context.Context, dollar_1 []string) ([]db.GetScopesByIATAsRow, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "GetScopesByIATAs", ctx, dollar_1) ret0, _ := ret[0].([]db.GetScopesByIATAsRow) @@ -492,7 +492,7 @@ func (mr *MockQuerierMockRecorder) GetScopesByIATAs(ctx, dollar_1 any) *gomock.C } // GetStatsNodeTypes mocks base method. -func (m *MockQuerier) GetStatsNodeTypes(ctx context.Context, dollar_1 string) ([]db.GetStatsNodeTypesRow, error) { +func (m *MockQuerier) GetStatsNodeTypes(ctx context.Context, dollar_1 []string) ([]db.GetStatsNodeTypesRow, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "GetStatsNodeTypes", ctx, dollar_1) ret0, _ := ret[0].([]db.GetStatsNodeTypesRow) @@ -507,7 +507,7 @@ func (mr *MockQuerierMockRecorder) GetStatsNodeTypes(ctx, dollar_1 any) *gomock. } // GetStatsOverview mocks base method. -func (m *MockQuerier) GetStatsOverview(ctx context.Context, dollar_1 string) (db.GetStatsOverviewRow, error) { +func (m *MockQuerier) GetStatsOverview(ctx context.Context, dollar_1 []string) (db.GetStatsOverviewRow, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "GetStatsOverview", ctx, dollar_1) ret0, _ := ret[0].(db.GetStatsOverviewRow) diff --git a/db/sqlc/querier.go b/db/sqlc/querier.go index 5fdc42f..f4effb4 100644 --- a/db/sqlc/querier.go +++ b/db/sqlc/querier.go @@ -47,13 +47,13 @@ type Querier interface { GetScopeByName(ctx context.Context, name string) (GetScopeByNameRow, error) GetScopeNames(ctx context.Context) ([]string, error) GetScopeStats(ctx context.Context) ([]GetScopeStatsRow, error) - GetScopesByIATAs(ctx context.Context, dollar_1 string) ([]GetScopesByIATAsRow, error) + GetScopesByIATAs(ctx context.Context, dollar_1 []string) ([]GetScopesByIATAsRow, error) // Returns node counts grouped by type, optionally filtered by IATA. - GetStatsNodeTypes(ctx context.Context, dollar_1 string) ([]GetStatsNodeTypesRow, error) + GetStatsNodeTypes(ctx context.Context, dollar_1 []string) ([]GetStatsNodeTypesRow, error) // ============================================================ // STATS // ============================================================ - GetStatsOverview(ctx context.Context, dollar_1 string) (GetStatsOverviewRow, error) + GetStatsOverview(ctx context.Context, dollar_1 []string) (GetStatsOverviewRow, error) // Returns observation counts grouped by payload type for the given window and IATA. GetStatsPayloadBreakdown(ctx context.Context, arg GetStatsPayloadBreakdownParams) ([]GetStatsPayloadBreakdownRow, error) // Returns the top N observers by observation count for the given window and IATA. diff --git a/db/sqlc/queries.sql.go b/db/sqlc/queries.sql.go index bdcb860..c0e7b4e 100644 --- a/db/sqlc/queries.sql.go +++ b/db/sqlc/queries.sql.go @@ -118,13 +118,13 @@ func (q *Queries) GetCrossIATANeighbors(ctx context.Context, arg GetCrossIATANei const getHourlyStats = `-- name: GetHourlyStats :many SELECT iata, hour, observation_count, unique_packets, active_observers FROM mv_hourly_iata_stats -WHERE ($1::text = '' OR iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR iata = ANY($1::bpchar[])) AND hour >= NOW() - $2::interval ORDER BY iata, hour ` type GetHourlyStatsParams struct { - Column1 string `json:"column_1"` + Column1 []string `json:"column_1"` Column2 pgtype.Interval `json:"column_2"` } @@ -823,13 +823,13 @@ const getRadioPresets = `-- name: GetRadioPresets :many SELECT preset, iata, source_type, count FROM mv_radio_presets WHERE ($1::text = '' OR preset = $1::text) - AND ($2::text = '' OR iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR iata = ANY($2::bpchar[])) ORDER BY preset, iata, source_type ` type GetRadioPresetsParams struct { - Column1 string `json:"column_1"` - Column2 string `json:"column_2"` + Column1 string `json:"column_1"` + Column2 []string `json:"column_2"` } func (q *Queries) GetRadioPresets(ctx context.Context, arg GetRadioPresetsParams) ([]MvRadioPreset, error) { @@ -1067,7 +1067,7 @@ LEFT JOIN observer_scopes os ON os.scope_id = ts.id LEFT JOIN observers o ON o.id = os.observer_id LEFT JOIN packet_observations po ON po.observer_id = o.id LEFT JOIN nodes n ON n.default_scope_id = ts.id -WHERE ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])) GROUP BY ts.name ORDER BY ts.name ` @@ -1079,7 +1079,7 @@ type GetScopesByIATAsRow struct { IataCount int64 `json:"iata_count"` } -func (q *Queries) GetScopesByIATAs(ctx context.Context, dollar_1 string) ([]GetScopesByIATAsRow, error) { +func (q *Queries) GetScopesByIATAs(ctx context.Context, dollar_1 []string) ([]GetScopesByIATAsRow, error) { rows, err := q.db.Query(ctx, getScopesByIATAs, dollar_1) if err != nil { return nil, err @@ -1110,7 +1110,7 @@ SELECT COUNT(DISTINCT n.id)::bigint AS count FROM nodes n LEFT JOIN node_iatas ni ON ni.node_id = n.id -WHERE ($1::text = '' OR ni.iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR ni.iata = ANY($1::bpchar[])) GROUP BY n.node_type ORDER BY count DESC ` @@ -1121,7 +1121,7 @@ type GetStatsNodeTypesRow struct { } // Returns node counts grouped by type, optionally filtered by IATA. -func (q *Queries) GetStatsNodeTypes(ctx context.Context, dollar_1 string) ([]GetStatsNodeTypesRow, error) { +func (q *Queries) GetStatsNodeTypes(ctx context.Context, dollar_1 []string) ([]GetStatsNodeTypesRow, error) { rows, err := q.db.Query(ctx, getStatsNodeTypes, dollar_1) if err != nil { return nil, err @@ -1150,7 +1150,7 @@ SELECT COUNT(DISTINCT po.iata) AS active_iatas FROM packet_observations po WHERE po.heard_at > NOW() - INTERVAL '24 hours' - AND ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))) + AND (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])) ` type GetStatsOverviewRow struct { @@ -1163,7 +1163,7 @@ type GetStatsOverviewRow struct { // ============================================================ // STATS // ============================================================ -func (q *Queries) GetStatsOverview(ctx context.Context, dollar_1 string) (GetStatsOverviewRow, error) { +func (q *Queries) GetStatsOverview(ctx context.Context, dollar_1 []string) (GetStatsOverviewRow, error) { row := q.db.QueryRow(ctx, getStatsOverview, dollar_1) var i GetStatsOverviewRow err := row.Scan( @@ -1182,14 +1182,14 @@ SELECT FROM packet_observations po JOIN packets p ON p.packet_hash = po.packet_hash WHERE po.heard_at > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) GROUP BY p.payload_type ORDER BY count DESC ` type GetStatsPayloadBreakdownParams struct { HeardAt pgtype.Timestamptz `json:"heard_at"` - Column2 string `json:"column_2"` + Column2 []string `json:"column_2"` } type GetStatsPayloadBreakdownRow struct { @@ -1232,7 +1232,7 @@ SELECT FROM packet_observations po JOIN observers o ON o.id = po.observer_id WHERE po.heard_at > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) GROUP BY o.id ORDER BY observation_count DESC LIMIT $3 @@ -1240,7 +1240,7 @@ LIMIT $3 type GetStatsTopObserversParams struct { HeardAt pgtype.Timestamptz `json:"heard_at"` - Column2 string `json:"column_2"` + Column2 []string `json:"column_2"` Limit int32 `json:"limit"` } @@ -1281,14 +1281,14 @@ func (q *Queries) GetStatsTopObservers(ctx context.Context, arg GetStatsTopObser const getTopNodes = `-- name: GetTopNodes :many SELECT iata, node_id, name, node_type, observation_count, last_heard FROM mv_top_nodes_by_iata -WHERE ($1::text = '' OR iata = ANY(string_to_array($1::text, ','))) +WHERE (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR iata = ANY($1::bpchar[])) ORDER BY observation_count DESC LIMIT $2 ` type GetTopNodesParams struct { - Column1 string `json:"column_1"` - Limit int32 `json:"limit"` + Column1 []string `json:"column_1"` + Limit int32 `json:"limit"` } func (q *Queries) GetTopNodes(ctx context.Context, arg GetTopNodesParams) ([]MvTopNodesByIatum, error) { @@ -1530,7 +1530,7 @@ JOIN packet_observations po ON po.packet_hash = cm.packet_hash JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE ($1::timestamptz IS NULL OR cm.sent_at >= $1) - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) AND ($3::text = '' OR ts.name = $3::text) AND ($4 = 0 OR cm.id < $4) ORDER BY cm.id DESC @@ -1539,7 +1539,7 @@ LIMIT $5 type ListAllChannelMessagesParams struct { Column1 pgtype.Timestamptz `json:"column_1"` - Column2 string `json:"column_2"` + Column2 []string `json:"column_2"` Column3 string `json:"column_3"` Column4 interface{} `json:"column_4"` Limit int32 `json:"limit"` @@ -1608,7 +1608,7 @@ JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE cm.channel_id = $1 AND ($2::timestamptz IS NULL OR cm.sent_at >= $2) - AND ($3::text = '' OR po.iata = ANY(string_to_array($3::text, ','))) + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR po.iata = ANY($3::bpchar[])) AND ($4::text = '' OR ts.name = $4::text) AND ($5::bigint = 0 OR cm.id < $5::bigint) ORDER BY cm.id DESC @@ -1618,7 +1618,7 @@ LIMIT $6 type ListChannelMessagesParams struct { ChannelID int32 `json:"channel_id"` Column2 pgtype.Timestamptz `json:"column_2"` - Column3 string `json:"column_3"` + Column3 []string `json:"column_3"` Column4 string `json:"column_4"` Column5 int64 `json:"column_5"` Limit int32 `json:"limit"` @@ -1689,7 +1689,7 @@ JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE c.channel_hash = $1 AND ($2::timestamptz IS NULL OR cm.sent_at >= $2) - AND ($3::text = '' OR po.iata = ANY(string_to_array($3::text, ','))) + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR po.iata = ANY($3::bpchar[])) AND ($4::text = '' OR ts.name = $4::text) AND ($5::bigint = 0 OR cm.id < $5::bigint) ORDER BY cm.id DESC @@ -1699,7 +1699,7 @@ LIMIT $6 type ListChannelMessagesByHashParams struct { ChannelHash []byte `json:"channel_hash"` Column2 pgtype.Timestamptz `json:"column_2"` - Column3 string `json:"column_3"` + Column3 []string `json:"column_3"` Column4 string `json:"column_4"` Column5 int64 `json:"column_5"` Limit int32 `json:"limit"` @@ -1910,17 +1910,17 @@ JOIN packet_observations po ON po.packet_hash = cm.packet_hash JOIN packets p ON p.packet_hash = cm.packet_hash LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE cm.id > $1 - AND ($2::text = '' OR po.iata = ANY(string_to_array($2::text, ','))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR po.iata = ANY($2::bpchar[])) AND ($3::text = '' OR ts.name = $3::text) ORDER BY cm.id ASC LIMIT $4 ` type ListMessagesAfterIDParams struct { - ID int64 `json:"id"` - Column2 string `json:"column_2"` - Column3 string `json:"column_3"` - Limit int32 `json:"limit"` + ID int64 `json:"id"` + Column2 []string `json:"column_2"` + Column3 string `json:"column_3"` + Limit int32 `json:"limit"` } type ListMessagesAfterIDRow struct { @@ -2050,7 +2050,7 @@ LEFT JOIN node_iatas ni ON ni.node_id = n.id LEFT JOIN transport_scopes ts ON ts.id = n.default_scope_id WHERE ($1 = 0 OR n.node_type = $1) - AND ($2::text = '' OR n.id IN (SELECT node_id FROM node_iatas WHERE iata = ANY(string_to_array($2::text, ',')))) + AND (COALESCE(cardinality($2::bpchar[]), 0) = 0 OR n.id IN (SELECT node_id FROM node_iatas WHERE iata = ANY($2::bpchar[]))) AND ( $3::text = 'any' OR ($3::text = 'true' AND n.supports_multibyte_paths = TRUE) @@ -2072,7 +2072,7 @@ LIMIT $8 type ListNodesParams struct { Column1 interface{} `json:"column_1"` - Column2 string `json:"column_2"` + Column2 []string `json:"column_2"` Column3 string `json:"column_3"` Column4 string `json:"column_4"` Column5 []byte `json:"column_5"` @@ -2318,11 +2318,11 @@ LEFT JOIN observer_brokers ob ON ob.observer_id = o.id LEFT JOIN observer_scopes os ON os.observer_id = o.id LEFT JOIN transport_scopes ts ON ts.id = os.scope_id WHERE - ($1::text = '' OR ( + (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR ( SELECT po.iata FROM packet_observations po WHERE po.observer_id = o.id ORDER BY po.heard_at DESC LIMIT 1 - ) = ANY(string_to_array($1::text, ','))) + ) = ANY($1::bpchar[])) AND ($2 = '' OR o.observer_type = $2) AND ($3 = '' OR ob.broker_name = $3) AND ($4 = '' OR CASE @@ -2342,7 +2342,7 @@ LIMIT $7 ` type ListObserversParams struct { - Column1 string `json:"column_1"` + Column1 []string `json:"column_1"` Column2 interface{} `json:"column_2"` Column3 interface{} `json:"column_3"` Column4 interface{} `json:"column_4"` @@ -2433,10 +2433,10 @@ LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE ($1::smallint = -1 OR p.payload_type = $1::smallint) AND ($2::smallint = -1 OR p.route_type = $2::smallint) - AND ($3::text = '' OR EXISTS ( + AND (COALESCE(cardinality($3::bpchar[]), 0) = 0 OR EXISTS ( SELECT 1 FROM packet_observations po3 WHERE po3.packet_hash = p.packet_hash - AND po3.iata = ANY(string_to_array($3::text, ',')) + AND po3.iata = ANY($3::bpchar[]) )) AND ($4::timestamptz IS NULL OR p.first_heard_at >= $4) AND ($5::timestamptz IS NULL OR p.first_heard_at <= $5) @@ -2449,7 +2449,7 @@ LIMIT $7 type ListPacketsParams struct { Column1 int16 `json:"column_1"` Column2 int16 `json:"column_2"` - Column3 string `json:"column_3"` + Column3 []string `json:"column_3"` Column4 pgtype.Timestamptz `json:"column_4"` Column5 pgtype.Timestamptz `json:"column_5"` Column6 pgtype.Timestamptz `json:"column_6"` @@ -2533,19 +2533,19 @@ LEFT JOIN transport_scopes ts ON ts.id = p.scope_id WHERE po.id > $1 AND ($2::smallint = -1 OR p.payload_type = $2::smallint) AND ($3::smallint = -1 OR p.route_type = $3::smallint) - AND ($4::text = '' OR po.iata = ANY(string_to_array($4::text, ','))) + AND (COALESCE(cardinality($4::bpchar[]), 0) = 0 OR po.iata = ANY($4::bpchar[])) AND ($5::text = '' OR ts.name = $5::text) ORDER BY po.id ASC LIMIT $6 ` type ListPacketsAfterIDParams struct { - ID int64 `json:"id"` - Column2 int16 `json:"column_2"` - Column3 int16 `json:"column_3"` - Column4 string `json:"column_4"` - Column5 string `json:"column_5"` - Limit int32 `json:"limit"` + ID int64 `json:"id"` + Column2 int16 `json:"column_2"` + Column3 int16 `json:"column_3"` + Column4 []string `json:"column_4"` + Column5 string `json:"column_5"` + Limit int32 `json:"limit"` } type ListPacketsAfterIDRow struct { @@ -2657,7 +2657,7 @@ LEFT JOIN LATERAL ( LIMIT 1 ) best ON true WHERE p.trace_tag IS NOT NULL - AND ($1::text = '' OR po.iata = ANY(string_to_array($1::text, ','))) + AND (COALESCE(cardinality($1::bpchar[]), 0) = 0 OR po.iata = ANY($1::bpchar[])) AND ($2::text = '' OR p.scope_id = (SELECT id FROM transport_scopes WHERE name = $2)) AND ($3::timestamptz IS NULL OR p.first_heard_at >= $3) AND ($4::timestamptz IS NULL OR p.first_heard_at <= $4) @@ -2669,7 +2669,7 @@ LIMIT $6 ` type ListTraceTagsParams struct { - Column1 string `json:"column_1"` + Column1 []string `json:"column_1"` Column2 string `json:"column_2"` Column3 pgtype.Timestamptz `json:"column_3"` Column4 pgtype.Timestamptz `json:"column_4"` diff --git a/db/stats.go b/db/stats.go index f6a8382..8da0074 100644 --- a/db/stats.go +++ b/db/stats.go @@ -5,7 +5,6 @@ package db import ( "context" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -14,7 +13,7 @@ import ( ) func (s *Store) GetStatsOverview(ctx context.Context, iatas []string) (*api.StatsOverview, error) { - row, err := s.q.GetStatsOverview(ctx, strings.Join(iatas, ",")) + row, err := s.q.GetStatsOverview(ctx, iatas) if err != nil { return nil, err } @@ -33,7 +32,7 @@ func (s *Store) GetStatsObservations(ctx context.Context, iatas []string, since } interval := time.Since(since) rows, err := s.q.GetHourlyStats(ctx, sqlc.GetHourlyStatsParams{ - Column1: strings.Join(iatas, ","), + Column1: iatas, Column2: pgtype.Interval{Microseconds: int64(interval.Hours()) * 3600 * 1e6, Valid: true}, }) if err != nil { @@ -58,7 +57,7 @@ func (s *Store) GetStatsPayloadBreakdown(ctx context.Context, iatas []string, si } rows, err := s.q.GetStatsPayloadBreakdown(ctx, sqlc.GetStatsPayloadBreakdownParams{ HeardAt: pgtype.Timestamptz{Time: since, Valid: true}, - Column2: strings.Join(iatas, ","), + Column2: iatas, }) if err != nil { return nil, err @@ -76,7 +75,7 @@ func (s *Store) GetStatsPayloadBreakdown(ctx context.Context, iatas []string, si func (s *Store) GetStatsTopNodes(ctx context.Context, iatas []string, limit int32) ([]api.TopNode, error) { rows, err := s.q.GetTopNodes(ctx, sqlc.GetTopNodesParams{ - Column1: strings.Join(iatas, ","), + Column1: iatas, Limit: limit, }) if err != nil { @@ -107,7 +106,7 @@ func (s *Store) GetStatsTopObservers(ctx context.Context, iatas []string, since } rows, err := s.q.GetStatsTopObservers(ctx, sqlc.GetStatsTopObserversParams{ HeardAt: pgtype.Timestamptz{Time: since, Valid: true}, - Column2: strings.Join(iatas, ","), + Column2: iatas, Limit: limit, }) if err != nil { @@ -130,7 +129,7 @@ func (s *Store) GetStatsTopObservers(ctx context.Context, iatas []string, since func (s *Store) GetRadioPresets(ctx context.Context, preset string, iatas []string) ([]api.RadioPreset, error) { rows, err := s.q.GetRadioPresets(ctx, sqlc.GetRadioPresetsParams{ Column1: preset, - Column2: strings.Join(iatas, ","), + Column2: iatas, }) if err != nil { return nil, err @@ -165,7 +164,7 @@ func (s *Store) GetScopeStats(ctx context.Context) ([]api.ScopeStats, error) { } func (s *Store) GetStatsNodeTypes(ctx context.Context, iatas []string) ([]api.NodeTypeCount, error) { - rows, err := s.q.GetStatsNodeTypes(ctx, strings.Join(iatas, ",")) + rows, err := s.q.GetStatsNodeTypes(ctx, iatas) if err != nil { return nil, err } diff --git a/db/stats_test.go b/db/stats_test.go index 685f7b8..fec74d1 100644 --- a/db/stats_test.go +++ b/db/stats_test.go @@ -20,7 +20,7 @@ func TestGetStatsOverview(t *testing.T) { mock := mockdb.NewMockQuerier(ctrl) mock.EXPECT(). - GetStatsOverview(gomock.Any(), "YVR"). + GetStatsOverview(gomock.Any(), []string{"YVR"}). Return(sqlc.GetStatsOverviewRow{ TotalPackets: 100, TotalObservations: 500, @@ -51,7 +51,7 @@ func TestGetStatsTopNodes_NilObservationCount(t *testing.T) { mock.EXPECT(). GetTopNodes(gomock.Any(), sqlc.GetTopNodesParams{ - Column1: "YVR", + Column1: []string{"YVR"}, Limit: 5, }). Return([]sqlc.MvTopNodesByIatum{ @@ -119,7 +119,7 @@ func TestGetStatsNodeTypes_Mapping(t *testing.T) { mock := mockdb.NewMockQuerier(ctrl) mock.EXPECT(). - GetStatsNodeTypes(gomock.Any(), "YVR"). + GetStatsNodeTypes(gomock.Any(), []string{"YVR"}). Return([]sqlc.GetStatsNodeTypesRow{ {NodeType: 1, Count: 10}, {NodeType: 2, Count: 5}, diff --git a/db/traces.go b/db/traces.go index fa8b1b6..1b63c2f 100644 --- a/db/traces.go +++ b/db/traces.go @@ -7,7 +7,6 @@ import ( "context" "encoding/hex" "encoding/json" - "strings" "time" sqlc "github.com/MeshCore-Beacon/beacon-server/db/sqlc" @@ -22,7 +21,6 @@ type tracePayload struct { } func (s *Store) ListTraceTags(ctx context.Context, iatas []string, scope, traceType string, since, until time.Time, cursor time.Time, limit int32) ([]api.TraceTagSummary, error) { - iataFilter := strings.Join(iatas, ",") var sinceTS, untilTS, cursorTS pgtype.Timestamptz if !since.IsZero() { sinceTS = pgtype.Timestamptz{Time: since, Valid: true} @@ -34,7 +32,7 @@ func (s *Store) ListTraceTags(ctx context.Context, iatas []string, scope, traceT cursorTS = pgtype.Timestamptz{Time: cursor, Valid: true} } rows, err := s.q.ListTraceTags(ctx, sqlc.ListTraceTagsParams{ - Column1: iataFilter, + Column1: iatas, Column2: scope, Column3: sinceTS, Column4: untilTS, diff --git a/sqlc.yaml b/sqlc.yaml index 892a782..7943109 100644 --- a/sqlc.yaml +++ b/sqlc.yaml @@ -20,3 +20,5 @@ sql: go_type: "github.com/google/uuid.UUID" - db_type: "jsonb" go_type: "encoding/json.RawMessage" + - db_type: "bpchar" + go_type: "string"