rewrite enumerate children async to be less fragile
License: MIT Signed-off-by: Jeromy <why@ipfs.io>
Jeromy committed
Jan 7, 2017 at 05:46 UTC
8fc425787694c48c0a9542e7814afa5cb0f7b5d2
2 files changed
+109
-77
merkledag/merkledag.go
+66
-77
@@ -389,103 +389,92 @@ func EnumerateChildren(ctx context.Context, ds LinkService, root *cid.Cid, visit
389
return nil
390
}
391
392
-func EnumerateChildrenAsync(ctx context.Context, ds DAGService, c *cid.Cid, visit func(*cid.Cid) bool) error {
393
- toprocess := make(chan []*cid.Cid, 8)
394
- nodes := make(chan *NodeOption, 8)
395
-
396
- ctx, cancel := context.WithCancel(ctx)
397
- defer cancel()
398
- defer close(toprocess)
392
+// FetchGraphConcurrency is total number of concurrent fetches that
393
+// 'fetchNodes' will start at a time
394
+var FetchGraphConcurrency = 8
395
400
- go fetchNodes(ctx, ds, toprocess, nodes)
396
+func EnumerateChildrenAsync(ctx context.Context, ds DAGService, c *cid.Cid, visit func(*cid.Cid) bool) error {
397
+ if !visit(c) {
398
+ return nil
399
+ }
400
401
root, err := ds.Get(ctx, c)
402
if err != nil {
403
return err
404
}
405
407
- nodes <- &NodeOption{Node: root}
408
- live := 1
409
-
410
- for {
411
- select {
412
- case opt, ok := <-nodes:
413
- if !ok {
414
- return nil
415
- }
416
-
417
- if opt.Err != nil {
418
- return opt.Err
419
- }
420
-
421
- nd := opt.Node
422
-
423
- // a node has been fetched
424
- live--
425
-
426
- var cids []*cid.Cid
427
- for _, lnk := range nd.Links() {
428
- c := lnk.Cid
429
- if visit(c) {
430
- live++
431
- cids = append(cids, c)
406
+ feed := make(chan node.Node)
407
+ out := make(chan *NodeOption)
408
+ done := make(chan struct{})
409
+
410
+ var setlk sync.Mutex
411
+
412
+ for i := 0; i < FetchGraphConcurrency; i++ {
413
+ go func() {
414
+ for n := range feed {
415
+ links := n.Links()
416
+ cids := make([]*cid.Cid, 0, len(links))
417
+ for _, l := range links {
418
+ setlk.Lock()
419
+ unseen := visit(l.Cid)
420
+ setlk.Unlock()
421
+ if unseen {
422
+ cids = append(cids, l.Cid)
423
+ }
424
}
433
- }
434
-
435
- if live == 0 {
436
- return nil
437
- }
425
439
- if len(cids) > 0 {
426
+ for nopt := range ds.GetMany(ctx, cids) {
427
+ select {
428
+ case out <- nopt:
429
+ case <-ctx.Done():
430
+ return
431
+ }
432
+ }
433
select {
441
- case toprocess <- cids:
434
+ case done <- struct{}{}:
435
case <-ctx.Done():
443
- return ctx.Err()
436
}
437
}
446
- case <-ctx.Done():
447
- return ctx.Err()
448
- }
438
+ }()
439
}
450
-}
440
+ defer close(feed)
441
452
-// FetchGraphConcurrency is total number of concurrenct fetches that
453
-// 'fetchNodes' will start at a time
454
-var FetchGraphConcurrency = 8
455
-
456
-func fetchNodes(ctx context.Context, ds DAGService, in <-chan []*cid.Cid, out chan<- *NodeOption) {
457
- var wg sync.WaitGroup
458
- defer func() {
459
- // wait for all 'get' calls to complete so we don't accidentally send
460
- // on a closed channel
461
- wg.Wait()
462
- close(out)
463
- }()
442
+ send := feed
443
+ var todobuffer []node.Node
444
+ var inProgress int
445
465
- rateLimit := make(chan struct{}, FetchGraphConcurrency)
446
+ next := root
447
+ for {
448
+ select {
449
+ case send <- next:
450
+ inProgress++
451
+ if len(todobuffer) > 0 {
452
+ next = todobuffer[0]
453
+ todobuffer = todobuffer[1:]
454
+ } else {
455
+ next = nil
456
+ send = nil
457
+ }
458
+ case <-done:
459
+ inProgress--
460
+ if inProgress == 0 && next == nil {
461
+ return nil
462
+ }
463
+ case nc := <-out:
464
+ if nc.Err != nil {
465
+ return nc.Err
466
+ }
467
467
- get := func(ks []*cid.Cid) {
468
- defer wg.Done()
469
- defer func() {
470
- <-rateLimit
471
- }()
472
- nodes := ds.GetMany(ctx, ks)
473
- for opt := range nodes {
474
- select {
475
- case out <- opt:
476
- case <-ctx.Done():
477
- return
468
+ if next == nil {
469
+ next = nc.Node
470
+ send = feed
471
+ } else {
472
+ todobuffer = append(todobuffer, nc.Node)
473
}
479
- }
480
- }
474
482
- for ks := range in {
483
- select {
484
- case rateLimit <- struct{}{}:
475
case <-ctx.Done():
486
- return
476
+ return ctx.Err()
477
}
488
- wg.Add(1)
489
- go get(ks)
478
}
479
+
480
}
merkledag/merkledag_test.go
+43
@@ -504,3 +504,46 @@ func TestCidRawDoesnNeedData(t *testing.T) {
504
t.Fatal("raw node shouldn't have any links")
505
}
506
}
507
+
508
+func TestEnumerateAsyncFailsNotFound(t *testing.T) {
509
+ a := NodeWithData([]byte("foo1"))
510
+ b := NodeWithData([]byte("foo2"))
511
+ c := NodeWithData([]byte("foo3"))
512
+ d := NodeWithData([]byte("foo4"))
513
+
514
+ ds := dstest.Mock()
515
+ for _, n := range []node.Node{a, b, c} {
516
+ _, err := ds.Add(n)
517
+ if err != nil {
518
+ t.Fatal(err)
519
+ }
520
+ }
521
+
522
+ parent := new(ProtoNode)
523
+ if err := parent.AddNodeLinkClean("a", a); err != nil {
524
+ t.Fatal(err)
525
+ }
526
+
527
+ if err := parent.AddNodeLinkClean("b", b); err != nil {
528
+ t.Fatal(err)
529
+ }
530
+
531
+ if err := parent.AddNodeLinkClean("c", c); err != nil {
532
+ t.Fatal(err)
533
+ }
534
+
535
+ if err := parent.AddNodeLinkClean("d", d); err != nil {
536
+ t.Fatal(err)
537
+ }
538
+
539
+ pcid, err := ds.Add(parent)
540
+ if err != nil {
541
+ t.Fatal(err)
542
+ }
543
+
544
+ cset := cid.NewSet()
545
+ err = EnumerateChildrenAsync(context.Background(), ds, pcid, cset.Visit)
546
+ if err == nil {
547
+ t.Fatal("this should have failed")
548
+ }
549
+}