@cryptotaxi247 / kubo / commits / 8fc425787

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 +}