diff --git a/ctl/broadcast.go b/ctl/broadcast.go new file mode 100644 index 000000000..9ef644356 --- /dev/null +++ b/ctl/broadcast.go @@ -0,0 +1,70 @@ +package ctl + +import ( + "io" +) + +type Broadcaster struct { + source io.Reader + Readers []io.Reader +} + +func (b *Broadcaster) Consume() { + if len(b.Readers) == 1 { + return + } + buf := make([]byte, 65535) + for { + n, err := b.source.Read(buf) + if err == io.EOF { + // close and shutdown + b.shutdown() + return + } + buf = buf[:n] + for _, s := range b.Readers { + subscriber := s.(*Subscriber) + // possibly a timeout channel here + subscriber.src <- buf + } + + } +} + +func (b *Broadcaster) shutdown() { + for _, s := range b.Readers { + subscriber := s.(*Subscriber) + close(subscriber.src) + } +} + +type Subscriber struct { + src chan []byte + local []byte + i int +} + +func (s *Subscriber) Read(b []byte) (n int, err error) { + if len(s.local) == 0 { + s.local = <-s.src + } + if len(s.local) == 0 { + return 0, io.EOF + } + n = copy(b, s.local) + s.local = s.local[n:] + return n, nil +} + +func NewBroadcaster(reader io.Reader, numSubscribers int) *Broadcaster { + b := &Broadcaster{Readers: make([]io.Reader, numSubscribers)} + if numSubscribers == 1 { + b.Readers[0] = reader + return b + } + for i := 0; i < len(b.Readers); i++ { + b.Readers[i] = &Subscriber{src: make(chan []byte)} + } + b.source = reader + return b +} diff --git a/ctl/restore_tar.go b/ctl/restore_tar.go index 80bf9c5e7..7432a3b75 100644 --- a/ctl/restore_tar.go +++ b/ctl/restore_tar.go @@ -269,25 +269,20 @@ func (cmd *RestoreTarCommand) Run(ctx context.Context) (err error) { switch action := record[4]; action { case "translate": logger.Printf("field keys %v %v", indexName, fieldName) - // needs to go to all nodes - - _, err = io.Copy(buf, tarReader) - if err != nil { - return errors.Wrap(err, "copying") - } - + bc := NewBroadcaster(tarReader, len(nodes)) g, _ := errgroup.WithContext(ctx) - for _, node := range nodes { + for i, node := range nodes { + i := i node := node g.Go(func() error { - // rd := bytes.NewReader(shardBytes) rd := func() (io.Reader, error) { - return bytes.NewReader(buf.Bytes()), nil + return bc.Readers[i], nil } return cmd.client.ImportFieldKeys(ctx, &node.URI, indexName, fieldName, false, rd) }) } + bc.Consume() if err := g.Wait(); err != nil { return err }