diff --git a/api.go b/api.go index 9f4190535..dfbd33e01 100644 --- a/api.go +++ b/api.go @@ -378,37 +378,42 @@ func importWorker(importWork chan importJob) { } } - tx, finisher := j.qcx.GetTx(Txo{Write: writable, Index: j.field.idx, Shard: j.shard}) - defer finisher(&err0) + if err := func() (err1 error) { + tx, finisher := j.qcx.GetTx(Txo{Write: writable, Index: j.field.idx, Shard: j.shard}) + defer finisher(&err1) - var doClear bool - switch doAction { - case RequestActionOverwrite: - err := j.field.importRoaringOverwrite(j.ctx, tx, viewData, j.shard, viewName, j.req.Block) - if err != nil { - return errors.Wrap(err, "importing roaring as overwrite") - } - case RequestActionClear: - doClear = true - fallthrough - case RequestActionSet: - fileMagic := uint32(binary.LittleEndian.Uint16(viewData[0:2])) - if fileMagic == roaring.MagicNumber { // if pilosa roaring format - err := j.field.importRoaring(j.ctx, tx, viewData, j.shard, viewName, doClear) + var doClear bool + switch doAction { + case RequestActionOverwrite: + err := j.field.importRoaringOverwrite(j.ctx, tx, viewData, j.shard, viewName, j.req.Block) if err != nil { - return errors.Wrap(err, "importing pilosa roaring") + return errors.Wrap(err, "importing roaring as overwrite") } - } else { - // must make a copy of data to operate on locally on standard roaring format. - // field.importRoaring changes the standard roaring run format to pilosa roaring - data := make([]byte, len(viewData)) - copy(data, viewData) - err := j.field.importRoaring(j.ctx, tx, data, j.shard, viewName, doClear) + case RequestActionClear: + doClear = true + fallthrough + case RequestActionSet: + fileMagic := uint32(binary.LittleEndian.Uint16(viewData[0:2])) + if fileMagic == roaring.MagicNumber { // if pilosa roaring format + err := j.field.importRoaring(j.ctx, tx, viewData, j.shard, viewName, doClear) + if err != nil { + return errors.Wrap(err, "importing pilosa roaring") + } + } else { + // must make a copy of data to operate on locally on standard roaring format. + // field.importRoaring changes the standard roaring run format to pilosa roaring + data := make([]byte, len(viewData)) + copy(data, viewData) + err := j.field.importRoaring(j.ctx, tx, data, j.shard, viewName, doClear) - if err != nil { - return errors.Wrap(err, "importing standard roaring") + if err != nil { + return errors.Wrap(err, "importing standard roaring") + } } } + return nil + }(); err != nil { + return err } } return nil diff --git a/http/client_test.go b/http/client_test.go index cfe51bde4..b1b154647 100644 --- a/http/client_test.go +++ b/http/client_test.go @@ -588,6 +588,42 @@ func TestClient_ImportRoaring(t *testing.T) { } } +// Ensure client can bulk import data with multiple views and not deadlock. +func TestClient_ImportRoaring_MultiView(t *testing.T) { + cluster := test.MustNewCluster(t, 2) + for _, c := range cluster.Nodes { + c.Config.Cluster.ReplicaN = 2 + } + err := cluster.Start() + if err != nil { + t.Fatalf("starting cluster: %v", err) + } + defer cluster.Close() + + _, err = cluster.GetNode(0).API.CreateIndex(context.Background(), "i", pilosa.IndexOptions{}) + if err != nil { + t.Fatalf("creating index: %v", err) + } + _, err = cluster.GetNode(0).API.CreateField(context.Background(), "i", "f", pilosa.OptFieldTypeSet(pilosa.CacheTypeRanked, 100)) + if err != nil { + t.Fatalf("creating field: %v", err) + } + _, err = cluster.GetNode(0).API.Query(context.Background(), &pilosa.QueryRequest{Index: "i", Query: "Set(0, f=1)"}) + if err != nil { + t.Fatalf("querying: %v", err) + } + + // Send import request. + host := cluster.GetNode(0).URL() + c := MustNewClient(host, http.GetHTTPClient(nil)) + req := &pilosa.ImportRoaringRequest{Views: map[string][]byte{}} + req.Views["a"], _ = hex.DecodeString("3B3001000100000900010000000100010009000100") + req.Views["b"], _ = hex.DecodeString("3B3001000100000900010000000100010009000100") + if err := c.ImportRoaring(context.Background(), &cluster.GetNode(0).API.Node().URI, "i", "f", 0, false, req); err != nil { + t.Fatal(err) + } +} + // Ensure client can bulk import data. func TestClient_ImportKeys(t *testing.T) { t.Run("SingleNode", func(t *testing.T) {