mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: improve coordinator peer mapping performance (#23696)
- Skipping DB querying entirely for peers that aren't actually connected to our coordinator - Opportunistically batching the queries for peers
This commit is contained in:
@@ -3646,18 +3646,18 @@ func (q *querier) GetTailnetPeers(ctx context.Context, id uuid.UUID) ([]database
|
||||
return q.db.GetTailnetPeers(ctx, id)
|
||||
}
|
||||
|
||||
func (q *querier) GetTailnetTunnelPeerBindings(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsRow, error) {
|
||||
func (q *querier) GetTailnetTunnelPeerBindingsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsBatchRow, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceTailnetCoordinator); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetTailnetTunnelPeerBindings(ctx, srcID)
|
||||
return q.db.GetTailnetTunnelPeerBindingsBatch(ctx, ids)
|
||||
}
|
||||
|
||||
func (q *querier) GetTailnetTunnelPeerIDs(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerIDsRow, error) {
|
||||
func (q *querier) GetTailnetTunnelPeerIDsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerIDsBatchRow, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceTailnetCoordinator); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetTailnetTunnelPeerIDs(ctx, srcID)
|
||||
return q.db.GetTailnetTunnelPeerIDsBatch(ctx, ids)
|
||||
}
|
||||
|
||||
func (q *querier) GetTaskByID(ctx context.Context, id uuid.UUID) (database.Task, error) {
|
||||
|
||||
@@ -3750,13 +3750,11 @@ func (s *MethodTestSuite) TestTailnetFunctions() {
|
||||
check.Args(uuid.New()).
|
||||
Asserts(rbac.ResourceTailnetCoordinator, policy.ActionRead)
|
||||
}))
|
||||
s.Run("GetTailnetTunnelPeerBindings", s.Subtest(func(_ database.Store, check *expects) {
|
||||
check.Args(uuid.New()).
|
||||
Asserts(rbac.ResourceTailnetCoordinator, policy.ActionRead)
|
||||
s.Run("GetTailnetTunnelPeerBindingsBatch", s.Subtest(func(_ database.Store, check *expects) {
|
||||
check.Args([]uuid.UUID{uuid.New()}).Asserts(rbac.ResourceTailnetCoordinator, policy.ActionRead)
|
||||
}))
|
||||
s.Run("GetTailnetTunnelPeerIDs", s.Subtest(func(_ database.Store, check *expects) {
|
||||
check.Args(uuid.New()).
|
||||
Asserts(rbac.ResourceTailnetCoordinator, policy.ActionRead)
|
||||
s.Run("GetTailnetTunnelPeerIDsBatch", s.Subtest(func(_ database.Store, check *expects) {
|
||||
check.Args([]uuid.UUID{uuid.New()}).Asserts(rbac.ResourceTailnetCoordinator, policy.ActionRead)
|
||||
}))
|
||||
s.Run("GetAllTailnetCoordinators", s.Subtest(func(_ database.Store, check *expects) {
|
||||
check.Args().
|
||||
|
||||
@@ -2216,19 +2216,19 @@ func (m queryMetricsStore) GetTailnetPeers(ctx context.Context, id uuid.UUID) ([
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetTailnetTunnelPeerBindings(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsRow, error) {
|
||||
func (m queryMetricsStore) GetTailnetTunnelPeerBindingsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsBatchRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetTailnetTunnelPeerBindings(ctx, srcID)
|
||||
m.queryLatencies.WithLabelValues("GetTailnetTunnelPeerBindings").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetTailnetTunnelPeerBindings").Inc()
|
||||
r0, r1 := m.s.GetTailnetTunnelPeerBindingsBatch(ctx, ids)
|
||||
m.queryLatencies.WithLabelValues("GetTailnetTunnelPeerBindingsBatch").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetTailnetTunnelPeerBindingsBatch").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetTailnetTunnelPeerIDs(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerIDsRow, error) {
|
||||
func (m queryMetricsStore) GetTailnetTunnelPeerIDsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerIDsBatchRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetTailnetTunnelPeerIDs(ctx, srcID)
|
||||
m.queryLatencies.WithLabelValues("GetTailnetTunnelPeerIDs").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetTailnetTunnelPeerIDs").Inc()
|
||||
r0, r1 := m.s.GetTailnetTunnelPeerIDsBatch(ctx, ids)
|
||||
m.queryLatencies.WithLabelValues("GetTailnetTunnelPeerIDsBatch").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetTailnetTunnelPeerIDsBatch").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
|
||||
@@ -4113,34 +4113,34 @@ func (mr *MockStoreMockRecorder) GetTailnetPeers(ctx, id any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTailnetPeers", reflect.TypeOf((*MockStore)(nil).GetTailnetPeers), ctx, id)
|
||||
}
|
||||
|
||||
// GetTailnetTunnelPeerBindings mocks base method.
|
||||
func (m *MockStore) GetTailnetTunnelPeerBindings(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsRow, error) {
|
||||
// GetTailnetTunnelPeerBindingsBatch mocks base method.
|
||||
func (m *MockStore) GetTailnetTunnelPeerBindingsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerBindingsBatchRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetTailnetTunnelPeerBindings", ctx, srcID)
|
||||
ret0, _ := ret[0].([]database.GetTailnetTunnelPeerBindingsRow)
|
||||
ret := m.ctrl.Call(m, "GetTailnetTunnelPeerBindingsBatch", ctx, ids)
|
||||
ret0, _ := ret[0].([]database.GetTailnetTunnelPeerBindingsBatchRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetTailnetTunnelPeerBindings indicates an expected call of GetTailnetTunnelPeerBindings.
|
||||
func (mr *MockStoreMockRecorder) GetTailnetTunnelPeerBindings(ctx, srcID any) *gomock.Call {
|
||||
// GetTailnetTunnelPeerBindingsBatch indicates an expected call of GetTailnetTunnelPeerBindingsBatch.
|
||||
func (mr *MockStoreMockRecorder) GetTailnetTunnelPeerBindingsBatch(ctx, ids any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTailnetTunnelPeerBindings", reflect.TypeOf((*MockStore)(nil).GetTailnetTunnelPeerBindings), ctx, srcID)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTailnetTunnelPeerBindingsBatch", reflect.TypeOf((*MockStore)(nil).GetTailnetTunnelPeerBindingsBatch), ctx, ids)
|
||||
}
|
||||
|
||||
// GetTailnetTunnelPeerIDs mocks base method.
|
||||
func (m *MockStore) GetTailnetTunnelPeerIDs(ctx context.Context, srcID uuid.UUID) ([]database.GetTailnetTunnelPeerIDsRow, error) {
|
||||
// GetTailnetTunnelPeerIDsBatch mocks base method.
|
||||
func (m *MockStore) GetTailnetTunnelPeerIDsBatch(ctx context.Context, ids []uuid.UUID) ([]database.GetTailnetTunnelPeerIDsBatchRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetTailnetTunnelPeerIDs", ctx, srcID)
|
||||
ret0, _ := ret[0].([]database.GetTailnetTunnelPeerIDsRow)
|
||||
ret := m.ctrl.Call(m, "GetTailnetTunnelPeerIDsBatch", ctx, ids)
|
||||
ret0, _ := ret[0].([]database.GetTailnetTunnelPeerIDsBatchRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetTailnetTunnelPeerIDs indicates an expected call of GetTailnetTunnelPeerIDs.
|
||||
func (mr *MockStoreMockRecorder) GetTailnetTunnelPeerIDs(ctx, srcID any) *gomock.Call {
|
||||
// GetTailnetTunnelPeerIDsBatch indicates an expected call of GetTailnetTunnelPeerIDsBatch.
|
||||
func (mr *MockStoreMockRecorder) GetTailnetTunnelPeerIDsBatch(ctx, ids any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTailnetTunnelPeerIDs", reflect.TypeOf((*MockStore)(nil).GetTailnetTunnelPeerIDs), ctx, srcID)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTailnetTunnelPeerIDsBatch", reflect.TypeOf((*MockStore)(nil).GetTailnetTunnelPeerIDsBatch), ctx, ids)
|
||||
}
|
||||
|
||||
// GetTaskByID mocks base method.
|
||||
|
||||
@@ -477,8 +477,8 @@ type sqlcQuerier interface {
|
||||
// Used for recovery after coderd crashes or long hangs.
|
||||
GetStaleChats(ctx context.Context, staleThreshold time.Time) ([]Chat, error)
|
||||
GetTailnetPeers(ctx context.Context, id uuid.UUID) ([]TailnetPeer, error)
|
||||
GetTailnetTunnelPeerBindings(ctx context.Context, srcID uuid.UUID) ([]GetTailnetTunnelPeerBindingsRow, error)
|
||||
GetTailnetTunnelPeerIDs(ctx context.Context, srcID uuid.UUID) ([]GetTailnetTunnelPeerIDsRow, error)
|
||||
GetTailnetTunnelPeerBindingsBatch(ctx context.Context, ids []uuid.UUID) ([]GetTailnetTunnelPeerBindingsBatchRow, error)
|
||||
GetTailnetTunnelPeerIDsBatch(ctx context.Context, ids []uuid.UUID) ([]GetTailnetTunnelPeerIDsBatchRow, error)
|
||||
GetTaskByID(ctx context.Context, id uuid.UUID) (Task, error)
|
||||
GetTaskByOwnerIDAndName(ctx context.Context, arg GetTaskByOwnerIDAndNameParams) (Task, error)
|
||||
GetTaskByWorkspaceID(ctx context.Context, workspaceID uuid.UUID) (Task, error)
|
||||
|
||||
@@ -19309,43 +19309,44 @@ func (q *sqlQuerier) GetTailnetPeers(ctx context.Context, id uuid.UUID) ([]Tailn
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getTailnetTunnelPeerBindings = `-- name: GetTailnetTunnelPeerBindings :many
|
||||
SELECT id AS peer_id, coordinator_id, updated_at, node, status
|
||||
FROM tailnet_peers
|
||||
WHERE id IN (
|
||||
SELECT dst_id as peer_id
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.src_id = $1
|
||||
const getTailnetTunnelPeerBindingsBatch = `-- name: GetTailnetTunnelPeerBindingsBatch :many
|
||||
SELECT tp.id AS peer_id, tp.coordinator_id, tp.updated_at, tp.node, tp.status,
|
||||
tunnels.lookup_id
|
||||
FROM (
|
||||
SELECT dst_id AS peer_id, src_id AS lookup_id
|
||||
FROM tailnet_tunnels WHERE src_id = ANY($1 :: uuid[])
|
||||
UNION
|
||||
SELECT src_id as peer_id
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.dst_id = $1
|
||||
)
|
||||
SELECT src_id AS peer_id, dst_id AS lookup_id
|
||||
FROM tailnet_tunnels WHERE dst_id = ANY($1 :: uuid[])
|
||||
) tunnels
|
||||
INNER JOIN tailnet_peers tp ON tp.id = tunnels.peer_id
|
||||
`
|
||||
|
||||
type GetTailnetTunnelPeerBindingsRow struct {
|
||||
type GetTailnetTunnelPeerBindingsBatchRow struct {
|
||||
PeerID uuid.UUID `db:"peer_id" json:"peer_id"`
|
||||
CoordinatorID uuid.UUID `db:"coordinator_id" json:"coordinator_id"`
|
||||
UpdatedAt time.Time `db:"updated_at" json:"updated_at"`
|
||||
Node []byte `db:"node" json:"node"`
|
||||
Status TailnetStatus `db:"status" json:"status"`
|
||||
LookupID uuid.UUID `db:"lookup_id" json:"lookup_id"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetTailnetTunnelPeerBindings(ctx context.Context, srcID uuid.UUID) ([]GetTailnetTunnelPeerBindingsRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getTailnetTunnelPeerBindings, srcID)
|
||||
func (q *sqlQuerier) GetTailnetTunnelPeerBindingsBatch(ctx context.Context, ids []uuid.UUID) ([]GetTailnetTunnelPeerBindingsBatchRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getTailnetTunnelPeerBindingsBatch, pq.Array(ids))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetTailnetTunnelPeerBindingsRow
|
||||
var items []GetTailnetTunnelPeerBindingsBatchRow
|
||||
for rows.Next() {
|
||||
var i GetTailnetTunnelPeerBindingsRow
|
||||
var i GetTailnetTunnelPeerBindingsBatchRow
|
||||
if err := rows.Scan(
|
||||
&i.PeerID,
|
||||
&i.CoordinatorID,
|
||||
&i.UpdatedAt,
|
||||
&i.Node,
|
||||
&i.Status,
|
||||
&i.LookupID,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -19360,32 +19361,36 @@ func (q *sqlQuerier) GetTailnetTunnelPeerBindings(ctx context.Context, srcID uui
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getTailnetTunnelPeerIDs = `-- name: GetTailnetTunnelPeerIDs :many
|
||||
SELECT dst_id as peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.src_id = $1
|
||||
UNION
|
||||
SELECT src_id as peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.dst_id = $1
|
||||
const getTailnetTunnelPeerIDsBatch = `-- name: GetTailnetTunnelPeerIDsBatch :many
|
||||
SELECT src_id AS lookup_id, dst_id AS peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels WHERE src_id = ANY($1 :: uuid[])
|
||||
UNION ALL
|
||||
SELECT dst_id AS lookup_id, src_id AS peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels WHERE dst_id = ANY($1 :: uuid[])
|
||||
`
|
||||
|
||||
type GetTailnetTunnelPeerIDsRow struct {
|
||||
type GetTailnetTunnelPeerIDsBatchRow struct {
|
||||
LookupID uuid.UUID `db:"lookup_id" json:"lookup_id"`
|
||||
PeerID uuid.UUID `db:"peer_id" json:"peer_id"`
|
||||
CoordinatorID uuid.UUID `db:"coordinator_id" json:"coordinator_id"`
|
||||
UpdatedAt time.Time `db:"updated_at" json:"updated_at"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetTailnetTunnelPeerIDs(ctx context.Context, srcID uuid.UUID) ([]GetTailnetTunnelPeerIDsRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getTailnetTunnelPeerIDs, srcID)
|
||||
func (q *sqlQuerier) GetTailnetTunnelPeerIDsBatch(ctx context.Context, ids []uuid.UUID) ([]GetTailnetTunnelPeerIDsBatchRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getTailnetTunnelPeerIDsBatch, pq.Array(ids))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetTailnetTunnelPeerIDsRow
|
||||
var items []GetTailnetTunnelPeerIDsBatchRow
|
||||
for rows.Next() {
|
||||
var i GetTailnetTunnelPeerIDsRow
|
||||
if err := rows.Scan(&i.PeerID, &i.CoordinatorID, &i.UpdatedAt); err != nil {
|
||||
var i GetTailnetTunnelPeerIDsBatchRow
|
||||
if err := rows.Scan(
|
||||
&i.LookupID,
|
||||
&i.PeerID,
|
||||
&i.CoordinatorID,
|
||||
&i.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, i)
|
||||
|
||||
@@ -96,28 +96,6 @@ DELETE
|
||||
FROM tailnet_tunnels
|
||||
WHERE coordinator_id = $1 and src_id = $2;
|
||||
|
||||
-- name: GetTailnetTunnelPeerIDs :many
|
||||
SELECT dst_id as peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.src_id = $1
|
||||
UNION
|
||||
SELECT src_id as peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.dst_id = $1;
|
||||
|
||||
-- name: GetTailnetTunnelPeerBindings :many
|
||||
SELECT id AS peer_id, coordinator_id, updated_at, node, status
|
||||
FROM tailnet_peers
|
||||
WHERE id IN (
|
||||
SELECT dst_id as peer_id
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.src_id = $1
|
||||
UNION
|
||||
SELECT src_id as peer_id
|
||||
FROM tailnet_tunnels
|
||||
WHERE tailnet_tunnels.dst_id = $1
|
||||
);
|
||||
|
||||
-- For PG Coordinator HTMLDebug
|
||||
|
||||
-- name: GetAllTailnetCoordinators :many
|
||||
@@ -128,3 +106,22 @@ SELECT * FROM tailnet_peers;
|
||||
|
||||
-- name: GetAllTailnetTunnels :many
|
||||
SELECT * FROM tailnet_tunnels;
|
||||
|
||||
-- name: GetTailnetTunnelPeerIDsBatch :many
|
||||
SELECT src_id AS lookup_id, dst_id AS peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels WHERE src_id = ANY(@ids :: uuid[])
|
||||
UNION ALL
|
||||
SELECT dst_id AS lookup_id, src_id AS peer_id, coordinator_id, updated_at
|
||||
FROM tailnet_tunnels WHERE dst_id = ANY(@ids :: uuid[]);
|
||||
|
||||
-- name: GetTailnetTunnelPeerBindingsBatch :many
|
||||
SELECT tp.id AS peer_id, tp.coordinator_id, tp.updated_at, tp.node, tp.status,
|
||||
tunnels.lookup_id
|
||||
FROM (
|
||||
SELECT dst_id AS peer_id, src_id AS lookup_id
|
||||
FROM tailnet_tunnels WHERE src_id = ANY(@ids :: uuid[])
|
||||
UNION
|
||||
SELECT src_id AS peer_id, dst_id AS lookup_id
|
||||
FROM tailnet_tunnels WHERE dst_id = ANY(@ids :: uuid[])
|
||||
) tunnels
|
||||
INNER JOIN tailnet_peers tp ON tp.id = tunnels.peer_id;
|
||||
|
||||
+179
-108
@@ -3,6 +3,8 @@ package tailnet
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"math"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -807,7 +809,8 @@ type querier struct {
|
||||
newConnections chan *connIO
|
||||
closeConnections chan *connIO
|
||||
|
||||
workQ *workQ[querierWorkKey]
|
||||
peerUpdateQ *workQ[uuid.UUID]
|
||||
mappingQ *workQ[mKey]
|
||||
|
||||
wg sync.WaitGroup
|
||||
|
||||
@@ -840,7 +843,8 @@ func newQuerier(ctx context.Context,
|
||||
store: store,
|
||||
newConnections: newConnections,
|
||||
closeConnections: closeConnections,
|
||||
workQ: newWorkQ[querierWorkKey](ctx),
|
||||
peerUpdateQ: newWorkQ[uuid.UUID](ctx),
|
||||
mappingQ: newWorkQ[mKey](ctx),
|
||||
heartbeats: newHeartbeats(ctx, logger, ps, store, self, updates, firstHeartbeat, clk),
|
||||
mappers: make(map[mKey]*mapper),
|
||||
updates: updates,
|
||||
@@ -848,14 +852,21 @@ func newQuerier(ctx context.Context,
|
||||
}
|
||||
q.subscribe()
|
||||
|
||||
q.wg.Add(2 + numWorkers)
|
||||
// For an odd number of workers we allocate more to the mapping workers since they're busier.
|
||||
mappingWorkers := int(math.Ceil(float64(numWorkers) / 2))
|
||||
peerWorkers := numWorkers - mappingWorkers
|
||||
|
||||
q.wg.Add(2 + mappingWorkers + peerWorkers)
|
||||
go func() {
|
||||
<-firstHeartbeat
|
||||
go q.handleIncoming()
|
||||
for i := 0; i < numWorkers; i++ {
|
||||
go q.worker()
|
||||
}
|
||||
go q.handleUpdates()
|
||||
for range mappingWorkers {
|
||||
go q.mappingWorker()
|
||||
}
|
||||
for range peerWorkers {
|
||||
go q.peerUpdateWorker()
|
||||
}
|
||||
}()
|
||||
return q
|
||||
}
|
||||
@@ -913,9 +924,7 @@ func (q *querier) newConn(c *connIO) {
|
||||
}
|
||||
}
|
||||
q.mappers[mk] = mpr
|
||||
q.workQ.enqueue(querierWorkKey{
|
||||
mappingQuery: mk,
|
||||
})
|
||||
q.mappingQ.enqueue(mk)
|
||||
q.logger.Debug(q.ctx, "added new mapper", slog.F("peer_id", c.UniqueID()))
|
||||
}
|
||||
|
||||
@@ -947,87 +956,144 @@ func (q *querier) cleanupConn(c *connIO) {
|
||||
q.logger.Debug(q.ctx, "removed mapper", slog.F("peer_id", c.UniqueID()))
|
||||
}
|
||||
|
||||
func (q *querier) worker() {
|
||||
// maxBatchSize is the maximum number of keys to process in a single batch
|
||||
// query.
|
||||
const maxBatchSize = 50
|
||||
|
||||
func (q *querier) peerUpdateWorker() {
|
||||
defer q.wg.Done()
|
||||
defer q.logger.Debug(q.ctx, "worker exited")
|
||||
defer q.logger.Debug(q.ctx, "peerUpdate worker exited")
|
||||
eb := backoff.NewExponentialBackOff()
|
||||
eb.MaxElapsedTime = 0 // retry indefinitely
|
||||
eb.MaxInterval = dbMaxBackoff
|
||||
bkoff := backoff.WithContext(eb, q.ctx)
|
||||
for {
|
||||
qk, err := q.workQ.acquire()
|
||||
allKeys, err := q.peerUpdateQ.acquireBatch(maxBatchSize)
|
||||
if err != nil {
|
||||
// context expired
|
||||
return
|
||||
}
|
||||
peers := make([]uuid.UUID, 0, len(allKeys))
|
||||
peers = append(peers, allKeys...)
|
||||
err = backoff.Retry(func() error {
|
||||
return q.query(qk)
|
||||
return q.peerUpdate(peers)
|
||||
}, bkoff)
|
||||
if err != nil {
|
||||
bkoff.Reset()
|
||||
}
|
||||
q.workQ.done(qk)
|
||||
q.peerUpdateQ.done(allKeys...)
|
||||
}
|
||||
}
|
||||
|
||||
func (q *querier) query(qk querierWorkKey) error {
|
||||
if uuid.UUID(qk.mappingQuery) != uuid.Nil {
|
||||
return q.mappingQuery(qk.mappingQuery)
|
||||
func (q *querier) mappingWorker() {
|
||||
defer q.wg.Done()
|
||||
defer q.logger.Debug(q.ctx, "mapping worker exited")
|
||||
eb := backoff.NewExponentialBackOff()
|
||||
eb.MaxElapsedTime = 0 // retry indefinitely
|
||||
eb.MaxInterval = dbMaxBackoff
|
||||
bkoff := backoff.WithContext(eb, q.ctx)
|
||||
for {
|
||||
allKeys, err := q.mappingQ.acquireBatch(maxBatchSize)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
mkeys := make([]mKey, 0, len(allKeys))
|
||||
mkeys = append(mkeys, allKeys...)
|
||||
err = backoff.Retry(func() error {
|
||||
return q.mappingQuery(mkeys)
|
||||
}, bkoff)
|
||||
if err != nil {
|
||||
bkoff.Reset()
|
||||
}
|
||||
q.mappingQ.done(allKeys...)
|
||||
}
|
||||
if qk.peerUpdate != uuid.Nil {
|
||||
return q.peerUpdate(qk.peerUpdate)
|
||||
}
|
||||
q.logger.Critical(q.ctx, "bad querierWorkKey", slog.F("work_key", qk))
|
||||
return backoff.Permanent(xerrors.Errorf("bad querierWorkKey %v", qk))
|
||||
}
|
||||
|
||||
// peerUpdate is work scheduled in response to a new peer->binding. We need to find out all the
|
||||
// other peers that share a tunnel with the indicated peer, and then schedule a mapping update on
|
||||
// each, so that they can find out about the new binding.
|
||||
func (q *querier) peerUpdate(peer uuid.UUID) error {
|
||||
logger := q.logger.With(slog.F("peer_id", peer))
|
||||
logger.Debug(q.ctx, "querying peers that share a tunnel")
|
||||
others, err := q.store.GetTailnetTunnelPeerIDs(q.ctx, peer)
|
||||
func (q *querier) peerUpdate(peers []uuid.UUID) error {
|
||||
q.logger.Debug(q.ctx, "batch querying peers that share tunnels",
|
||||
slog.F("num_peers", len(peers)))
|
||||
others, err := q.store.GetTailnetTunnelPeerIDsBatch(q.ctx, peers)
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
return xerrors.Errorf("get tunnel peer IDs batch: %w", err)
|
||||
}
|
||||
logger.Debug(q.ctx, "queried peers that share a tunnel", slog.F("num_peers", len(others)))
|
||||
q.logger.Debug(q.ctx, "batch queried tunnel peers",
|
||||
slog.F("num_results", len(others)))
|
||||
q.mu.Lock()
|
||||
for _, other := range others {
|
||||
logger.Debug(q.ctx, "got tunnel peer", slog.F("other_id", other.PeerID))
|
||||
q.workQ.enqueue(querierWorkKey{mappingQuery: mKey(other.PeerID)})
|
||||
mk := mKey(other.PeerID)
|
||||
if _, ok := q.mappers[mk]; ok {
|
||||
q.mappingQ.enqueue(mk)
|
||||
}
|
||||
}
|
||||
q.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
// mappingQuery queries the database for all the mappings that the given peers should know about,
|
||||
// that is, all the peers that it shares a tunnel with and their current node mappings (if they
|
||||
// exist). It then sends the mapping snapshot to the corresponding mapper, where it will get
|
||||
// transmitted to the peer.
|
||||
func (q *querier) mappingQuery(peers []mKey) error {
|
||||
// Filter to peers with active mappers before hitting the DB.
|
||||
q.mu.Lock()
|
||||
active := make([]uuid.UUID, 0, len(peers))
|
||||
activeKeys := make([]mKey, 0, len(peers))
|
||||
for _, p := range peers {
|
||||
if _, ok := q.mappers[p]; ok {
|
||||
active = append(active, uuid.UUID(p))
|
||||
activeKeys = append(activeKeys, p)
|
||||
}
|
||||
}
|
||||
q.mu.Unlock()
|
||||
if len(active) == 0 {
|
||||
q.logger.Debug(q.ctx, "batch mapping query: no active mappers")
|
||||
return nil
|
||||
}
|
||||
|
||||
q.logger.Debug(q.ctx, "batch querying mappings",
|
||||
slog.F("num_peers", len(active)))
|
||||
bindings, err := q.store.GetTailnetTunnelPeerBindingsBatch(q.ctx, active)
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return xerrors.Errorf("get tunnel peer bindings batch: %w", err)
|
||||
}
|
||||
q.logger.Debug(q.ctx, "batch queried mappings",
|
||||
slog.F("num_bindings", len(bindings)))
|
||||
|
||||
// Group bindings by lookup_id (the peer that needs the mapping).
|
||||
grouped := make(map[uuid.UUID][]database.GetTailnetTunnelPeerBindingsBatchRow)
|
||||
for _, b := range bindings {
|
||||
grouped[b.LookupID] = append(grouped[b.LookupID], b)
|
||||
}
|
||||
|
||||
// Dispatch each peer's mappings to its mapper.
|
||||
for _, mk := range activeKeys {
|
||||
peerID := uuid.UUID(mk)
|
||||
rows := grouped[peerID]
|
||||
mappings, err := q.bindingsToMappings(rows)
|
||||
if err != nil {
|
||||
q.logger.Error(q.ctx, "failed to convert batch mappings",
|
||||
slog.F("peer_id", peerID), slog.Error(err))
|
||||
continue
|
||||
}
|
||||
q.mu.Lock()
|
||||
mpr, ok := q.mappers[mk]
|
||||
q.mu.Unlock()
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if err := agpl.SendCtx(mpr.ctx, mpr.mappings, mappings); err != nil {
|
||||
q.logger.Debug(q.ctx, "failed to send mappings to peer",
|
||||
slog.F("peer_id", peerID), slog.Error(err))
|
||||
continue
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// mappingQuery queries the database for all the mappings that the given peer should know about,
|
||||
// that is, all the peers that it shares a tunnel with and their current node mappings (if they
|
||||
// exist). It then sends the mapping snapshot to the corresponding mapper, where it will get
|
||||
// transmitted to the peer.
|
||||
func (q *querier) mappingQuery(peer mKey) error {
|
||||
logger := q.logger.With(slog.F("peer_id", uuid.UUID(peer)))
|
||||
logger.Debug(q.ctx, "querying mappings")
|
||||
bindings, err := q.store.GetTailnetTunnelPeerBindings(q.ctx, uuid.UUID(peer))
|
||||
logger.Debug(q.ctx, "queried mappings", slog.F("num_mappings", len(bindings)))
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
mappings, err := q.bindingsToMappings(bindings)
|
||||
if err != nil {
|
||||
logger.Debug(q.ctx, "failed to convert mappings", slog.Error(err))
|
||||
return err
|
||||
}
|
||||
q.mu.Lock()
|
||||
mpr, ok := q.mappers[peer]
|
||||
q.mu.Unlock()
|
||||
if !ok {
|
||||
logger.Debug(q.ctx, "query for missing mapper")
|
||||
return nil
|
||||
}
|
||||
logger.Debug(q.ctx, "sending mappings", slog.F("mapping_len", len(mappings)))
|
||||
return agpl.SendCtx(mpr.ctx, mpr.mappings, mappings)
|
||||
}
|
||||
|
||||
func (q *querier) bindingsToMappings(bindings []database.GetTailnetTunnelPeerBindingsRow) ([]mapping, error) {
|
||||
// bindingsToMappings converts binding rows to mappings.
|
||||
func (q *querier) bindingsToMappings(bindings []database.GetTailnetTunnelPeerBindingsBatchRow) ([]mapping, error) {
|
||||
slog.Helper()
|
||||
mappings := make([]mapping, 0, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
@@ -1162,7 +1228,7 @@ func (q *querier) listenPeer(_ context.Context, msg []byte, err error) {
|
||||
// we know that this peer has an updated node mapping, but we don't yet know who to send that
|
||||
// update to. We need to query the database to find all the other peers that share a tunnel with
|
||||
// this one, and then run mapping queries against all of them.
|
||||
q.workQ.enqueue(querierWorkKey{peerUpdate: peer})
|
||||
q.peerUpdateQ.enqueue(peer)
|
||||
}
|
||||
|
||||
func (q *querier) listenTunnel(_ context.Context, msg []byte, err error) {
|
||||
@@ -1192,13 +1258,17 @@ func (q *querier) listenTunnel(_ context.Context, msg []byte, err error) {
|
||||
slog.F("peer_id", peer))
|
||||
continue
|
||||
}
|
||||
q.workQ.enqueue(querierWorkKey{mappingQuery: mk})
|
||||
q.mappingQ.enqueue(mk)
|
||||
}
|
||||
}
|
||||
|
||||
func (q *querier) listenReadyForHandshake(_ context.Context, msg []byte, err error) {
|
||||
if err != nil && !xerrors.Is(err, pubsub.ErrDroppedMessages) {
|
||||
q.logger.Warn(q.ctx, "unhandled pubsub error", slog.Error(err))
|
||||
if err != nil {
|
||||
if xerrors.Is(err, pubsub.ErrDroppedMessages) {
|
||||
q.logger.Warn(q.ctx, "pubsub dropped ready-for-handshake messages")
|
||||
} else {
|
||||
q.logger.Warn(q.ctx, "unhandled pubsub error", slog.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1230,7 +1300,7 @@ func (q *querier) resyncPeerMappings() {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
for mk := range q.mappers {
|
||||
q.workQ.enqueue(querierWorkKey{mappingQuery: mk})
|
||||
q.mappingQ.enqueue(mk)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1347,17 +1417,8 @@ type mapping struct {
|
||||
kind proto.CoordinateResponse_PeerUpdate_Kind
|
||||
}
|
||||
|
||||
// querierWorkKey describes two kinds of work the querier needs to do. If peerUpdate
|
||||
// is not uuid.Nil, then the querier needs to find all tunnel peers of the given peer and
|
||||
// mark them for a mapping query. If mappingQuery is not uuid.Nil, then the querier has to
|
||||
// query the mappings of the tunnel peers of the given peer.
|
||||
type querierWorkKey struct {
|
||||
peerUpdate uuid.UUID
|
||||
mappingQuery mKey
|
||||
}
|
||||
|
||||
type queueKey interface {
|
||||
bKey | tKey | querierWorkKey
|
||||
bKey | tKey | uuid.UUID | mKey
|
||||
}
|
||||
|
||||
// workQ allows scheduling work based on a key. Multiple enqueue requests for the same key are coalesced, and
|
||||
@@ -1387,59 +1448,69 @@ func newWorkQ[K queueKey](ctx context.Context) *workQ[K] {
|
||||
}
|
||||
|
||||
// enqueue adds the key to the workQ if it is not already pending.
|
||||
func (q *workQ[K]) enqueue(key K) {
|
||||
func (q *workQ[K]) enqueue(keys ...K) {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
for _, mk := range q.pending {
|
||||
if mk == key {
|
||||
// already pending, no-op
|
||||
return
|
||||
for _, key := range keys {
|
||||
if slices.Contains(q.pending, key) {
|
||||
continue
|
||||
}
|
||||
q.pending = append(q.pending, key)
|
||||
}
|
||||
q.pending = append(q.pending, key)
|
||||
q.cond.Signal()
|
||||
}
|
||||
|
||||
// acquire gets a new key to begin working on. This call blocks until work is available. After acquiring a key, the
|
||||
// worker MUST call done() with the same key to mark it complete and allow new pending work to be acquired for the key.
|
||||
// acquireBatch blocks until at least one pending key is available, then
|
||||
// returns up to limit keys, moving them to inProgress. Caller must call
|
||||
// done() for each returned key.
|
||||
// An error is returned if the workQ context is canceled to unblock waiting workers.
|
||||
func (q *workQ[K]) acquire() (key K, err error) {
|
||||
func (q *workQ[K]) acquireBatch(limit int) ([]K, error) {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
for !q.workAvailable() && q.ctx.Err() == nil {
|
||||
for {
|
||||
if q.ctx.Err() != nil {
|
||||
return nil, q.ctx.Err()
|
||||
}
|
||||
var batch []K
|
||||
remaining := make([]K, 0, len(q.pending))
|
||||
for _, k := range q.pending {
|
||||
if len(batch) >= limit {
|
||||
remaining = append(remaining, k)
|
||||
continue
|
||||
}
|
||||
if _, inProg := q.inProgress[k]; inProg {
|
||||
remaining = append(remaining, k)
|
||||
continue
|
||||
}
|
||||
batch = append(batch, k)
|
||||
q.inProgress[k] = true
|
||||
}
|
||||
q.pending = remaining
|
||||
if len(batch) > 0 {
|
||||
return batch, nil
|
||||
}
|
||||
q.cond.Wait()
|
||||
}
|
||||
if q.ctx.Err() != nil {
|
||||
return key, q.ctx.Err()
|
||||
}
|
||||
for i, mk := range q.pending {
|
||||
_, ok := q.inProgress[mk]
|
||||
if !ok {
|
||||
q.pending = append(q.pending[:i], q.pending[i+1:]...)
|
||||
q.inProgress[mk] = true
|
||||
return mk, nil
|
||||
}
|
||||
}
|
||||
// this should not be possible because we are holding the lock when we exit the loop that waits
|
||||
panic("woke with no work available")
|
||||
}
|
||||
|
||||
// workAvailable returns true if there is work we can do. Must be called while holding q.cond.L
|
||||
func (q workQ[K]) workAvailable() bool {
|
||||
for _, mk := range q.pending {
|
||||
_, ok := q.inProgress[mk]
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
// acquire blocks until a work item is available and returns it. After
|
||||
// acquiring a key, the worker MUST call done() with the same key to mark
|
||||
// it complete and allow new pending work to be acquired for the key.
|
||||
func (q *workQ[K]) acquire() (key K, err error) {
|
||||
items, err := q.acquireBatch(1)
|
||||
if err != nil {
|
||||
return key, err
|
||||
}
|
||||
return false
|
||||
return items[0], nil
|
||||
}
|
||||
|
||||
// done marks the key completed; MUST be called after acquire() for each key.
|
||||
func (q *workQ[K]) done(key K) {
|
||||
func (q *workQ[K]) done(keys ...K) {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
delete(q.inProgress, key)
|
||||
for _, key := range keys {
|
||||
delete(q.inProgress, key)
|
||||
}
|
||||
q.cond.Signal()
|
||||
}
|
||||
|
||||
|
||||
@@ -433,3 +433,86 @@ func TestPGCoordinatorUnhealthy(t *testing.T) {
|
||||
_ = coordinator.Close()
|
||||
require.Eventually(t, ctrl.Satisfied, testutil.WaitShort, testutil.IntervalFast)
|
||||
}
|
||||
|
||||
func TestWorkQ_AcquireBatch_RespectsMax(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
q := newWorkQ[uuid.UUID](ctx)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
q.enqueue(uuid.New())
|
||||
}
|
||||
|
||||
batch, err := q.acquireBatch(3)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, batch, 3, "should respect max parameter")
|
||||
|
||||
for _, k := range batch {
|
||||
q.done(k)
|
||||
}
|
||||
|
||||
// Remaining 2 should be available.
|
||||
batch, err = q.acquireBatch(10)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, batch, 2)
|
||||
|
||||
for _, k := range batch {
|
||||
q.done(k)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWorkQ_AcquireBatch_SkipsInProgress(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
q := newWorkQ[uuid.UUID](ctx)
|
||||
|
||||
peer1 := uuid.New()
|
||||
peer2 := uuid.New()
|
||||
q.enqueue(peer1)
|
||||
q.enqueue(peer2)
|
||||
|
||||
// Acquire one item.
|
||||
key, err := q.acquire()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, peer1, key)
|
||||
|
||||
// Re-enqueue peer1 (simulating a new update while in progress).
|
||||
q.enqueue(peer1)
|
||||
|
||||
// acquireBatch should only return peer2 (peer1 is in progress).
|
||||
batch, err := q.acquireBatch(10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, batch, 1)
|
||||
assert.Equal(t, peer2, batch[0])
|
||||
|
||||
q.done(key)
|
||||
for _, k := range batch {
|
||||
q.done(k)
|
||||
}
|
||||
|
||||
// Now peer1 (re-enqueued) should be available.
|
||||
batch, err = q.acquireBatch(10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, batch, 1)
|
||||
assert.Equal(t, peer1, batch[0])
|
||||
}
|
||||
|
||||
func TestWorkQ_Acquire_WrapsAcquireBatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
q := newWorkQ[uuid.UUID](ctx)
|
||||
|
||||
peer := uuid.New()
|
||||
q.enqueue(peer)
|
||||
|
||||
key, err := q.acquire()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, peer, key)
|
||||
q.done(key)
|
||||
}
|
||||
|
||||
@@ -50,7 +50,7 @@ func TestPGCoordinatorSingle_ClientWithoutAgent(t *testing.T) {
|
||||
defer client.Close(ctx)
|
||||
client.UpdateDERP(10)
|
||||
require.Eventually(t, func() bool {
|
||||
clients, err := store.GetTailnetTunnelPeerBindings(ctx, agentID)
|
||||
clients, err := store.GetTailnetTunnelPeerBindingsBatch(ctx, []uuid.UUID{agentID})
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
t.Fatalf("database error: %v", err)
|
||||
}
|
||||
@@ -590,9 +590,8 @@ func TestPGCoordinator_Unhealthy(t *testing.T) {
|
||||
mStore.EXPECT().CleanTailnetCoordinators(gomock.Any()).AnyTimes().Return(nil)
|
||||
mStore.EXPECT().CleanTailnetLostPeers(gomock.Any()).AnyTimes().Return(nil)
|
||||
mStore.EXPECT().CleanTailnetTunnels(gomock.Any()).AnyTimes().Return(nil)
|
||||
mStore.EXPECT().GetTailnetTunnelPeerIDs(gomock.Any(), gomock.Any()).AnyTimes().Return(nil, nil)
|
||||
mStore.EXPECT().GetTailnetTunnelPeerBindings(gomock.Any(), gomock.Any()).
|
||||
AnyTimes().Return(nil, nil)
|
||||
mStore.EXPECT().GetTailnetTunnelPeerIDsBatch(gomock.Any(), gomock.Any()).AnyTimes().Return(nil, nil)
|
||||
mStore.EXPECT().GetTailnetTunnelPeerBindingsBatch(gomock.Any(), gomock.Any()).AnyTimes().Return(nil, nil)
|
||||
mStore.EXPECT().DeleteTailnetPeer(gomock.Any(), gomock.Any()).
|
||||
AnyTimes().Return(database.DeleteTailnetPeerRow{}, nil)
|
||||
mStore.EXPECT().DeleteAllTailnetTunnels(gomock.Any(), gomock.Any()).AnyTimes().Return(nil)
|
||||
@@ -934,7 +933,7 @@ func assertEventuallyLost(ctx context.Context, t *testing.T, store database.Stor
|
||||
func assertEventuallyNoClientsForAgent(ctx context.Context, t *testing.T, store database.Store, agentID uuid.UUID) {
|
||||
t.Helper()
|
||||
assert.Eventually(t, func() bool {
|
||||
clients, err := store.GetTailnetTunnelPeerIDs(ctx, agentID)
|
||||
clients, err := store.GetTailnetTunnelPeerIDsBatch(ctx, []uuid.UUID{agentID})
|
||||
if xerrors.Is(err, sql.ErrNoRows) {
|
||||
return true
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user