diff --git a/api.go b/api.go index 44ea0a7b0..910bc582f 100644 --- a/api.go +++ b/api.go @@ -321,13 +321,19 @@ func (api *API) ImportRoaring(ctx context.Context, indexName, fieldName string, nodes := api.cluster.shardNodes(indexName, shard) var eg errgroup.Group + field := api.holder.Field(indexName, fieldName) + if field == nil { + return newNotFoundError(ErrFieldNotFound) + } + + // only set fields are supported + if field.Type() != FieldTypeSet { + return NewBadRequestError(errors.New("roaring import is only supported for set fields")) + } + for _, node := range nodes { node := node if node.ID == api.server.nodeID { - field := api.holder.Field(indexName, fieldName) - if field == nil { - return newNotFoundError(ErrFieldNotFound) - } // must make a copy of data to operate on locally. field.importRoaring changes data d2 := make([]byte, len(data)) copy(d2, data) diff --git a/http/handler.go b/http/handler.go index 47f5cf374..2980da12b 100644 --- a/http/handler.go +++ b/http/handler.go @@ -1494,11 +1494,16 @@ func (h *Handler) handlePostImportRoaring(w http.ResponseWriter, r *http.Request return } + resp := &pilosa.ImportResponse{} // TODO give meaningful stats for import err = h.api.ImportRoaring(r.Context(), urlVars["index"], urlVars["field"], shard, remote, body) - resp := &pilosa.ImportResponse{} if err != nil { resp.Err = err.Error() + if _, ok := err.(pilosa.BadRequestError); ok { + w.WriteHeader(http.StatusBadRequest) + } else { + w.WriteHeader(http.StatusInternalServerError) + } } // Marshal response object. buf, err := h.api.Serializer.Marshal(resp) diff --git a/server/handler_test.go b/server/handler_test.go index 2c1aa9868..9e76d77b9 100644 --- a/server/handler_test.go +++ b/server/handler_test.go @@ -104,6 +104,22 @@ func TestHandler_Endpoints(t *testing.T) { }) + t.Run("ImportRoaringFieldTypeFail", func(t *testing.T) { + // Roaring import into a non-set field should fail. + if _, err := i0.CreateFieldIfNotExists("int-field", pilosa.OptFieldTypeInt(0, 1)); err != nil { + t.Fatal(err) + } + w := httptest.NewRecorder() + roaringData, _ := hex.DecodeString("3B3001000100000900010000000100010009000100") + req := test.MustNewHTTPRequest("POST", "/index/i0/field/int-field/import-roaring/0", bytes.NewBuffer(roaringData)) + req.Header.Set("Content-Type", "application/x-binary") + h.ServeHTTP(w, req) + if w.Code != gohttp.StatusBadRequest { + t.Fatalf("unexpected status code: %d", w.Code) + } + + }) + t.Run("Status", func(t *testing.T) { w := httptest.NewRecorder() h.ServeHTTP(w, test.MustNewHTTPRequest("GET", "/status", nil))