mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-01 06:28:12 +08:00
Create hosted user annotations [1685] (#1726)
* add function to retrieve latest annotation from db, db updates * read and write tiledb arrays * adding tests
This commit is contained in:
@@ -1,6 +1,11 @@
|
||||
import json
|
||||
from os import path, listdir
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import tiledb
|
||||
from flask import Flask
|
||||
|
||||
import server.test.unit.decode_fbs as decode_fbs
|
||||
import shutil
|
||||
|
||||
@@ -8,8 +13,134 @@ import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from server.common.rest import schema_get_helper, annotations_put_fbs_helper
|
||||
from server.test import data_with_tmp_annotations, make_fbs
|
||||
from server.db.cellxgene_orm import CellxGeneDataset, Annotation
|
||||
from server.test import data_with_tmp_annotations, make_fbs, data_with_tmp_tiledb_annotations
|
||||
from server.data_common.matrix_loader import MatrixDataType
|
||||
from server.common.errors import AnnotationCategoryNameError
|
||||
|
||||
|
||||
class auth(object):
|
||||
def get_user_id():
|
||||
return "1234"
|
||||
|
||||
|
||||
class WritableTileDBStoredAnnotationTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.user_id = '1234'
|
||||
self.data, self.tmp_dir, self.annotations = data_with_tmp_tiledb_annotations(MatrixDataType.H5AD)
|
||||
self.data.dataset_config.user_annotations = self.annotations
|
||||
self.db = self.annotations.db
|
||||
self.n_rows = self.data.get_shape()[0]
|
||||
self.test_dict = {
|
||||
"cat_A": pd.Series(["label_A"] * self.n_rows, dtype="category"),
|
||||
"cat_B": pd.Series(["label_B"] * self.n_rows, dtype="category"),
|
||||
}
|
||||
self.fbs = make_fbs(self.test_dict)
|
||||
self.df = pd.DataFrame(self.test_dict)
|
||||
self.app = Flask('fake_app')
|
||||
self.app.__setattr__("auth", auth)
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.tmp_dir)
|
||||
|
||||
def annotation_put_fbs(self, fbs):
|
||||
annotations_put_fbs_helper(self.data, fbs)
|
||||
res = json.dumps({"status": "OK"})
|
||||
return res
|
||||
|
||||
def test_category_name_throws_errors_for_categories_that_cant_be_converted_to_filenames(self):
|
||||
with self.app.test_request_context():
|
||||
bad_category_names = make_fbs(
|
||||
{
|
||||
"cat_A": pd.Series(["label_A"] * self.n_rows, dtype="category"),
|
||||
"cat/B": pd.Series(["label_B"] * self.n_rows, dtype="category"),
|
||||
}
|
||||
)
|
||||
with self.assertRaises(AnnotationCategoryNameError):
|
||||
self.annotation_put_fbs(bad_category_names)
|
||||
|
||||
def test_convert_to_pandas__converts_tiledb_to_pandas_df(self):
|
||||
with self.app.test_request_context():
|
||||
self.annotations.write_labels(self.df, self.data)
|
||||
dataset_id = self.db.query([CellxGeneDataset], [CellxGeneDataset.name == self.data.get_location()])[0].id
|
||||
annotation = self.db.query_for_most_recent(
|
||||
Annotation,
|
||||
[Annotation.user_id == self.user_id, Annotation.dataset_id == str(dataset_id)]
|
||||
)
|
||||
# retrieve tiledb array
|
||||
df = tiledb.open(annotation.tiledb_uri)
|
||||
self.assertEqual(type(df), tiledb.array.SparseArray)
|
||||
|
||||
# convert to pandas df
|
||||
pandas_df = self.annotations.convert_to_pandas_df(df)
|
||||
self.assertEqual(type(pandas_df), pd.DataFrame)
|
||||
|
||||
def test_write_labels_creates_a_dataset_if_it_doesnt_exist(self):
|
||||
with self.app.test_request_context():
|
||||
|
||||
new_name = 'new_dataset/location'
|
||||
self.data.get_location = MagicMock(return_value=new_name)
|
||||
num_datasets = len(self.db.query([CellxGeneDataset]))
|
||||
self.annotation_put_fbs(self.fbs)
|
||||
more_datasets = len(self.db.query([CellxGeneDataset]))
|
||||
self.assertGreater(more_datasets, num_datasets)
|
||||
|
||||
self.assertGreater(len(self.db.query([CellxGeneDataset], [CellxGeneDataset.name == new_name])), 0)
|
||||
|
||||
def test_write_labels_links_to_existing_dataset(self):
|
||||
with self.app.test_request_context():
|
||||
# add dataset to to db
|
||||
self.annotation_put_fbs(self.fbs)
|
||||
|
||||
num_datasets = len(self.db.query([CellxGeneDataset]))
|
||||
|
||||
# create another annotation with the same dataset
|
||||
self.annotation_put_fbs(self.fbs)
|
||||
|
||||
same_num_datasets = len(self.db.query([CellxGeneDataset]))
|
||||
|
||||
self.assertEqual(num_datasets, same_num_datasets)
|
||||
|
||||
def test_read_labels_returns_pandas_df(self):
|
||||
with self.app.test_request_context():
|
||||
self.annotation_put_fbs(self.fbs)
|
||||
pandas_df = self.annotations.read_labels(self.data)
|
||||
self.assertEqual(type(pandas_df), pd.DataFrame)
|
||||
|
||||
def test_read_labels_returns_df_matching_original(self):
|
||||
with self.app.test_request_context():
|
||||
self.annotation_put_fbs(self.fbs)
|
||||
pandas_df = self.annotations.read_labels(self.data)
|
||||
|
||||
self.assertEqual(pandas_df.shape, (self.n_rows, 2))
|
||||
self.assertEqual(set(pandas_df.columns), {"cat_A", "cat_B"})
|
||||
self.assertTrue(self.data.original_obs_index.equals(pandas_df.index))
|
||||
self.assertTrue(np.all(pandas_df["cat_A"] == ["label_A"] * self.n_rows))
|
||||
self.assertTrue(np.all(pandas_df["cat_B"] == ["label_B"] * self.n_rows))
|
||||
|
||||
def test_error_checks(self):
|
||||
# verify that the expected errors are generated
|
||||
with self.app.test_request_context():
|
||||
n_rows = self.data.get_shape()[0]
|
||||
fbs_bad = make_fbs({"louvain": pd.Series(["undefined"] * n_rows, dtype="category")})
|
||||
|
||||
# ensure we catch attempt to overwrite non-writable data
|
||||
with self.assertRaises(KeyError):
|
||||
self.annotation_put_fbs(fbs_bad)
|
||||
|
||||
@patch('server.common.annotations.hosted_tiledb.current_app')
|
||||
def test_write_labels_stores_df_as_tiledb_array(self, mock_user_id):
|
||||
mock_user_id.auth.get_user_id.return_value = '1234'
|
||||
self.annotations.write_labels(self.df, self.data)
|
||||
# get uri
|
||||
dataset_id = self.db.query([CellxGeneDataset], [CellxGeneDataset.name == self.data.get_location()])[0].id
|
||||
annotation = self.db.query_for_most_recent(
|
||||
Annotation,
|
||||
[Annotation.user_id == '1234', Annotation.dataset_id == str(dataset_id)]
|
||||
)
|
||||
|
||||
df = tiledb.open(annotation.tiledb_uri)
|
||||
self.assertEqual(type(df), tiledb.array.SparseArray)
|
||||
|
||||
|
||||
class WritableAnnotationTest(unittest.TestCase):
|
||||
|
||||
Reference in New Issue
Block a user