diff --git a/server/common/annotations/hosted_tiledb.py b/server/common/annotations/hosted_tiledb.py index 54f7dac2..235f8a16 100644 --- a/server/common/annotations/hosted_tiledb.py +++ b/server/common/annotations/hosted_tiledb.py @@ -34,6 +34,12 @@ class AnnotationsHostedTileDB(Annotations): f"{unsanitary_original_category_names} are not valid category names, please resubmit" ) + def get_user_name(self): + return current_app.auth.get_user_name() + + def get_user_id(self): + return current_app.auth.get_user_id() + def is_safe_collection_name(self, name): """ return true if this is a safe collection name @@ -48,7 +54,7 @@ class AnnotationsHostedTileDB(Annotations): self.CXG_ANNO_COLLECTION = name def read_labels(self, data_adaptor): - user_id = current_app.auth.get_user_id() + user_id = self.get_user_id() if user_id is None: return dataset_name = data_adaptor.get_location() @@ -103,8 +109,8 @@ class AnnotationsHostedTileDB(Annotations): return new_df def write_labels(self, df, data_adaptor): - auth_user_id = current_app.auth.get_user_id() - user_name = current_app.auth.get_user_name() + auth_user_id = self.get_user_id() + user_name = self.get_user_name() timestamp = time.time() dataset_location = data_adaptor.get_location() dataset_id = self.db.get_or_create_dataset(dataset_location) diff --git a/server/requirements.txt b/server/requirements.txt index 7aeec71a..537abb04 100644 --- a/server/requirements.txt +++ b/server/requirements.txt @@ -12,6 +12,7 @@ flatbuffers>=1.11.0 flatten-dict>=0.2.0 fsspec>=0.4.4,<0.8.0 gunicorn>=20.0.4 +h5py<3.0.0 # h5py returns bytes instead of str, which breaks many assumptions numba>=0.49.1 numpy>=1.15.0 packaging>=20.0 diff --git a/server/test/unit/auth/test_oauth.py b/server/test/unit/auth/test_oauth.py index be6ab850..caa13968 100644 --- a/server/test/unit/auth/test_oauth.py +++ b/server/test/unit/auth/test_oauth.py @@ -62,22 +62,44 @@ def jwks(): data = dict(alg="RS256", kty="RSA", use="sig", kid="fake_kid",) return make_response(jsonify(dict(keys=[data]))) -# The port that the mock oauth server will listen on -PORT = random.randint(10000, 12000) - # function to launch the mock oauth server -def launch_mock_oauth(): - mock_oauth_app.run(port=PORT) +def launch_mock_oauth(mock_port): + mock_oauth_app.run(port=mock_port) class AuthTest(unittest.TestCase): @classmethod def setUpClass(cls): + # The port that the mock oauth server will listen on + cls.mock_port = random.randint(10000, 12000) cls.dataset_dataroot = FIXTURES_ROOT - cls.mock_oauth_process = Process(target=launch_mock_oauth) + cls.mock_oauth_process = Process(target=launch_mock_oauth, args=(cls.mock_port,)) cls.mock_oauth_process.start() + # Verify that the mock oauth server is ready (accepting requests) before starting the tests. + + # The following lines are polling until the mock server is ready. + # The issue is we are starting a mock oauth server, then we are starting a cellxgene server, + # which will start making requests to the mock oauth server. + # So there is a race condition because the mock oauth server needs to be ready before it gets requests. + # We check to see if it is ready, and if not we wait 1 second, then try again. + # If it gets to 5 seconds, which is shouldn't, we assume something has gone wrong and fail the test. + server_okay = False + for _ in range(5): + try: + response = requests.get(f"http://localhost:{cls.mock_port}/.well-known/jwks.json") + if response.status_code == 200: + server_okay = True + break + except: # noqa: E722 + pass + + # wait one second and try again + time.sleep(1) + + assert(server_okay) + @classmethod def tearDownClass(cls): cls.mock_oauth_process.terminate() @@ -87,7 +109,7 @@ class AuthTest(unittest.TestCase): app_config.update_server_config( app__api_base_url="local", authentication__type="oauth", - authentication__params_oauth__oauth_api_base_url=f"http://localhost:{PORT}", + authentication__params_oauth__oauth_api_base_url=f"http://localhost:{self.mock_port}", authentication__params_oauth__client_id="mock_client_id", authentication__params_oauth__client_secret="mock_client_secret", authentication__params_oauth__jwt_decode_options={"verify_signature": False, "verify_iss": False}, diff --git a/server/test/unit/common/test_writable_annotation.py b/server/test/unit/common/test_writable_annotation.py index 0e369e56..fda910aa 100644 --- a/server/test/unit/common/test_writable_annotation.py +++ b/server/test/unit/common/test_writable_annotation.py @@ -129,9 +129,11 @@ class WritableTileDBStoredAnnotationTest(unittest.TestCase): 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" + @patch("server.common.annotations.hosted_tiledb.AnnotationsHostedTileDB.get_user_id") + @patch("server.common.annotations.hosted_tiledb.AnnotationsHostedTileDB.get_user_name") + def test_write_labels_stores_df_as_tiledb_array(self, mock_user_name, mock_user_id): + mock_user_id.return_value = "1234" + mock_user_name.return_value = "user1234" self.annotations.write_labels(self.df, self.data) # get uri dataset_id = self.db.query([CellxGeneDataset], [CellxGeneDataset.name == self.data.get_location()])[0].id