lib/connections: Dial devices in parallel (#7783)
This commit is contained in:
@@ -66,6 +66,8 @@ const (
|
||||
worstDialerPriority = math.MaxInt32
|
||||
recentlySeenCutoff = 7 * 24 * time.Hour
|
||||
shortLivedConnectionThreshold = 5 * time.Second
|
||||
dialMaxParallel = 64
|
||||
dialMaxParallelPerDevice = 8
|
||||
)
|
||||
|
||||
// From go/src/crypto/tls/cipher_suites.go
|
||||
@@ -490,14 +492,40 @@ func (s *service) dialDevices(ctx context.Context, now time.Time, cfg config.Con
|
||||
// Perform dials according to the queue, stopping when we've reached the
|
||||
// allowed additional number of connections (if limited).
|
||||
numConns := 0
|
||||
for _, entry := range queue {
|
||||
if conn, ok := s.dialParallel(ctx, entry.id, entry.targets); ok {
|
||||
s.conns <- conn
|
||||
numConns++
|
||||
if allowAdditional > 0 && numConns >= allowAdditional {
|
||||
break
|
||||
}
|
||||
var numConnsMut stdsync.Mutex
|
||||
dialSemaphore := util.NewSemaphore(dialMaxParallel)
|
||||
dialWG := new(stdsync.WaitGroup)
|
||||
dialCtx, dialCancel := context.WithCancel(ctx)
|
||||
defer func() {
|
||||
dialWG.Wait()
|
||||
dialCancel()
|
||||
}()
|
||||
for i := range queue {
|
||||
select {
|
||||
case <-dialCtx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
dialWG.Add(1)
|
||||
go func(entry dialQueueEntry) {
|
||||
defer dialWG.Done()
|
||||
conn, ok := s.dialParallel(dialCtx, entry.id, entry.targets, dialSemaphore)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
numConnsMut.Lock()
|
||||
if allowAdditional == 0 || numConns < allowAdditional {
|
||||
select {
|
||||
case s.conns <- conn:
|
||||
numConns++
|
||||
if allowAdditional > 0 && numConns >= allowAdditional {
|
||||
dialCancel()
|
||||
}
|
||||
case <-dialCtx.Done():
|
||||
}
|
||||
}
|
||||
numConnsMut.Unlock()
|
||||
}(queue[i])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -959,7 +987,7 @@ func IsAllowedNetwork(host string, allowed []string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *service) dialParallel(ctx context.Context, deviceID protocol.DeviceID, dialTargets []dialTarget) (internalConn, bool) {
|
||||
func (s *service) dialParallel(ctx context.Context, deviceID protocol.DeviceID, dialTargets []dialTarget, parentSema *util.Semaphore) (internalConn, bool) {
|
||||
// Group targets into buckets by priority
|
||||
dialTargetBuckets := make(map[int][]dialTarget, len(dialTargets))
|
||||
for _, tgt := range dialTargets {
|
||||
@@ -975,13 +1003,19 @@ func (s *service) dialParallel(ctx context.Context, deviceID protocol.DeviceID,
|
||||
// Sort the priorities so that we dial lowest first (which means highest...)
|
||||
sort.Ints(priorities)
|
||||
|
||||
sema := util.MultiSemaphore{util.NewSemaphore(dialMaxParallelPerDevice), parentSema}
|
||||
for _, prio := range priorities {
|
||||
tgts := dialTargetBuckets[prio]
|
||||
res := make(chan internalConn, len(tgts))
|
||||
wg := stdsync.WaitGroup{}
|
||||
for _, tgt := range tgts {
|
||||
sema.Take(1)
|
||||
wg.Add(1)
|
||||
go func(tgt dialTarget) {
|
||||
defer func() {
|
||||
wg.Done()
|
||||
sema.Give(1)
|
||||
}()
|
||||
conn, err := tgt.Dial(ctx)
|
||||
if err == nil {
|
||||
// Closes the connection on error
|
||||
@@ -994,7 +1028,6 @@ func (s *service) dialParallel(ctx context.Context, deviceID protocol.DeviceID,
|
||||
l.Debugln("dialing", deviceID, tgt.uri, "success:", conn)
|
||||
res <- conn
|
||||
}
|
||||
wg.Done()
|
||||
}(tgt)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user