Compare commits

..
Author SHA1 Message Date
Colin Megill 87ba3ff870 Merge branch 'main' into colinmegill/#632 2021-08-16 12:14:34 -07:00
Colin Megill b553da0264 async 1 2021-08-05 16:39:22 -07:00
Colin Megill b67142e98f embedding to tsx 2021-08-04 16:02:52 -07:00
Colin Megill 22a0921147 Merge branch 'main' into colinmegill/#632 2021-08-04 14:43:09 -07:00
Colin Megill 020e562f5c merge typescript changes 2021-08-04 14:42:31 -07:00
Colin Megill 60d89b9478 flip scale 2021-07-30 10:59:42 -07:00
Colin Megill c2b12abe2b toggle and scale dotplot 2021-07-26 13:23:21 -07:00
Colin Megill 2462d4afb1 color scale, button 2021-07-23 13:29:27 -07:00
Colin Megill 02e79d502f add classnames to canvases 2021-07-23 13:29:27 -07:00
Colin Megill 9c2323b7bc queries 2021-07-23 13:29:27 -07:00
Colin Megill e7200ce6c3 todo comment 2021-07-23 13:29:27 -07:00
Colin Megill b7ffb2748d colorby type 2021-07-23 13:29:27 -07:00
Colin Megill 7ab8a8894d no overlap 2021-07-23 13:29:27 -07:00
Colin Megill 7b01e9e67b remove hardcoded geneset 2021-07-23 13:29:27 -07:00
Colin Megill 72ee670620 remove logs 2021-07-23 13:29:27 -07:00
Colin Megill e21cac65bf set row and column 2021-07-23 13:29:27 -07:00
Colin Megill bb5bbaac8a reducer 2021-07-23 13:29:27 -07:00
Colin Megill e29a6f72c2 d3 scale for dot size 2021-07-23 13:29:10 -07:00
Colin Megill ed97013277 dotplot button 2021-07-23 13:29:10 -07:00
Colin Megill 2b29a152b9 metadata as var, maxsize todo 2021-07-23 13:28:17 -07:00
Colin Megill f0e9b1ab91 dotplot proto full 2021-07-23 13:28:17 -07:00
Colin Megill c489221296 geneset iterate 2021-07-23 13:28:17 -07:00
Colin Megill 0a69af98c5 break out load and err 2021-07-23 13:28:17 -07:00
Colin Megill 11570273e0 dotplot 1 2021-07-23 13:28:17 -07:00
Colin Megill 5f9d0a6b34 logging out values 2021-07-23 13:28:17 -07:00
750 changed files with 51865 additions and 15479 deletions
+2 -2
View File
@@ -1,5 +1,5 @@
[bumpversion]
current_version = 1.0.0
current_version = 0.17.0
commit = True
parse = (?P<major>\d+)\.(?P<minor>\d+)\.(?P<patch>\d+)(?:-(?P<prerel>rc)\.(?P<prerelversion>\d+))?
serialize =
@@ -20,6 +20,6 @@ replace = version="{new_version}"
search = "version": "{current_version}"
replace = "version": "{new_version}"
[bumpversion:file:server/__init__.py]
[bumpversion:file:backend/server/__init__.py]
search = __version__ = "{current_version}"
replace = __version__ = "{new_version}"
+1 -1
View File
@@ -2,4 +2,4 @@ bin
client
dist
docs
server
backend
+13
View File
@@ -0,0 +1,13 @@
name: Deploy canary via single cell infra repo
on:
push:
branches: main-canary
jobs:
deploy:
runs-on: ubuntu-latest
steps:
- name: repository dispatch
run: |
curl -XPOST -u czi-sci-single-cell-eng:${{secrets.SCI_GITHUB_TOKEN}} -H "Accept: application/vnd.github.everest-preview+json" -H "Content-Type: application/json" https://api.github.com/repos/chanzuckerberg/single-cell-infra/dispatches --data '{"event_type": "canary-hook"}'
+73 -89
View File
@@ -9,6 +9,7 @@ on:
env:
JEST_ENV: prod
CXG_AUTH_TYPE: none
jobs:
docker-build:
@@ -22,102 +23,85 @@ jobs:
- name: Build docker image
run: docker build .
matrix-compatibility-test:
name: cxg:${{ matrix.cellxgene_build }} os:${{ matrix.os }} py:${{ matrix.python-version }} anndata:${{ matrix.anndata_version || 'latest' }}
runs-on: ${{ matrix.os }}
cellxgene-main-with-python-and-anndata-versions:
name: python versions x anndata versions
runs-on: ubuntu-latest
continue-on-error: true
strategy:
fail-fast: false
matrix:
# note: The `macos-latest` is latest Catalina version, and not Big Sur. So we explicitly ask for Big Sur (`macos-11`)
os: [ubuntu-latest, macos-latest, macos-11]
python-version: [3.6, 3.7, 3.8, 3.9]
cellxgene_build: [main, latest]
exclude:
# 3.6 no longer avail on Big Sur (`macos-11`)
- os: macos-11
python-version: 3.6
# no pypi build exists for macos+py3.9 and source install fails to
# install `tables` py pkg (a `scanpy` dependency), so we test py3.9
# only on ubuntu
- os: macos-11
python-version: 3.9
- os: macos-latest
python-version: 3.9
# add anndata pinned version test for subset of matrix configurations,
# in order to reduce matrix cross-product explosion
include:
- python-version: 3.8
cellxgene_build: latest
# TODO: dynamically use the literal version in requirements.txt,
# to avoid having to update this in manually in the future
# TODO: Do not bother running this if anndata latest version
# matches this pinned version, to avoid a redundant test
anndata_version: '==0.7.6'
python-version: [3.6, 3.7, 3.8]
anndata-version: [0.7.6]
test-suite: [smoke-test, smoke-test-annotations]
steps:
- uses: actions/checkout@v2
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v1
with:
python-version: ${{ matrix.python-version }}
- name: Cache env vars
run: echo "PIP_CACHE=`python -m pip cache dir`" >> $GITHUB_ENV
- name: Cache env vars (MacOS)
if: startsWith(matrix.os, 'macos')
run: echo "BREW_CACHE=`brew --cache`" >> $GITHUB_ENV
# FIXME: Only working for Linux
- name: Python cache
uses: actions/cache@v1
with:
path: ${{ env.PIP_CACHE }}
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements*.txt') }}
restore-keys: |
${{ runner.os }}-pip-
- name: Node cache
uses: actions/cache@v1
with:
path: ~/.npm
key: ${{ runner.os }}-node-${{ hashFiles('**/package-lock.json') }}
restore-keys: |
${{ runner.os }}-node-
- name: Brew cache (MacOS)
if: startsWith(matrix.os, 'macos')
uses: actions/cache@v1
with:
path: ${{ env.BREW_CACHE }}
key: ${{ runner.os }}-brew-
- name: Install dependencies (Ubuntu Linux)
if: startsWith(matrix.os, 'ubuntu')
- name: Install dependencies
run: |
sudo apt-get update
sudo apt-get install -y libhdf5-serial-dev
- name: Install dependencies (MacOS)
if: startsWith(matrix.os, 'macos')
run: brew install hdf5
- name: Install cellxgene from `main` branch
if: matrix.cellxgene_build == 'main'
run: |
pip install -r server/requirements-dev.txt
# 1. only install the dev requirements on top of what is in the cellxgene pip package
sudo apt-get update && sudo apt-get install -y libhdf5-serial-dev
sed -i 's/-r requirements.txt//' backend/server/requirements-dev.txt
pip install -r backend/server/requirements-dev.txt
# 2. install cellxgene
make pydist install-dist
- name: Install cellxgene from latest release (pypi.org)
if: matrix.cellxgene_build == 'latest'
run: |
pip install --upgrade cellxgene
# install the additional dev requirements on top of what is in the
# cellxgene pip package, which are needed for testing, but otherwise
# keep same pip pkg versions as in the cxg release
sed -i'' -e 's/-r requirements.txt//' server/requirements-dev.txt
pip install -r server/requirements-dev.txt
- name: Install anndata version per matrix variable
run: pip install anndata${{ matrix.anndata_version }}
- name: Install node
run: make dev-env-client
# Run different types of test separately, to facilitate troubleshooting
- name: Unit Tests - client
run: make unit-test-client
- name: Unit Tests - server
run: make unit-test-server
- name: Smoke Tests
run: make smoke-test
# FIXME: Fails intermittently. See https://app.zenhub.com/workspaces/single-cell-5e2a191dad828d52cc78b028/issues/chanzuckerberg/cellxgene/2415
# - name: Smoke Tests with Annotations
# run: make smoke-test-annotations
# 3. install anndata
pip install anndata==${{ matrix.anndata-version }}
- name: Tests
run: make unit-test ${{ matrix.test-suite }}
cellxgene-release-with-anndata-master:
name: cellxgene release with anndata master
runs-on: ubuntu-latest
strategy:
matrix:
test-suite: [smoke-test, smoke-test-annotations]
steps:
- uses: actions/checkout@v2
- name: Set up Python 3.7
uses: actions/setup-python@v1
with:
python-version: 3.7
- name: Checkout
uses: actions/checkout@v2
with:
path: cellxgene
- name: Install dependencies
run: |
cd cellxgene
# 1. only install the dev requirements on top of what is in the cellxgene pip package
make dev-env-client
sed -i 's/-r requirements.txt//' backend/server/requirements-dev.txt
pip install -r backend/server/requirements-dev.txt
# 2. install cellxgene
pip install --upgrade cellxgene
# 3. install anndata
pip install git+https://github.com/theislab/anndata
- name: Tests
run: cd cellxgene && make unit-test ${{ matrix.test-suite }}
cellxgene-main-with-anndata-master:
name: cellxgene main with anndata master
runs-on: ubuntu-latest
strategy:
matrix:
test-suite: [smoke-test, smoke-test-annotations]
steps:
- uses: actions/checkout@v2
- name: Set up Python 3.7
uses: actions/setup-python@v1
with:
python-version: 3.7
- name: Checkout
uses: actions/checkout@v2
with:
path: cellxgene
- name: Install dependencies
run: |
cd cellxgene
sed -i -E 's/^anndata[>=]=[0-9]+.[0-9]+.[0-9]+$/anndata/g' backend/server/requirements.txt
make pydist install-dist dev-env
pip install git+https://github.com/theislab/anndata
- name: Tests
run: cd cellxgene && make unit-test ${{ matrix.test-suite }}
+13
View File
@@ -0,0 +1,13 @@
name: Deploy via single cell infra repo
on:
push:
branches: main
jobs:
deploy:
runs-on: ubuntu-latest
steps:
- name: repository dispatch
run: |
curl -XPOST -u czi-sci-single-cell-eng:${{secrets.SCI_GITHUB_TOKEN}} -H "Accept: application/vnd.github.everest-preview+json" -H "Content-Type: application/json" https://api.github.com/repos/chanzuckerberg/single-cell-infra/dispatches --data '{"event_type": "cellxgene-hook"}'
+34 -4
View File
@@ -36,7 +36,7 @@ jobs:
npm install
- name: Format with black and lint with flake8
run: |
make lint-server
make lint-servers
- name: Lint src with eslint
working-directory: ./client
run: |
@@ -68,8 +68,38 @@ jobs:
run: make pydist install-dist dev-env-server
- name: Unit tests
run: |
make unit-test-server unit-test-client
bash <(curl -s https://codecov.io/bash) -y .codecov.yml -k server -cF server,python,unitTest
make unit-test-server
bash <(curl -s https://codecov.io/bash) -y .codecov.yml -k backend/server -cF backend,python,unitTest
cd client && ./node_modules/codecov/bin/codecov --yml=../.codecov.yml --root=../ --gcov-root=../ -C -F frontend,javascript,unitTest
unit-test-czi-hosted:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- name: Set up Python 3.7
uses: actions/setup-python@v1
with:
python-version: 3.7
- name: Python cache
uses: actions/cache@v1
with:
path: ~/.cache/pip
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements*.txt') }}
restore-keys: |
${{ runner.os }}-pip-
- name: Node cache
uses: actions/cache@v1
with:
path: ~/.npm
key: ${{ runner.os }}-node-${{ hashFiles('**/package-lock.json') }}
restore-keys: |
${{ runner.os }}-node-
- name: Install dependencies
run: make pydist-czi-hosted install-dist dev-env-czi-hosted
- name: Unit tests
run: |
make unit-test-czi-hosted
bash <(curl -s https://codecov.io/bash) -y .codecov.yml -k backend/czi-hosted -cF backend,python,unitTest
cd client && ./node_modules/codecov/bin/codecov --yml=../.codecov.yml --root=../ --gcov-root=../ -C -F frontend,javascript,unitTest
smoke-tests:
@@ -96,7 +126,7 @@ jobs:
restore-keys: |
${{ runner.os }}-node-
- name: Install dependencies
run: make pydist install-dist
run: make pydist-czi-hosted install-dist
- name: Smoke tests (without annotations feature)
run: |
cd client && make smoke-test
+27
View File
@@ -0,0 +1,27 @@
name: Run SASTisfaction
on:
- pull_request
jobs:
sastisfaction:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- uses: actions/checkout@v2
with:
repository: chanzuckerberg/sastisfaction
ref: main
path: .github/actions/sastisfaction
ssh-key: ${{ secrets.SASTISFACTION_READ_KEY }}
- name: Login to GitHub Container Registry
uses: docker/login-action@v1
with:
registry: ghcr.io
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Docker pull
run: docker pull ghcr.io/chanzuckerberg/sastisfaction:main
- name: Run SASTisfaction
uses: ./.github/actions/sastisfaction
with:
snowflake_private_key: ${{ secrets.SASTISFACTION_RSA_KEY }}
+30
View File
@@ -0,0 +1,30 @@
name: "Scale test cellxgene APIs for initial loading"
on:
schedule:
- cron: "0 0 * * Sun"
jobs:
locust-build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- name: Set up Python 3.7
uses: actions/setup-python@v1
with:
python-version: 3.7
- name: Install dependencies
run: |
pip install -r backend/test/test_czi_hosted/locust/requirements-locust.txt
- name: Dev Scale Test
run: |
locust -f backend/test/test_czi_hosted/locust/locustfile.py --headless -u 30 -r 10 --host https://api.cellxgene.dev.single-cell.czi.technology/cellxgene/e/ --run-time 5m 2>&1 | tee locust_dev_stats.txt
- name: Slack success webhook
env:
SLACK_WEBHOOK: ${{ secrets.SLACK_WEBHOOK }}
run: |
DEV_STATS=$(tail -n 15 locust_dev_stats.txt)
DEV_MSG="\`\`\`CELLXGENE EXPLORER DEV SCALE TEST RESULTS: ${DEV_STATS}\`\`\`"
curl -X POST -H 'Content-type: application/json' --data "{'text':'${DEV_MSG}'}" $SLACK_WEBHOOK
+8 -4
View File
@@ -15,13 +15,17 @@ dist/
*.egg-info
# Environments
venv*/
venv/
cellxgene/
# client build
server/common/web/static/*
server/common/web/templates/
server/common/web/csp-hashes.json
backend/server/common/web/static/*
backend/server/common/web/templates/
backend/server/common/web/csp-hashes.json
backend/czi_hosted/common/web/static/*
backend/czi_hosted/common/web/templates/
backend/czi_hosted/common/web/csp-hashes.json
# eb build
artifact.dir
+1 -1
View File
@@ -1,3 +1,3 @@
We warmly welcome contributions from the community!
Whether you want to contribute ideas, requests, documentation, or code, you can get started by visiting our [contribution guide](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/contribute.md).
Whether you want to contribute ideas, requests, documentation, or code, you can get started by visiting our [contribution guide](https://chanzuckerberg.github.io/cellxgene/posts/contribute).
+6 -6
View File
@@ -1,7 +1,7 @@
recursive-include server/common/web/templates *
recursive-include server/common/web/static *
recursive-include backend/server/common/web/templates *
recursive-include backend/server/common/web/static *
include server/requirements.txt
include server/requirements-prepare.txt
include server/converters/schema/hgnc_complete_set.txt.gz
include server/converters/schema/schema_definitions/*
include backend/server/requirements.txt
include backend/server/requirements-prepare.txt
include backend/server/converters/schema/hgnc_complete_set.txt.gz
include backend/server/converters/schema/schema_definitions/*
+7
View File
@@ -0,0 +1,7 @@
recursive-include backend/czi_hosted/common/web/templates *
recursive-include backend/czi_hosted/common/web/static *
include backend/czi_hosted/requirements.txt
include backend/czi_hosted/requirements-prepare.txt
include backend/czi_hosted/converters/schema/hgnc_complete_set.txt.gz
include backend/czi_hosted/converters/schema/schema_definitions/*
+81 -30
View File
@@ -2,14 +2,15 @@ include common.mk
BUILDDIR := build
CLIENTBUILD := $(BUILDDIR)/client
SERVERBUILD := $(BUILDDIR)/server
CZIHOSTEDBUILD := $(BUILDDIR)/backend/czi_hosted
SERVERBUILD := $(BUILDDIR)/backend/server
CLEANFILES := $(BUILDDIR)/ client/build build dist cellxgene.egg-info
PART ?= patch
# CLEANING
.PHONY: clean
clean: clean-lite clean-server clean-client
clean: clean-lite clean-czi-hosted clean-server clean-client
# cleaning the client's node_modules is the longest one, so we avoid that if possible
.PHONY: clean-lite
@@ -22,8 +23,11 @@ clean-client:
.PHONY: clean-server
clean-server:
cd server && $(MAKE) clean
cd backend/server && $(MAKE) clean
.PHONY: clean-czi-hosted
clean-czi-hosted:
cd backend/czi_hosted && $(MAKE) clean
# BUILDING PACKAGE
@@ -33,43 +37,71 @@ build-client:
.PHONY: build
build: clean build-client
git ls-files server/ | cpio -pdm $(BUILDDIR)
git ls-files backend/server/ | grep -v 'backend/server/test/' | cpio -pdm $(BUILDDIR)
cp -r client/build/ $(CLIENTBUILD)
$(call copy_client_assets,$(CLIENTBUILD),$(SERVERBUILD))
cp backend/__init__.py $(BUILDDIR)
cp backend/__init__.py $(BUILDDIR)/backend
cp -r backend/common $(BUILDDIR)/backend/common
cp MANIFEST.in README.md setup.cfg setup.py $(BUILDDIR)
.PHONY: build-czi-hosted
build-czi-hosted: clean build-client
git ls-files backend/czi_hosted/ | grep -v 'backend/czi_hosted/test/' | cpio -pdm $(BUILDDIR)
cp -r client/build/ $(CLIENTBUILD)
$(call copy_client_assets,$(CLIENTBUILD),$(CZIHOSTEDBUILD))
cp -r backend/common $(BUILDDIR)/backend/common
cp backend/__init__.py $(BUILDDIR)
cp backend/__init__.py $(BUILDDIR)/backend
cp MANIFEST_hosted.in README.md setup.cfg setup_hosted.py $(BUILDDIR)
mv $(BUILDDIR)/setup_hosted.py $(BUILDDIR)/setup.py
mv $(BUILDDIR)/MANIFEST_hosted.in $(BUILDDIR)/MANIFEST.in
# If you are actively developing in the server folder use this, dirties the source tree
.PHONY: build-for-server-dev
build-for-server-dev: clean-server build-client copy-client-assets
build-for-server-dev: clean-server build-client
$(call copy_client_assets,client/build,backend/server)
.PHONY: build-for-czi-hosted-dev
build-for-czi-hosted-dev: clean-czi-hosted build-client
$(call copy_client_assets,client/build,backend/czi_hosted)
.PHONY: copy-client-assets
copy-client-assets:
$(call copy_client_assets,client/build,server)
$(call copy_client_assets,client/build,backend/server)
.PHONY: copy-client-assets-czi-hosted
copy-client-assets-czi-hosted:
$(call copy_client_assets,client/build,backend/czi_hosted)
# TESTING
.PHONY: test
test: unit-test smoke-test
.PHONY: unit-test
unit-test: unit-test-server unit-test-client
unit-test: unit-test-server unit-test-client unit-test-common
.PHONY: test-server
test-server: unit-test-server smoke-test
.PHONY: test-czi-hosted
test-czi-hosted: unit-test-czi-hosted smoke-test
.PHONY: unit-test-client
unit-test-client:
cd client && $(MAKE) unit-test
.PHONY: unit-test-czi-hosted
unit-test-czi-hosted:
cd backend/czi_hosted && $(MAKE) unit-test
.PHONY: unit-test-server
unit-test-server:
PYTHONWARNINGS=ignore:ResourceWarning coverage run \
--source=server \
--omit=.coverage,venv \
-m unittest discover \
--start-directory test/unit \
--verbose; test_result=$$?; \
exit $$test_result \
cd backend/server && $(MAKE) unit-test
.PHONY: unit-test-common
unit-test-common:
cd backend/common && $(MAKE) unit-test
.PHONY: smoke-test
smoke-test:
@@ -79,6 +111,10 @@ smoke-test:
smoke-test-annotations:
cd client && $(MAKE) smoke-test-annotations
.PHONY: test-db
test-db:
cd backend/czi_hosted && $(MAKE) test-db
# FORMATTING CODE
.PHONY: fmt
@@ -93,12 +129,18 @@ fmt-py:
black .
.PHONY: lint
lint: lint-server lint-client
lint: lint-servers lint-client
.PHONY: lint-servers
lint-servers: lint-server lint-czi-hosted-server
.PHONY: lint-server
lint-server: fmt-py
flake8 server --per-file-ignores='test/fixtures/dataset_config_outline.py:F821 test/fixtures/server_config_outline.py:F821 test/performance/scale_test_annotations.py:E501'
flake8 backend/server --per-file-ignores='backend/test/fixtures/dataset_config_outline.py:F821 backend/test/fixtures/server_config_outline.py:F821 backend/server/test/performance/scale_test_annotations.py:E501'
.PHONY: lint-czi-hosted-server
lint-czi-hosted-server: fmt-py
flake8 backend/czi_hosted --per-file-ignores='backend/test/fixtures/czi_hosted_dataset_config_outline.py:F821 backend/test/fixtures/czi_hosted_server_config_outline.py:F821 backend/test/performance/scale_test_annotations.py:E501'
.PHONY: lint-client
lint-client:
@@ -111,34 +153,34 @@ pydist: build
cd $(BUILDDIR); python setup.py sdist -d ../dist
@echo "done"
# RELEASE HELPERS
.PHONY: pydist-czi-hosted
pydist-czi-hosted: build-czi-hosted
cd $(BUILDDIR); python setup.py sdist -d ../dist
@echo "done"
# Set PART=[major, minor, patch] as param to make bump.
# This will create a release candidate. (i.e. 0.16.1 -> 0.16.2-rc.0 for a patch bump)
.PHONY: bump-version
bump-version:
bumpversion --config-file .bumpversion.cfg $(PART)
# RELEASE HELPERS
# Create new version to commit to main
.PHONY: create-release-candidate
create-release-candidate: bump-version clean-lite gen-package-lock
create-release-candidate: dev-env bump-version clean-lite gen-package-lock
@echo "Version bumped part:$(PART) and client built. Ready to commit and push"
# Bump the release candidate version if needed (i.e. the previous release candidate had errors).
.PHONY: recreate-release-candidate
recreate-release-candidate: bump-release-candidate clean-lite gen-package-lock
recreate-release-candidate: dev-env bump-release-candidate clean-lite gen-package-lock
@echo "Version bumped part:$(PART) and client built. Ready to commit and push"
# Build dist and release to Test PyPI
.PHONY: release-candidate-to-test-pypi
release-candidate-to-test-pypi: pydist twine
release-candidate-to-test-pypi: dev-env pydist twine
@echo "Dist built and uploaded to test.pypi.org"
@echo "Test the install:"
@echo " make install-release-test"
# Build final dist (gets rid of the rc tag) and release final candidate to TestPyPI
.PHONY: release-final-to-test-pypi
release-final-to-test-pypi: bump-release clean-lite gen-package-lock pydist twine
release-final-to-test-pypi: dev-env bump-release clean-lite gen-package-lock pydist twine
@echo "Final release dist built and uploaded to test.pypi.org"
@echo "Test the install:"
@echo " make install-release-test"
@@ -148,9 +190,9 @@ release-final: twine-prod
@echo "Release uploaded to pypi.org"
# DANGER: releases directly to prod
# use this if you accidentally burned a test release version number,
# use this if you accidently burned a test release version number,
.PHONY: release-directly-to-prod
release-directly-to-prod: pydist twine-prod
release-directly-to-prod: dev-env pydist twine-prod
@echo "Dist built and uploaded to pypi.org"
@echo "Test the install:"
@echo " make install-release"
@@ -164,7 +206,16 @@ dev-env-client:
.PHONY: dev-env-server
dev-env-server:
pip install -r server/requirements-dev.txt
pip install -r backend/server/requirements-dev.txt
.PHONY: dev-env-czi-hosted
dev-env-czi-hosted:
pip install -r backend/czi_hosted/requirements-dev.txt
# Set PART=[major, minor, patch] as param to make bump.
# This will create a release candidate. (i.e. 0.16.1 -> 0.16.2-rc.0 for a patch bump)
.PHONY: bump-version
bump-version:
bumpversion --config-file .bumpversion.cfg $(PART)
# Increments the release candidate version (i.e. 0.16.2-rc.1 -> 0.16.2-rc.2)
.PHONY: bump-release-candidate
@@ -200,7 +251,7 @@ install-dev: uninstall
# install from test.pypi to test your release
.PHONY: install-release-test
install-release-test: uninstall
pip install --no-cache-dir --index-url https://test.pypi.org/simple/ --extra-index-url https://pypi.org/simple cellxgene==$(VERSION)
pip install --no-cache-dir --index-url https://test.pypi.org/simple/ --extra-index-url https://pypi.org/simple cellxgene
@echo "Installed cellxgene from test.pypi.org, now run and smoke test"
# install from pypi to test your release
+19 -19
View File
@@ -7,27 +7,27 @@ _an interactive explorer for single-cell transcriptomics data_
[![Compatibility Tests](https://github.com/chanzuckerberg/cellxgene/workflows/Compatibility%20Tests/badge.svg)](https://github.com/chanzuckerberg/cellxgene/actions?query=workflow%3A%22Compatibility+Tests%22)
![Code Coverage](https://codecov.io/gh/chanzuckerberg/cellxgene/branch/main/graph/badge.svg)
cellxgene Desktop (pronounced "cell-by-gene") is an interactive data explorer for single-cell datasets, such as those coming from the [Human Cell Atlas](https://humancellatlas.org). Leveraging modern web development techniques to enable fast visualizations of at least 1 million cells, we hope to enable biologists and computational researchers to explore their data.
cellxgene (pronounced "cell-by-gene") is an interactive data explorer for single-cell transcriptomics datasets, such as those coming from the [Human Cell Atlas](https://humancellatlas.org). Leveraging modern web development techniques to enable fast visualizations of at least 1 million cells, we hope to enable biologists and computational researchers to explore their data.
Whether you need to visualize one thousand cells or one million, cellxgene Desktop helps you gain insight into your single-cell data.
Whether you need to visualize one thousand cells or one million, cellxgene helps you gain insight into your single-cell data.
<img src="https://github.com/chanzuckerberg/cellxgene/raw/main/docs/images/crossfilter.gif" width="350" height="200" hspace="30"><img src="https://github.com/chanzuckerberg/cellxgene/raw/main/docs/images/category-breakdown.gif" width="350" height="200" hspace="30">
# Getting started
### The comprehensive guide to cellxgene Desktop
### The comprehensive guide to cellxgene
[The cellxgene documentation is your one-stop-shop for information about cellxgene Desktop](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/README.md)! You may be particularly interested in:
[The cellxgene documentation is your one-stop-shop for information about cellxgene](https://chanzuckerberg.github.io/cellxgene/)! You may be particularly interested in:
- Seeing [what cellxgene Desktop can do](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/explore-data/explorer-tutorials.md)
- Learning more about cellxgene [installation](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/install.md) and [usage](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/quick-start.md#quick-start-1)
- [Preparing your own data](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/data-reqs.md) for use in cellxgene Desktop
- Checking out [our roadmap](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/roadmap.md) for future development
- [Contributing](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/contribute.md) to cellxgene Desktop
- Seeing [what cellxgene can do](https://chanzuckerberg.github.io/cellxgene/posts/gallery)
- Learning more about cellxgene [installation](https://chanzuckerberg.github.io/cellxgene/posts/install) and [usage](https://chanzuckerberg.github.io/cellxgene/posts/launch)
- [Preparing your own data](https://chanzuckerberg.github.io/cellxgene/posts/prepare) for use in cellxgene
- Checking out [our roadmap](https://chanzuckerberg.github.io/cellxgene/posts/roadmap) for future development
- [Contributing](https://chanzuckerberg.github.io/cellxgene/posts/contribute) to cellxgene
### Quick start
To install cellxgene Desktop you need Python 3.6+. We recommend [installing cellxgene Desktop into a conda or virtual environment.](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/install.md)
To install cellxgene you need Python 3.6+. We recommend [installing cellxgene into a conda or virtual environment.](https://chanzuckerberg.github.io/cellxgene/posts/install)
Install the package.
@@ -35,19 +35,19 @@ Install the package.
pip install cellxgene
```
Launch cellxgene Desktop with an example [anndata](https://anndata.readthedocs.io/en/latest/) file
Launch cellxgene with an example [anndata](https://anndata.readthedocs.io/en/latest/) file
```bash
cellxgene launch https://cellxgene-example-data.czi.technology/pbmc3k.h5ad
```
To explore more datasets already formatted for cellxgene Desktop, check out the [Demo data](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/quick-start.md#example-datasets) or
see [Preparing your data](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/desktop/data-reqs.md) to learn more about formatting your own
data for cellxgene Desktop.
To explore more datasets already formatted for cellxgene, check out the [Demo data](https://chanzuckerberg.github.io/cellxgene/posts/demo-data) or
see [Preparing your data](https://chanzuckerberg.github.io/cellxgene/posts/prepare) to learn more about formatting your own
data for cellxgene.
### Supported browsers
cellxgene Desktop currently supports the following browsers:
cellxgene currently supports the following browsers:
- Google Chrome 61+
- Edge 15+
@@ -62,11 +62,11 @@ For questions, suggestions, or accolades, [join the `#cellxgene-users` channel o
For any errors, [report bugs on Github](https://github.com/chanzuckerberg/cellxgene/issues).
# Developing with cellxgene Desktop
# Developing with cellxgene
### Contributing
We warmly welcome contributions from the community! Please see our [contributing guide](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/contribute.md) and don't hesitate to open an issue or send a pull request to improve cellxgene Desktop. Please see the [dev_docs](https://github.com/chanzuckerberg/cellxgene/tree/main/dev_docs) for pull request suggestions, unit test details, local documentation preview, and other development specifics.
We warmly welcome contributions from the community! Please see our [contributing guide](https://chanzuckerberg.github.io/cellxgene/posts/contribute) and don't hesitate to open an issue or send a pull request to improve cellxgene. Please see the [dev_docs](https://github.com/chanzuckerberg/cellxgene/tree/main/dev_docs) for pull request suggestions, unit test details, local documentation preview, and other development specifics.
This project adheres to the Contributor Covenant [code of conduct](https://github.com/chanzuckerberg/.github/blob/master/CODE_OF_CONDUCT.md). By participating, you are expected to uphold this code. Please report unacceptable behavior to opensource@chanzuckerberg.com.
@@ -79,9 +79,9 @@ this project. All code is freely available for reuse under the [MIT license](htt
Before extending cellxgene, we encourage you to reach out to us with ideas or questions. It might be possible that an
extension could be directly contributed, which would make it available for a wider audience, or that it's on our
[roadmap](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/roadmap.md) and under active development.
[roadmap](./docs/posts/roadmap.md) and under active development.
See the [cellxgene extensions](https://github.com/chanzuckerberg/cellxgene-documentation/blob/main/community-extensions.md) section of our documentation for examples of community use and cellxgene extensions.
See the [cellxgene extensions](./docs/posts/extensions.md) section of our documentation for examples of community use and cellxgene extensions.
### Security
+11
View File
@@ -0,0 +1,11 @@
.PHONY: unit-test
unit-test:
PYTHONWARNINGS=ignore:ResourceWarning coverage run \
--source=fbs,utils \
--omit=.coverage,data_common/fbs/NetEncoding,venv \
-m unittest discover \
--start-directory ../test/test_common/unit \
--top-level-directory ../../ \
--verbose; test_result=$$?; \
exit $$test_result \
@@ -1,6 +1,6 @@
import re
from server.common.errors import ColorFormatException
from backend.common.errors import ColorFormatException
HEX_COLOR_FORMAT = re.compile("^#[a-fA-F0-9]{6,6}$")
@@ -1,6 +1,6 @@
import numpy as np
from scipy import sparse, stats
from server.common.constants import XApproximateDistribution
from backend.common.constants import XApproximateDistribution
def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
@@ -27,8 +27,7 @@ def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
:param top_n: number of variables to return stats for
:param diffexp_lfc_cutoff: minimum
absolute value returning [ varindex, logfoldchange, pval, pval_adj ] for top N genes
:return: for top N genes, {"positive": for top N genes, [ varindex, foldchange, pval, pval_adj ],
"negative": for top N genes, [ varindex, foldchange, pval, pval_adj ]}
:return: for top N genes, {"positive": for top N genes, [ varindex, foldchange, pval, pval_adj ], "negative": for top N genes, [ varindex, foldchange, pval, pval_adj ]}
"""
X_approximate_distribution = adaptor.get_X_approximate_distribution()
@@ -1,13 +1,12 @@
from typing import Tuple
import numba
import concurrent.futures
import numpy as np
from scipy import sparse
from server.common.constants import XApproximateDistribution
from backend.common.constants import XApproximateDistribution
@numba.njit(error_model="numpy", nogil=True)
def min_max_fast(arr: np.ndarray) -> Tuple[float, float]:
def min_max(arr: np.ndarray):
"""Return (min, max) values for the ndarray."""
# initialize to first finite value in array. Normally,
@@ -48,24 +47,6 @@ def min_max_fast(arr: np.ndarray) -> Tuple[float, float]:
return min_val, max_val
def min_max_numpy(arr: np.ndarray) -> Tuple[float, float]:
return arr.min(), arr.max()
def numba_has_support_for_scalar_type(arr: np.ndarray) -> bool:
"""Numba does not support half-floats, 128 bit floats, ints > 64 bit or non-scalars."""
if arr.dtype == np.float32 or arr.dtype == np.float64:
return True
if np.issubdtype(arr.dtype, np.integer) and arr.dtype <= np.int64:
return True
if arr.dtype == np.bool_:
return True
return False
def estimate_approximate_distribution(X) -> XApproximateDistribution:
"""
Estimate the distribution (normal, count) of the X matrix.
@@ -91,8 +72,6 @@ def estimate_approximate_distribution(X) -> XApproximateDistribution:
else:
raise TypeError(f"Unsupported matrix format: {str(type(X))}")
min_max = min_max_fast if numba_has_support_for_scalar_type(Xdata) else min_max_numpy
CHUNKSIZE = 1 << 24
if Xdata.size > CHUNKSIZE:
min_val = max_val = Xdata[0]
@@ -41,6 +41,9 @@ define_request_exception(
)
define_request_exception("ExceedsLimitError", "Raised when an HTTP request exceeds a limit/quota")
define_request_exception("ColorFormatException", "Raised when color helper functions encounter an unknown color format")
define_request_exception(
"AuthenticationError", "Raised when there is an authentication error", default_status_code=HTTPStatus.UNAUTHORIZED
)
define_request_exception(
"AnnotationCategoryNameError",
@@ -50,5 +53,6 @@ define_request_exception(
define_exception("ConfigurationError", "Raised when checking configuration errors")
define_exception("PrepareError", "Raised when data is misprepared")
define_exception("SecretKeyRetrievalError", "Raised when get_secret_key from AWS fails")
define_exception("ObsoleteRequest", "Raised when the request is no longer valid.")
define_exception("UnsupportedSummaryMethod", "Raised when a gene set summary method is unknown or unsupported.")
@@ -5,16 +5,16 @@ import pandas as pd
from flatbuffers import Builder
from scipy import sparse
from server.common.utils.type_conversion_utils import get_encoding_dtype_of_array
from backend.common.utils.type_conversion_utils import get_encoding_dtype_of_array
import server.common.fbs.NetEncoding.Column as Column
import server.common.fbs.NetEncoding.Float32Array as Float32Array
import server.common.fbs.NetEncoding.Float64Array as Float64Array
import server.common.fbs.NetEncoding.Int32Array as Int32Array
import server.common.fbs.NetEncoding.JSONEncodedArray as JSONEncodedArray
import server.common.fbs.NetEncoding.Matrix as Matrix
import server.common.fbs.NetEncoding.TypedArray as TypedArray
import server.common.fbs.NetEncoding.Uint32Array as Uint32Array
import backend.common.fbs.NetEncoding.Column as Column
import backend.common.fbs.NetEncoding.Float32Array as Float32Array
import backend.common.fbs.NetEncoding.Float64Array as Float64Array
import backend.common.fbs.NetEncoding.Int32Array as Int32Array
import backend.common.fbs.NetEncoding.JSONEncodedArray as JSONEncodedArray
import backend.common.fbs.NetEncoding.Matrix as Matrix
import backend.common.fbs.NetEncoding.TypedArray as TypedArray
import backend.common.fbs.NetEncoding.Uint32Array as Uint32Array
# Serialization helper
@@ -75,7 +75,7 @@ def read_gene_sets_tidycsv(gs_locator, context=None):
# if this is the first non-comment row, assume it is a header and validate
# column names. OK if the user has extra columns after our initial set.
if not haveReadHeader:
if row[0 : len(GENESETS_TIDYCSV_HEADER)] != GENESETS_TIDYCSV_HEADER:
if row[0:len(GENESETS_TIDYCSV_HEADER)] != GENESETS_TIDYCSV_HEADER:
raise AnnotationsError("Gene set CSV file missing the required column header.")
haveReadHeader = True
continue
+23
View File
@@ -0,0 +1,23 @@
import logging
import boto3
from flask import json
from backend.common.errors import SecretKeyRetrievalError
def get_secret_key(region_name, secret_name):
session = boto3.session.Session()
client = session.client(service_name="secretsmanager", region_name=region_name)
try:
get_secret_value_response = client.get_secret_value(SecretId=secret_name)
if "SecretString" in get_secret_value_response:
var = get_secret_value_response["SecretString"]
secret = json.loads(var)
return secret
except Exception as e:
logging.critical(f"Caught exception during get_secret_key, {e}", exc_info=True)
raise SecretKeyRetrievalError(str(e))
return None
@@ -6,7 +6,7 @@ import pandas as pd
"""
These routines drive all type inference for the schema generation and the
FBS (REST OTA) encoding.
FBS (REST OTA) encoding. They are also used for CXG generation.
H5AD Type REST REST
@@ -10,7 +10,7 @@ from urllib.parse import urlsplit, urljoin
import numpy as np
from flask import json
from server.common.errors import ConfigurationError
from backend.common.errors import ConfigurationError
def find_available_port(host, port=5005):
@@ -65,13 +65,7 @@ def path_join(base, *urls):
return btpl._replace(path=path).geturl()
class StrictJSONEncoder(json.JSONEncoder):
"""
Custom JSON encoder set-up performing two tasks:
1. Strict JSON conformance with non-finite floats (NaN, +/-Inf) via allow_nan=False
2. Convert various Numpy types into python types so the encoder will correctly encode.
"""
class Float32JSONEncoder(json.JSONEncoder):
def __init__(self, *args, **kwargs):
"""
NaN/Infinities are illegal in standard JSON. Python extends JSON with
@@ -84,11 +78,9 @@ class StrictJSONEncoder(json.JSONEncoder):
super().__init__(*args, **kwargs)
def default(self, obj):
"""This helps us convert types not supported by the native JSON encoder into
standard python types, eg, np.int64."""
if isinstance(obj, np.floating):
if isinstance(obj, np.float32):
return float(obj)
if isinstance(obj, np.integer):
elif isinstance(obj, np.integer):
return int(obj)
return json.JSONEncoder.default(self, obj)
@@ -97,8 +89,8 @@ def custom_format_warning(msg, *args, **kwargs):
return f"[cellxgene] Warning: {msg} \n"
def jsonify_strict(data):
return json.dumps(data, cls=StrictJSONEncoder, allow_nan=False)
def jsonify_numpy(data):
return json.dumps(data, cls=Float32JSONEncoder, allow_nan=False)
def import_plugins(plugin_module):
+49
View File
@@ -0,0 +1,49 @@
include ../../common.mk
.PHONY: clean
clean:
rm -f common/web/templates/index.html
rm -rf common/web/static
rm -f common/web/csp-hashes.json
.PHONY: unit-test
unit-test: create-test-db
PYTHONWARNINGS=ignore:ResourceWarning coverage run \
--source=app,auth,cli,common,compute,converters,data_anndata,data_common,data_cxg,eb \
--omit=.coverage,venv \
-m unittest discover \
--start-directory ../test/test_czi_hosted/unit \
--top-level-directory ../.. \
--verbose; test_result=$$?; \
$(MAKE) clean-test-db; \
exit $$test_result \
.PHONY: test-db
test-db: create-test-db
PYTHONWARNINGS=ignore:ResourceWarning coverage run \
--source=db \
--omit=.coverage,venv \
-m unittest discover \
--start-directory ../test/test_czi_hosted/test_database \
--top-level-directory ../.. \
--verbose; test_result=$$?; \
$(MAKE) clean-test-db; \
exit $$test_result
.PHONY: create-test-db
create-test-db:
-docker run -d -p 5432:5432 --name test_db -e POSTGRES_PASSWORD=test_pw postgres
.PHONY: clean-test-db
clean-test-db:
-docker stop test_db
-docker rm test_db
.PHONY: test-annotations-performance
test-annotations-performance:
python ../test/test_czi_hosted/performance/performance_test_annotations_backend.py
.PHONY: test-annotations-scale
test-annotations-scale:
locust -f ../test/test_czi_hosted/performance/scale_test_annotations.py --headless -u 30 -r 10 --host https://api.cellxgene.dev.single-cell.czi.technology/cellxgene/e/ --run-time 5m 2>&1 | tee locust_dev_stats.txt
+15
View File
@@ -0,0 +1,15 @@
import logging
import sys
from backend.common.utils.utils import import_plugins
__version__ = "0.16.7"
display_version = "cellxgene v" + __version__
try:
import_plugins("backend.czi_hosted.plugins")
except Exception as e:
# Make sure to exit in this case, as the server may not be configured as expected.
logging.critical(f"Error in import_plugins: {str(e)}")
sys.exit(1)
+14
View File
@@ -0,0 +1,14 @@
# Work around bug https://github.com/pallets/werkzeug/issues/461
if __package__ is None:
import sys
from pathlib import Path
PKG_PATH = Path(__file__).parent
sys.path.insert(0, str(PKG_PATH.parent))
import backend.czi_hosted # noqa F401
__package__ = PKG_PATH.name
# Main thing
from .cli.cli import cli # noqa F402
cli()
+475
View File
@@ -0,0 +1,475 @@
import datetime
import logging
from functools import wraps
from http import HTTPStatus
from urllib.parse import urlparse
import hashlib
import os
from flask import (
Flask,
redirect,
current_app,
make_response,
render_template,
abort,
Blueprint,
request,
send_from_directory,
)
from flask_restful import Api, Resource
from server_timing import Timing as ServerTiming
import backend.czi_hosted.common.rest as common_rest
from backend.common.utils.data_locator import DataLocator
from backend.common.errors import DatasetAccessError, RequestException
from backend.czi_hosted.common.health import health_check
from backend.common.utils.utils import path_join, Float32JSONEncoder
from backend.czi_hosted.data_common.matrix_loader import MatrixDataLoader
webbp = Blueprint("webapp", "backend.czi_hosted.common.web", template_folder="templates")
ONE_WEEK = 7 * 24 * 60 * 60
def _cache_control(always, **cache_kwargs):
"""
Used to easily manage cache control headers on responses.
See Werkzeug for attributes that can be set, eg, no_cache, private, max_age, etc.
https://werkzeug.palletsprojects.com/en/1.0.x/datastructures/#werkzeug.datastructures.ResponseCacheControl
"""
def inner_cache_control(f):
@wraps(f)
def wrapper(*args, **kwargs):
response = make_response(f(*args, **kwargs))
if not always and not current_app.app_config.server_config.app__generate_cache_control_headers:
return response
if response.status_code >= 400:
return response
for k, v in cache_kwargs.items():
setattr(response.cache_control, k, v)
return response
return wrapper
return inner_cache_control
def cache_control(**cache_kwargs):
""" config driven """
return _cache_control(False, **cache_kwargs)
def cache_control_always(**cache_kwargs):
""" always generate headers, regardless of the config """
return _cache_control(True, **cache_kwargs)
# tell the client not to cache the index.html page so that changes to the app work on redeployment
# note that the bulk of the data needed by the client (datasets) will still be cached
@webbp.route("/", methods=["GET"])
@cache_control_always(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True)
def dataset_index(url_dataroot=None, dataset=None):
app_config = current_app.app_config
server_config = app_config.server_config
if dataset is None:
if app_config.is_multi_dataset():
return dataroot_index()
else:
location = server_config.single_dataset__datapath
else:
dataroot = None
for key, dataroot_dict in server_config.multi_dataset__dataroot.items():
if dataroot_dict["base_url"] == url_dataroot:
dataroot = dataroot_dict["dataroot"]
break
if dataroot is None:
abort(HTTPStatus.NOT_FOUND)
location = path_join(dataroot, dataset)
dataset_config = app_config.get_dataset_config(url_dataroot)
scripts = dataset_config.app__scripts
inline_scripts = dataset_config.app__inline_scripts
try:
cache_manager = current_app.matrix_data_cache_manager
with cache_manager.data_adaptor(url_dataroot, location, app_config) as data_adaptor:
data_adaptor.set_uri_path(f"{url_dataroot}/{dataset}")
args = {"SCRIPTS": scripts, "INLINE_SCRIPTS": inline_scripts}
return render_template("index.html", **args)
except DatasetAccessError as e:
return common_rest.abort_and_log(
e.status_code, f"Invalid dataset {dataset}: {e.message}", loglevel=logging.INFO, include_exc_info=True
)
@webbp.errorhandler(RequestException)
def handle_request_exception(error):
return common_rest.abort_and_log(error.status_code, error.message, loglevel=logging.INFO, include_exc_info=True)
def get_data_adaptor(url_dataroot=None, dataset=None):
config = current_app.app_config
server_config = config.server_config
dataset_key = None
if dataset is None:
datapath = server_config.single_dataset__datapath
else:
dataroot = None
for key, dataroot_dict in server_config.multi_dataset__dataroot.items():
if dataroot_dict["base_url"] == url_dataroot:
dataroot = dataroot_dict["dataroot"]
dataset_key = key
break
if dataroot is None:
raise DatasetAccessError(f"Invalid dataset {url_dataroot}/{dataset}")
datapath = path_join(dataroot, dataset)
# path_join returns a normalized path. Therefore it is
# sufficient to check that the datapath starts with the
# dataroot to determine that the datapath is under the dataroot.
if not datapath.startswith(dataroot):
raise DatasetAccessError(f"Invalid dataset {url_dataroot}/{dataset}")
if datapath is None:
return common_rest.abort_and_log(HTTPStatus.BAD_REQUEST, "Invalid dataset NONE", loglevel=logging.INFO)
cache_manager = current_app.matrix_data_cache_manager
return cache_manager.data_adaptor(dataset_key, datapath, config)
def requires_authentication(func):
@wraps(func)
def wrapped_function(self, *args, **kwargs):
auth = current_app.auth
if auth.is_user_authenticated():
return func(self, *args, **kwargs)
else:
return make_response("not authenticated", HTTPStatus.UNAUTHORIZED)
return wrapped_function
def rest_get_data_adaptor(func):
@wraps(func)
def wrapped_function(self, dataset=None):
try:
with get_data_adaptor(self.url_dataroot, dataset) as data_adaptor:
data_adaptor.set_uri_path(f"{self.url_dataroot}/{dataset}")
return func(self, data_adaptor)
except DatasetAccessError as e:
return common_rest.abort_and_log(
e.status_code, f"Invalid dataset {dataset}: {e.message}", loglevel=logging.INFO, include_exc_info=True
)
return wrapped_function
def dataroot_test_index():
# the following index page is meant for testing/debugging purposes
data = '<!doctype html><html lang="en">'
data += "<head><title>Hosted Cellxgene</title></head>"
data += "<body><H1>Welcome to cellxgene</H1>"
config = current_app.app_config
server_config = config.server_config
auth = server_config.auth
if auth.is_valid_authentication_type():
if server_config.auth.is_user_authenticated():
data += f"<p>Logged in as {auth.get_user_id()} / {auth.get_user_name()} / {auth.get_user_email()}</p>"
if auth.requires_client_login():
if server_config.auth.is_user_authenticated():
data += f"<p><a href='{auth.get_logout_url(None)}'>Logout</a></p>"
else:
data += f"<p><a href='{auth.get_login_url(None)}'>Login</a></p>"
datasets = []
for dataroot_dict in server_config.multi_dataset__dataroot.values():
dataroot = dataroot_dict["dataroot"]
url_dataroot = dataroot_dict["base_url"]
locator = DataLocator(dataroot, region_name=server_config.data_locator__s3__region_name)
for fname in locator.ls():
location = path_join(dataroot, fname)
try:
MatrixDataLoader(location, app_config=config)
datasets.append((url_dataroot, fname))
except DatasetAccessError:
# skip over invalid datasets
pass
data += "<br/>Select one of these datasets...<br/>"
data += "<ul>"
datasets.sort()
for url_dataroot, dataset in datasets:
data += f"<li><a href={url_dataroot}/{dataset}/>{dataset}</a></li>"
data += "</ul>"
data += "</body></html>"
return make_response(data)
def dataroot_index():
# Handle the base url for the cellxgene server when running in multi dataset mode
config = current_app.app_config
if not config.server_config.multi_dataset__index:
abort(HTTPStatus.NOT_FOUND)
elif config.server_config.multi_dataset__index is True:
return dataroot_test_index()
else:
return redirect(config.server_config.multi_dataset__index)
class HealthAPI(Resource):
@cache_control(no_store=True)
def get(self):
config = current_app.app_config
return health_check(config)
class DatasetResource(Resource):
"""Base class for all Resources that act on datasets."""
def __init__(self, url_dataroot):
super().__init__()
self.url_dataroot = url_dataroot
class SchemaAPI(DatasetResource):
# TODO @mdunitz separate dataset schema and user schema
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.schema_get(data_adaptor)
class ConfigAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.config_get(current_app.app_config, data_adaptor)
class UserInfoAPI(DatasetResource):
@cache_control_always(no_store=True)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.userinfo_get(current_app.app_config, data_adaptor)
class AnnotationsObsAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.annotations_obs_get(request, data_adaptor)
@requires_authentication
@cache_control(no_store=True)
@rest_get_data_adaptor
def put(self, data_adaptor):
return common_rest.annotations_obs_put(request, data_adaptor)
class AnnotationsVarAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.annotations_var_get(request, data_adaptor)
class DataVarAPI(DatasetResource):
@cache_control(no_store=True)
@rest_get_data_adaptor
def put(self, data_adaptor):
return common_rest.data_var_put(request, data_adaptor)
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.data_var_get(request, data_adaptor)
class ColorsAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.colors_get(data_adaptor)
class DiffExpObsAPI(DatasetResource):
@cache_control(no_store=True)
@rest_get_data_adaptor
def post(self, data_adaptor):
return common_rest.diffexp_obs_post(request, data_adaptor)
class LayoutObsAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.layout_obs_get(request, data_adaptor)
class GenesetsAPI(DatasetResource):
@cache_control(public=True, max_age=ONE_WEEK)
@rest_get_data_adaptor
def get(self, data_adaptor):
return common_rest.genesets_get(request, data_adaptor)
class SummarizeVarAPI(DatasetResource):
@rest_get_data_adaptor
@cache_control(public=True, max_age=ONE_WEEK)
def get(self, data_adaptor):
return common_rest.summarize_var_get(request, data_adaptor)
@rest_get_data_adaptor
@cache_control(no_store=True)
def post(self, data_adaptor):
return common_rest.summarize_var_post(request, data_adaptor)
def get_api_base_resources(bp_base):
"""Add resources that are accessed from the api_base_url"""
api = Api(bp_base)
# Diagnostics routes
api.add_resource(HealthAPI, "/health")
return api
def get_api_dataroot_resources(bp_dataroot, url_dataroot=None):
"""Add resources that refer to a dataset"""
api = Api(bp_dataroot)
def add_resource(resource, url):
"""convenience function to make the outer function less verbose"""
api.add_resource(resource, url, resource_class_args=(url_dataroot,))
# Initialization routes
add_resource(SchemaAPI, "/schema")
add_resource(ConfigAPI, "/config")
add_resource(UserInfoAPI, "/userinfo")
# Data routes
add_resource(AnnotationsObsAPI, "/annotations/obs")
add_resource(AnnotationsVarAPI, "/annotations/var")
add_resource(DataVarAPI, "/data/var")
add_resource(GenesetsAPI, "/genesets")
add_resource(SummarizeVarAPI, "/summarize/var")
# Display routes
add_resource(ColorsAPI, "/colors")
# Computation routes
add_resource(DiffExpObsAPI, "/diffexp/obs")
add_resource(LayoutObsAPI, "/layout/obs")
return api
def handle_api_base_url(app, app_config):
"""If an api_base_url is provided, then an inline script is generated to
handle the new API prefix"""
api_base_url = app_config.server_config.get_api_base_url()
if not api_base_url:
return
sha256 = hashlib.sha256(api_base_url.encode()).hexdigest()
script_name = f"api_base_url-{sha256}.js"
script_path = os.path.join(app.root_path, "../common/web/templates", script_name)
with open(script_path, "w") as fout:
fout.write("window.CELLXGENE.API.prefix = `" + api_base_url + "${location.pathname}api/`;\n")
dataset_configs = [app_config.default_dataset_config] + list(app_config.dataroot_config.values())
for dataset_config in dataset_configs:
inline_scripts = dataset_config.app__inline_scripts
inline_scripts.append(script_name)
class Server:
@staticmethod
def _before_adding_routes(app, app_config):
""" will be called before routes are added, during __init__. Subclass protocol """
pass
def __init__(self, app_config):
self.app = Flask(__name__, static_folder=None)
handle_api_base_url(self.app, app_config)
self._before_adding_routes(self.app, app_config)
self.app.json_encoder = Float32JSONEncoder
server_config = app_config.server_config
if server_config.app__server_timing_headers:
ServerTiming(self.app, force_debug=True)
# enable session data
self.app.permanent_session_lifetime = datetime.timedelta(days=50 * 365)
# Config
secret_key = server_config.app__flask_secret_key
self.app.config.update(SECRET_KEY=secret_key)
self.app.register_blueprint(webbp)
api_version = "/api/v0.2"
api_base_url = server_config.get_api_base_url()
api_path = "/"
if api_base_url:
parse = urlparse(api_base_url)
api_path = parse.path
bp_base = Blueprint("bp_base", __name__, url_prefix=api_path)
base_resources = get_api_base_resources(bp_base)
self.app.register_blueprint(base_resources.blueprint)
if app_config.is_multi_dataset():
# NOTE: These routes only allow the dataset to be in the directory
# of the dataroot, and not a subdirectory. We may want to change
# the route format at some point
for dataroot_dict in server_config.multi_dataset__dataroot.values():
url_dataroot = dataroot_dict["base_url"]
bp_dataroot = Blueprint(
f"api_dataset_{url_dataroot}",
__name__,
url_prefix=f"{api_path}/{url_dataroot}/<dataset>" + api_version,
)
dataroot_resources = get_api_dataroot_resources(bp_dataroot, url_dataroot)
self.app.register_blueprint(dataroot_resources.blueprint)
self.app.add_url_rule(
f"/{url_dataroot}/<dataset>",
f"dataset_index_{url_dataroot}",
lambda dataset, url_dataroot=url_dataroot: dataset_index(url_dataroot, dataset),
methods=["GET"],
)
self.app.add_url_rule(
f"/{url_dataroot}/<dataset>/",
f"dataset_index_{url_dataroot}/",
lambda dataset, url_dataroot=url_dataroot: dataset_index(url_dataroot, dataset),
methods=["GET"],
)
self.app.add_url_rule(
f"/{url_dataroot}/<dataset>/static/<path:filename>",
f"static_assets_{url_dataroot}",
view_func=lambda dataset, filename: send_from_directory("../common/web/static", filename),
methods=["GET"],
)
else:
bp_api = Blueprint("api", __name__, url_prefix=f"{api_path}{api_version}")
resources = get_api_dataroot_resources(bp_api)
self.app.register_blueprint(resources.blueprint)
self.app.add_url_rule(
"/static/<path:filename>",
"static_assets",
view_func=lambda filename: send_from_directory("../common/web/static", filename),
methods=["GET"],
)
self.app.matrix_data_cache_manager = server_config.matrix_data_cache_manager
self.app.app_config = app_config
auth = server_config.auth
self.app.auth = auth
if auth and auth.requires_client_login():
auth.add_url_rules(self.app)
auth.complete_setup(self.app)
+6
View File
@@ -0,0 +1,6 @@
# import the built in auth types so they can be registered
import backend.czi_hosted.auth.auth_test # noqa: F401
import backend.czi_hosted.auth.auth_session # noqa: F401
import backend.czi_hosted.auth.auth_oauth # noqa: F401
import backend.czi_hosted.auth.auth_none # noqa: F401
+91
View File
@@ -0,0 +1,91 @@
from abc import ABC, abstractmethod
class AuthTypeBase(ABC):
"""Base type for all authentication types."""
def __init__(self):
super().__init__()
@abstractmethod
def is_valid_authentication_type(self):
"""Return True if the auth type is valid, e.g. it can return userinfo and username.
(AuthTypeNone is the only one type that returns False)"""
pass
def requires_client_login(self):
"""Return True if the user needs to login from the client (e.g. Login button is shown)"""
return False
@abstractmethod
def complete_setup(self, app):
"""complete any setup that may be needed by this auth type. The Flask app is passed in.
This is the last auth function called before the server starts to run."""
pass
@abstractmethod
def is_user_authenticated(self):
"""Return True if the user is authenticated"""
pass
@abstractmethod
def get_user_id(self):
"""Return the id for this user (string)"""
pass
@abstractmethod
def get_user_name(self):
"""Return the name of the user (string)"""
pass
@abstractmethod
def get_user_email(self):
"""Return the name of the user (string)"""
pass
def get_user_picture(self):
"""Return the location to the user's picture"""
return None
class AuthTypeClientBase(AuthTypeBase):
"""Base type for all authentication types that require the client to login"""
def __init__(self):
super().__init__()
def requires_client_login(self):
return True
@abstractmethod
def add_url_rules(self, selfapp):
"""Add url rules to the app (like /login, /logout, etc)"""
pass
@abstractmethod
def get_login_url(self, data_adaptor):
"""Return the url for the login route"""
pass
@abstractmethod
def get_logout_url(self, data_adaptor):
"""Return the url for the logout route"""
pass
class AuthTypeFactory:
"""Factory class to create an authentication type"""
auth_types = {}
@staticmethod
def register(name, auth_type):
assert issubclass(auth_type, AuthTypeBase)
AuthTypeFactory.auth_types[name] = auth_type
@staticmethod
def create(name, app_config):
auth_type = AuthTypeFactory.auth_types.get(name)
if auth_type is None:
return None
return auth_type(app_config)
+27
View File
@@ -0,0 +1,27 @@
from backend.czi_hosted.auth.auth import AuthTypeBase, AuthTypeFactory
class AuthTypeNone(AuthTypeBase):
def __init__(self, app_config):
super().__init__()
def is_valid_authentication_type(self):
return False
def complete_setup(self, app):
pass
def is_user_authenticated(self):
return True
def get_user_id(self):
return None
def get_user_name(self):
return None
def get_user_email(self):
return None
AuthTypeFactory.register(None, AuthTypeNone)
+385
View File
@@ -0,0 +1,385 @@
from flask import session, request, redirect, current_app, after_this_request, has_request_context, g
from backend.czi_hosted.auth.auth import AuthTypeClientBase, AuthTypeFactory
from backend.common.errors import AuthenticationError, ConfigurationError
from urllib.parse import urlencode, urlparse
import json
import requests
import base64
# It is not required to have authlib or jose.
# However, it is a configuration error to use this auth type if they are not installed.
missingimport = []
try:
from authlib.integrations.flask_client import OAuth
except ModuleNotFoundError:
missingimport.append("authlib")
try:
from jose import jwt
from jose.exceptions import ExpiredSignatureError, JWTError, JWTClaimsError
except ModuleNotFoundError:
missingimport.append("jose")
class Tokens:
"""Simple class to represent the tokens that are saved/restored from the cookie"""
def __init__(self, access_token, id_token, refresh_token, expires_at, **kwargs):
self.access_token = access_token
self.id_token = id_token
self.refresh_token = refresh_token
self.expires_at = expires_at
# expires_at may be None after a token refresh, and so it is not checked here
if not (access_token and id_token and refresh_token):
raise KeyError(str(self.__dict__))
class AuthTypeOAuth(AuthTypeClientBase):
"""An authentication type for oauth2 logins."""
CXG_TOKENS = "auth_tokens"
def __init__(self, server_config):
super().__init__()
if missingimport:
raise ConfigurationError(f"oauth requires these modules: {', '.join(missingimport)}")
self.algorithms = ["RS256"]
self.oauth_api_base_url = server_config.authentication__params_oauth__oauth_api_base_url
self.client_id = server_config.authentication__params_oauth__client_id
self.client_secret = server_config.authentication__params_oauth__client_secret
self.session_cookie = server_config.authentication__params_oauth__session_cookie
self.cookie_params = server_config.authentication__params_oauth__cookie
self.jwt_decode_options = server_config.authentication__params_oauth__jwt_decode_options
self._validate_cookie_params()
self._validate_jwt_decode_options()
self.api_base_url = server_config.get_api_base_url()
self.web_base_url = server_config.get_web_base_url()
if self.api_base_url is None:
raise ConfigurationError("oauth requires the app__api_base_url to be set")
# set the audience
self.audience = self.client_id
# load the jwks (JSON Web Key Set).
# The JSON Web Key Set (JWKS) is a set of keys which contains the public keys used to verify
# any JSON Web Token (JWT) issued by the authorization server and signed using the RS256
try:
jwksloc = f"{self.oauth_api_base_url}/.well-known/jwks.json"
jwksurl = requests.get(jwksloc)
self.jwks = jwksurl.json()
except Exception:
raise ConfigurationError(
f"error in oauth, api_url_base: {self.oauth_api_base_url}, cannot access {jwksloc}"
)
def _validate_cookie_params(self):
"""check the cookie_params, and raise a ConfigurationError if there is something wrong"""
if self.session_cookie:
return
if not isinstance(self.cookie_params, dict):
raise ConfigurationError("either session_cookie or cookie must be set")
valid_keys = {"key", "max_age", "expires", "path", "domain", "secure", "httponly", "samesite"}
keys = set(self.cookie_params.keys())
unknown = keys - valid_keys
if unknown:
raise ConfigurationError(f"unexpected key in cookie params: {', '.join(unknown)}")
if "key" not in keys:
raise ConfigurationError("must have a key (name) in the cookie params")
def _validate_jwt_decode_options(self):
"""check the jwt_decode_options, and raise a ConfigurationError if there is something wrong"""
if self.jwt_decode_options is None:
self.jwt_decode_options = {}
return
valid_keys = {
"verify_signature",
"verify_aud",
"verify_iat",
"verify_exp",
"verify_nbf",
"verify_iss",
"verify_sub",
"verify_jti",
"verify_at_hash",
"leeway",
}
keys = set(self.jwt_decode_options.keys())
unknown = keys - valid_keys
if unknown:
raise ConfigurationError(f"unexpected key in jwt_decode_options: {', '.join(unknown)}")
def is_valid_authentication_type(self):
return True
def requires_client_login(self):
return True
def add_url_rules(self, app):
parse = urlparse(self.api_base_url)
app.add_url_rule(f"{parse.path}/login", "login", self.login, methods=["GET"])
app.add_url_rule(f"{parse.path}/logout", "logout", self.logout, methods=["GET"])
app.add_url_rule(f"{parse.path}/logout_redirect", "logout_redirect", self.logout_redirect, methods=["GET"])
app.add_url_rule(f"{parse.path}/oauth2/callback", "callback", self.callback, methods=["GET"])
def complete_setup(self, flask_app):
self.oauth = OAuth(flask_app)
self.client = self.oauth.register(
"auth0",
client_id=self.client_id,
client_secret=self.client_secret,
api_base_url=self.oauth_api_base_url,
refresh_token_url=f"{self.oauth_api_base_url}/oauth/token",
access_token_url=f"{self.oauth_api_base_url}/oauth/token",
authorize_url=f"{self.oauth_api_base_url}/authorize",
client_kwargs={"scope": "openid profile email offline_access"},
)
def is_user_authenticated(self):
payload = self.get_userinfo()
return payload is not None
def get_user_id(self):
payload = self.get_userinfo()
return payload.get("sub") if payload else None
def get_user_name(self):
payload = self.get_userinfo()
return payload.get("name") if payload else None
def get_user_email(self):
payload = self.get_userinfo()
return payload.get("email") if payload else None
def get_user_picture(self):
payload = self.get_userinfo()
return payload.get("picture") if payload else None
def update_response(self, response):
response.cache_control.update(dict(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True))
def login(self):
callbackurl = f"{self.api_base_url}/oauth2/callback"
return_path = request.args.get("dataset", "")
return_to = f"{self.web_base_url}/{return_path}"
# save the return path in the session cookie, accessed in the callback function
session["oauth_callback_redirect"] = return_to
response = self.client.authorize_redirect(redirect_uri=callbackurl)
self.update_response(response)
return response
def logout(self):
"""
We would like for the user to remain on the same dataset after logout. oauth requires that
the redirect `returnTo` path be whitelisted by the oauth server, therefore a level of
indirection is used. We first redirect to a single path "logout_redirect", and logout_redirect
will redirect the user's browser back to the current page.
"""
self.remove_tokens()
redirect_path = request.args.get("dataset", "")
redirect_to = f"{self.web_base_url}/{redirect_path}"
session["oauth_logout_redirect"] = redirect_to
return_to = f"{self.api_base_url}/logout_redirect"
params = {"returnTo": return_to, "client_id": self.client_id}
response = redirect(self.client.api_base_url + "/v2/logout?" + urlencode(params))
self.update_response(response)
return response
def logout_redirect(self):
oauth_logout_redirect = session.pop("oauth_logout_redirect", "/")
response = redirect(oauth_logout_redirect)
self.update_response(response)
return response
def callback(self):
data = self.client.authorize_access_token()
tokens = Tokens(
access_token=data.get("access_token"),
id_token=data.get("id_token"),
refresh_token=data.get("refresh_token"),
expires_at=data.get("expires_at"),
)
self.save_tokens(tokens)
oauth_callback_redirect = session.pop("oauth_callback_redirect", "/")
response = redirect(oauth_callback_redirect)
self.update_response(response)
return response
def get_tokens(self):
"""Extract the tokens from the cookie, and store them in the flask global context"""
if "tokens" in g:
return g.tokens
try:
if self.session_cookie:
value = session.get(self.CXG_TOKENS)
if value:
g.tokens = Tokens(**value)
else:
return None
else:
value = request.cookies.get(self.cookie_params["key"])
if value is None:
return None
value = base64.b64decode(value)
value = json.loads(value)
g.tokens = Tokens(**value)
except Exception:
# there are many types of exceptions that can be raise in the above section.
# It is impractical to list all the exceptions here, since that would be brittle.
# If an exception occurs, then return None, meaning that no token could be retrieved.
current_app.logger.warning(f"auth cookie is in the wrong format: {str(value)}")
g.pop("tokens", None)
return None
return g.tokens
def save_tokens(self, tokens):
g.tokens = tokens
if self.session_cookie:
session[self.CXG_TOKENS] = tokens.__dict__
else:
@after_this_request
def set_cookie(response):
args = self.cookie_params.copy()
value = base64.b64encode(json.dumps(tokens.__dict__).encode("utf-8"))
del args["key"]
try:
response.set_cookie(self.cookie_params["key"], value, **args)
except Exception as e:
raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e
return response
def remove_tokens(self):
g.pop("tokens", None)
if self.session_cookie:
if self.CXG_TOKENS in session:
del session[self.CXG_TOKENS]
else:
@after_this_request
def remove_cookie(response):
response.set_cookie(self.cookie_params["key"], "", expires=0)
self.update_response(response)
return response
def get_login_url(self, data_adaptor):
"""Return the url for the login route"""
if data_adaptor and current_app.app_config.is_multi_dataset():
return f"{self.api_base_url}/login?dataset={data_adaptor.uri_path}/"
else:
return f"{self.api_base_url}/login"
def get_logout_url(self, data_adaptor):
"""Return the url for the logout route"""
if data_adaptor and current_app.app_config.is_multi_dataset():
return f"{self.api_base_url}/logout?dataset={data_adaptor.uri_path}/"
else:
return f"{self.api_base_url}/logout"
def check_jwt_payload(self, id_token):
try:
unverified_header = jwt.get_unverified_header(id_token)
except JWTError:
return None
rsa_key = {}
for key in self.jwks["keys"]:
if key["kid"] == unverified_header["kid"]:
rsa_key = {
"kty": key["kty"],
"kid": key["kid"],
"use": key["use"],
"n": key.get("n"),
"e": key.get("e"),
}
if rsa_key:
try:
payload = jwt.decode(
id_token,
rsa_key,
algorithms=self.algorithms,
audience=self.audience,
issuer=self.oauth_api_base_url + "/",
options=self.jwt_decode_options,
)
return payload
except ExpiredSignatureError:
# This exception is handled in get_userinfo
raise
except JWTClaimsError as e:
raise AuthenticationError(f"invalid claims {str(e)}") from e
except JWTError as e:
raise AuthenticationError(f"invalid signature: {str(e)}") from e
raise AuthenticationError("Unable to find the appropriate key")
def get_userinfo(self):
if not has_request_context():
return None
# check if the userinfo has been retrieved already in this request
if "userinfo" in g:
return g.get("userinfo")
# if there is no id_token, return None (user is not authenticated)
tokens = self.get_tokens()
if tokens is None or tokens.id_token is None:
return None
try:
# check the jwt payload. This raises an AuthenticationError if the token is not valid.
# It the token has expired, we attempt to refresh the token
g.userinfo = self.check_jwt_payload(tokens.id_token)
return g.userinfo
except ExpiredSignatureError:
tokens = self.refresh_expired_token(tokens.refresh_token)
if tokens is None or tokens.id_token is None:
return None
else:
try:
g.userinfo = self.check_jwt_payload(tokens.id_token)
return g.userinfo
except JWTError as e:
raise AuthenticationError(f"error during token refresh: {str(e)}") from e
except AuthenticationError:
self.remove_tokens()
raise
def refresh_expired_token(self, refresh_token):
params = {
"grant_type": "refresh_token",
"client_id": self.client_id,
"refresh_token": refresh_token,
"client_secret": self.client_secret,
}
headers = {"content-type": "application/x-www-form-urlencoded"}
request = requests.post(f"{self.oauth_api_base_url}/oauth/token", urlencode(params), headers=headers)
if request.status_code != 200:
# unable to refresh the token, log the user out
self.remove_tokens()
return None
data = request.json()
tokens = Tokens(
access_token=data.get("access_token"),
id_token=data.get("id_token"),
refresh_token=data.get("refresh_token", refresh_token),
expires_at=data.get("expires_at"),
)
self.save_tokens(tokens)
return tokens
AuthTypeFactory.register("oauth", AuthTypeOAuth)
+40
View File
@@ -0,0 +1,40 @@
from flask import session
from uuid import uuid4
from backend.czi_hosted.auth.auth import AuthTypeBase, AuthTypeFactory
class AuthTypeSession(AuthTypeBase):
"""Session based authentication. The user is always logged. The user id is a random number
associated with the session. This is a good choice for desktop servers."""
# key in the session token for userid
CXGUID = "cxguid"
def __init__(self, app_config):
super().__init__()
def is_valid_authentication_type(self):
return True
def complete_setup(self, app):
pass
def is_user_authenticated(self):
# always authenticated
return True
def get_user_id(self):
if self.CXGUID not in session:
session[self.CXGUID] = uuid4().hex
session.permanent = True
return session[self.CXGUID]
def get_user_name(self):
return "anonymous"
def get_user_email(self):
return None
AuthTypeFactory.register("session", AuthTypeSession)
+80
View File
@@ -0,0 +1,80 @@
from flask import session, request, redirect, current_app
from backend.czi_hosted.auth.auth import AuthTypeClientBase, AuthTypeFactory
class AuthTypeTest(AuthTypeClientBase):
"""An authentication type for testing client based logins. When the login route is accessed
the user is automatically logged in with a default or configured username"""
# key in session token with userid and username
CXGUID = "cxguid_test"
CXGUNAME = "cxguname_test"
CXGUEMAIL = "cxguemail_test"
CXGUPICTURE = "cxgupicture_test"
def __init__(self, app_config):
super().__init__()
self.user_name = "test_account"
self.user_id = "id0001"
self.user_email = "test_account@test.com"
self.user_picture = None
def is_valid_authentication_type(self):
return True
def requires_client_login(self):
return True
def add_url_rules(self, app):
app.add_url_rule("/login", "login", self.login, methods=["GET"])
app.add_url_rule("/logout", "logout", self.logout, methods=["GET"])
def complete_setup(self, app):
pass
def is_user_authenticated(self):
return self.CXGUID in session
def get_user_id(self):
return session.get(self.CXGUID)
def get_user_name(self):
return session.get(self.CXGUNAME)
def get_user_email(self):
return session.get(self.CXGUEMAIL)
def get_user_picture(self):
return session.get(self.CXGUPICTURE)
def login(self):
args = request.args
return_to = args.get("dataset", "/")
session[self.CXGUID] = args.get("userid", self.user_id)
session[self.CXGUNAME] = args.get("username", self.user_name)
session[self.CXGUEMAIL] = args.get("email", self.user_email)
session[self.CXGUPICTURE] = args.get("picture", self.user_picture)
return redirect(return_to)
def logout(self):
session.clear()
return_to = request.args.get("dataset", "/")
return redirect(return_to)
def get_login_url(self, data_adaptor):
"""Return the url for the login route"""
if current_app.app_config.is_multi_dataset():
return f"/login?dataset={data_adaptor.uri_path}"
else:
return "/login"
def get_logout_url(self, data_adaptor):
"""Return the url for the logout route"""
if current_app.app_config.is_multi_dataset():
return f"/logout?dataset={data_adaptor.uri_path}"
else:
return "/logout"
AuthTypeFactory.register("test", AuthTypeTest)
+35
View File
@@ -0,0 +1,35 @@
import click
from .convert_to_cxg import convert_to_cxg
from .launch import launch
from .prepare import prepare
from .upgrade import log_upgrade_check
from .schema import schema_cli
from .. import __version__
@click.group(
name="cellxgene",
subcommand_metavar="COMMAND <args>",
options_metavar="<options>",
context_settings=dict(max_content_width=85, help_option_names=["-h", "--help"]),
)
@click.help_option("--help", "-h", help="Show this message and exit.")
@click.version_option(
version=__version__,
prog_name="cellxgene",
message="[%(prog)s] Version %(version)s",
help="Show the software version and exit.",
)
@click.option(
"--upgrade-check/--no-upgrade-check", default=True, show_default=True, help="Check for release upgrades on start.",
)
def cli(upgrade_check):
if upgrade_check:
log_upgrade_check()
cli.add_command(launch)
cli.add_command(prepare)
cli.add_command(convert_to_cxg)
cli.add_command(schema_cli)
+133
View File
@@ -0,0 +1,133 @@
from os import path
import click
from backend.czi_hosted.converters.h5ad_data_file import H5ADDataFile
@click.command(
name="convert",
short_help="Converts an H5AD dataset to the CXG format.",
help="Converts an H5AD dataset to the CXG format. The CXG format is a cellxgene-private data format "
"that has performance and access characteristics amenable to a multi-dataset, multi-user serving "
"environment. You will be able to launch the cellxgene using the `cellxgene launch` command as "
"usually with the generated CXG file.",
)
@click.argument(
"input-file", nargs=1, type=click.Path(exists=True, dir_okay=False),
)
@click.option(
"-o",
"--output-directory",
help="Name of the output CXG directory. If not provided, will default to be the input filename with a "
"CXG extension.",
)
@click.option(
"-b",
"--backed",
help="When true, loads the H5AD in file backed mode. This will cause the conversion to be slower, "
"but will use less memory.",
default=False,
show_default=True,
is_flag=True,
)
@click.option(
"-t",
"--title",
help="Human readable dataset title that will be included as metadata about the CXG file. If omitted, "
"the dataset title will be the filename.",
)
@click.option(
"-a",
"--about",
help="A fully qualified URL that provides more information about the dataset and will be included as "
"metadata about the CXG file.",
)
@click.option(
"-s",
"--sparse-threshold",
help="If the dataset's percent of non-zero values falls belows the specified threshold, then the X "
"array of the dataset will be sparse. Since the default value is 0.0, the default will be to "
"convert to dense array.",
default=0.0,
show_default=True,
)
@click.option(
"--obs-names",
help="Name to a column in the obs dataframe that will be used as the index for the dataframe instead of "
"the one designated by the dataframe generated-index.",
)
@click.option(
"--var-names",
help="Name to a column in the var dataframe that will be used as the index for the dataframe instead of "
"the one designated by the dataframe generated-index.",
)
@click.option(
"--disable-custom-colors",
help="When set, conversion process will not extract scanpy-compatible category colors from the H5AD file.",
default=False,
show_default=True,
is_flag=True,
)
@click.option(
"--disable-corpora-schema",
help="When set, conversion process will neither extract nor store Corpora schema information. See "
"https://github.com/chanzuckerberg/corpora-data-portal/blob/main/backend/schema/corpora_schema.md for more "
"information.",
default=False,
show_default=True,
is_flag=True,
)
@click.option(
"--overwrite",
help="When set to true, will overwrite the output file if the output file already exists.",
default=False,
show_default=True,
is_flag=True,
)
@click.help_option("--help", "-h", help="Show this message and exit.")
def convert_to_cxg(
input_file,
output_directory,
backed,
title,
about,
sparse_threshold,
obs_names,
var_names,
disable_custom_colors,
disable_corpora_schema,
overwrite,
):
"""
Convert a dataset file into CXG.
"""
h5ad_data_file = H5ADDataFile(
input_file, backed, title, about, obs_names, var_names, use_corpora_schema=not disable_corpora_schema
)
# Get the directory that will hold all the CXG files
cxg_output_container = get_output_directory(input_file, output_directory, overwrite)
h5ad_data_file.to_cxg(
cxg_output_container, sparse_threshold, convert_anndata_colors_to_cxg_colors=not disable_custom_colors
)
def get_output_directory(input_filename, output_directory, should_overwrite):
"""
Get the name of the CXG output directory to be created/populated during the dataset conversion.
"""
if output_directory and (not path.isdir(output_directory) or (path.isdir(output_directory) and should_overwrite)):
if output_directory.endswith(".cxg"):
return output_directory
return output_directory + ".cxg"
if output_directory and path.isdir(output_directory) and not should_overwrite:
raise click.BadParameter(
f"Output directory {output_directory} already exists. If you'd like to overwrite, then run the command "
f"with the --overwrite flag."
)
return path.splitext(input_filename)[0] + ".cxg"
+432
View File
@@ -0,0 +1,432 @@
import errno
import functools
import logging
import sys
import webbrowser
import os
import click
from flask_compress import Compress
from flask_cors import CORS
from backend.czi_hosted.default_config import default_config
from backend.czi_hosted.app.app import Server
from backend.czi_hosted.common.config.app_config import AppConfig
from backend.common.errors import DatasetAccessError, ConfigurationError
from backend.common.utils.utils import sort_options
DEFAULT_CONFIG = AppConfig()
def annotation_args(func):
@click.option(
"--disable-annotations",
is_flag=True,
default=not DEFAULT_CONFIG.default_dataset_config.user_annotations__enable,
show_default=True,
help="Disable user annotation of data.",
)
@click.option(
"--annotations-file",
default=DEFAULT_CONFIG.default_dataset_config.user_annotations__local_file_csv__file,
show_default=True,
multiple=False,
metavar="<path>",
help="CSV file to initialize editing of existing annotations; will be altered in-place. "
"Incompatible with --annotations-dir.",
)
@click.option(
"--annotations-dir",
default=DEFAULT_CONFIG.default_dataset_config.user_annotations__local_file_csv__directory,
show_default=False,
multiple=False,
metavar="<directory path>",
help="Directory of where to save output annotations; filename will be specified in the application. "
"Incompatible with --annotations-file.",
)
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
return wrapper
def config_args(func):
@click.option(
"--max-category-items",
default=DEFAULT_CONFIG.default_dataset_config.presentation__max_categories,
metavar="<integer>",
show_default=True,
help="Will not display categories with more distinct values than specified.",
)
@click.option(
"--disable-custom-colors",
is_flag=True,
default=False,
show_default=False,
help="Disable user-defined category-label colors drawn from source data file.",
)
@click.option(
"--diffexp-lfc-cutoff",
"-de",
default=DEFAULT_CONFIG.default_dataset_config.diffexp__lfc_cutoff,
show_default=True,
metavar="<float>",
help="Minimum log fold change threshold for differential expression.",
)
@click.option(
"--disable-diffexp",
is_flag=True,
default=not DEFAULT_CONFIG.default_dataset_config.diffexp__enable,
show_default=False,
help="Disable on-demand differential expression.",
)
@click.option(
"--embedding",
"-e",
default=DEFAULT_CONFIG.default_dataset_config.embeddings__names,
multiple=True,
show_default=False,
metavar="<text>",
help="Embedding name, eg, 'umap'. Repeat option for multiple embeddings. Defaults to all.",
)
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
return wrapper
def dataset_args(func):
@click.option(
"--obs-names",
"-obs",
default=DEFAULT_CONFIG.server_config.single_dataset__obs_names,
metavar="<text>",
help="Name of annotation field to use for observations. If not specified cellxgene will use the the obs index.",
)
@click.option(
"--var-names",
"-var",
default=DEFAULT_CONFIG.server_config.single_dataset__var_names,
metavar="<text>",
help="Name of annotation to use for variables. If not specified cellxgene will use the the var index.",
)
@click.option(
"--backed",
"-b",
is_flag=True,
default=DEFAULT_CONFIG.server_config.adaptor__anndata_adaptor__backed,
show_default=False,
help="Load anndata in file-backed mode. " "This may save memory, but may result in slower overall performance.",
)
@click.option(
"--title",
"-t",
default=DEFAULT_CONFIG.server_config.single_dataset__title,
metavar="<text>",
help="Title to display. If omitted will use file name.",
)
@click.option(
"--about",
default=DEFAULT_CONFIG.server_config.single_dataset__about,
metavar="<URL>",
help="URL providing more information about the dataset (hint: must be a fully specified absolute URL).",
)
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
return wrapper
def server_args(func):
@click.option(
"--debug",
"-d",
is_flag=True,
default=DEFAULT_CONFIG.server_config.app__debug,
show_default=True,
help="Run in debug mode. This is helpful for cellxgene developers, "
"or when you want more information about an error condition.",
)
@click.option(
"--verbose",
"-v",
is_flag=True,
default=DEFAULT_CONFIG.server_config.app__verbose,
show_default=True,
help="Provide verbose output, including warnings and all server requests.",
)
@click.option(
"--port",
"-p",
metavar="<port>",
default=DEFAULT_CONFIG.server_config.app__port,
type=int,
show_default=True,
help="Port to run server on. If not specified cellxgene will find an available port.",
)
@click.option(
"--host",
metavar="<IP address>",
default=DEFAULT_CONFIG.server_config.app__host,
show_default=False,
help="Host IP address. By default cellxgene will use localhost (e.g. 127.0.0.1).",
)
@click.option(
"--scripts",
"-s",
default=DEFAULT_CONFIG.default_dataset_config.app__scripts,
multiple=True,
metavar="<text>",
help="Additional script files to include in HTML page. If not specified, "
"no additional script files will be included.",
show_default=False,
)
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
return wrapper
def launch_args(func):
@annotation_args
@config_args
@dataset_args
@server_args
@click.option(
"--dataroot",
default=DEFAULT_CONFIG.server_config.multi_dataset__dataroot,
metavar="<data directory>",
help="Enable cellxgene to serve multiple files. Supply path (local directory or URL)"
" to folder containing H5AD and/or CXG datasets.",
hidden=True,
) # TODO, unhide when dataroot is supported)
@click.argument("datapath", required=False, metavar="<path to data file>")
@click.option(
"--open",
"-o",
"open_browser",
is_flag=True,
default=DEFAULT_CONFIG.server_config.app__open_browser,
show_default=True,
help="Open web browser after launch.",
)
@click.option(
"--config-file",
"-c",
"config_file",
default=None,
show_default=True,
help="Location to yaml file with configuration settings",
)
@click.option(
"--dump-default-config",
"dump_default_config",
is_flag=True,
default=False,
show_default=True,
help="Print default configuration settings and exit",
)
@click.help_option("--help", "-h", help="Show this message and exit.")
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
return wrapper
def handle_scripts(scripts):
if scripts:
click.echo(
r"""
/ / /\ \ \__ _ _ __ _ __ (_)_ __ __ _
\ \/ \/ / _` | '__| '_ \| | '_ \ / _` |
\ /\ / (_| | | | | | | | | | | (_| |
\/ \/ \__,_|_| |_| |_|_|_| |_|\__, |
|___/
The --scripts flag is intended for developers to include google analytics etc. You could be opening yourself to a
security risk by including the --scripts flag. Make sure you trust the scripts that you are including.
"""
)
scripts_pretty = ", ".join(scripts)
click.confirm(f"Are you sure you want to inject these scripts: {scripts_pretty}?", abort=True)
class CliLaunchServer(Server):
"""
the CLI runs a local web server, and needs to enable a few more features.
"""
def __init__(self, app_config):
super().__init__(app_config)
@staticmethod
def _before_adding_routes(app, app_config):
app.config["COMPRESS_MIMETYPES"] = [
"text/html",
"text/css",
"text/xml",
"application/json",
"application/javascript",
"application/octet-stream",
]
Compress(app)
if app_config.server_config.app__debug:
CORS(app, supports_credentials=True)
@sort_options
@click.command(
short_help="Launch the cellxgene data viewer. " "Run `cellxgene launch --help` for more information.",
options_metavar="<options>",
)
@launch_args
def launch(
datapath,
dataroot,
verbose,
debug,
open_browser,
port,
host,
embedding,
obs_names,
var_names,
max_category_items,
disable_custom_colors,
diffexp_lfc_cutoff,
title,
scripts,
about,
disable_annotations,
annotations_file,
annotations_dir,
backed,
disable_diffexp,
config_file,
dump_default_config,
):
"""Launch the cellxgene data viewer.
This web app lets you explore single-cell expression data.
Data must be in a format that cellxgene expects.
Read the "getting started" guide to learn more:
https://chanzuckerberg.github.io/cellxgene/getting-started.html
Examples:
> cellxgene launch example-dataset/pbmc3k.h5ad --title pbmc3k
> cellxgene launch <your data file> --title <your title>
> cellxgene launch <url>"""
# TODO Examples to provide when "--dataroot" is unhidden
# > cellxgene launch --dataroot example-dataset/
#
# > cellxgene launch --dataroot <url>
if dump_default_config:
print(default_config)
sys.exit(0)
# Startup message
click.echo("[cellxgene] Starting the CLI...")
# app config
app_config = AppConfig()
server_config = app_config.server_config
try:
if config_file:
app_config.update_from_config_file(config_file)
# Determine which config options were give on the command line.
# Those will override the ones provided in the config file (if provided).
cli_config = AppConfig()
cli_config.update_server_config(
app__verbose=verbose,
app__debug=debug,
app__host=host,
app__port=port,
app__open_browser=open_browser,
single_dataset__datapath=datapath,
single_dataset__title=title,
single_dataset__about=about,
single_dataset__obs_names=obs_names,
single_dataset__var_names=var_names,
multi_dataset__dataroot=dataroot,
adaptor__anndata_adaptor__backed=backed,
)
cli_config.update_default_dataset_config(
app__scripts=scripts,
user_annotations__enable=not disable_annotations,
user_annotations__local_file_csv__file=annotations_file,
user_annotations__local_file_csv__directory=annotations_dir,
presentation__max_categories=max_category_items,
presentation__custom_colors=not disable_custom_colors,
embeddings__names=embedding,
diffexp__enable=not disable_diffexp,
diffexp__lfc_cutoff=diffexp_lfc_cutoff,
)
diff = cli_config.server_config.changes_from_default()
changes = {key: val for key, val, _ in diff}
app_config.update_server_config(**changes)
diff = cli_config.default_dataset_config.changes_from_default()
changes = {key: val for key, val, _ in diff}
app_config.update_default_dataset_config(**changes)
# process the configuration
# any errors will be thrown as an exception.
# any info messages will be passed to the messagefn function.
def messagefn(message):
click.echo("[cellxgene] " + message)
# Use a default secret if one is not provided
if not server_config.app__flask_secret_key:
app_config.update_server_config(app__flask_secret_key="SparkleAndShine")
app_config.complete_config(messagefn)
except (ConfigurationError, DatasetAccessError) as e:
raise click.ClickException(e)
handle_scripts(scripts)
# create the server
server = CliLaunchServer(app_config)
if not server_config.app__verbose:
log = logging.getLogger("werkzeug")
log.setLevel(logging.ERROR)
cellxgene_url = f"http://{app_config.server_config.app__host}:{app_config.server_config.app__port}"
if server_config.app__open_browser:
click.echo(f"[cellxgene] Launching! Opening your browser to {cellxgene_url} now.")
webbrowser.open(cellxgene_url)
else:
click.echo(f"[cellxgene] Launching! Please go to {cellxgene_url} in your browser.")
click.echo("[cellxgene] Type CTRL-C at any time to exit.")
if not server_config.app__verbose:
f = open(os.devnull, "w")
sys.stdout = f
try:
server.app.run(
host=server_config.app__host,
debug=server_config.app__debug,
port=server_config.app__port,
threaded=not server_config.app__debug,
use_debugger=False,
use_reloader=False,
)
except OSError as e:
if e.errno == errno.EADDRINUSE:
raise click.ClickException("Port is in use, please specify an open port using the --port flag.") from e
raise
@@ -5,7 +5,7 @@ import pandas as pd
from numpy import ndarray, unique
from scipy.sparse.csc import csc_matrix
from server.common.utils.utils import sort_options
from backend.common.utils.utils import sort_options
@sort_options
@@ -24,11 +24,7 @@ from server.common.utils.utils import sort_options
show_default=True,
)
@click.option(
"--recipe",
"-r",
default="none",
type=click.Choice(["none", "seurat", "zheng17"]),
show_default=True,
"--recipe", "-r", default="none", type=click.Choice(["none", "seurat", "zheng17"]), show_default=True,
)
@click.option("--output", "-o", default="", help="Save a new file to filename.", metavar="<filename>")
@click.option("--plotting", "-p", default=False, is_flag=True, help="Generate plots.", show_default=True)
+72
View File
@@ -0,0 +1,72 @@
import click
from backend.czi_hosted.converters.schema import remix, validate
@click.group(
name="schema",
subcommand_metavar="COMMAND <args>",
short_help="Apply and validate the cellxgene data integration schema to an h5ad file.",
context_settings=dict(max_content_width=85, help_option_names=["-h", "--help"]),
)
def schema_cli():
try:
import scanpy # noqa: F401
except ImportError:
raise click.ClickException(
"[cellxgene] cellxgene schema requires scanpy"
)
@click.command(
name="apply",
short_help="(experimental) Apply the cellxgene data integration schema to an h5ad.",
help="(experimental) Using a yaml file that describes schema values to insert or convert and in input "
"h5ad file, apply the schema changes and create a new, conforming h5ad.",
)
@click.option(
"--source-h5ad",
help="Input h5ad file.",
nargs=1,
required=True,
type=click.Path(exists=True, dir_okay=False),
)
@click.option(
"--remix-config",
help="Config yaml with information on how to apply the schema.",
nargs=1,
required=True,
type=click.Path(exists=True, dir_okay=False),
)
@click.option(
"--output-filename",
help="Filename for the new, schema-conforming h5ad file.",
required=True,
nargs=1
)
def schema_apply(source_h5ad, remix_config, output_filename):
remix.apply_schema(source_h5ad, remix_config, output_filename)
@click.command(
name="validate",
short_help="(experimental) Check that an h5ad follows the cellxgene data integration schema.",
)
@click.argument(
"h5ad",
nargs=1,
type=click.Path(exists=True, dir_okay=False),
)
@click.option(
"--shallow",
help="When true, just check that the correct version information is present.",
default=False,
show_default=True,
is_flag=True,
)
def schema_validate(h5ad, shallow):
validate.validate(h5ad, shallow)
schema_cli.add_command(schema_apply)
schema_cli.add_command(schema_validate)
@@ -0,0 +1,110 @@
import os
from flask import current_app, has_request_context
from backend.common.errors import DisabledFeatureError
from backend.common.utils.type_conversion_utils import get_schema_type_hint_of_array
from backend.common.genesets import write_gene_sets_tidycsv, read_gene_sets_tidycsv, validate_gene_sets
from backend.common.utils.data_locator import DataLocator
from backend.common.utils.utils import path_join
class Annotations:
"""baseclass for annotations and genesets"""
def __init__(self, config={}):
self.config = config
def user_annotations_enabled(self):
return self.config.get("user-annotations", False)
def check_user_annotations_enabled(self):
if not self.user_annotations_enabled():
raise DisabledFeatureError("User annotations are disabled.")
def get_schema(self, data_adaptor):
schema = []
labels = self.read_labels(data_adaptor)
if labels is not None and not labels.empty:
for col in labels.columns:
col_schema = dict(name=col, writable=True)
col_schema.update(get_schema_type_hint_of_array(labels[col]))
schema.append(col_schema)
return schema
def set_collection(self, name):
"""set or create a new annotation collection"""
raise NotImplementedError
def read_labels(self, data_adaptor):
"""Return the labels as a pandas.DataFrame"""
raise NotImplementedError
def write_labels(self, df, data_adaptor):
"""Write the labels (df) to a persistent storage such that it can later be read"""
raise NotImplementedError
def update_parameters(self, parameters, data_adaptor):
"""Update configuration parameters that describe information about the annotations feature"""
params = {}
params["annotations_genesets_readonly"] = True
params["annotations_genesets_name_is_read_only"] = True
parameters.update(params)
@staticmethod
def gene_sets_to_csv(genesets):
"""
Convert the internal genesets format (returned by read_gene_set) into
the simple Tidy CSV.
"""
from io import StringIO
if isinstance(genesets, dict):
genesets = genesets.values()
with StringIO() as sio:
write_gene_sets_tidycsv(sio, genesets)
return sio.getvalue()
@staticmethod
def gene_sets_to_response(genesets):
"""
Convert the internal genesets format (returned by read_gene_set) into
the dict expected by the JSON REST API
"""
return list(genesets.values())
def read_gene_sets(self, data_adaptor, context=None):
if has_request_context():
if not current_app.auth.is_user_authenticated():
return ({}, 0)
gene_sets_uri_or_path = dataset_uri_to_geneset_uri(data_adaptor.data_locator.uri_or_path)
server_config = data_adaptor.server_config
region_name = None if server_config is None else server_config.data_locator__s3__region_name
gene_sets_locator = DataLocator(gene_sets_uri_or_path, region_name=region_name)
if not gene_sets_locator.exists():
return ({}, 0)
gene_sets = read_gene_sets_tidycsv(gene_sets_locator, context)
schema = data_adaptor.get_schema()
var_index = schema["annotations"]["var"].get("index", "index")
var_names = set(data_adaptor.query_var_array(var_index))
gene_sets = validate_gene_sets(gene_sets, var_names)
return (gene_sets, 0)
def dataset_uri_to_geneset_uri(data_uri_or_path):
"""given a dataset URI, return the associated gene set URI"""
data_basename = os.path.basename(data_uri_or_path)
base, ext = os.path.splitext(data_basename)
if ext is not None: # strip extension, if any
data_basename = base
genesets_basename = f"{data_basename}-genesets.csv"
gene_sets_uri_or_path = path_join(data_uri_or_path, "..", genesets_basename)
return gene_sets_uri_or_path
@@ -0,0 +1,167 @@
import json
import os
import re
import time
import pandas as pd
import tiledb
from flask import current_app
from backend.czi_hosted.common.annotations.annotations import Annotations
from backend.common.errors import AnnotationCategoryNameError
from backend.czi_hosted.common.utils.sanitization_utils import sanitize_values_in_list
from backend.common.utils.type_conversion_utils import get_dtypes_and_schemas_of_dataframe, get_encoding_dtype_of_array
from backend.czi_hosted.db.cellxgene_orm import Annotation
class AnnotationsHostedTileDB(Annotations):
CXG_ANNO_COLLECTION = "cxg_anno_collection"
def __init__(self, config, directory_path, db):
super().__init__(config)
self.db = db
if directory_path[-1] == "/":
self.directory_path = directory_path
else:
self.directory_path = directory_path + "/"
def check_category_names(self, df):
original_category_names = df.keys().to_list()
sanitized_category_names = set(sanitize_values_in_list(original_category_names).values())
unsanitary_original_category_names = set(original_category_names).difference(sanitized_category_names)
if unsanitary_original_category_names:
raise AnnotationCategoryNameError(
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
this is ultra conservative. If we want to allow full legal file name syntax,
we could look at modules like `pathvalidate`
"""
if name is None:
return False
return re.match(r"^[\w\-]+$", name) is not None
def set_collection(self, name):
self.CXG_ANNO_COLLECTION = name
def read_labels(self, data_adaptor):
user_id = self.get_user_id()
if user_id is None:
return
dataset_name = data_adaptor.get_location()
dataset_id = self.db.get_or_create_dataset(dataset_name)
annotation_object = self.db.query_for_most_recent(
Annotation, [Annotation.user_id == user_id, Annotation.dataset_id == dataset_id]
)
if annotation_object:
if annotation_object.tiledb_uri == "":
# this mean the user has removed all the categories.
return None
try:
df = tiledb.open(annotation_object.tiledb_uri)
except tiledb.TileDBError:
# don't crash if the annotations file is missing or can't be read.
current_app.logger.warning(f"Cannot read annotation file: {annotation_object.tiledb_uri}")
return None
pandas_df = self.convert_to_pandas_df(df, annotation_object.schema_hints)
return pandas_df
else:
return None
def convert_to_pandas_df(self, tileDBArray, schema_hints):
repr_meta = None
index_dims = None
schema_hints = json.loads(schema_hints)
if "__pandas_attribute_repr" in tileDBArray.meta:
# backwards compatibility... unsure if necessary at this point
repr_meta = json.loads(tileDBArray.meta["__pandas_attribute_repr"])
if "__pandas_index_dims" in tileDBArray.meta:
index_dims = json.loads(tileDBArray.meta["__pandas_index_dims"])
data = tileDBArray[:]
indexes = list()
for col_name, col_val in data.items():
# If the column values are byte literals, decode them
if isinstance(col_val[0], bytes):
col_val = [value.decode("utf-8") for value in col_val]
if schema_hints and col_name in schema_hints:
type = schema_hints.get(col_name).get("type")
if type and type == "categorical":
new_col = pd.Series(col_val, dtype="category")
data[col_name] = new_col
elif repr_meta and col_name in repr_meta:
new_col = pd.Series(col_val, dtype=repr_meta[col_name])
data[col_name] = new_col
elif index_dims and col_name in index_dims:
new_col = pd.Series(col_val, dtype=index_dims[col_name])
data[col_name] = new_col
indexes.append(col_name)
new_df = pd.DataFrame.from_dict(data)
if len(indexes) > 0:
new_df.set_index(indexes, inplace=True)
return new_df
def write_labels(self, df, data_adaptor):
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)
dataset_name = data_adaptor.get_title()
user_id = self.db.get_or_create_user(auth_user_id)
"""
NOTE: The uri contains the dataset name, user name and a timestamp as a convenience for debugging purposes.
People may have the same name and time.time() can be server dependent.
See - https://docs.python.org/2/library/time.html#time.time
The annotations objects in the database should be used as the source of truth about who an annotation belongs
to (for authorization purposes) and what time it was created (for garbage collection).
"""
uri = f"{self.directory_path}{dataset_name}/{user_name}/{timestamp}"
if uri.startswith("s3://"):
pass
else:
os.makedirs(uri, exist_ok=True)
_, dataframe_schema_type_hints = get_dtypes_and_schemas_of_dataframe(df)
if not df.empty:
self.check_category_names(df)
# convert to tiledb datatypes
for col in df:
df[col] = df[col].astype(get_encoding_dtype_of_array(df[col]))
tiledb.from_pandas(uri, df, sparse=True)
else:
uri = ""
annotation = Annotation(
tiledb_uri=uri,
user_id=user_id,
dataset_id=str(dataset_id),
schema_hints=json.dumps(dataframe_schema_type_hints),
)
self.db.session.add(annotation)
self.db.session.commit()
def update_parameters(self, parameters, data_adaptor):
super().update_parameters(parameters, data_adaptor)
params = {}
params["annotations"] = True
params["user_annotation_collection_name_enabled"] = False
parameters.update(params)
@@ -0,0 +1,192 @@
import base64
import os
import re
import threading
from datetime import datetime
from hashlib import blake2b
import pandas as pd
from flask import session, has_request_context, current_app
from backend.czi_hosted import __version__ as cellxgene_version
from backend.czi_hosted.common.annotations.annotations import Annotations
from backend.common.errors import AnnotationsError
class AnnotationsLocalFile(Annotations):
CXG_ANNO_COLLECTION = "cxg_anno_collection"
def __init__(self, config, output_dir, output_file):
super().__init__(config)
self.output_dir = output_dir
self.output_file = output_file
# lock used to protect label file write ops
self.label_lock = threading.RLock()
# cache the most recent annotations
self.last_fname = None
self.last_labels = None
def is_safe_collection_name(self, name):
"""
return true if this is a safe collection name
this is ultra conservative. If we want to allow full legal file name syntax,
we could look at modules like `pathvalidate`
"""
if name is None:
return False
return re.match(r"^[\w\-]+$", name) is not None
def set_collection(self, name):
session[self.CXG_ANNO_COLLECTION] = name
session.permanent = True
def get_collection(self):
if session is None:
return None
return session.get(self.CXG_ANNO_COLLECTION)
def read_labels(self, data_adaptor):
if has_request_context():
if not current_app.auth.is_user_authenticated():
return pd.DataFrame()
fname = self._get_filename(data_adaptor)
with self.label_lock:
if fname is not None and os.path.exists(fname) and os.path.getsize(fname) > 0:
# returned the cached labels if possible, otherwise read them from the file
if fname == self.last_fname:
return self.last_labels
else:
labels = pd.read_csv(
fname, dtype="category", index_col=0, header=0, comment="#", keep_default_na=False
)
# update the cache
self.last_fname = fname
self.last_labels = labels
return labels
else:
return pd.DataFrame()
def write_labels(self, df, data_adaptor):
# update our internal state and save it. Multi-threading often enabled,
# so treat this as a critical section.
with self.label_lock:
lastmod = data_adaptor.get_last_mod_time()
lastmodstr = "'unknown'" if lastmod is None else lastmod.isoformat(timespec="seconds")
header = (
f"# Annotations generated on {datetime.now().isoformat(timespec='seconds')} "
f"using cellxgene version {cellxgene_version}\n"
f"# Input data file was {data_adaptor.get_location()}, "
f"which was last modified on {lastmodstr}\n"
)
fname = self._get_filename(data_adaptor)
self._backup(fname)
if not df.empty:
with open(fname, "w", newline="") as f:
if header is not None:
f.write(header)
df.to_csv(f)
else:
open(fname, "w").close()
# update the cache
self.last_fname = fname
self.last_labels = df
def _get_userdata_idhash(self, data_adaptor):
"""
Return a short hash that weakly identifies the user and dataset.
Used to create safe annotations output file names.
"""
uid = current_app.auth.get_user_id()
id = (uid + data_adaptor.get_location()).encode()
idhash = base64.b32encode(blake2b(id, digest_size=5).digest()).decode("utf-8")
return idhash
def _get_output_dir(self):
if self.output_dir:
return self.output_dir
if self.output_file:
return os.path.dirname(self.path.abspath(self.output_dir))
return os.getcwd()
def _get_filename(self, data_adaptor):
"""return the current annotation file name"""
if self.output_file:
return self.output_file
# we need to generate a file name, which we can only do if we have a UID and collection name
if session is None:
raise AnnotationsError("unable to determine file name for annotations")
collection = self.get_collection()
if collection is None:
return None
if data_adaptor is None:
raise AnnotationsError("unable to determine file name for annotations")
idhash = self._get_userdata_idhash(data_adaptor)
return os.path.join(self._get_output_dir(), f"{collection}-{idhash}.csv")
def _backup(self, fname, max_backups=9):
"""
save N backups of file to backup_dir.
1. fname -> backup_dir/fname-TIME
2. delete excess files in backup_dir
"""
root, ext = os.path.splitext(fname)
backup_dir = f"{root}-backups"
# Make sure there is work to do
if not os.path.exists(fname):
return
# Ensure backup_dir exists
if not os.path.exists(backup_dir):
os.mkdir(backup_dir)
# Save current file to backup_dir
fname_base = os.path.basename(fname)
fname_base_root, fname_base_ext = os.path.splitext(fname_base)
# don't use ISO standard time format, as it contains characters illegal on some filesytems.
nowish = datetime.now().strftime("%Y-%m-%dT%H-%M-%S")
backup_fname = os.path.join(backup_dir, f"{fname_base_root}-{nowish}{fname_base_ext}")
if os.path.exists(backup_fname):
os.remove(backup_fname)
os.rename(fname, backup_fname)
# prune the backup_dir to max number of backup files, keeping the most recent backups
backups = list(filter(lambda s: s.startswith(fname_base_root), os.listdir(backup_dir)))
excess_count = len(backups) - max_backups
if excess_count > 0:
backups.sort()
for bu in backups[0:excess_count]:
os.remove(os.path.join(backup_dir, bu))
def update_parameters(self, parameters, data_adaptor):
super().update_parameters(parameters, data_adaptor)
params = {}
params["annotations"] = True
params["user_annotation_collection_name_enabled"] = True
if self.output_file is not None:
# user has hard-wired the name of the annotation data collection
fname = os.path.basename(self.output_file)
collection_fname = os.path.splitext(fname)[0]
params["annotations-data-collection-is-read-only"] = True
params["annotations-data-collection-name"] = collection_fname
elif session is not None:
collection = self.get_collection()
if current_app.auth.is_user_authenticated():
params["annotations-user-data-idhash"] = self._get_userdata_idhash(data_adaptor)
params["annotations-data-collection-is-read-only"] = not self.user_annotations_enabled()
params["annotations-data-collection-name"] = collection
parameters.update(params)
@@ -0,0 +1,4 @@
from backend.common.utils.aws_secret_utils import get_secret_key # noqa F504
DEFAULT_SERVER_PORT = 5005
BIG_FILE_SIZE_THRESHOLD = 100 * 2 ** 20 # 100MB
@@ -0,0 +1,247 @@
import yaml
from flatten_dict import unflatten
from backend.czi_hosted.common.config.external_config import ExternalConfig
from backend.czi_hosted.common.config.dataset_config import DatasetConfig
from backend.czi_hosted.common.config.server_config import ServerConfig
from backend.common.errors import ConfigurationError
from backend.czi_hosted.default_config import get_default_config
class AppConfig(object):
"""
AppConfig stores all the configuration for cellxgene.
AppConfig contains one or more DatasetConfig(s) and one ServerConfig.
The server_config contains attributes that refer to the server process as a whole.
The default_dataset_config refers to attributes that are associated with the features and
presentations of a dataset.
The dataset config attributes can be overridden depending on the url by which the
dataset was accessed. These are stored in dataroot_config.
AppConfig has methods to initialize, modify, and access the configuration.
"""
def __init__(self):
# the default configuration (see default_config.py)
# TODO @madison -- if we always read from the default config (hard coded path) can we set those values as
# defaults within the config class?
self.default_config = get_default_config()
# the server configuration
self.server_config = ServerConfig(self, self.default_config["server"])
# the dataset config, unless overridden by an entry in dataroot_config
self.default_dataset_config = DatasetConfig(None, self, self.default_config["dataset"])
# a dictionary of keys to DatasetConfig objects. Each key must exist in the multi_dataset__dataroot
# attribute of the server_config. The default dataset config will apply to all datasets unless a different set
# of config vars was passed for a specific dataset under the multidataset config. For example:
"""
per_dataset_config:
d1:
user_annotations:
enable: false
d2:
user_annotations:
enable: true
"""
# dataroot config
self.dataroot_config = {}
# external config
self.external_config = ExternalConfig(self, self.default_config["external"])
# Set to true when config_completed is called
self.is_completed = False
def get_dataset_config(self, dataroot_key):
if self.server_config.single_dataset__datapath:
return self.default_dataset_config
else:
return self.dataroot_config.get(dataroot_key, self.default_dataset_config)
def check_config(self):
"""Verify all the attributes in the config have been type checked"""
if not self.is_completed:
raise ConfigurationError("The configuration has not been completed")
self.server_config.check_config()
self.default_dataset_config.check_config()
for dataset_config in self.dataroot_config.values():
dataset_config.check_config()
self.external_config.check_config()
def update_server_config(self, **kw):
self.server_config.update(**kw)
self.is_completed = False
def update_default_dataset_config(self, **kw):
self.default_dataset_config.update(**kw)
# update all the other dataset configs, if any
for value in self.dataroot_config.values():
value.update(**kw)
self.is_completed = False
def update_single_config_from_path_and_value(self, path, value):
"""Update a single config parameter with the value.
Path is a list of string, that gives a path to the config parameter to be updated.
For example, path may be ["server","app","port"].
"""
self.is_completed = False
if not isinstance(path, list):
raise ConfigurationError(f"path must be a list of strings, got '{str(path)}'")
for part in path:
if not isinstance(part, str):
raise ConfigurationError(f"path must be a list of strings, got '{str(path)}'")
if len(path) < 1 or path[0] not in ("server", "dataset", "per_dataset_config"):
raise ConfigurationError("path must start with 'server', 'dataset', or 'per_dataset_config'")
if path[0] == "server":
attr = "__".join(path[1:])
try:
self.update_server_config(**{attr: value})
except ConfigurationError:
raise ConfigurationError(f"unknown config parameter at path: '{str(path)}'")
elif path[0] == "dataset":
attr = "__".join(path[1:])
try:
self.update_default_dataset_config(**{attr: value})
except ConfigurationError:
raise ConfigurationError(f"unknown config parameter at path: '{str(path)}'")
elif path[0] == "per_dataset_config":
if len(path) < 2:
raise ConfigurationError(f"missing dataroot when using per_dataset_config: got '{path}'")
dataroot = path[1]
if dataroot not in self.dataroot_config:
dataroots = str(list(self.dataroot_config.keys()))
raise ConfigurationError(
f"unknown dataroot when using per_dataset_config: got '{path}',"
f" dataroots specified in config are {dataroots}"
)
attr = "__".join(path[2:])
try:
self.dataroot_config[dataroot].update(**{attr: value})
except ConfigurationError:
raise ConfigurationError(f"unknown config parameter at path: '{str(path)}'")
def update_from_config_file(self, config_file):
try:
with open(config_file) as yml_file:
config = yaml.safe_load(yml_file)
except yaml.YAMLError as e:
raise ConfigurationError(f"The specified config file contained an error: {e}")
except OSError as e:
raise ConfigurationError(f"Issue retrieving the specified config file: {e}")
if config.get("server"):
self.server_config.update_from_config(config["server"], "server")
if config.get("dataset"):
self.default_dataset_config.update_from_config(config["dataset"], "dataset")
per_dataset_config = config.get("per_dataset_config", {})
for key, dataroot_config in per_dataset_config.items():
# first create and initialize the dataroot with the default config
self.add_dataroot_config(key, **config["dataset"])
# then apply the per dataset configuration
self.dataroot_config[key].update_from_config(dataroot_config, f"per_dataset_config__{key}")
if config.get("external"):
self.external_config.update_from_config(config["external"], "external")
self.is_completed = False
def config_to_dict(self):
"""return the configuration as an unflattened dict"""
server = self.server_config.create_mapping(self.server_config.default_config)
dataset = self.default_dataset_config.create_mapping(self.default_dataset_config.default_config)
external = self.external_config.create_mapping(self.external_config.default_config)
config = dict(server={}, dataset={})
for attrname in server.keys():
config["server__" + attrname] = getattr(self.server_config, attrname)
for attrname in dataset.keys():
config["dataset__" + attrname] = getattr(self.default_dataset_config, attrname)
if self.dataroot_config:
config["per_dataset_config"] = {}
for dataroot_tag, dataroot_config in self.dataroot_config.items():
dataset = dataroot_config.create_mapping(dataroot_config.default_config)
for attrname in dataset.keys():
config[f"per_dataset_config__{dataroot_tag}__" + attrname] = getattr(dataroot_config, attrname)
for attrname in external.keys():
config["external__" + attrname] = getattr(self.external_config, attrname)
config = unflatten(config, splitter=lambda key: key.split("__"))
return config
def write_config(self, config_file):
"""output the config to a yaml file"""
config = self.config_to_dict()
yaml.dump(config, open(config_file, "w"))
def changes_from_default(self):
"""Return all the attribute that are different from the default"""
diff_server = self.server_config.changes_from_default()
diff_dataset = self.default_dataset_config.changes_from_default()
diff_external = self.external.changes_from_default()
diff = dict(server=diff_server, dataset=diff_dataset, external=diff_external)
return diff
def add_dataroot_config(self, dataroot_tag, **kw):
"""Create a new dataset config object based on the default dataset config, and kw parameters"""
if dataroot_tag in self.dataroot_config:
raise ConfigurationError(f"dataroot config already exists: {dataroot_tag}")
if type(self.server_config.multi_dataset__dataroot) != dict:
raise ConfigurationError("The server__multi_dataset__dataroot must be a dictionary")
if dataroot_tag not in self.server_config.multi_dataset__dataroot:
raise ConfigurationError(f"The dataroot_tag ({dataroot_tag}) not found in server__multi_dataset__dataroot")
self.is_completed = False
self.dataroot_config[dataroot_tag] = DatasetConfig(dataroot_tag, self, self.default_config["dataset"])
flat_config = self.default_dataset_config.create_mapping(self.default_dataset_config.default_config)
config = {key: value[1] for key, value in flat_config.items()}
self.dataroot_config[dataroot_tag].update(**config)
self.dataroot_config[dataroot_tag].update_from_config(kw, dataroot_tag)
def complete_config(self, messagefn=None):
"""The configure options are checked, and any additional setup based on the config
parameters is done"""
if messagefn is None:
def noop(message):
pass
messagefn = noop
# TODO: to give better error messages we can add a mapping between where each config
# attribute originated (e.g. command line argument or config file), then in the error
# messages we can give correct context for attributes with bad value.
context = dict(messagefn=messagefn)
# complete config for external_config first, since this may update values in the other sections
self.external_config.complete_config(context)
self.server_config.complete_config(context)
self.default_dataset_config.complete_config(context)
for dataroot_config in self.dataroot_config.values():
dataroot_config.complete_config(context)
self.is_completed = True
self.check_config()
def get_matrix_data_cache_manager(self):
return self.server_config.matrix_data_cache_manager
def is_multi_dataset(self):
return self.server_config.multi_dataset__dataroot is not None
def get_title(self, data_adaptor):
return (
self.server_config.single_dataset__title
if self.server_config.single_dataset__title
else data_adaptor.get_title()
)
def get_about(self, data_adaptor):
return (
self.server_config.single_dataset__about
if self.server_config.single_dataset__about
else data_adaptor.get_about()
)
@@ -0,0 +1,132 @@
import copy
from flatten_dict import flatten
from backend.common.errors import ConfigurationError
class BaseConfig(object):
"""
This class handles the mechanics of updating and checking attributes.
Derived classes are expected to store the actual attributes
Currently DatasetConfig and ServerConfig both inherit from BaseConfig.
"""
def __init__(self, app_config, default_config, dictval_cases={}):
# reference back to the app_config
self.app_config = app_config
# the complete set of attributes and their default values (unflattened)
self.default_config = default_config
# attributes where the value may be a dict (and therefore are not flattened)
self.dictval_cases = dictval_cases
# used to make sure every attribute value is checked
self.attr_checked = {key_name: False for key_name in self.create_mapping(default_config).keys()}
def create_mapping(self, config):
"""
Create a dictionary where the keys are the name of attributes (using double underscore convention)
For example: authentication__type
The values are a tuple,
- the first item of the tuple is a tuple of path elements (location in config 'tree')
- the second item is the value of the config parameter
For example: (('authentication', 'type'), 'session'))
"""
config_copy = copy.deepcopy(config)
mapping = {}
# special cases where the value could be a dict.
# If its value is not None, the entry is added to the mapping, and not included
# in the flattening below.
for dictval_case in self.dictval_cases:
cur = config_copy
for part in dictval_case[:-1]:
cur = cur.get(part, {})
val = cur.get(dictval_case[-1])
if val is not None:
key = "__".join(dictval_case)
mapping[key] = (dictval_case, val)
del cur[dictval_case[-1]]
flat_config = flatten(config_copy)
for key, value in flat_config.items():
# name of the attribute
attr = "__".join(key)
mapping[attr] = (key, value)
return mapping
def validate_correct_type_of_configuration_attribute(self, attrname, vtype):
val = getattr(self, attrname)
if type(vtype) in (list, tuple):
if type(val) not in vtype:
tnames = ",".join([x.__name__ for x in vtype])
raise ConfigurationError(
f"Invalid type for attribute: {attrname}, expected types ({tnames}), got {type(val).__name__}"
)
else:
if type(val) != vtype:
raise ConfigurationError(
f"Invalid type for attribute: {attrname}, "
f"expected type {vtype.__name__}, got {type(val).__name__}"
)
self.attr_checked[attrname] = True
def check_config(self):
mapping = self.create_mapping(self.default_config)
for key in mapping.keys():
if not self.attr_checked[key]:
raise ConfigurationError(f"The attr '{key}' has not been checked")
def update(self, **kw):
"""Update the attributes defined in kw with their new values."""
for key, value in kw.items():
if not hasattr(self, key):
# check if the key is setting into a dictval entry.
found_dictval = False
for dictval in self.dictval_cases:
dictvalname = "__".join(dictval)
if dictvalname + "__" in key:
dictkey = key[len(dictvalname) + 2 :]
curdictval = getattr(self, dictvalname)
if curdictval is None:
setattr(self, dictvalname, dict(dictkey=value))
else:
curdictval[dictkey] = value
found_dictval = True
break
if found_dictval:
continue
raise ConfigurationError(f"unknown config parameter {key}.")
try:
if type(value) == tuple:
# convert tuple values to list values
value = list(value)
setattr(self, key, value)
except KeyError:
raise ConfigurationError(f"Unable to set config parameter {key}.")
self.attr_checked[key] = False
def update_from_config(self, config, prefix):
mapping = self.create_mapping(config)
for attr, (key, value) in mapping.items():
if not hasattr(self, attr):
raise ConfigurationError(f"Unknown key from config file: {prefix}__{attr}")
setattr(self, attr, value)
self.attr_checked[attr] = False
def changes_from_default(self):
"""Return all the attribute that are different from the default"""
mapping = self.create_mapping(self.default_config)
diff = []
for attrname, (key, defval) in mapping.items():
curval = getattr(self, attrname)
if curval != defval:
diff.append((attrname, curval, defval))
return diff
@@ -0,0 +1,121 @@
from backend.czi_hosted import display_version as cellxgene_display_version
def get_client_config(app_config, data_adaptor):
"""
Return the configuration as required by the /config REST route
"""
server_config = app_config.server_config
dataset_config = data_adaptor.dataset_config
annotation = dataset_config.user_annotations
auth = server_config.auth
# FIXME The current set of config is not consistently presented:
# we have camalCase, hyphen-text, and underscore_text
# make sure the configuration has been checked.
app_config.check_config()
# display_names
title = app_config.get_title(data_adaptor)
about = app_config.get_about(data_adaptor)
display_names = dict(engine=data_adaptor.get_name(), dataset=title)
# library_versions
library_versions = {}
library_versions.update(data_adaptor.get_library_versions())
library_versions["cellxgene"] = cellxgene_display_version
# links
links = {"about-dataset": about}
# parameters
parameters = {
"layout": dataset_config.embeddings__names,
"max-category-items": dataset_config.presentation__max_categories,
"obs_names": server_config.single_dataset__obs_names,
"var_names": server_config.single_dataset__var_names,
"diffexp_lfc_cutoff": dataset_config.diffexp__lfc_cutoff,
"backed": server_config.adaptor__anndata_adaptor__backed,
"disable-diffexp": not dataset_config.diffexp__enable,
"annotations": False,
"annotations_file": None,
"annotations_dir": None,
"annotations_genesets": True, # feature flag
"annotations_genesets_readonly": True,
"annotations_genesets_summary_methods": ["mean"],
"custom_colors": dataset_config.presentation__custom_colors,
"diffexp-may-be-slow": False,
"about_legal_tos": dataset_config.app__about_legal_tos,
"about_legal_privacy": dataset_config.app__about_legal_privacy,
}
# corpora dataset_props
# TODO/Note: putting info from the dataset into the /config is not ideal.
# However, it is definitely not part of /schema, and we do not have a top-level
# route for data properties. Consider creating one at some point.
corpora_props = data_adaptor.get_corpora_props()
if corpora_props and "default_embedding" in corpora_props:
default_embedding = corpora_props["default_embedding"]
if isinstance(default_embedding, str) and default_embedding.startswith("X_"):
default_embedding = default_embedding[2:] # drop X_ prefix
if default_embedding in data_adaptor.get_embedding_names():
parameters["default_embedding"] = default_embedding
data_adaptor.update_parameters(parameters)
if annotation:
annotation.update_parameters(parameters, data_adaptor)
# gather it all together
client_config = {}
config = client_config["config"] = {}
config["displayNames"] = display_names
config["library_versions"] = library_versions
config["links"] = links
config["parameters"] = parameters
config["corpora_props"] = corpora_props
config["limits"] = {
"column_request_max": server_config.limits__column_request_max,
"diffexp_cellcount_max": server_config.limits__diffexp_cellcount_max,
}
if dataset_config.app__authentication_enable and auth.is_valid_authentication_type():
config["authentication"] = {
"requires_client_login": auth.requires_client_login(),
}
if auth.requires_client_login():
config["authentication"].update(
{
# Todo why are these stored on the data_adaptor?
"login": auth.get_login_url(data_adaptor),
"logout": auth.get_logout_url(data_adaptor),
}
)
return client_config
def get_client_userinfo(app_config, data_adaptor):
"""
Return the userinfo as required by the /userinfo REST route
"""
server_config = app_config.server_config
dataset_config = data_adaptor.dataset_config
auth = server_config.auth
# make sure the configuration has been checked.
app_config.check_config()
if dataset_config.app__authentication_enable and auth.is_valid_authentication_type():
userinfo = {}
userinfo["userinfo"] = {
"is_authenticated": auth.is_user_authenticated(),
"username": auth.get_user_name(),
"user_id": auth.get_user_id(),
"email": auth.get_user_email(),
"picture": auth.get_user_picture(),
}
return userinfo
@@ -0,0 +1,211 @@
import os
from os.path import splitext, isdir
from backend.czi_hosted.common.annotations.annotations import Annotations
from backend.czi_hosted.common.annotations.hosted_tiledb import AnnotationsHostedTileDB
from backend.czi_hosted.common.annotations.local_file_csv import AnnotationsLocalFile
from backend.czi_hosted.common.config.base_config import BaseConfig
from backend.common.errors import ConfigurationError
from backend.czi_hosted.db.db_utils import DbUtils
class DatasetConfig(BaseConfig):
"""Manages the config attribute associated with a dataset."""
def __init__(self, tag, app_config, default_config):
super().__init__(app_config, default_config)
self.tag = tag
try:
self.app__scripts = default_config["app"]["scripts"]
self.app__inline_scripts = default_config["app"]["inline_scripts"]
self.app__about_legal_tos = default_config["app"]["about_legal_tos"]
self.app__about_legal_privacy = default_config["app"]["about_legal_privacy"]
self.app__authentication_enable = default_config["app"]["authentication_enable"]
self.presentation__max_categories = default_config["presentation"]["max_categories"]
self.presentation__custom_colors = default_config["presentation"]["custom_colors"]
self.user_annotations__enable = default_config["user_annotations"]["enable"]
self.user_annotations__type = default_config["user_annotations"]["type"]
self.user_annotations__local_file_csv__directory = default_config["user_annotations"]["local_file_csv"][
"directory"
]
self.user_annotations__local_file_csv__file = default_config["user_annotations"]["local_file_csv"]["file"]
self.user_annotations__hosted_tiledb_array__db_uri = default_config["user_annotations"][
"hosted_tiledb_array"
]["db_uri"]
self.user_annotations__hosted_tiledb_array__hosted_file_directory = default_config["user_annotations"][
"hosted_tiledb_array"
]["hosted_file_directory"]
self.embeddings__names = default_config["embeddings"]["names"]
self.diffexp__enable = default_config["diffexp"]["enable"]
self.diffexp__lfc_cutoff = default_config["diffexp"]["lfc_cutoff"]
self.diffexp__top_n = default_config["diffexp"]["top_n"]
self.X_approximate_distribution = default_config["X_approximate_distribution"]
except KeyError as e:
raise ConfigurationError(f"Unexpected config: {str(e)}")
# Create the default annotation, which supports gene set reading without
# further configuration. Depending on configuration options, `complete_config`
# may create a more specialized annotation object and replace this default.
self.user_annotations = Annotations()
def complete_config(self, context):
self.handle_app()
self.handle_presentation()
self.handle_user_annotations(context)
self.handle_embeddings()
self.handle_diffexp(context)
self.handle_X_approximate_distribution()
def handle_app(self):
self.validate_correct_type_of_configuration_attribute("app__scripts", list)
self.validate_correct_type_of_configuration_attribute("app__inline_scripts", list)
self.validate_correct_type_of_configuration_attribute("app__about_legal_tos", (type(None), str))
self.validate_correct_type_of_configuration_attribute("app__about_legal_privacy", (type(None), str))
self.validate_correct_type_of_configuration_attribute("app__authentication_enable", bool)
# scripts can be string (filename) or dict (attributes). Convert string to dict.
scripts = []
for script in self.app__scripts:
try:
if isinstance(script, str):
scripts.append({"src": script})
elif isinstance(script, dict) and isinstance(script["src"], str):
scripts.append(script)
else:
raise Exception
except Exception as e:
raise ConfigurationError(f"Scripts must be string or a dict containing an src key: {e}")
self.app__scripts = scripts
def handle_presentation(self):
self.validate_correct_type_of_configuration_attribute("presentation__max_categories", int)
self.validate_correct_type_of_configuration_attribute("presentation__custom_colors", bool)
def handle_user_annotations(self, context):
self.validate_correct_type_of_configuration_attribute("user_annotations__enable", bool)
self.validate_correct_type_of_configuration_attribute("user_annotations__type", str)
self.validate_correct_type_of_configuration_attribute(
"user_annotations__local_file_csv__directory", (type(None), str)
)
self.validate_correct_type_of_configuration_attribute(
"user_annotations__local_file_csv__file", (type(None), str)
)
self.validate_correct_type_of_configuration_attribute(
"user_annotations__hosted_tiledb_array__db_uri", (type(None), str)
)
self.validate_correct_type_of_configuration_attribute(
"user_annotations__hosted_tiledb_array__hosted_file_directory", (type(None), str)
)
if self.user_annotations__enable:
server_config = self.app_config.server_config
if not self.app__authentication_enable:
raise ConfigurationError("user annotations requires authentication to be enabled")
if not server_config.auth.is_valid_authentication_type():
auth_type = server_config.authentication__type
raise ConfigurationError(f"authentication method {auth_type} is not compatible with user annotations")
if self.user_annotations__type == "local_file_csv":
self.handle_local_file_csv_annotations()
elif self.user_annotations__type == "hosted_tiledb_array":
self.handle_hosted_tiledb_annotations()
else:
raise ConfigurationError('The only annotation type support is "local_file_csv" or "hosted_tiledb_array')
else:
self.check_annotation_config_vars_not_set(context)
def handle_local_file_csv_annotations(self):
dirname = self.user_annotations__local_file_csv__directory
filename = self.user_annotations__local_file_csv__file
if filename is not None and dirname is not None:
raise ConfigurationError("'annotations-file' and 'annotations-dir' may not be used together.")
if filename is not None:
lf_name, lf_ext = splitext(filename)
if lf_ext and lf_ext != ".csv":
raise ConfigurationError(f"annotation file type must be .csv: {filename}")
if dirname is not None and not isdir(dirname):
try:
os.mkdir(dirname)
except OSError:
raise ConfigurationError("Unable to create directory specified by --annotations-dir")
anno_config = {
"user-annotations": self.user_annotations__enable,
"genesets-save": False,
}
self.user_annotations = AnnotationsLocalFile(anno_config, dirname, filename)
# if the user has specified a fixed label file, go ahead and validate it
# so that we can remove errors early in the process.
server_config = self.app_config.server_config
if server_config.single_dataset__datapath and self.user_annotations__local_file_csv__file:
with server_config.matrix_data_cache_manager.data_adaptor(
self.tag, server_config.single_dataset__datapath, self.app_config
) as data_adaptor:
data_adaptor.check_new_labels(self.user_annotations.read_labels(data_adaptor))
def handle_hosted_tiledb_annotations(self):
self.validate_correct_type_of_configuration_attribute("user_annotations__hosted_tiledb_array__db_uri", str)
self.validate_correct_type_of_configuration_attribute(
"user_annotations__hosted_tiledb_array__hosted_file_directory", str
)
anno_config = {
"user-annotations": self.user_annotations__enable,
"genesets-save": False,
}
self.user_annotations = AnnotationsHostedTileDB(
anno_config,
directory_path=self.user_annotations__hosted_tiledb_array__hosted_file_directory,
db=DbUtils(self.user_annotations__hosted_tiledb_array__db_uri),
)
def check_annotation_config_vars_not_set(self, context):
if self.user_annotations__type is not None:
dirname = self.user_annotations__local_file_csv__directory
filename = self.user_annotations__local_file_csv__file
db_uri = self.user_annotations__hosted_tiledb_array__db_uri
hosted_file_dirname = self.user_annotations__hosted_tiledb_array__hosted_file_directory
if filename is not None:
context["messagefn"]("Warning: --annotations-file ignored as annotations are disabled.")
if dirname is not None:
context["messagefn"]("Warning: --annotations-dir ignored as annotations are disabled.")
if db_uri is not None:
context["messagefn"]("Warning: db_uri ignored as annotations are disabled.")
if hosted_file_dirname is not None:
context["messagefn"](
"Warning: hosted_file_directory for hosted_tiledb_array ignored as annotations are disabled."
)
def handle_embeddings(self):
self.validate_correct_type_of_configuration_attribute("embeddings__names", list)
def handle_diffexp(self, context):
self.validate_correct_type_of_configuration_attribute("diffexp__enable", bool)
self.validate_correct_type_of_configuration_attribute("diffexp__lfc_cutoff", float)
self.validate_correct_type_of_configuration_attribute("diffexp__top_n", int)
server_config = self.app_config.server_config
if server_config.single_dataset__datapath:
with server_config.matrix_data_cache_manager.data_adaptor(
self.tag, server_config.single_dataset__datapath, self.app_config
) as data_adaptor:
if self.diffexp__enable and data_adaptor.parameters.get("diffexp_may_be_slow", False):
context["messagefn"](
"CAUTION: due to the size of your dataset, "
"running differential expression may take longer or fail."
)
def handle_X_approximate_distribution(self):
self.validate_correct_type_of_configuration_attribute("X_approximate_distribution", str)
if self.X_approximate_distribution not in ["normal", "count"]:
raise ConfigurationError(
"X_approximate_distribution has unknown value -- must be 'normal' or 'count'."
)
@@ -0,0 +1,95 @@
import os
from backend.czi_hosted.common.config.base_config import BaseConfig
from backend.common.errors import ConfigurationError, SecretKeyRetrievalError
from backend.common.utils.aws_secret_utils import get_secret_key
from backend.common.utils.type_conversion_utils import convert_string_to_value
class ExternalConfig(BaseConfig):
"""Manages the config attribute associated with external configuration sources, such as
environment variables or the AWS Secrets Manager."""
def __init__(self, app_config, default_config):
super().__init__(app_config, default_config)
try:
self.environment = default_config["environment"]
self.aws_secrets_manager__region = default_config["aws_secrets_manager"]["region"]
self.aws_secrets_manager__secrets = default_config["aws_secrets_manager"]["secrets"]
except KeyError as e:
raise ConfigurationError(f"Unexpected config: {str(e)}")
def complete_config(self, context):
self.handle_environment(context)
self.handle_aws_secrets_manager(context)
def handle_environment(self, context):
"""For each environment variable defined, get the value (if it is set),
and set the specified config parameter"""
self.validate_correct_type_of_configuration_attribute("environment", list)
for envdict in self.environment:
name = envdict.get("name")
if name is None:
raise ConfigurationError("environment: 'name' is missing")
required = envdict.get("required", False)
if type(required) != bool:
raise ConfigurationError("environment: 'required' must be a bool")
path = envdict.get("path")
if path is None:
raise ConfigurationError("environment: 'path' is missing")
value = os.environ.get(name)
if value is None:
if required:
raise ConfigurationError(f"required environment variable '{name}' not set")
else:
value = convert_string_to_value(value)
self.app_config.update_single_config_from_path_and_value(path, value)
def handle_aws_secrets_manager(self, context):
"""For each aws secret defined, get the key/values, and set the specified config parameter"""
self.validate_correct_type_of_configuration_attribute("aws_secrets_manager__region", (type(None), str))
self.validate_correct_type_of_configuration_attribute("aws_secrets_manager__secrets", list)
if not self.aws_secrets_manager__secrets:
return
self.validate_correct_type_of_configuration_attribute("aws_secrets_manager__region", str)
for secret in self.aws_secrets_manager__secrets:
secret_name = secret.get("name")
if secret_name is None:
raise ConfigurationError("aws_secrets_manager: 'name' is missing")
if not isinstance(secret_name, str):
raise ConfigurationError("aws_secrets_manager: 'name' must be a string")
try:
secret_dict = get_secret_key(self.aws_secrets_manager__region, secret_name)
except SecretKeyRetrievalError as e:
raise ConfigurationError(f"Unable to retrieve secret {secret_name}: {str(e)}")
values = secret.get("values")
if values is None:
raise ConfigurationError("aws_secrets_manager: 'values' is missing")
if not isinstance(values, list):
raise ConfigurationError("aws_secrets_manager: 'values' must be a list")
for value in values:
key = value.get("key")
if key is None:
raise ConfigurationError(f"missing 'key' in secret values: {secret_name}")
path = value.get("path")
if path is None:
raise ConfigurationError(f"missing 'path' in secret values: {secret_name}")
required = value.get("required", False)
if type(required) != bool:
raise ConfigurationError(f"wrong type for 'required' in secret values: {secret_name}")
secret_value = secret_dict.get(key)
if secret_value is None:
if required:
raise ConfigurationError(f"required secret '{secret_name}:{key}' not set")
else:
secret_value = convert_string_to_value(secret_value)
self.app_config.update_single_config_from_path_and_value(path, secret_value)
@@ -0,0 +1,387 @@
import os
import sys
import warnings
from os.path import basename
from urllib.parse import urlparse, quote_plus
from backend.czi_hosted.auth.auth import AuthTypeFactory
from backend.czi_hosted.common.config import DEFAULT_SERVER_PORT, BIG_FILE_SIZE_THRESHOLD
from backend.czi_hosted.common.config.base_config import BaseConfig
from backend.common.utils.data_locator import discover_s3_region_name
from backend.common.errors import ConfigurationError, DatasetAccessError
from backend.common.utils.utils import is_port_available, find_available_port, custom_format_warning
from backend.czi_hosted.compute import diffexp_cxg as diffexp_tiledb
from backend.czi_hosted.data_common.matrix_loader import MatrixDataCacheManager, MatrixDataLoader, MatrixDataType
class ServerConfig(BaseConfig):
"""Manages the config attribute associated with the server."""
def __init__(self, app_config, default_config):
dictval_cases = [
("app", "csp_directives"),
("authentication", "params_oauth", "cookie"),
("authentication", "params_oauth", "jwt_decode_options"),
("adaptor", "cxg_adaptor", "tiledb_ctx"),
("multi_dataset", "dataroot"),
]
super().__init__(app_config, default_config, dictval_cases)
try:
self.app__verbose = default_config["app"]["verbose"]
self.app__debug = default_config["app"]["debug"]
self.app__host = default_config["app"]["host"]
self.app__port = default_config["app"]["port"]
self.app__open_browser = default_config["app"]["open_browser"]
self.app__force_https = default_config["app"]["force_https"]
self.app__flask_secret_key = default_config["app"]["flask_secret_key"]
self.app__generate_cache_control_headers = default_config["app"]["generate_cache_control_headers"]
self.app__server_timing_headers = default_config["app"]["server_timing_headers"]
self.app__csp_directives = default_config["app"]["csp_directives"]
self.app__api_base_url = default_config["app"]["api_base_url"]
self.app__web_base_url = default_config["app"]["web_base_url"]
self.authentication__type = default_config["authentication"]["type"]
self.authentication__insecure_test_environment = default_config["authentication"][
"insecure_test_environment"
]
self.authentication__params_oauth__oauth_api_base_url = default_config["authentication"]["params_oauth"][
"oauth_api_base_url"
]
self.authentication__params_oauth__client_id = default_config["authentication"]["params_oauth"]["client_id"]
self.authentication__params_oauth__client_secret = default_config["authentication"]["params_oauth"][
"client_secret"
]
self.authentication__params_oauth__jwt_decode_options = default_config["authentication"]["params_oauth"][
"jwt_decode_options"
]
self.authentication__params_oauth__session_cookie = default_config["authentication"]["params_oauth"][
"session_cookie"
]
self.authentication__params_oauth__cookie = default_config["authentication"]["params_oauth"]["cookie"]
self.multi_dataset__dataroot = default_config["multi_dataset"]["dataroot"]
self.multi_dataset__index = default_config["multi_dataset"]["index"]
self.multi_dataset__allowed_matrix_types = default_config["multi_dataset"]["allowed_matrix_types"]
self.multi_dataset__matrix_cache__max_datasets = default_config["multi_dataset"]["matrix_cache"][
"max_datasets"
]
self.multi_dataset__matrix_cache__timelimit_s = default_config["multi_dataset"]["matrix_cache"][
"timelimit_s"
]
self.single_dataset__datapath = default_config["single_dataset"]["datapath"]
self.single_dataset__obs_names = default_config["single_dataset"]["obs_names"]
self.single_dataset__var_names = default_config["single_dataset"]["var_names"]
self.single_dataset__about = default_config["single_dataset"]["about"]
self.single_dataset__title = default_config["single_dataset"]["title"]
self.diffexp__alg_cxg__max_workers = default_config["diffexp"]["alg_cxg"]["max_workers"]
self.diffexp__alg_cxg__cpu_multiplier = default_config["diffexp"]["alg_cxg"]["cpu_multiplier"]
self.diffexp__alg_cxg__target_workunit = default_config["diffexp"]["alg_cxg"]["target_workunit"]
self.data_locator__s3__region_name = default_config["data_locator"]["s3"]["region_name"]
self.adaptor__cxg_adaptor__tiledb_ctx = default_config["adaptor"]["cxg_adaptor"]["tiledb_ctx"]
self.adaptor__anndata_adaptor__backed = default_config["adaptor"]["anndata_adaptor"]["backed"]
self.limits__diffexp_cellcount_max = default_config["limits"]["diffexp_cellcount_max"]
self.limits__column_request_max = default_config["limits"]["column_request_max"]
except KeyError as e:
raise ConfigurationError(f"Unexpected config: {str(e)}")
# The matrix data cache manager is created during the complete_config and stored here.
self.matrix_data_cache_manager = None
# The authentication object
self.auth = None
def complete_config(self, context):
self.handle_app(context)
self.handle_data_source()
self.handle_authentication()
self.handle_data_locator()
self.handle_adaptor() # may depend on data_locator
self.handle_single_dataset(context) # may depend on adaptor
self.handle_multi_dataset() # may depend on adaptor
self.handle_diffexp()
self.handle_limits()
self.check_config()
def handle_app(self, context):
self.validate_correct_type_of_configuration_attribute("app__verbose", bool)
self.validate_correct_type_of_configuration_attribute("app__debug", bool)
self.validate_correct_type_of_configuration_attribute("app__host", str)
self.validate_correct_type_of_configuration_attribute("app__port", (type(None), int))
self.validate_correct_type_of_configuration_attribute("app__open_browser", bool)
self.validate_correct_type_of_configuration_attribute("app__force_https", bool)
self.validate_correct_type_of_configuration_attribute("app__flask_secret_key", str)
self.validate_correct_type_of_configuration_attribute("app__generate_cache_control_headers", bool)
self.validate_correct_type_of_configuration_attribute("app__server_timing_headers", bool)
self.validate_correct_type_of_configuration_attribute("app__csp_directives", (type(None), dict))
self.validate_correct_type_of_configuration_attribute("app__api_base_url", (type(None), str))
self.validate_correct_type_of_configuration_attribute("app__web_base_url", (type(None), str))
if self.app__port:
try:
if not is_port_available(self.app__host, self.app__port):
raise ConfigurationError(
f"The port selected {self.app__port} is in use, please configure an open port."
)
except OverflowError:
raise ConfigurationError(f"Invalid port: {self.app__port}")
else:
try:
default_server_port = int(os.environ.get("CXG_SERVER_PORT", DEFAULT_SERVER_PORT))
except ValueError:
raise ConfigurationError(
"Invalid port from environment variable CXG_SERVER_PORT: " + os.environ.get("CXG_SERVER_PORT")
)
try:
self.app__port = find_available_port(self.app__host, default_server_port)
except OverflowError:
raise ConfigurationError(f"Invalid port: {default_server_port}")
if self.app__debug:
context["messagefn"]("in debug mode, setting verbose=True and open_browser=False")
self.app__verbose = True
self.app__open_browser = False
else:
warnings.formatwarning = custom_format_warning
if not self.app__verbose:
sys.tracebacklimit = 0
# CSP Directives are a dict of string: list(string) or string: string
if self.app__csp_directives is not None:
for k, v in self.app__csp_directives.items():
if not isinstance(k, str):
raise ConfigurationError("CSP directive names must be a string.")
if isinstance(v, list):
for policy in v:
if not isinstance(policy, str):
raise ConfigurationError("CSP directive value must be a string or list of strings.")
elif not isinstance(v, str):
raise ConfigurationError("CSP directive value must be a string or list of strings.")
if self.app__web_base_url is None:
self.app__web_base_url = self.app__api_base_url
def handle_authentication(self):
self.validate_correct_type_of_configuration_attribute("authentication__type", (type(None), str))
self.validate_correct_type_of_configuration_attribute("authentication__insecure_test_environment", bool)
if self.authentication__type == "test" and not self.authentication__insecure_test_environment:
raise ConfigurationError("Test auth can only be used in an insecure test environment")
# oauth
ptypes = str if self.authentication__type == "oauth" else (type(None), str)
self.validate_correct_type_of_configuration_attribute(
"authentication__params_oauth__oauth_api_base_url", ptypes
)
self.validate_correct_type_of_configuration_attribute("authentication__params_oauth__client_id", ptypes)
self.validate_correct_type_of_configuration_attribute("authentication__params_oauth__client_secret", ptypes)
self.validate_correct_type_of_configuration_attribute(
"authentication__params_oauth__jwt_decode_options", (type(None), dict)
)
self.validate_correct_type_of_configuration_attribute("authentication__params_oauth__session_cookie", bool)
if self.authentication__params_oauth__session_cookie:
self.validate_correct_type_of_configuration_attribute(
"authentication__params_oauth__cookie", (type(None), dict)
)
else:
self.validate_correct_type_of_configuration_attribute("authentication__params_oauth__cookie", dict)
self.auth = AuthTypeFactory.create(self.authentication__type, self)
if self.auth is None:
raise ConfigurationError(f"Unknown authentication type: {self.authentication__type}")
def handle_data_locator(self):
self.validate_correct_type_of_configuration_attribute("data_locator__s3__region_name", (type(None), bool, str))
if self.data_locator__s3__region_name is True:
path = self.single_dataset__datapath or self.multi_dataset__dataroot
if type(path) == dict:
# if multi_dataset__dataroot is a dict, then use the first key
# that is in s3. NOTE: it is not supported to have dataroots
# in different regions.
paths = [val.get("dataroot") for val in path.values()]
for path in paths:
if path.startswith("s3://"):
break
if path.startswith("s3://"):
region_name = discover_s3_region_name(path)
if region_name is None:
raise ConfigurationError(f"Unable to discover s3 region name from {path}")
else:
region_name = None
self.data_locator__s3__region_name = region_name
def handle_data_source(self):
self.validate_correct_type_of_configuration_attribute("single_dataset__datapath", (str, type(None)))
self.validate_correct_type_of_configuration_attribute("multi_dataset__dataroot", (type(None), dict, str))
if self.single_dataset__datapath and self.multi_dataset__dataroot:
raise ConfigurationError(
"You must supply either a datapath (for single datasets) or a dataroot (for multidatasets). Not both"
)
if self.single_dataset__datapath is None and self.multi_dataset__dataroot is None:
raise ConfigurationError("You must specify a datapath for a single dataset or a dataroot for multidatasets")
def handle_single_dataset(self, context):
self.validate_correct_type_of_configuration_attribute("single_dataset__datapath", (str, type(None)))
self.validate_correct_type_of_configuration_attribute("single_dataset__title", (str, type(None)))
self.validate_correct_type_of_configuration_attribute("single_dataset__about", (str, type(None)))
self.validate_correct_type_of_configuration_attribute("single_dataset__obs_names", (str, type(None)))
self.validate_correct_type_of_configuration_attribute("single_dataset__var_names", (str, type(None)))
if self.single_dataset__datapath is None:
return
# create the matrix data cache manager:
if self.matrix_data_cache_manager is None:
self.matrix_data_cache_manager = MatrixDataCacheManager(max_cached=1, timelimit_s=None)
# preload this data set
matrix_data_loader = MatrixDataLoader(self.single_dataset__datapath, app_config=self.app_config)
try:
matrix_data_loader.pre_load_validation()
except DatasetAccessError as e:
raise ConfigurationError(str(e))
file_size = matrix_data_loader.file_size()
file_basename = basename(self.single_dataset__datapath)
if file_size > BIG_FILE_SIZE_THRESHOLD:
context["messagefn"](f"Loading data from {file_basename}, this may take a while...")
else:
context["messagefn"](f"Loading data from {file_basename}.")
if self.single_dataset__about:
def url_check(url):
try:
result = urlparse(url)
if all([result.scheme, result.netloc]):
return True
else:
return False
except ValueError:
return False
if not url_check(self.single_dataset__about):
raise ConfigurationError(
"Must provide an absolute URL for --about. (Example format: http://example.com)"
)
def handle_multi_dataset(self):
self.validate_correct_type_of_configuration_attribute("multi_dataset__dataroot", (type(None), dict, str))
self.validate_correct_type_of_configuration_attribute("multi_dataset__index", (type(None), bool, str))
self.validate_correct_type_of_configuration_attribute("multi_dataset__allowed_matrix_types", list)
self.validate_correct_type_of_configuration_attribute("multi_dataset__matrix_cache__max_datasets", int)
self.validate_correct_type_of_configuration_attribute(
"multi_dataset__matrix_cache__timelimit_s", (type(None), int, float)
)
if self.multi_dataset__dataroot is None:
return
if type(self.multi_dataset__dataroot) == str:
default_dict = dict(base_url="d", dataroot=self.multi_dataset__dataroot)
self.multi_dataset__dataroot = dict(d=default_dict)
for tag, dataroot_dict in self.multi_dataset__dataroot.items():
if "base_url" not in dataroot_dict:
raise ConfigurationError(f"error in multi_dataset__dataroot: missing base_url for tag {tag}")
if "dataroot" not in dataroot_dict:
raise ConfigurationError(f"error in multi_dataset__dataroot: missing dataroot, for tag {tag}")
base_url = dataroot_dict["base_url"]
# sanity check for well formed base urls
bad = False
if type(base_url) != str:
bad = True
elif os.path.normpath(base_url) != base_url:
bad = True
else:
base_url_parts = base_url.split("/")
if [quote_plus(part) for part in base_url_parts] != base_url_parts:
bad = True
if ".." in base_url_parts:
bad = True
if bad:
raise ConfigurationError(f"error in multi_dataset__dataroot base_url {base_url} for tag {tag}")
# verify all the base_urls are unique
base_urls = [d["base_url"] for d in self.multi_dataset__dataroot.values()]
if len(base_urls) > len(set(base_urls)):
raise ConfigurationError("error in multi_dataset__dataroot: base_urls must be unique")
# error checking
for mtype in self.multi_dataset__allowed_matrix_types:
try:
MatrixDataType(mtype)
except ValueError:
raise ConfigurationError(f'Invalid matrix type in "allowed_matrix_types": {mtype}')
# create the matrix data cache manager:
if self.matrix_data_cache_manager is None:
self.matrix_data_cache_manager = MatrixDataCacheManager(
max_cached=self.multi_dataset__matrix_cache__max_datasets,
timelimit_s=self.multi_dataset__matrix_cache__timelimit_s,
)
def handle_diffexp(self):
self.validate_correct_type_of_configuration_attribute("diffexp__alg_cxg__max_workers", (str, int))
self.validate_correct_type_of_configuration_attribute("diffexp__alg_cxg__cpu_multiplier", int)
self.validate_correct_type_of_configuration_attribute("diffexp__alg_cxg__target_workunit", int)
max_workers = self.diffexp__alg_cxg__max_workers
cpu_multiplier = self.diffexp__alg_cxg__cpu_multiplier
cpu_count = os.cpu_count()
max_workers = min(max_workers, cpu_multiplier * cpu_count)
diffexp_tiledb.set_config(max_workers, self.diffexp__alg_cxg__target_workunit)
def handle_adaptor(self):
# cxg
self.validate_correct_type_of_configuration_attribute("adaptor__cxg_adaptor__tiledb_ctx", dict)
regionkey = "vfs.s3.region"
if regionkey not in self.adaptor__cxg_adaptor__tiledb_ctx:
if type(self.data_locator__s3__region_name) == str:
self.adaptor__cxg_adaptor__tiledb_ctx[regionkey] = self.data_locator__s3__region_name
from backend.czi_hosted.data_cxg.cxg_adaptor import CxgAdaptor
CxgAdaptor.set_tiledb_context(self.adaptor__cxg_adaptor__tiledb_ctx)
# anndata
self.validate_correct_type_of_configuration_attribute("adaptor__anndata_adaptor__backed", bool)
def handle_limits(self):
self.validate_correct_type_of_configuration_attribute("limits__diffexp_cellcount_max", (type(None), int))
self.validate_correct_type_of_configuration_attribute("limits__column_request_max", (type(None), int))
def exceeds_limit(self, limit_name, value):
limit_value = getattr(self, "limits__" + limit_name, None)
if limit_value is None: # disabled
return False
return value > limit_value
def get_api_base_url(self):
if self.app__api_base_url == "local":
return f"http://{self.app__host}:{self.app__port}"
if self.app__api_base_url and self.app__api_base_url.endswith("/"):
return self.app__api_base_url[:-1]
return self.app__api_base_url
def get_web_base_url(self):
if self.app__web_base_url == "local":
return f"http://{self.app__host}:{self.app__port}"
if self.app__web_base_url is None:
return self.get_api_base_url()
if self.app__web_base_url.endswith("/"):
return self.app__web_base_url[:-1]
return self.app__web_base_url
+78
View File
@@ -0,0 +1,78 @@
"""
Corpora schema conventions support. Helper functions for reading.
https://github.com/chanzuckerberg/corpora-data-portal/blob/main/backend/schema/corpora_schema.md
https://github.com/chanzuckerberg/corpora-data-portal/blob/main/backend/schema/corpora_schema_h5ad_implementation.md
"""
import collections
import json
from backend.czi_hosted.cli.upgrade import validate_version_str
from backend.czi_hosted.common.utils.corpora_constants import CorporaConstants
def corpora_get_versions_from_anndata(adata):
"""
Given an AnnData object, return:
* None - if not a Corpora object
* [ corpora_schema_version, corpora_encoding_version ] - if a Corpora object
Implements the identification protocol defined in the specification.
"""
# per Corpora AnnData spec, this is a corpora file if the following is true
if "version" not in adata.uns_keys():
return None
version = adata.uns["version"]
if not isinstance(version, collections.abc.Mapping) or "corpora_schema_version" not in version:
return None
corpora_schema_version = version.get("corpora_schema_version")
corpora_encoding_version = version.get("corpora_encoding_version")
# TODO: spec says these must be SEMVER values, so check.
if validate_version_str(corpora_schema_version) and validate_version_str(corpora_encoding_version):
return [corpora_schema_version, corpora_encoding_version]
def corpora_is_version_supported(corpora_schema_version, corpora_encoding_version):
return (
corpora_schema_version
and corpora_encoding_version
and corpora_schema_version.startswith("1.")
and corpora_encoding_version.startswith("0.1.")
)
def corpora_get_props_from_anndata(adata):
"""
Get Corpora dataset properties from an AnnData
"""
versions = corpora_get_versions_from_anndata(adata)
if versions is None:
return None
[corpora_schema_version, corpora_encoding_version] = versions
version_is_supported = corpora_is_version_supported(corpora_schema_version, corpora_encoding_version)
if not version_is_supported:
raise ValueError("Unsupported Corpora schema version")
corpora_props = {}
for key in CorporaConstants.REQUIRED_SIMPLE_METADATA_FIELDS:
if key not in adata.uns:
raise KeyError(f"missing Corpora schema field {key}")
corpora_props[key] = adata.uns[key]
for key in CorporaConstants.OPTIONAL_JSON_ENCODED_METADATA_FIELD:
if key not in adata.uns:
continue
try:
corpora_props[key] = json.loads(adata.uns[key])
except json.JSONDecodeError:
raise json.JSONDecodeError(f"Corpora schema field {key} is expected to be a valid JSON string")
for key in CorporaConstants.OPTIONAL_SIMPLE_METADATA_FIELDS:
if key in adata.uns:
corpora_props[key] = adata.uns[key]
return corpora_props
+38
View File
@@ -0,0 +1,38 @@
from http import HTTPStatus
from flask import make_response, jsonify
from backend.czi_hosted import __version__ as cellxgene_version
from backend.common.utils.data_locator import DataLocator
def _is_accessible(path, config):
if path is None:
return True
try:
dl = DataLocator(path, region_name=config.data_locator__s3__region_name)
return dl.exists()
except RuntimeError:
return False
def health_check(config):
"""
simple health check - return HTTP response.
See https://tools.ietf.org/id/draft-inadarei-api-health-check-01.html
"""
health = {"status": None, "version": "1", "releaseID": cellxgene_version}
checks = False
server_config = config.server_config
if config.is_multi_dataset():
dataroots = [datapath_dict["dataroot"] for datapath_dict in server_config.multi_dataset__dataroot.values()]
checks = all([_is_accessible(dataroot, server_config) for dataroot in dataroots])
else:
checks = _is_accessible(server_config.single_dataset__datapath, server_config)
health["status"] = "pass" if checks else "fail"
code = HTTPStatus.OK if health["status"] == "pass" else HTTPStatus.BAD_REQUEST
response = make_response(jsonify(health), code)
response.headers["Content-Type"] = "application/health+json"
return response
@@ -0,0 +1,71 @@
import threading
from collections.abc import MutableMapping
class ImmutableKVCache(MutableMapping):
"""
Guarantees that the factory will be called for each key once, and
only once.
"""
def __init__(self, factory):
self.factory = factory # user-provided factory function
self.lock = threading.Lock() # guards factory_calls
self.factory_calls = {} # per-key factory condition variables
self.cache = {} # result cache, indexed by key
super().__init__()
def __getitem__(self, key):
if key in self.cache:
return self.cache[key]
# we need to call factory. First grab the main lock and the per-key CV.
factory_calls = None
creation_thr = False
with self.lock:
if key in self.cache:
return self.cache[key]
if key not in self.factory_calls:
creation_thr = True
self.factory_calls[key] = {"cv": threading.Condition(), "is_done": False, "error": None}
factory_calls = self.factory_calls[key]
# with the CV, create the value (or wait for it to be created)
cv = factory_calls["cv"]
with cv:
if creation_thr:
try:
self.cache[key] = self.factory(key)
except Exception as e:
factory_calls["error"] = e
factory_calls["is_done"] = True
cv.notify_all()
else:
""" wait for the value to be available """
while not factory_calls["is_done"]:
cv.wait()
with self.lock:
if key in self.factory_calls:
del self.factory_calls[key]
return self.cache[key]
def __iter__(self):
""" weak iter, don't call factory """
return self.cache.__iter__()
def __len__(self):
return self.cache.__len__()
def __contains__(self, key):
""" weak contain - don't call factory """
return self.cache.__contains__(key)
def __delitem__(self, key):
del self.cache[key]
def __setitem__(self, key, value):
""" unsupported """
raise NotImplementedError
+381
View File
@@ -0,0 +1,381 @@
import copy
import logging
import sys
from http import HTTPStatus
import zlib
import json
from flask import make_response, jsonify, current_app, abort
from werkzeug.urls import url_unquote
from backend.czi_hosted.common.config.client_config import get_client_config, get_client_userinfo
from backend.common.constants import Axis, DiffExpMode, JSON_NaN_to_num_warning_msg
from backend.common.errors import (
FilterError,
JSONEncodingValueError,
PrepareError,
DisabledFeatureError,
ExceedsLimitError,
DatasetAccessError,
ColorFormatException,
AnnotationsError,
UnsupportedSummaryMethod,
)
from backend.common.genesets import summarizeQueryHash
from backend.common.fbs.matrix import decode_matrix_fbs
def abort_and_log(code, logmsg, loglevel=logging.DEBUG, include_exc_info=False):
"""
Log the message, then abort with HTTP code. If include_exc_info is true,
also include current exception via sys.exc_info().
"""
if include_exc_info:
exc_info = sys.exc_info()
else:
exc_info = False
current_app.logger.log(loglevel, logmsg, exc_info=exc_info)
# Do NOT send log message to HTTP response.
return abort(code)
def _query_parameter_to_filter(args):
"""
Convert an annotation value filter, if present in the query args,
into the standard dict filter format used by internal code.
Query param filters look like: <axis>:name=value, where value
may be one of:
- a range, min,max, where either may be an open range by using an asterisc, eg, 10,*
- a value
Eg,
...?tissue=lung&obs:tissue=heart&obs:num_reads=1000,*
"""
filters = {
"obs": {},
"var": {},
}
# args has already been url-unquoted once. We assume double escaping
# on name and value.
try:
for key, value in args.items(multi=True):
axis, name = key.split(":")
if axis not in ("obs", "var"):
raise FilterError("unknown filter axis")
name = url_unquote(name)
current = filters[axis].setdefault(name, {"name": name})
val_split = value.split(",")
if len(val_split) == 1:
if "min" in current or "max" in current:
raise FilterError("do not mix range and value filters")
value = url_unquote(value)
values = current.setdefault("values", [])
values.append(value)
elif len(val_split) == 2:
if len(current) > 1:
raise FilterError("duplicate range specification")
min = url_unquote(val_split[0])
max = url_unquote(val_split[1])
if min != "*":
current["min"] = float(min)
if max != "*":
current["max"] = float(max)
if len(current) < 2:
raise FilterError("must specify at least min or max in range filter")
else:
raise FilterError("badly formated filter value")
except ValueError as e:
raise FilterError(str(e))
result = {}
for axis in ("obs", "var"):
axis_filter = filters[axis]
if len(axis_filter) > 0:
result[axis] = {"annotation_value": [val for val in axis_filter.values()]}
return result
def schema_get_helper(data_adaptor):
"""helper function to gather the schema from the data source and annotations"""
schema = data_adaptor.get_schema()
schema = copy.deepcopy(schema)
# add label obs annotations as needed
annotations = data_adaptor.dataset_config.user_annotations
if annotations.user_annotations_enabled():
label_schema = annotations.get_schema(data_adaptor)
schema["annotations"]["obs"]["columns"].extend(label_schema)
return schema
def schema_get(data_adaptor):
schema = schema_get_helper(data_adaptor)
return make_response(jsonify({"schema": schema}), HTTPStatus.OK)
def config_get(app_config, data_adaptor):
config = get_client_config(app_config, data_adaptor)
return make_response(jsonify(config), HTTPStatus.OK)
def userinfo_get(app_config, data_adaptor):
config = get_client_userinfo(app_config, data_adaptor)
return make_response(jsonify(config), HTTPStatus.OK)
def annotations_obs_get(request, data_adaptor):
fields = request.args.getlist("annotation-name", None)
num_columns_requested = len(data_adaptor.get_obs_keys()) if len(fields) == 0 else len(fields)
if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested):
return abort(HTTPStatus.BAD_REQUEST)
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
try:
labels = None
annotations = data_adaptor.dataset_config.user_annotations
if annotations.user_annotations_enabled():
labels = annotations.read_labels(data_adaptor)
fbs = data_adaptor.annotation_to_fbs_matrix(Axis.OBS, fields, labels)
return make_response(fbs, HTTPStatus.OK, {"Content-Type": "application/octet-stream"})
except KeyError as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
def annotations_put_fbs_helper(data_adaptor, fbs):
"""helper function to write annotations from fbs"""
annotations = data_adaptor.dataset_config.user_annotations
if not annotations.user_annotations_enabled():
raise DisabledFeatureError("Writable annotations are not enabled")
new_label_df = decode_matrix_fbs(fbs)
if not new_label_df.empty:
new_label_df = data_adaptor.check_new_labels(new_label_df)
annotations.write_labels(new_label_df, data_adaptor)
def inflate(data):
return zlib.decompress(data)
def annotations_obs_put(request, data_adaptor):
annotations = data_adaptor.dataset_config.user_annotations
if not annotations.user_annotations_enabled():
return abort(HTTPStatus.NOT_IMPLEMENTED)
anno_collection = request.args.get("annotation-collection-name", default=None)
fbs = inflate(request.get_data())
if anno_collection is not None:
if not annotations.is_safe_collection_name(anno_collection):
return abort(HTTPStatus.BAD_REQUEST, "Bad annotation collection name")
annotations.set_collection(anno_collection)
try:
annotations_put_fbs_helper(data_adaptor, fbs)
res = json.dumps({"status": "OK"})
return make_response(res, HTTPStatus.OK, {"Content-Type": "application/json"})
except (ValueError, DisabledFeatureError, KeyError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
def annotations_var_get(request, data_adaptor):
fields = request.args.getlist("annotation-name", None)
num_columns_requested = len(data_adaptor.get_var_keys()) if len(fields) == 0 else len(fields)
if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested):
return abort(HTTPStatus.BAD_REQUEST)
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
try:
labels = None
annotations = data_adaptor.dataset_config.user_annotations
if annotations.user_annotations_enabled():
labels = annotations.read_labels(data_adaptor)
return make_response(
data_adaptor.annotation_to_fbs_matrix(Axis.VAR, fields, labels),
HTTPStatus.OK,
{"Content-Type": "application/octet-stream"},
)
except KeyError as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
def data_var_put(request, data_adaptor):
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
filter_json = request.get_json()
filter = filter_json["filter"] if filter_json else None
try:
return make_response(
data_adaptor.data_frame_to_fbs_matrix(filter, axis=Axis.VAR),
HTTPStatus.OK,
{"Content-Type": "application/octet-stream"},
)
except (FilterError, ValueError, ExceedsLimitError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
def data_var_get(request, data_adaptor):
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
try:
filter = _query_parameter_to_filter(request.args)
return make_response(
data_adaptor.data_frame_to_fbs_matrix(filter, axis=Axis.VAR),
HTTPStatus.OK,
{"Content-Type": "application/octet-stream"},
)
except (FilterError, ValueError, ExceedsLimitError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
def colors_get(data_adaptor):
if not data_adaptor.dataset_config.presentation__custom_colors:
return make_response(jsonify({}), HTTPStatus.OK)
try:
return make_response(jsonify(data_adaptor.get_colors()), HTTPStatus.OK)
except ColorFormatException as e:
return abort_and_log(HTTPStatus.NOT_FOUND, str(e), include_exc_info=True)
def diffexp_obs_post(request, data_adaptor):
if not data_adaptor.dataset_config.diffexp__enable:
return abort(HTTPStatus.NOT_IMPLEMENTED)
args = request.get_json()
try:
# TODO: implement varfilter mode
mode = DiffExpMode(args["mode"])
if mode == DiffExpMode.VAR_FILTER or "varFilter" in args:
return abort_and_log(HTTPStatus.NOT_IMPLEMENTED, "varFilter not enabled")
set1_filter = args.get("set1", {"filter": {}})["filter"]
set2_filter = args.get("set2", {"filter": {}})["filter"]
# TODO(#1281): When we simplify the config, we should actually use the config to determine this number,
# this will also require an update in the client
count = 15
if set1_filter is None or set2_filter is None or count is None:
return abort_and_log(HTTPStatus.BAD_REQUEST, "missing required parameter")
if Axis.VAR in set1_filter or Axis.VAR in set2_filter:
return abort_and_log(HTTPStatus.BAD_REQUEST, "var axis filter not enabled")
except (KeyError, TypeError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
try:
diffexp = data_adaptor.diffexp_topN(set1_filter, set2_filter, count)
return make_response(diffexp, HTTPStatus.OK, {"Content-Type": "application/json"})
except (ValueError, DisabledFeatureError, FilterError, ExceedsLimitError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
except JSONEncodingValueError:
# JSON encoding failure, usually due to bad data. Just let it ripple up
# to default exception handler.
current_app.logger.warning(JSON_NaN_to_num_warning_msg)
raise
def layout_obs_get(request, data_adaptor):
fields = request.args.getlist("layout-name", None)
num_columns_requested = len(data_adaptor.get_embedding_names()) if len(fields) == 0 else len(fields)
if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested):
return abort(HTTPStatus.BAD_REQUEST)
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
try:
return make_response(
data_adaptor.layout_to_fbs_matrix(fields), HTTPStatus.OK, {"Content-Type": "application/octet-stream"}
)
except (KeyError, DatasetAccessError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
except PrepareError:
return abort_and_log(
HTTPStatus.NOT_IMPLEMENTED,
f"No embedding available {request.path}",
loglevel=logging.ERROR,
include_exc_info=True,
)
def genesets_get(request, data_adaptor):
preferred_mimetype = request.accept_mimetypes.best_match(["application/json", "text/csv"])
if preferred_mimetype not in ("application/json", "text/csv"):
return abort(HTTPStatus.NOT_ACCEPTABLE)
try:
annotations = data_adaptor.dataset_config.user_annotations
(genesets, tid) = annotations.read_gene_sets(data_adaptor)
if preferred_mimetype == "text/csv":
return make_response(
annotations.gene_sets_to_csv(genesets),
HTTPStatus.OK,
{
"Content-Type": "text/csv",
"Content-Disposition": "attachment; filename=genesets.csv",
},
)
else:
return make_response(
jsonify({"genesets": annotations.gene_sets_to_response(genesets), "tid": tid}), HTTPStatus.OK
)
except (ValueError, KeyError, AnnotationsError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e))
def summarize_var_helper(request, data_adaptor, key, raw_query):
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
summary_method = request.values.get("method", default="mean")
query_hash = summarizeQueryHash(raw_query)
if key and query_hash != key:
return abort(HTTPStatus.BAD_REQUEST, description="query key did not match")
args_filter_only = request.values.copy()
args_filter_only.poplist("method")
args_filter_only.poplist("key")
try:
filter = _query_parameter_to_filter(args_filter_only)
return make_response(
data_adaptor.summarize_var(summary_method, filter, query_hash),
HTTPStatus.OK,
{"Content-Type": "application/octet-stream"},
)
except (ValueError) as e:
return abort(HTTPStatus.NOT_FOUND, description=str(e))
except (UnsupportedSummaryMethod, FilterError) as e:
return abort(HTTPStatus.BAD_REQUEST, description=str(e))
def summarize_var_get(request, data_adaptor):
return summarize_var_helper(request, data_adaptor, None, request.query_string)
def summarize_var_post(request, data_adaptor):
if not request.content_type or "application/x-www-form-urlencoded" not in request.content_type:
return abort(HTTPStatus.UNSUPPORTED_MEDIA_TYPE)
if request.content_length > 1_000_000: # just a sanity check to avoid memory exhaustion
return abort(HTTPStatus.BAD_REQUEST)
key = request.args.get("key", default=None)
return summarize_var_helper(request, data_adaptor, key, request.get_data())
@@ -0,0 +1,4 @@
class CxgConstants(object):
# The CXG container version number. Must be a semver string (major.minor.patch)
# DO NOT UPDATE THIS WITHOUT ALSO UPDATING CXG SPECIFICATION.
CXG_VERSION = "0.2.0"
@@ -0,0 +1,178 @@
import json
import numpy as np
import tiledb
from backend.common.utils.type_conversion_utils import get_encoding_dtype_of_array, get_dtype_and_schema_of_array
def convert_dictionary_to_cxg_group(cxg_container, metadata_dict, group_metadata_name="cxg_group_metadata"):
"""
Saves the contents of the dictionary to the CXG output directory specified.
This function is primarily used to save metadata about a dataset to the CXG directory. At some point, tiledb will
have support for metadata on groups at which point the utility of this function should be revisited. Until such
feature exists, this function create an empty array and annotate that array.
For more information, visit https://github.com/TileDB-Inc/TileDB-Py/issues/254.
"""
array_name = f"{cxg_container}/{group_metadata_name}"
# Because TileDB does not allow one to attach metadata directly to a CXG group, we need to have a workaround
# where we create an empty array and attached the metadata onto to this empty array. Below we construct this empty
# array.
tiledb.from_numpy(array_name, np.zeros((1,)))
with tiledb.DenseArray(array_name, mode="w") as metadata_array:
for key, value in metadata_dict.items():
metadata_array.meta[key] = value
def convert_dataframe_to_cxg_array(cxg_container, dataframe_name, dataframe, index_column_name, ctx):
"""
Saves the contents of the dataframe to the CXG output directory specified.
Current access patterns are oriented toward reading very large slices of the dataframe, one attribute at a time.
Attribute data also tends to be (often) repetitive (bools, categories, strings). Given this, we use a large tile
size (1000) and very aggressive compression levels.
"""
def create_dataframe_array(array_name, dataframe):
tiledb_filter = tiledb.FilterList(
[
# Attempt aggressive compression as many of these dataframes are very repetitive strings, bools and
# other non-float data.
tiledb.ZstdFilter(level=22),
]
)
attrs = [
tiledb.Attr(name=column, dtype=get_encoding_dtype_of_array(dataframe[column]), filters=tiledb_filter)
for column in dataframe
]
domain = tiledb.Domain(
tiledb.Dim(domain=(0, dataframe.shape[0] - 1), tile=min(dataframe.shape[0], 1000), dtype=np.uint32)
)
schema = tiledb.ArraySchema(
domain=domain, sparse=False, attrs=attrs, cell_order="row-major", tile_order="row-major"
)
tiledb.DenseArray.create(array_name, schema)
array_name = f"{cxg_container}/{dataframe_name}"
create_dataframe_array(array_name, dataframe)
with tiledb.DenseArray(array_name, mode="w", ctx=ctx) as array:
value = {}
schema_hints = {}
for column_name, column_values in dataframe.items():
dtype, hints = get_dtype_and_schema_of_array(column_values)
value[column_name] = column_values.to_numpy(dtype=dtype)
if hints:
schema_hints.update({column_name: hints})
schema_hints.update({"index": index_column_name})
array[:] = value
array.meta["cxg_schema"] = json.dumps(schema_hints)
tiledb.consolidate(array_name, ctx=ctx)
def convert_ndarray_to_cxg_dense_array(ndarray_name, ndarray, ctx):
"""
Saves contents of ndarray to the CXG output directory specified.
Generally this function is used to convert dataset embeddings. Because embeddings are typically accessed with
very large slices (or all of the embedding), they do not benefit from overly aggressive compression due to their
format. Given this, we use a large tile size (1000) but only default compression level.
"""
def create_ndarray_array(ndarray_name, ndarray):
filters = tiledb.FilterList([tiledb.ZstdFilter()])
attrs = [tiledb.Attr(dtype=ndarray.dtype, filters=filters)]
dimensions = [
tiledb.Dim(
domain=(0, ndarray.shape[dimension] - 1), tile=min(ndarray.shape[dimension], 1000), dtype=np.uint32
)
for dimension in range(ndarray.ndim)
]
domain = tiledb.Domain(*dimensions)
schema = tiledb.ArraySchema(
domain=domain, sparse=False, attrs=attrs, capacity=1_000_000, cell_order="row-major", tile_order="row-major"
)
tiledb.DenseArray.create(ndarray_name, schema)
create_ndarray_array(ndarray_name, ndarray)
with tiledb.DenseArray(ndarray_name, mode="w", ctx=ctx) as array:
array[:] = ndarray
tiledb.consolidate(ndarray_name, ctx=ctx)
def convert_matrix_to_cxg_array(
matrix_name, matrix, encode_as_sparse_array, ctx, column_shift_for_sparse_encoding=None
):
"""
Converts a numpy array matrix into a TileDB SparseArray of DenseArray based on whether `encode_as_sparse_array`
is true or not. Note that when the matrix is encoded as a SparseArray, it only writes the values that are
nonzero. This means that if you count the number of elements in the SparseArray, it will not equal the total
number of elements in the matrix, only the number of nonzero elements.
Furthermore, if the `column_shift_for_sparse_encoding` matrix is not None, this function will subtract the sparse
encoding from the original given matrix and as previously stated, only write the nonzero values to the TileDB
SparseArray.
"""
def create_matrix_array(matrix_name, number_of_rows, number_of_columns, encode_as_sparse_array):
filters = tiledb.FilterList([tiledb.ZstdFilter()])
attrs = [tiledb.Attr(dtype=np.float32, filters=filters)]
if encode_as_sparse_array:
domain = tiledb.Domain(
tiledb.Dim(name="obs", domain=(0, number_of_rows - 1), tile=min(number_of_rows, 512), dtype=np.uint32),
tiledb.Dim(
name="var", domain=(0, number_of_columns - 1), tile=min(number_of_columns, 2048), dtype=np.uint32
),
)
else:
domain = tiledb.Domain(
tiledb.Dim(name="obs", domain=(0, number_of_rows - 1), tile=min(number_of_rows, 50), dtype=np.uint32),
tiledb.Dim(
name="var", domain=(0, number_of_columns - 1), tile=min(number_of_columns, 100), dtype=np.uint32
),
)
schema = tiledb.ArraySchema(
domain=domain, sparse=encode_as_sparse_array, attrs=attrs, cell_order="row-major", tile_order="col-major"
)
if encode_as_sparse_array:
tiledb.SparseArray.create(matrix_name, schema)
else:
tiledb.DenseArray.create(matrix_name, schema)
number_of_rows = matrix.shape[0]
number_of_columns = matrix.shape[1]
stride = min(int(np.power(10, np.around(np.log10(1e9 / number_of_columns)))), 10_000)
create_matrix_array(matrix_name, number_of_rows, number_of_columns, encode_as_sparse_array)
if encode_as_sparse_array:
with tiledb.SparseArray(matrix_name, mode="w", ctx=ctx) as array:
for start_row_index in range(0, number_of_rows, stride):
end_row_index = min(start_row_index + stride, number_of_rows)
matrix_subset = matrix[start_row_index:end_row_index, :]
if not isinstance(matrix_subset, np.ndarray):
matrix_subset = matrix_subset.toarray()
if column_shift_for_sparse_encoding is not None:
matrix_subset = matrix_subset - column_shift_for_sparse_encoding
indices = np.nonzero(matrix_subset)
trow = indices[0] + start_row_index
array[trow, indices[1]] = matrix_subset[indices[0], indices[1]]
else:
with tiledb.DenseArray(matrix_name, mode="w", ctx=ctx) as array:
for start_row_index in range(0, number_of_rows, stride):
end_row_index = min(start_row_index + stride, number_of_rows)
matrix_subset = matrix[start_row_index:end_row_index, :]
if not isinstance(matrix_subset, np.ndarray):
matrix_subset = matrix_subset.toarray()
array[start_row_index:end_row_index, :] = matrix_subset
@@ -0,0 +1,115 @@
import logging
import numpy as np
from scipy.stats import mode
def is_matrix_sparse(matrix: np.ndarray, sparse_threshold):
"""
Returns whether `matrix` is sparse or not (i.e. dense). This is determined by figuring out whether the matrix has
a sparsity percentage below the sparse_threshold, returning the number of non-zeros encountered and number of
elements evaluated. This function may return before evaluating the whole matrix if it can be determined that matrix
is not sparse enough.
"""
if sparse_threshold == 100.0:
return True
if sparse_threshold == 0.0:
return False
total_number_of_rows = matrix.shape[0]
total_number_of_columns = matrix.shape[1]
total_number_of_matrix_elements = total_number_of_rows * total_number_of_columns
# For efficiency, we count the number of non-zero elements in chunks of the matrix at a time until we hit the
# maximum number of non zero values allowed before the matrix is deemed "dense." This allows the function the
# quit early for large dense matrices.
row_stride = min(int(np.power(10, np.around(np.log10(1e9 / total_number_of_columns)))), 10_000)
maximum_number_of_non_zero_elements_in_matrix = int(
total_number_of_rows * total_number_of_columns * sparse_threshold / 100
)
number_of_non_zero_elements = 0
for start_row_index in range(0, total_number_of_rows, row_stride):
end_row_index = min(start_row_index + row_stride, total_number_of_rows)
matrix_subset = matrix[start_row_index:end_row_index, :]
if not isinstance(matrix_subset, np.ndarray):
matrix_subset = matrix_subset.toarray()
number_of_non_zero_elements += np.count_nonzero(matrix_subset)
if number_of_non_zero_elements > maximum_number_of_non_zero_elements_in_matrix:
if end_row_index != total_number_of_rows:
percentage_of_non_zero_elements = (
100 * number_of_non_zero_elements / (end_row_index * total_number_of_columns)
)
logging.info(
f"Matrix is not sparse. Percentage of non-zero elements (estimate): "
f"{percentage_of_non_zero_elements:6.2f}"
)
else:
percentage_of_non_zero_elements = 100 * number_of_non_zero_elements / total_number_of_matrix_elements
logging.info(
f"Matrix is not sparse. Percentage of non-zero elements (exact): "
f"{percentage_of_non_zero_elements:6.2f}"
)
return False
is_sparse = (100.0 * number_of_non_zero_elements / total_number_of_matrix_elements) < sparse_threshold
return is_sparse
def get_column_shift_encode_for_matrix(matrix, sparse_threshold):
"""
Returns a column shift if there is a column shift that allows the given matrix to be considered as sparse. Column
shift encoding works by taking the most common value in each column, then subtracting that value from each element
of the column. If each column mostly contains its most common value, then the resulting matrix can be very sparse.
This function determines if column shift encoding can be used to transform the matrix into a sparse matrix with a
sparsity below the sparse_threshold. If so, returns the array that stores this encoding. This function also returns
the number of non-zeros encountered and number of elements evaluated. This function may return before evaluating
the whole matrix if it can be determined that the matrix cannot benefit from column shift encoding.
"""
total_number_of_rows = matrix.shape[0]
total_number_of_columns = matrix.shape[1]
total_number_of_matrix_elements = total_number_of_rows * total_number_of_columns
stride = max(1, 128_000_000 // total_number_of_rows)
column_shift = np.zeros(total_number_of_columns)
maximum_number_of_non_zero_elements_in_matrix = int(
total_number_of_rows * total_number_of_columns * sparse_threshold / 100
)
number_of_non_zero_elements = 0
for start_column_index in range(0, total_number_of_columns, stride):
end_column_index = min(start_column_index + stride, total_number_of_columns)
matrix_subset = matrix[:, start_column_index:end_column_index]
if not isinstance(matrix_subset, np.ndarray):
matrix_subset = matrix_subset.toarray()
matrix_subset_mode = mode(matrix_subset)
column_shift[start_column_index:end_column_index] = matrix_subset_mode.mode
number_of_non_zero_elements += total_number_of_rows * (end_column_index - start_column_index) - np.sum(
matrix_subset_mode.count
)
if number_of_non_zero_elements > maximum_number_of_non_zero_elements_in_matrix:
if end_column_index != total_number_of_columns:
logging.info(
"Matrix is not sparse even with column shift. Percentage of non-zero elements (estimate): %6.2f"
% (100 * number_of_non_zero_elements / end_column_index * total_number_of_rows)
)
else:
logging.info(
"Matrix is not sparse even with column shift. Percentage of non-zero elements (exact): %6.2f"
% (100 * number_of_non_zero_elements / total_number_of_matrix_elements)
)
return None
is_sparse = (100.0 * number_of_non_zero_elements / total_number_of_matrix_elements) < sparse_threshold
return column_shift if is_sparse else None
@@ -0,0 +1,40 @@
import re
def sanitize_values_in_list(list_of_keys: list):
"""
Returns a dictionary mapping of the old keys in the list of `list_of_keys` to its new, clean name that is both
safe and unique.
"""
if not all([isinstance(key, str) for key in list_of_keys]):
raise Exception("List of keys to sanitize must contain all strings.")
# Mask out [~/.] and anything outside the ASCII range.
mask = re.compile(r"[^ -\-0-\[\]-\}]")
clean_keys_list = [mask.sub("_", key) for key in list_of_keys]
# Dedupe the clean keys list
deduped_clean_keys_list = []
for index, clean_key in enumerate(clean_keys_list):
total_occurrences_of_clean_key = clean_keys_list.count(clean_key)
total_occurrences_up_until_current_index = clean_keys_list[:index].count(clean_key)
deduped_clean_keys_list.append(
clean_key + "_" + str(total_occurrences_up_until_current_index + 1)
if total_occurrences_of_clean_key > 1
else clean_key
)
return dict(zip(list_of_keys, deduped_clean_keys_list))
def sanitize_keys_in_dictionary(dict_to_sanitize: dict):
"""
Clean and dedupe the keys in the given dictionary.
"""
clean_keys = sanitize_values_in_list(dict_to_sanitize.keys())
for original_key, sanitized_key in clean_keys.items():
if original_key != sanitized_key:
dict_to_sanitize[sanitized_key] = dict_to_sanitize[original_key]
del dict_to_sanitize[original_key]
+199
View File
@@ -0,0 +1,199 @@
import concurrent.futures
import numpy as np
from numba import jit
from backend.czi_hosted.data_cxg.cxg_util import pack_selector_from_indices
from backend.common.compute.diffexp_generic import diffexp_ttest_from_mean_var, mean_var_n
from backend.common.errors import ComputeError
"""
See the comments in diffexp_generic for a description of this algorithm
This implementation runs directly in-process. It is multi- threaded, but not particularly scalable.
Longer term, will likely move to a distributed framework for this.
There are currently no global throttles on simultaneous workers.
"""
diffexp_thread_executor = None
max_workers = None
target_workunit = None
def set_config(config_max_workers, config_target_workunit):
global max_workers
global target_workunit
max_workers = config_max_workers
target_workunit = config_target_workunit
def get_thread_executor():
global diffexp_thread_executor
if diffexp_thread_executor is None:
diffexp_thread_executor = concurrent.futures.ThreadPoolExecutor(max_workers=max_workers)
return diffexp_thread_executor
def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
matrix = adaptor.open_array("X")
row_selector_A = np.where(maskA)[0]
row_selector_B = np.where(maskB)[0]
nA = len(row_selector_A)
nB = len(row_selector_B)
dtype = matrix.dtype
cols = matrix.shape[1]
tile_extent = [dim.tile for dim in matrix.schema.domain]
is_sparse = matrix.schema.sparse
if is_sparse:
row_selector_A = pack_selector_from_indices(row_selector_A)
row_selector_B = pack_selector_from_indices(row_selector_B)
else:
# The rows from both row_selector_A and row_selector_B are gathered at the
# same time, then the mean and variance are computed by subsetting on that
# combined submatrix. Combining the gather reduces number of requests/bandwidth
# to the data source.
row_selector_AB = np.union1d(row_selector_A, row_selector_B)
row_selector_A_in_AB = np.in1d(row_selector_AB, row_selector_A, assume_unique=True)
row_selector_B_in_AB = np.in1d(row_selector_AB, row_selector_B, assume_unique=True)
row_selector_AB = pack_selector_from_indices(row_selector_AB)
# because all IO is done per-tile, and we are always col-major,
# use the tile column size as the unit of partition. Possibly access
# more than one column tile at a time based on the target_workunit.
# Revisit partitioning if we change the X layout, or start using a non-local execution environment
# which may have other constraints.
# TODO: If the number of row selections is large enough, then the cells_per_coltile will exceed
# the target_workunit. A potential improvement would be to partition by both columns and rows.
# However partitioning the rows is slightly more complex due to the arbitrary distribution
# of row selections that are passed into this algorithm.
cells_per_coltile = (nA + nB) * tile_extent[1]
cols_per_partition = max(1, int(target_workunit / cells_per_coltile)) * tile_extent[1]
col_partitions = [(c, min(c + cols_per_partition, cols)) for c in range(0, cols, cols_per_partition)]
meanA = np.zeros((cols,), dtype=np.float64)
varA = np.zeros((cols,), dtype=np.float64)
meanB = np.zeros((cols,), dtype=np.float64)
varB = np.zeros((cols,), dtype=np.float64)
executor = get_thread_executor()
futures = []
if is_sparse:
for cols in col_partitions:
futures.append(executor.submit(_mean_var_sparse_ab, matrix, row_selector_A, nA, row_selector_B, nB, cols))
else:
for cols in col_partitions:
futures.append(
executor.submit(_mean_var_ab, matrix, row_selector_AB, row_selector_A_in_AB, row_selector_B_in_AB, cols)
)
for future in futures:
# returns tuple: (meanA, varA, meanB, varB, cols)
try:
result = future.result()
part_meanA, part_varA, part_meanB, part_varB, cols = result
meanA[cols[0] : cols[1]] += part_meanA
varA[cols[0] : cols[1]] += part_varA
meanB[cols[0] : cols[1]] += part_meanB
varB[cols[0] : cols[1]] += part_varB
except Exception as e:
for future in futures:
future.cancel()
raise ComputeError(str(e))
if is_sparse:
if adaptor.has_array("X_col_shift"):
X_col_shift = adaptor.open_array("X_col_shift")[:]
meanA += X_col_shift
meanB += X_col_shift
r = diffexp_ttest_from_mean_var(
meanA=meanA.astype(dtype),
varA=varA.astype(dtype),
nA=nA,
meanB=meanB.astype(dtype),
varB=varB.astype(dtype),
nB=nB,
top_n=top_n,
diffexp_lfc_cutoff=diffexp_lfc_cutoff
)
return r
def _mean_var_ab(matrix, row_selector_AB, row_selector_A_in_AB, row_selector_B_in_AB, col_range):
X = matrix.multi_index[row_selector_AB, col_range[0] : col_range[1] - 1][""]
meanA, varA, n = mean_var_n(X[row_selector_A_in_AB])
meanB, varB, n = mean_var_n(X[row_selector_B_in_AB])
return (meanA, varA, meanB, varB, col_range)
def _mean_var_sparse_ab(matrix, row_selector_A, nrows_A, row_selector_B, nrows_B, col_range):
meanA, varA = _mean_var_sparse(matrix, row_selector_A, nrows_A, col_range)
meanB, varB = _mean_var_sparse(matrix, row_selector_B, nrows_B, col_range)
return (meanA, varA, meanB, varB, col_range)
@jit(nopython=True)
def _mean_var_sparse_numba(x, var, nrows, ncols):
"""Kernel to compute the mean and variance. It was not clear if this function
could be written using numpy, thus avoiding the loops. Therefore numba is
used here to speed things up. With numba, this function takes a negligible amount
of time compared to reading in the sparse matrix"""
mean = np.zeros((ncols,), dtype=np.float64)
for col, val in zip(var, x):
mean[col] += val
mean /= nrows
# optimize the sumsq computation.
# since most entries in a sparse matrix are 0, then start by assuming
# all values are 0, so fill the sumsq array with nrows * (0 - mean)**2.
# as non-zero values are encountered, subtract off the (mean*mean) value
# and replace with (val-mean)**2. Simplifying the expression
# gives the following code.
sumsq = nrows * np.multiply(mean, mean)
for col, val in zip(var, x):
sumsq[col] += val * (val - 2 * mean[col])
v = sumsq / (nrows - 1)
return mean, v
def _mean_var_sparse(matrix, selector, nrows, col_range):
data = matrix.multi_index[selector, col_range[0] : col_range[1] - 1]
x = data[""]
# tiledb < 0.6.0 and >= 0.6.0 have slightly different interfaces.
# the following takes care of both cases:
# older: data["coords]["var"]
# newer: data["var"]
var = data.get("coords", data)["var"]
# shift the column indices to start at 0, this
# will become the index into the mean and var arrays.
var -= col_range[0]
fp_err_occurred = False
def fp_err_set(err, flag):
nonlocal fp_err_occurred
fp_err_occurred = True
ncols = col_range[1] - col_range[0]
with np.errstate(divide="call", invalid="call", call=fp_err_set):
mean, v = _mean_var_sparse_numba(x, var, nrows, ncols)
if fp_err_occurred:
mean[np.isfinite(mean) == False] = 0 # noqa: E712
v[np.isfinite(v) == False] = 0 # noqa: E712
else:
mean[np.isnan(mean)] = 0
v[np.isnan(v)] = 0
return mean, v
@@ -0,0 +1,250 @@
import json
import logging
from os import path
import anndata
import numpy as np
import tiledb
from backend.common.colors import convert_anndata_category_colors_to_cxg_category_colors
from backend.czi_hosted.common.corpora import corpora_get_props_from_anndata
from backend.common.errors import ColorFormatException
from backend.czi_hosted.common.utils.cxg_constants import CxgConstants
from backend.czi_hosted.common.utils.cxg_generation_utils import (
convert_dictionary_to_cxg_group,
convert_dataframe_to_cxg_array,
convert_ndarray_to_cxg_dense_array,
convert_matrix_to_cxg_array,
)
from backend.czi_hosted.common.utils.matrix_utils import is_matrix_sparse, get_column_shift_encode_for_matrix
class H5ADDataFile:
""" Class encapsulating required information about an H5AD datafile that ultimately will be transformed into
another format (currently just CXG is supported). """
def __init__(
self,
input_filename,
backed=False,
dataset_title=None,
dataset_about=None,
obs_index_column_name=None,
vars_index_column_name=None,
use_corpora_schema=True,
):
self.input_filename = input_filename
self.backed = backed
self.dataset_title = dataset_title
self.dataset_about = dataset_about
self.obs_index_column_name = obs_index_column_name
self.vars_index_column_name = vars_index_column_name
self.use_corpora_schema = use_corpora_schema
self.validate_input_file_type()
self.extract_anndata_elements_from_file()
self.extract_metadata_about_dataset()
self.validate_anndata()
def to_cxg(self, output_cxg_directory, sparse_threshold, convert_anndata_colors_to_cxg_colors=True):
"""
Writes the following attributes of the anndata to CXG: 1) the metadata as metadata attached to an empty
DenseArray, 2) the obs DataFrame as a DenseArray, 3) the var DataFrame as a DenseArray, 4) all valid
embeddings stored in obsm, each one as a DenseArray, 5) the main X matrix of the anndata as either a
SparseArray or DenseArray based on the `sparse_threshold`, and optionally 6) the column shift of the main X
matrix that might turn an otherwise Dense matrix into a Sparse matrix.
"""
logging.info("Beginning writing to CXG.")
ctx = tiledb.Ctx(
{
"sm.num_reader_threads": 32,
"sm.num_writer_threads": 32,
"sm.consolidation.buffer_size": 1 * 1024 * 1024 * 1024,
}
)
tiledb.group_create(output_cxg_directory, ctx=ctx)
logging.info(f"\t...group created, with name {output_cxg_directory}")
convert_dictionary_to_cxg_group(
output_cxg_directory, self.generate_cxg_metadata(convert_anndata_colors_to_cxg_colors)
)
logging.info("\t...dataset metadata saved")
convert_dataframe_to_cxg_array(output_cxg_directory, "obs", self.obs, self.obs_index_column_name, ctx)
logging.info("\t...dataset obs dataframe saved")
convert_dataframe_to_cxg_array(output_cxg_directory, "var", self.var, self.var_index_column_name, ctx)
logging.info("\t...dataset var dataframe saved")
self.write_anndata_embeddings_to_cxg(output_cxg_directory, ctx)
logging.info("\t...dataset embeddings saved")
self.write_anndata_x_matrix_to_cxg(output_cxg_directory, ctx, sparse_threshold)
logging.info("\t...dataset X matrix saved")
logging.info("Completed writing to CXG.")
def write_anndata_x_matrix_to_cxg(self, output_cxg_directory, ctx, sparse_threshold):
matrix_container = f"{output_cxg_directory}/X"
x_matrix_data = self.anndata.X
is_sparse = is_matrix_sparse(x_matrix_data, sparse_threshold)
if not is_sparse:
col_shift = get_column_shift_encode_for_matrix(x_matrix_data, sparse_threshold)
is_sparse = col_shift is not None
else:
col_shift = None
if col_shift is not None:
logging.info("Converting matrix X as sparse matrix with column shift encoding")
x_col_shift_name = f"{output_cxg_directory}/X_col_shift"
convert_ndarray_to_cxg_dense_array(x_col_shift_name, col_shift, ctx)
convert_matrix_to_cxg_array(matrix_container, x_matrix_data, is_sparse, ctx, col_shift)
tiledb.consolidate(matrix_container, ctx=ctx)
if hasattr(tiledb, "vacuum"):
tiledb.vacuum(matrix_container)
def write_anndata_embeddings_to_cxg(self, output_cxg_directory, ctx):
def is_valid_embedding(adata, embedding_name, embedding_array):
"""
Returns true if this layout data is a valid array for front-end presentation with the following criteria:
* ndarray, with shape (n_obs, >= 2), dtype float/int/uint
* follows ScanPy embedding naming conventions
* with all values finite or NaN (no +Inf or -Inf)
"""
is_valid = isinstance(embedding_name, str) and embedding_name.startswith("X_") and len(embedding_name) > 2
is_valid = is_valid and isinstance(embedding_array, np.ndarray) and embedding_array.dtype.kind in "fiu"
is_valid = is_valid and embedding_array.shape[0] == adata.n_obs and embedding_array.shape[1] >= 2
is_valid = is_valid and not np.any(np.isinf(embedding_array)) and not np.all(np.isnan(embedding_array))
return is_valid
embedding_container = f"{output_cxg_directory}/emb"
tiledb.group_create(embedding_container, ctx=ctx)
for embedding_name, embedding_values in self.anndata.obsm.items():
if is_valid_embedding(self.anndata, embedding_name, embedding_values):
embedding_name = f"{embedding_container}/{embedding_name[2:]}"
convert_ndarray_to_cxg_dense_array(embedding_name, embedding_values, ctx)
logging.info(f"\t\t...{embedding_name} embedding created")
def generate_cxg_metadata(self, convert_anndata_colors_to_cxg_colors):
"""
Return a dictionary containing metadata about CXG dataset. This include data about the version as well as
Corpora schema properties if they exist, among other pieces of metadata.
"""
cxg_group_metadata = {
"cxg_version": CxgConstants.CXG_VERSION,
"cxg_properties": json.dumps({"title": self.dataset_title, "about": self.dataset_about}),
}
if self.corpora_properties is not None:
cxg_group_metadata["corpora"] = json.dumps(self.corpora_properties)
if convert_anndata_colors_to_cxg_colors:
try:
cxg_group_metadata["cxg_category_colors"] = json.dumps(
convert_anndata_category_colors_to_cxg_category_colors(self.anndata)
)
except ColorFormatException:
logging.warning(
"Failed to extract colors from H5AD file! Fix the H5AD file or rerun with "
"--disable-custom-colors. See help for more details."
)
return cxg_group_metadata
def validate_input_file_type(self):
"""
Validate that the input file is of a type that we can handle. Currently the only valid file type is `.h5ad`.
"""
if not self.input_filename.endswith(".h5ad"):
raise Exception(f"Cannot process input file {self.input_filename}. File must be an H5AD.")
if self.dataset_title or self.dataset_about:
logging.warning(
"If you convert this dataset into CXG and you explicit specify values for the dataset title metadata "
"or the dataset about metadata, it will override any metadata that is extracted as part of the "
"Corpora schema fields."
)
def validate_anndata(self):
if not self.var.index.is_unique:
raise ValueError("Variable index in AnnData object is not unique.")
if not self.obs.index.is_unique:
raise ValueError("Observation index in AnnData object is not unique.")
def extract_anndata_elements_from_file(self):
logging.info(f"Reading in AnnData dataset: {path.basename(self.input_filename)}")
self.anndata = anndata.read_h5ad(self.input_filename, backed="r" if self.backed else None)
logging.info("Completed reading in AnnData dataset!")
self.obs = self.transform_dataframe_index_into_column(self.anndata.obs, "obs", self.obs_index_column_name)
self.var = self.transform_dataframe_index_into_column(self.anndata.var, "var", self.vars_index_column_name)
def extract_metadata_about_dataset(self):
"""
Extract metadata information about the dataset that upon conversion will be saved as group metadata with the
CXG that is generated. This metadata information includes Corpora schema properties, the dataset title and
a link that details more information about the dataset.
"""
self.corpora_properties = corpora_get_props_from_anndata(self.anndata) if self.use_corpora_schema else None
if self.corpora_properties is None and self.use_corpora_schema:
# If the return value is None, this means that we were not able to figure out what version of the Corpora
# schema the object is using and therefore cannot extract any properties.
raise ValueError("Unknown source file schema version is unsupported.")
# The title and about properties of the dataset are set by the following order: if they are explicitly defined
# then use the explicit value. If the dataset is a Corpora-schema based schema, then extract the title and about
# from the corpora_properties. Otherwise, use the input filename (only for title, about will be blank).
if self.corpora_properties:
corpora_project_links = self.corpora_properties.get("project_links", [])
corpora_about_link = next(
(link for link in corpora_project_links if (link.get("link_type", None) == "SUMMARY")), {}
)
else:
corpora_about_link = {}
filename = path.splitext(path.basename(self.input_filename))[0]
self.dataset_title = self.dataset_title if self.dataset_title else corpora_about_link.get("link_name", filename)
self.dataset_about = self.dataset_about if self.dataset_about else corpora_about_link.get("link_url")
def transform_dataframe_index_into_column(self, dataframe, dataframe_name, index_column_name):
"""
Convert the dataframe's index into another column in the dataframe. If an index_column_name is specified,
use that column as the index instead.
"""
if index_column_name is None:
# Create a unique column name for the index.
suffix = 0
while f"name_{suffix}" in dataframe.columns:
suffix += 1
index_column_name = f"name_{suffix}"
# Turn the index into a normal column
dataframe.rename_axis(index_column_name, inplace=True)
dataframe.reset_index(inplace=True)
elif index_column_name in dataframe.columns:
# User has specified alternative column for unique names, and it exists
if not dataframe[index_column_name].is_unique:
raise KeyError(
f"Values in {dataframe_name}.{index_column_name} must be unique. Please prepare data to contain "
f"unique values."
)
else:
raise KeyError(f"Column {index_column_name} does not exist.")
setattr(self, f"{dataframe_name}_index_column_name", index_column_name)
return dataframe
@@ -0,0 +1,211 @@
"""Helpers for converting and checking HGNC gene symbols."""
import argparse
import enum
import logging
import os
import re
import numpy as np
import pandas as pd
def get_upgraded_var_index(var, hgnc_path=None):
"""Given an anndata var dataframe, return a new index for the dataframe
where human gene symbols have been upgraded to the current HGNC set.
"""
if not hgnc_path:
hgnc_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "hgnc_complete_set.txt.gz")
hgnc_symbol_checker = HGNCSymbolChecker.from_hgnc_records(hgnc_path)
return pd.Index([hgnc_symbol_checker.upgrade_symbol(s) for s in var.index])
class SymbolStatus(enum.Enum):
"""The status of a symbol in the HGNC database.
APPROVED: Currently a valid symbol
WITHDRAWN: A previously approved HGNC symbol for a gene that has since been shown
not to exist _unless_ that symbol is also approved
AMBIGUOUS: A symbol that is not approved but is an alias or previous symbol for
multiple approved symbols
UPGRADABLE: A symbol that is not approved but unambiguously maps to an approved
symbol
UNKNOWN: A symbol that does not appear in HGNC
"""
APPROVED = 1
WITHDRAWN = 2
AMBIGUOUS = 3
UPGRADABLE = 4
UNKNOWN = 5
class HGNCSymbolChecker:
"""Handle checking and correcting HGNC symbols."""
def __init__(self, approved_symbols, withdrawn_symbols, ambiguous_symbols, symbol_map):
self.approved_symbols = approved_symbols
self.withdrawn_symbols = withdrawn_symbols
self.ambiguous_symbols = ambiguous_symbols
self.symbol_map = symbol_map
def print_symbol_map(self):
"""Print out a map from old symbol to new symbol."""
for symbol_pair in self.symbol_map.items():
print("\t".join(symbol_pair))
def check_symbol(self, symbol):
"""See if a symbol if approved or something else."""
if symbol in self.approved_symbols:
return SymbolStatus.APPROVED
if symbol in self.withdrawn_symbols:
return SymbolStatus.WITHDRAWN
if symbol in self.ambiguous_symbols:
return SymbolStatus.AMBIGUOUS
if symbol in self.symbol_map:
return SymbolStatus.UPGRADABLE
return SymbolStatus.UNKNOWN
def upgrade_symbol(self, symbol):
"""Return the approved symbol for the given symbol.
If the symbol cannot be upgraded, just return the original symbol.
"""
fixed_symbol, stripped_symbol = format_symbol(symbol)
if fixed_symbol in self.approved_symbols:
return fixed_symbol
elif fixed_symbol in self.symbol_map:
return self.symbol_map[fixed_symbol]
elif stripped_symbol in self.approved_symbols:
return stripped_symbol
elif stripped_symbol in self.symbol_map:
return self.symbol_map[stripped_symbol]
return symbol
@classmethod
def from_hgnc_records(cls, hgnc_dataset_path):
"""Parse a hgnc database download into a HGNCSymbolChecker object."""
def all_symbols(record):
"""Get all the symbols associated with an HGNC record including previous, alias,
and approved."""
yield format_symbol(record["symbol"])[0]
for symbol in alias_and_previous_symbols(record):
yield symbol
def alias_and_previous_symbols(record):
"""Get alias and previous symbols from an HGNC record."""
for field in ("alias_symbol", "prev_symbol"):
if record[field] is not np.nan:
for symbol in record[field].split("|"):
yield format_symbol(symbol)[0]
# Sometimes something like HGNC:1234 appears in datasets, which we
# want to fix as well.
yield record["hgnc_id"]
hgnc_records = pd.read_csv(hgnc_dataset_path, sep="\t", header=0, low_memory=False).to_dict("records")
# Get all symbols that are currently approved.
approved_symbols = set()
for record in hgnc_records:
if record["status"] == "Approved":
approved_symbols.add(format_symbol(record["symbol"])[0])
# Get all symbols that have been withdrawn
withdrawn_symbols = set()
for record in hgnc_records:
if record["status"] == "Entry Withdrawn":
for symbol in all_symbols(record):
withdrawn_symbols.add(symbol)
# If a symbol is both approved and withdrawn, be optimistic and call it approved
logging.warning(
f"Some symbols are simulaneously withdrawn and approved\n"
f"We will treat them at approved:\n"
f"{withdrawn_symbols.intersection(approved_symbols)}"
)
withdrawn_symbols = withdrawn_symbols.difference(approved_symbols)
# Now try to map from symbols that are not approved but are an alias or previous symbol for an approved symbol
alias_previous_to_approved = {}
ambiguous_symbols = set()
for record in hgnc_records:
if record["status"] == "Approved":
# The approved symbol is what we'll map to
approved_symbol = format_symbol(record["symbol"])[0]
for symbol in alias_and_previous_symbols(record):
# If the alias or previous symbol is also an approved symbol,
# we'll just leave it alone
if symbol in approved_symbols:
continue
# If the alias or previous symbol maps to a different approved symbol, mark it as ambiguous
if symbol in alias_previous_to_approved and alias_previous_to_approved[symbol] != approved_symbol:
ambiguous_symbols.add(symbol)
else:
alias_previous_to_approved[symbol] = approved_symbol
# Remove all the ambiguous symbols from the map
for ambiguous_symbol in ambiguous_symbols:
alias_previous_to_approved.pop(ambiguous_symbol)
return HGNCSymbolChecker(approved_symbols, withdrawn_symbols, ambiguous_symbols, alias_previous_to_approved)
def format_symbol(symbol):
"""HGNC rules say symbols should all be upper case except for C#orf#. However, case is
variable in both alias and previous symbols as well as in the symbols we get in
submissions. So, upper case everything except for the one situation where mixed-case
is allowed, which are the genes like C2orf157.
Also, seurat and scanpy append ".1" or "-1" to duplicated gene names, and these altered
names persist throughout the life of the object. They won't match against the HGNC database
and we want to merge them, so we need to strip off the suffix and try matching again.
This function takes a symbol and returns the symbol with the fixed case and also with the
seurat/scanpy suffix stripped off.
"""
match = re.match(r"^(C)(\d+)(orf)(\d+)$", symbol, re.IGNORECASE)
if match:
fixed_case = f"C{match.group(2)}orf{match.group(4)}"
else:
fixed_case = symbol.upper()
suffix_stripped = re.sub(r"[\.\-]\d+$", "", fixed_case)
return fixed_case, suffix_stripped
def main():
"""When called as main, parse a given hgnc download and print out a map from old to new
symbol.
"""
parser = argparse.ArgumentParser()
parser.add_argument(
"hgnc_dataset", help="HGNC dataset tsv, available from www.genenames.org/download/statistics-and-files/"
)
args = parser.parse_args()
hgnc_symbol_checker = HGNCSymbolChecker.from_hgnc_records(args.hgnc_dataset)
hgnc_symbol_checker.print_symbol_map()
if __name__ == "__main__":
main()
@@ -0,0 +1,86 @@
"""Methods for working with ontologies and the OLS."""
from urllib.parse import quote_plus
import requests
OLS_API_ROOT = "http://www.ebi.ac.uk/ols/api"
# Curie means something like CL:0000001
def _ontology_name(curie):
"""Get the name of the ontology from the curie, CL or UBERON for example."""
return curie.split(":")[0]
def _ontology_value(curie):
"""Get the id component of the curie, 0000001 from CL:0000001 for example."""
return curie.split(":")[1]
def _double_encode(url):
"""Double url encode a url. This is required by the OLS API."""
return quote_plus(quote_plus(url))
def _iri(curie):
"""Get the iri from a curie. This is a bit hopeful that they all map to purl.obolibrary.org"""
if _ontology_name(curie) == "EFO":
return f"http://www.ebi.ac.uk/efo/EFO_{_ontology_value(curie)}"
return f"http://purl.obolibrary.org/obo/{_ontology_name(curie)}_{_ontology_value(curie)}"
class OntologyLookupError(Exception):
"""Exception for some problem with looking up ontology information."""
def _ontology_info_url(curie):
"""Get the to make a GET to to get information about an ontology term."""
# If the curie is empty, just return an empty string. This happens when there is no
# valid ontology value.
if not curie:
return ""
else:
return f"{OLS_API_ROOT}/ontologies/{_ontology_name(curie)}/terms/{_double_encode(_iri(curie))}"
def get_ontology_label(curie):
"""For a given curie like 'CL:1000413', get the label like 'endothelial cell of artery'"""
url = _ontology_info_url(curie)
if not url:
return ""
response = requests.get(url)
if not response.ok:
raise OntologyLookupError(
f"Curie {curie} lookup failed, got status code {response.status_code}: {response.text}"
)
return response.json()["label"]
def lookup_candidate_term(label, ontology="cl", method="select"):
"""Lookup candidate terms for a label. This is useful when there is an existing label in a
submitted dataset, and you want to find an appropriate ontology term.
Args:
label: the label to find ontology terms for
ontology: the ontology to search in, cl or uberon or efo for example
method: select or search. search provides much broader results
Returns:
list of (curie, label) tuples returned by OLS
"""
# using OLS REST API [https://www.ebi.ac.uk/ols/docs/api]
url = f"{OLS_API_ROOT}/{method}?q={quote_plus(label)}&ontology={ontology.lower()}"
response = requests.get(url)
if not response.ok:
raise OntologyLookupError(
f"Label {label} lookup failed, got status code {response.status_code}: {response.text}"
)
return [(r["obo_id"], r["label"]) for r in response.json()["response"]["docs"]]
@@ -0,0 +1,264 @@
import argparse
import collections
import json
import logging
import math
import string
import anndata
import numpy as np
import pandas as pd
import yaml
from . import gene_symbol
from . import ontology
from . import validate
REPLACE_SUFFIX = "_original"
ONTOLOGY_SUFFIX = "_ontology_term_id"
def is_curie(value):
"""Return True iff the value is an OBO-id CURIE like EFO:000001"""
return (value.count(":")
and all(len(part) > 0 for part in value.split(":"))
and all(c in string.digits for c in value.split(":")[1]))
def is_ontology_field(field_name):
"""Return True iff the field_name is an ontology field like tissue_ontology_term_id"""
return field_name.endswith(ONTOLOGY_SUFFIX)
def get_label_field_name(field_name):
"""Get the associated label field from an ontology field, assay_ontology_term_id --> assay"""
return field_name[: -len(ONTOLOGY_SUFFIX)]
def split_suffix(maybe_curie):
"""Split off the (cell culture) or (organoid) suffix."""
suffixes = [" (cell culture)", " (organoid)"]
for suffix in suffixes:
if maybe_curie.endswith(suffix):
return maybe_curie[:-len(suffix)], suffix
return maybe_curie, ""
def get_curie_and_label(maybe_curie):
"""Given a string that might be a curie, return a (curie, label) pair"""
maybe_curie, suffix = split_suffix(maybe_curie)
if not is_curie(maybe_curie):
return ("", maybe_curie + suffix)
return (maybe_curie + suffix, ontology.get_ontology_label(maybe_curie) + suffix)
def safe_add_field(adata_attr, field_name, field_value):
"""Add a field and value to an AnnData, but don't clobber an exising value."""
if (
isinstance(field_value, list)
and field_value
and isinstance(field_value[0], dict)
):
field_value = json.dumps(field_value)
if field_name in adata_attr:
adata_attr[field_name + REPLACE_SUFFIX] = adata_attr[field_name]
adata_attr[field_name] = field_value
def remix_uns(adata, uns_config):
"""Add fields from the config to adata.uns"""
for field_name, field_value in uns_config.items():
if is_ontology_field(field_name):
# If it's an ontology field, look it up
label_field_name = get_label_field_name(field_name)
ontology_term, ontology_label = get_curie_and_label(field_value)
safe_add_field(adata.uns, field_name, ontology_term)
safe_add_field(adata.uns, label_field_name, ontology_label)
else:
safe_add_field(adata.uns, field_name, field_value)
def remix_obs(adata, obs_config):
"""Add fields from the config to adata.obs"""
for field_name, field_value in obs_config.items():
if isinstance(field_value, dict):
# If the value is a dict, that means we are supposed to map from an
# existing column to the new one
source_column, column_map = next(iter(field_value.items()))
nan_value = None
for key in column_map:
if isinstance(key, float) and math.isnan(key):
nan_value = column_map[key]
if nan_value is not None:
column_map["nan"] = nan_value
for key in column_map:
if key not in adata.obs[source_column].unique():
logging.warning(f'Key {key} not in adata.obs["{source_column}"]')
for value in adata.obs[source_column].unique():
if value not in column_map:
logging.warning(f'Value {value} in adata.obs["{source_column}"] not in translation dict')
if is_ontology_field(field_name):
ontology_term_map, ontology_label_map = {}, {}
logging.info(f"Looking up labels for {field_name}")
for original_value, maybe_curie in column_map.items():
curie, label = get_curie_and_label(maybe_curie)
ontology_term_map[original_value] = curie
ontology_label_map[original_value] = label
logging.info(f"Mapping {original_value} -> {curie} -> {label}")
ontology_column = adata.obs[source_column].replace(
ontology_term_map, inplace=False
)
label_column = adata.obs[source_column].replace(
ontology_label_map, inplace=False
)
safe_add_field(adata.obs, field_name, ontology_column)
safe_add_field(
adata.obs, get_label_field_name(field_name), label_column
)
else:
label_column = adata.obs[source_column].replace(
column_map, inplace=False
)
safe_add_field(adata.obs, field_name, label_column)
else:
if is_ontology_field(field_name):
# If it's an ontology field, look it up
label_field_name = get_label_field_name(field_name)
ontology_term, ontology_label = get_curie_and_label(field_value)
safe_add_field(adata.obs, field_name, ontology_term)
safe_add_field(adata.obs, label_field_name, ontology_label)
else:
safe_add_field(adata.obs, field_name, field_value)
def merge_df(df, domain, index, columns):
"""
Given a dataframe with duplicate column labels, merge and return a dataframe where
the duplicates have been merged together, resulting in a dataframe with unique column
labels.
"merge" depends on the value of domain. If the domain is "raw", then duplicate columns
can just be summed. If it's "log1p" or "sqrt", it needs to be exp1m'd or squared, then
summed, and then logged or sqrt'd again.
"""
if not isinstance(df, np.ndarray):
to_merge = df.toarray()
else:
to_merge = df
if domain == "raw":
merged_df = pd.DataFrame(to_merge, index=index, columns=columns).sum(
axis=1, level=0, skipna=False
)
elif domain == "log1p":
merged_df = (
pd.DataFrame(np.expm1(to_merge, dtype=np.float128), index=index, columns=columns)
.sum(axis=1, level=0, skipna=False)
)
merged_df = pd.DataFrame(np.log1p(merged_df.to_numpy()), index=merged_df.index, columns=merged_df.columns)
elif domain == "sqrt":
merged_df = (
pd.DataFrame(np.square(to_merge), index=index, columns=columns)
.sum(axis=1, level=0, skipna=False)
)
merged_df = pd.DataFrame(np.sqrt(merged_df.to_numpy()), index=merged_df.index, columns=merged_df.columns)
return merged_df
def fixup_gene_symbols(adata, fixup_config):
"""Update the var index to hold a consistent set of HGNC gene symbols."""
upgraded_var_index = gene_symbol.get_upgraded_var_index(adata.var)
merged_X = merge_df(adata.X, fixup_config["X"], adata.obs.index, upgraded_var_index)
fixup_adata = anndata.AnnData(
X=merged_X,
obs=adata.obs,
var=merged_X.columns.to_frame(name="hgnc_gene_symbol"),
uns=adata.uns,
obsm=adata.obsm,
)
for layer, domain in fixup_config.items():
if layer == "X":
continue
if layer == "raw.X":
df = adata.raw.X
else:
df = adata.layers[layer]
merged_df = merge_df(df, domain, adata.obs.index, upgraded_var_index)
assert merged_df.index.equals(merged_X.index)
assert merged_df.columns.equals(merged_X.columns)
if domain == "raw":
fixup_raw = anndata.AnnData(
X=merged_df,
obs=adata.obs,
var=merged_X.columns.to_frame(name="hgnc_gene_symbol"),
)
fixup_adata.raw = fixup_raw
else:
fixup_adata.layers[layer] = merged_df
return fixup_adata
def _strip_version(adata):
"""Remove version information from the AnnData object."""
if "version" in adata.uns_keys():
del adata.uns["version"]
def apply_schema(source_h5ad, remix_config, output_filename):
try:
import scanpy
except ImportError:
raise ImportError("scanpy must be installed for cellxgene schema")
adata = scanpy.read_h5ad(source_h5ad)
config = yaml.load(open(remix_config), Loader=yaml.FullLoader)
remix_uns(adata, config["uns"])
remix_obs(adata, config["obs"])
if config.get("fixup_gene_symbols"):
adata = fixup_gene_symbols(adata, config["fixup_gene_symbols"])
if ("version" in adata.uns_keys()
and isinstance(adata.uns["version"], collections.Mapping)
and "corpora_schema_version" in adata.uns["version"]):
schema_version = adata.uns["version"]["corpora_schema_version"]
try:
validate.get_schema_definition(schema_version)
except ValueError:
logging.warning(f"Stripping version information out of AnnData because schema "
f"version {schema_version} is unknown.")
_strip_version(adata)
if not validate.validate_adata(adata, shallow=False):
logging.warning(f"Stripping version information out of AnnData because it does not "
f"follow schema version {schema_version} .")
_strip_version(adata)
adata.write_h5ad(output_filename, compression="gzip")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--source-h5ad", required=True)
parser.add_argument("--remix-config", required=True)
parser.add_argument("--output-filename", required=True)
args = parser.parse_args()
apply_schema(args.source_h5ad, args.remix_config, args.output_filename)
@@ -0,0 +1,95 @@
title: Corpora schema version 1.0.0
type: anndata
components:
uns:
type: dict
keys:
version:
type: dict
keys:
corpora_schema_version: null
corpora_encoding_version: null
title:
type: string
contributors:
type: stringified list of dicts
layer_descriptions:
type: dict
keys:
X: null
organism:
type: string
nullable: false
organism_ontology_term_id:
type: curie
prefixes:
- NCBITaxon
var:
type: dataframe
index:
type: human-readable string
unique: true
obs:
type: dataframe
index:
unique: true
columns:
tissue:
type: human-readable string
nullable: false
tissue_ontology_term_id:
type: suffixed curie
nullable: true
prefixes:
- UBERON
assay:
type: human-readable string
nullable: false
assay_ontology_term_id:
type: curie
nullable: true
prefixes:
- EFO
disease:
type: human-readable string
nullable: false
disease_ontology_term_id:
type: curie
nullable: true
prefixes:
- MONDO
- PATO
cell_type:
type: human-readable string
nullable: false
cell_type_ontology_term_id:
type: curie
nullable: true
prefixes:
- CL
- UBERON
sex:
type: string
enum:
- male
- female
- mixed
- unknown
- other
ethnicity:
type: human-readable string
nullable: false
ethnicity_ontology_term_id:
type: curie
nullable: true
prefixes:
- HANCESTRO
development_stage:
type: human-readable string
nullable: false
development_stage_ontology_term_id:
type: curie
nullable: true
prefixes:
- HsapDv
- EFO
@@ -0,0 +1,93 @@
title: Corpora schema version 1.1.0
type: anndata
components:
uns:
type: dict
keys:
version:
type: dict
keys:
corpora_schema_version: null
corpora_encoding_version: null
title:
type: string
layer_descriptions:
type: dict
keys:
X: null
organism:
type: string
nullable: false
organism_ontology_term_id:
type: curie
prefixes:
- NCBITaxon
var:
type: dataframe
index:
type: human-readable string
unique: true
obs:
type: dataframe
index:
unique: true
columns:
tissue:
type: human-readable string
nullable: false
tissue_ontology_term_id:
type: suffixed curie
nullable: true
prefixes:
- UBERON
assay:
type: human-readable string
nullable: false
assay_ontology_term_id:
type: curie
nullable: true
prefixes:
- EFO
disease:
type: human-readable string
nullable: false
disease_ontology_term_id:
type: curie
nullable: true
prefixes:
- MONDO
- PATO
cell_type:
type: human-readable string
nullable: false
cell_type_ontology_term_id:
type: curie
nullable: true
prefixes:
- CL
- UBERON
sex:
type: string
enum:
- male
- female
- mixed
- unknown
- other
ethnicity:
type: human-readable string
nullable: false
ethnicity_ontology_term_id:
type: curie
nullable: true
prefixes:
- HANCESTRO
development_stage:
type: human-readable string
nullable: false
development_stage_ontology_term_id:
type: curie
nullable: true
prefixes:
- HsapDv
- EFO
@@ -0,0 +1,236 @@
import json
import re
import os
import sys
import pandas as pd
import yaml
def _is_null(v):
"""Return True if v is null, for one of the multiple ways a "null" value shows up in an h5ad."""
return pd.isnull(v) or (hasattr(v, "__len__") and len(v) == 0)
def _validate_stringified_list_of_dicts(s):
"""Verify that a string can be parsed into a list.
We have some types that are lists of dicts. Those cannot be stored directly in an h5ad, so we have to
json.dumps them. This verifies that we can load them back.
"""
try:
list_ = json.loads(s)
if not isinstance(list_, list):
return False
for el in list_:
if not isinstance(el, dict):
return False
return True
except (json.JSONDecodeError, TypeError):
pass
return False
def _validate_human_readable_string(s):
"""Verify that a string is human-readable.
There are parts of the schema where a "human-readable" string is required. "Human-readable" is kind
of vague and subjective. I feel like I can read many strings. So here we just check for the main ways
that fails: someone puts in an ontology term id or and ensembl gene/transcript id.
Returns False if s is not a string or is one of those bad string types.
"""
return isinstance(s, str) and (not re.match(r"[A-Z]\w+:\d+", s)) and (not re.match(r"ENS[GT]\d+$", s))
def _validate_curie(c, prefixes):
"""Verify that a string is a valid compact URI, like EFO:000001. If prefixes is not empty, make sure the
prefix of the curies is in prefixes.
"""
if not c:
return True
match = re.match(r"([A-Z]\w+):\d+$", c)
if prefixes:
return match and match.group(1) in prefixes
else:
return match
def _validate_suffixed_curie(c, prefixes):
"""Verify that a string is a compact URI with an optional suffix like 'EFO:00001 (cell culture)'"""
# Pull off the suffix
suffix = re.findall(r"\ \(.*\)$", c)
if suffix:
c = c[: -len(suffix[0])]
return _validate_curie(c, prefixes)
def _validate_column(column, column_name, df_name, schema_def):
"""Given a schema definition and the column of a dataframe, verify that the column satifies
the schema.
"""
errors = []
if schema_def.get("unique"):
if column.nunique() != len(column):
errors.append(f"Column {column_name} in dataframe {df_name} is not unique.")
if "nullable" in schema_def and not schema_def["nullable"]:
if any(_is_null(v) for v in column):
errors.append(f"Column {column_name} in dataframe {df_name} contains empty values.")
if schema_def.get("type") == "human-readable string":
non_readables = [v for v in column if not _validate_human_readable_string(v)]
if non_readables:
errors.append(
f"Column {column_name} in dataframe {df_name} contains non-human-readable "
f"values like {non_readables[0]}"
)
if schema_def.get("type") in ("curie", "suffixed curie"):
validation_func = _validate_curie if schema_def.get("type") == "curie" else _validate_suffixed_curie
non_valid_curies = [v for v in column if not validation_func(v, schema_def.get("prefixes"))]
if non_valid_curies:
errors.append(
f"Column {column_name} in dataframe {df_name} contains invalid ontology values like "
f"{non_valid_curies[0]}."
)
if "prefixes" in schema_def:
errors[-1] += f" Values must be curies from one of these ontologies {schema_def['prefixes']}."
if "enum" in schema_def:
bad_enums = [v for v in column if v not in schema_def["enum"]]
if bad_enums:
errors.append(
f"Column {column_name} in dataframe {df_name} contains unpermitted values like "
f"{bad_enums[0]}. Values must be one of {schema_def['enum']}."
)
return errors
def _validate_dict(dict_, dict_name, schema_def):
"""Given a schema definition and dict, verify that the dict satifies the schema."""
errors = []
for key in schema_def.get("keys", []):
if key not in dict_:
errors.append(f"{dict_name} is missing key {key}.")
elif schema_def["keys"][key]:
if schema_def["keys"][key]["type"] == "stringified list of dicts":
if not _validate_stringified_list_of_dicts(dict_[key]):
errors.append(
f"Key {key} in {dict_name} should be a JSON-encoded list of dicts, but it is {dict_[key]}"
)
elif schema_def["keys"][key]["type"] == "dict":
errors.extend(_validate_dict(dict_[key], key, schema_def["keys"][key]))
elif schema_def["keys"][key]["type"] == "curie":
if not _validate_curie(dict_[key], schema_def["keys"][key]["prefixes"]):
errors.append(f"Key {key} in {dict_name} contains invalid ontology value.")
if "nullable" in schema_def["keys"][key] and not schema_def["keys"][key]["nullable"]:
if _is_null(dict_[key]):
errors.append(f"Key {key} in dict {dict_name} is an empty value.")
return errors
def _validate_dataframe(df, df_name, schema_def):
"""Given a dataframe and schema definition, verify that the dataframe follows the schema."""
errors = []
if "index" in schema_def:
errors.extend(_validate_column(df.index, "index", df_name, schema_def["index"]))
for column in schema_def.get("columns", []):
if column not in df.columns:
errors.append(f"Dataframe {df_name} is missing column {column}.")
else:
errors.extend(_validate_column(df[column], column, df_name, schema_def["columns"][column]))
return errors
def get_schema_definition(version):
"""Look up and read a schema definition based on a version number like "1.0.0"."""
path = os.path.join(
os.path.dirname(os.path.realpath(__file__)), "schema_definitions", version.replace(".", "_") + ".yaml"
)
if not os.path.isfile(path):
raise ValueError(f"No definition for version {version} found.")
return yaml.load(open(path), Loader=yaml.FullLoader)
def deep_check(adata, schema_def):
"""Perform a "deep" check of the AnnData object using the schema definition.
This checks all the columns and unstructured metadata rather than just the version.
Returns a list of error messages. If that list is empty, the object passed validation.
"""
errors = []
for component, component_def in schema_def["components"].items():
if component_def["type"] == "dataframe":
errors.extend(_validate_dataframe(getattr(adata, component), component, component_def))
elif component_def["type"] == "dict":
errors.extend(_validate_dict(getattr(adata, component), component, component_def))
else:
raise ValueError(f"Unexpected component type {component['type']}")
return errors
def validate_adata(adata, shallow):
"""Validate an AnnData object. If shallow, just check that the required version information is
present.
"""
# Does it have the version information written into uns?
if "version" not in adata.uns_keys() or "corpora_schema_version" not in adata.uns["version"]:
print("AnnData file is missing corpora version information")
return False
# We can stop here if it's a "shallow" check, that is, if we're just
# checking that version is present.
if shallow:
return True
schema_def = get_schema_definition(adata.uns["version"]["corpora_schema_version"])
errors = deep_check(adata, schema_def)
for error in errors:
print(error)
return not errors
def validate(h5ad_path, shallow=False):
"""Entry point for validation."""
try:
import scanpy
except ImportError:
raise ImportError("scanpy must be installed for cellxgene schema")
try:
adata = scanpy.read_h5ad(h5ad_path, backed="r")
except (OSError, TypeError):
print(f"Unable to open {h5ad_path} with scanpy.")
sys.exit(1)
if not validate_adata(adata, shallow):
sys.exit(1)
@@ -0,0 +1,81 @@
"""
Script to create a sparse dataset in CXG format based on an input dataset in CXG format.
The input dataset is not modified.
"""
import argparse
import os
import shutil
import sys
import tiledb
from backend.czi_hosted.common.utils.cxg_generation_utils import convert_ndarray_to_cxg_dense_array, \
convert_matrix_to_cxg_array
from backend.czi_hosted.common.utils.matrix_utils import is_matrix_sparse, get_column_shift_encode_for_matrix
def main():
parser = argparse.ArgumentParser()
parser.add_argument("input", help="input cxg directory")
parser.add_argument("output", help="output cxg directory")
parser.add_argument("--overwrite", action="store_true", help="replace output cxg directory")
parser.add_argument("--verbose", "-v", action="count", default=0, help="verbose output")
parser.add_argument(
"--sparse-threshold",
"-s",
type=float,
default=5.0, # default is 5% non-zero values
help="The X array will be sparse if the percent of non-zeros falls below this value",
)
args = parser.parse_args()
if os.path.exists(args.output):
print("output dir exists:", args.output)
if args.overwrite:
print("output dir removed:", args.output)
shutil.rmtree(args.output)
else:
print("use the overwrite option to remove the output directory")
sys.exit(1)
if not os.path.isdir(args.input):
print("input is not a directory", args.input)
sys.exit(1)
shutil.copytree(args.input, args.output, ignore=shutil.ignore_patterns("X", "X_col_shift"))
ctx = tiledb.Ctx(
{
"sm.num_reader_threads": 32,
"sm.num_writer_threads": 32,
"sm.consolidation.buffer_size": 1 * 1024 * 1024 * 1024,
}
)
with tiledb.DenseArray(os.path.join(args.input, "X"), mode="r", ctx=ctx) as X_in:
x_matrix_data = X_in[:, :]
matrix_container = args.output
is_sparse = is_matrix_sparse(x_matrix_data, args.sparse_threshold)
if not is_sparse:
col_shift = get_column_shift_encode_for_matrix(x_matrix_data, args.sparse_threshold)
is_sparse = col_shift is not None
else:
col_shift = None
if col_shift is not None:
x_col_shift_name = f"{args.output}/X_col_shift"
convert_ndarray_to_cxg_dense_array(x_col_shift_name, col_shift, ctx)
tiledb.consolidate(matrix_container, ctx=ctx)
if is_sparse:
convert_matrix_to_cxg_array(matrix_container, x_matrix_data, is_sparse, ctx, col_shift)
tiledb.consolidate(matrix_container, ctx=ctx)
if not is_sparse:
print("The array is not sparse, cleaning up, abort.")
shutil.rmtree(args.output)
sys.exit(1)
if __name__ == "__main__":
main()
@@ -0,0 +1,362 @@
import warnings
import anndata
import numpy as np
from packaging import version
from pandas.core.dtypes.dtypes import CategoricalDtype
from scipy import sparse
import backend.common.compute.diffexp_generic as diffexp_generic
from backend.common.colors import convert_anndata_category_colors_to_cxg_category_colors
from backend.common.constants import Axis, MAX_LAYOUTS, XApproximateDistribution
from backend.czi_hosted.common.corpora import corpora_get_props_from_anndata
from backend.common.errors import PrepareError, DatasetAccessError, ConfigurationError
from backend.common.utils.type_conversion_utils import get_schema_type_hint_of_array
from backend.czi_hosted.data_common.data_adaptor import DataAdaptor
from backend.common.fbs.matrix import encode_matrix_fbs
anndata_version = version.parse(str(anndata.__version__)).release
def anndata_version_is_pre_070():
major = anndata_version[0]
minor = anndata_version[1] if len(anndata_version) > 1 else 0
return major == 0 and minor < 7
class AnndataAdaptor(DataAdaptor):
def __init__(self, data_locator, app_config=None, dataset_config=None):
super().__init__(data_locator, app_config, dataset_config)
self.data = None
self.X_approximate_distribution = None
self._load_data(data_locator)
self._validate_and_initialize()
def cleanup(self):
pass
@staticmethod
def pre_load_validation(data_locator):
if data_locator.islocal():
# if data locator is local, apply file system conventions and other "cheap"
# validation checks. If a URI, defer until we actually fetch the data and
# try to read it. Many of these tests don't make sense for URIs (eg, extension-
# based typing).
if not data_locator.exists():
raise DatasetAccessError("does not exist")
if not data_locator.isfile():
raise DatasetAccessError("is not a file")
@staticmethod
def file_size(data_locator):
return data_locator.size() if data_locator.islocal() else 0
@staticmethod
def open(data_locator, app_config, dataset_config=None):
return AnndataAdaptor(data_locator, app_config, dataset_config)
def get_corpora_props(self):
return corpora_get_props_from_anndata(self.data)
def get_name(self):
return "cellxgene anndata adaptor version"
def get_library_versions(self):
return dict(anndata=str(anndata.__version__))
@staticmethod
def _create_unique_column_name(df, col_name_prefix):
"""given the columns of a dataframe, and a name prefix, return a column name which
does not exist in the dataframe, AND which is prefixed by `prefix`
The approach is to append a numeric suffix, starting at zero and increasing by
one, until an unused name is found (eg, prefix_0, prefix_1, ...).
"""
suffix = 0
while f"{col_name_prefix}{suffix}" in df:
suffix += 1
return f"{col_name_prefix}{suffix}"
def _alias_annotation_names(self):
"""
The front-end relies on the existance of a unique, human-readable
index for obs & var (eg, var is typically gene name, obs the cell name).
The user can specify these via the --obs-names and --var-names config.
If they are not specified, use the existing index to create them, giving
the resulting column a unique name (eg, "name").
In both cases, enforce that the result is unique, and communicate the
index column name to the front-end via the obs_names and var_names config
(which is incorporated into the schema).
"""
self.original_obs_index = self.data.obs.index
for (ax_name, var_name) in ((Axis.OBS, "obs"), (Axis.VAR, "var")):
config_name = f"single_dataset__{var_name}_names"
parameter_name = f"{var_name}_names"
name = getattr(self.server_config, config_name)
df_axis = getattr(self.data, str(ax_name))
if name is None:
# Default: create unique names from index
if not df_axis.index.is_unique:
raise KeyError(
f"Values in {ax_name}.index must be unique. "
"Please prepare data to contain unique index values, or specify an "
"alternative with --{ax_name}-name."
)
name = self._create_unique_column_name(df_axis.columns, "name_")
self.parameters[parameter_name] = name
# reset index to simple range; alias name to point at the
# previously specified index.
df_axis.rename_axis(name, inplace=True)
df_axis.reset_index(inplace=True)
elif name in df_axis.columns:
# User has specified alternative column for unique names, and it exists
if not df_axis[name].is_unique:
raise KeyError(
f"Values in {ax_name}.{name} must be unique. " "Please prepare data to contain unique values."
)
df_axis.reset_index(drop=True, inplace=True)
self.parameters[parameter_name] = name
else:
# user specified a non-existent column name
raise KeyError(f"Annotation name {name}, specified in --{ax_name}-name does not exist.")
def _create_schema(self):
self.schema = {
"dataframe": {
"nObs": self.cell_count,
"nVar": self.gene_count,
**get_schema_type_hint_of_array(self.data.X),
},
"annotations": {
"obs": {"index": self.parameters.get("obs_names"), "columns": []},
"var": {"index": self.parameters.get("var_names"), "columns": []},
},
"layout": {"obs": []},
}
for ax in Axis:
curr_axis = getattr(self.data, str(ax))
for ann in curr_axis:
ann_schema = {"name": ann, "writable": False}
ann_schema.update(get_schema_type_hint_of_array(curr_axis[ann]))
self.schema["annotations"][ax]["columns"].append(ann_schema)
for layout in self.get_embedding_names():
layout_schema = {"name": layout, "type": "float32", "dims": [f"{layout}_0", f"{layout}_1"]}
self.schema["layout"]["obs"].append(layout_schema)
def get_schema(self):
return self.schema
def _load_data(self, data_locator):
# as of AnnData 0.6.19, backed mode performs initial load fast, but at the
# cost of significantly slower access to X data.
try:
# there is no guarantee data_locator indicates a local file. The AnnData
# API will only consume local file objects. If we get a non-local object,
# make a copy in tmp, and delete it after we load into memory.
with data_locator.local_handle() as lh:
# as of AnnData 0.6.19, backed mode performs initial load fast, but at the
# cost of significantly slower access to X data.
backed = "r" if self.server_config.adaptor__anndata_adaptor__backed else None
self.data = anndata.read_h5ad(lh, backed=backed)
except ValueError:
raise DatasetAccessError(
"File must be in the .h5ad format. Please read "
"https://github.com/theislab/scanpy_usage/blob/master/170505_seurat/info_h5ad.md to "
"learn more about this format. You may be able to convert your file into this format "
"using `cellxgene prepare`, please run `cellxgene prepare --help` for more "
"information."
)
except MemoryError:
raise DatasetAccessError("Out of memory - file is too large for available memory.")
except Exception:
raise DatasetAccessError(
"File not found or is inaccessible. File must be an .h5ad object. "
"Please check your input and try again."
)
def _validate_and_initialize(self):
if anndata_version_is_pre_070():
warnings.warn(
"Use of anndata versions older than 0.7 will have serious issues. Please update to at "
"least anndata 0.7 or later."
)
# var and obs column names must be unique
if not self.data.obs.columns.is_unique or not self.data.var.columns.is_unique:
raise KeyError("All annotation column names must be unique.")
self._alias_annotation_names()
self._validate_data_types()
self.cell_count = self.data.shape[0]
self.gene_count = self.data.shape[1]
self._create_schema()
if self.dataset_config.X_approximate_distribution == "auto":
raise ConfigurationError("X-approximate-distribution 'auto' mode unsupported.")
self.X_approximate_distribution = self.dataset_config.X_approximate_distribution
# heuristic
n_values = self.data.shape[0] * self.data.shape[1]
if (n_values > 1e8 and self.server_config.adaptor__anndata_adaptor__backed is True) or (n_values > 5e8):
self.parameters.update({"diffexp_may_be_slow": True})
def _is_valid_layout(self, arr):
"""return True if this layout data is a valid array for front-end presentation:
* ndarray, dtype float/int/uint
* with shape (n_obs, >= 2)
* with all values finite or NaN (no +Inf or -Inf)
"""
is_valid = type(arr) == np.ndarray and arr.dtype.kind in "fiu"
is_valid = is_valid and arr.shape[0] == self.data.n_obs and arr.shape[1] >= 2
is_valid = is_valid and not np.any(np.isinf(arr)) and not np.all(np.isnan(arr))
return is_valid
def _validate_data_types(self):
# The backed API does not support interrogation of the underlying sparsity or sparse matrix type
# Fake it by asking for a small subarray and testing it. NOTE: if the user has ignored our
# anndata <= 0.7 warning, opted for the --backed option, and specified a large, sparse dataset,
# this "small" indexing request will load the entire X array. This is due to a bug in anndata<=0.7
# which will load the entire X matrix to fullfill any slicing request if X is sparse. See
# user warning in _load_data().
X0 = self.data.X[0, 0:1]
if sparse.isspmatrix(X0) and not sparse.isspmatrix_csc(X0):
warnings.warn(
"Anndata data matrix is sparse, but not a CSC (columnar) matrix. "
"Performance may be improved by using CSC."
)
if self.data.X.dtype != "float32":
warnings.warn(
f"Anndata data matrix is in {self.data.X.dtype} format not float32. " f"Precision may be truncated."
)
for ax in Axis:
curr_axis = getattr(self.data, str(ax))
for ann in curr_axis:
datatype = curr_axis[ann].dtype
downcast_map = {
"int64": "int32",
"uint32": "int32",
"uint64": "int32",
"float64": "float32",
}
if datatype in downcast_map:
warnings.warn(
f"Anndata annotation {ax}:{ann} is in unsupported format: {datatype}. "
f"Data will be downcast to {downcast_map[datatype]}."
)
if isinstance(datatype, CategoricalDtype):
category_num = len(curr_axis[ann].dtype.categories)
if category_num > 500 and category_num > self.dataset_config.presentation__max_categories:
warnings.warn(
f"{str(ax).title()} annotation '{ann}' has {category_num} categories, this may be "
f"cumbersome or slow to display. We recommend setting the "
f"--max-category-items option to 500, this will hide categorical "
f"annotations with more than 500 categories in the UI"
)
def annotation_to_fbs_matrix(self, axis, fields=None, labels=None):
if axis == Axis.OBS:
if labels is not None and not labels.empty:
df = self.data.obs.join(labels, self.parameters.get("obs_names"))
else:
df = self.data.obs
else:
df = self.data.var
if fields is not None and len(fields) > 0:
df = df[fields]
return encode_matrix_fbs(df, col_idx=df.columns)
def get_embedding_names(self):
"""
Return pre-computed embeddings.
function:
a) generate list of default layouts
b) validate layouts are legal. remove/warn on any that are not
c) cap total list of layouts at global const MAX_LAYOUTS
"""
# load default layouts from the data.
layouts = self.dataset_config.embeddings__names
if layouts is None or len(layouts) == 0:
layouts = [key[2:] for key in self.data.obsm_keys() if type(key) == str and key.startswith("X_")]
# remove invalid layouts
valid_layouts = []
obsm_keys = self.data.obsm_keys()
for layout in layouts:
layout_name = f"X_{layout}"
if layout_name not in obsm_keys:
warnings.warn(f"Ignoring unknown layout name: {layout}.")
elif not self._is_valid_layout(self.data.obsm[layout_name]):
warnings.warn(f"Ignoring layout due to malformed shape or data type: {layout}")
else:
valid_layouts.append(layout)
if len(valid_layouts) == 0:
raise PrepareError("No valid layout data.")
# cap layouts to MAX_LAYOUTS
return valid_layouts[0:MAX_LAYOUTS]
def get_embedding_array(self, ename, dims=2):
full_embedding = self.data.obsm[f"X_{ename}"]
return full_embedding[:, 0:dims]
def compute_diffexp_ttest(self, maskA, maskB, top_n=None, lfc_cutoff=None):
if top_n is None:
top_n = self.dataset_config.diffexp__top_n
if lfc_cutoff is None:
lfc_cutoff = self.dataset_config.diffexp__lfc_cutoff
return diffexp_generic.diffexp_ttest(self, maskA, maskB, top_n, lfc_cutoff)
def get_colors(self):
return convert_anndata_category_colors_to_cxg_category_colors(self.data)
def get_X_array(self, obs_mask=None, var_mask=None):
# H5Py does not support boolean indexing (masks), so convert to integer indexing
# when backed (ie, when AnnData is using H5Py indexing)
if obs_mask is None:
obs_mask = slice(None)
elif self.data.isbacked and obs_mask.dtype == bool:
obs_mask = obs_mask.nonzero()[0]
if var_mask is None:
var_mask = slice(None)
elif self.data.isbacked and var_mask.dtype == bool:
var_mask = var_mask.nonzero()[0]
X = self.data.X[obs_mask, var_mask]
return X
def get_X_approximate_distribution(self) -> XApproximateDistribution:
return self.X_approximate_distribution
def get_shape(self):
return self.data.shape
def query_var_array(self, term_name):
return getattr(self.data.var, term_name)
def query_obs_array(self, term_name):
return getattr(self.data.obs, term_name)
def get_obs_index(self):
name = self.server_config.single_dataset__obs_names
if name is None:
return self.original_obs_index
else:
return self.data.obs[name]
def get_obs_columns(self):
return self.data.obs.columns
def get_obs_keys(self):
# return list of keys
return self.data.obs.keys().to_list()
def get_var_keys(self):
# return list of keys
return self.data.var.keys().to_list()
@@ -0,0 +1,428 @@
from abc import ABCMeta, abstractmethod
from os.path import basename, splitext
import numpy as np
import pandas as pd
from scipy import sparse
from server_timing import Timing as ServerTiming
from backend.czi_hosted.common.config.app_config import AppConfig
from backend.common.constants import Axis, XApproximateDistribution
from backend.common.errors import (
FilterError,
JSONEncodingValueError,
ExceedsLimitError,
UnsupportedSummaryMethod,
DatasetAccessError,
)
from backend.common.utils.utils import jsonify_numpy
from backend.common.fbs.matrix import encode_matrix_fbs
class DataAdaptor(metaclass=ABCMeta):
"""Base class for loading and accessing matrix data"""
def __init__(self, data_locator, app_config, dataset_config=None):
if type(app_config) != AppConfig:
raise TypeError("config expected to be of type AppConfig")
# location to the dataset
self.data_locator = data_locator
# config is the application configuration
self.app_config = app_config
self.server_config = self.app_config.server_config
self.dataset_config = dataset_config or app_config.default_dataset_config
# parameters set by this data adaptor based on the data.
self.parameters = {}
self.uri_path = None
def set_uri_path(self, path):
# uri path to the dataset, e.g. /d/<datasetname>
self.uri_path = path
@staticmethod
@abstractmethod
def pre_load_validation(data_locator):
pass
@staticmethod
@abstractmethod
def open(data_locator, app_config, dataset_config):
pass
@staticmethod
@abstractmethod
def file_size(data_locator):
pass
@abstractmethod
def get_name(self):
"""return a string name for this data adaptor"""
pass
@abstractmethod
def get_library_versions(self):
"""return a dictionary of library name to library versions"""
pass
@abstractmethod
def get_embedding_names(self):
"""return a list of pre-computed embedding names"""
pass
@abstractmethod
def get_embedding_array(self, ename, dims=2):
"""return an numpy array for the given pre-computed embedding name."""
pass
@abstractmethod
def get_X_array(self, obs_mask=None, var_mask=None):
"""return the X array, possibly filtered by obs_mask or var_mask.
the return type is either ndarray or scipy.sparse.spmatrix."""
pass
@abstractmethod
def get_X_approximate_distribution(self) -> XApproximateDistribution:
"""return the approximate distribution of the X matrix."""
pass
@abstractmethod
def get_shape(self):
pass
@abstractmethod
def query_var_array(self, term_var):
pass
@abstractmethod
def query_obs_array(self, term_var):
pass
@abstractmethod
def get_colors(self):
pass
@abstractmethod
def get_obs_index(self):
pass
@abstractmethod
def get_obs_columns(self):
pass
@abstractmethod
def get_obs_keys(self):
# return list of keys
pass
@abstractmethod
def get_var_keys(self):
# return list of keys
pass
@abstractmethod
def cleanup(self):
pass
def get_data_locator(self):
return self.data_locator
def get_location(self):
return self.data_locator.uri_or_path
def get_about(self):
return None
def get_title(self):
# default to file name
location = self.get_location()
if location.endswith("/"):
location = location[:-1]
return splitext(basename(location))[0]
def get_corpora_props(self):
return None
@abstractmethod
def get_schema(self):
"""
Return current schema
"""
pass
@abstractmethod
def annotation_to_fbs_matrix(self, axis, field=None, uid=None):
"""
Gets annotation value for each observation
:param axis: string obs or var
:param fields: list of keys for annotation to return, returns all annotation values if not set.
:return: flatbuffer: in fbs/matrix.fbs encoding
"""
pass
def update_parameters(self, parameters):
parameters.update(self.parameters)
def _index_filter_to_mask(self, filter, count):
mask = np.zeros((count,), dtype=np.bool)
for i in filter:
if type(i) == list:
mask[i[0] : i[1]] = True
else:
mask[i] = True
return mask
def _axis_filter_to_mask(self, axis, filter, count):
mask = np.ones((count,), dtype=np.bool)
if "index" in filter:
mask = np.logical_and(mask, self._index_filter_to_mask(filter["index"], count))
if "annotation_value" in filter:
mask = np.logical_and(mask, self._annotation_filter_to_mask(axis, filter["annotation_value"], count))
return mask
def _annotation_filter_to_mask(self, axis, filter, count):
mask = np.ones((count,), dtype=np.bool)
for v in filter:
name = v["name"]
if axis == Axis.VAR:
anno_data = self.query_var_array(name)
elif axis == Axis.OBS:
anno_data = self.query_obs_array(name)
if anno_data.dtype.name in ["boolean", "category", "object"]:
values = v.get("values", [])
key_idx = np.in1d(anno_data, values)
mask = np.logical_and(mask, key_idx)
else:
min_ = v.get("min", None)
max_ = v.get("max", None)
if min_ is not None:
key_idx = (anno_data >= min_).ravel()
mask = np.logical_and(mask, key_idx)
if max_ is not None:
key_idx = (anno_data <= max_).ravel()
mask = np.logical_and(mask, key_idx)
return mask
def _filter_to_mask(self, filter):
"""
Return the filter as a row and column selection list.
No filter on a dimension means 'all'
"""
shape = self.get_shape()
var_selector = None
obs_selector = None
if filter is not None:
if Axis.OBS in filter:
obs_selector = self._axis_filter_to_mask(Axis.OBS, filter["obs"], shape[0])
if Axis.VAR in filter:
var_selector = self._axis_filter_to_mask(Axis.VAR, filter["var"], shape[1])
return (obs_selector, var_selector)
def check_new_labels(self, labels_df):
"""Check the new annotations labels, then set the labels_df index"""
if labels_df is None or labels_df.empty:
return
labels_df.index = self.get_obs_index()
if labels_df.index.name is None:
labels_df.index.name = "index"
# all labels must have a name, which must be unique and not used in obs column names
if not labels_df.columns.is_unique:
raise KeyError("All column names specified in user annotations must be unique.")
# the label index must be unique, and must have same values the anndata obs index
if not labels_df.index.is_unique:
raise KeyError("All row index values specified in user annotations must be unique.")
obs_columns = self.get_obs_columns()
duplicate_columns = list(set(labels_df.columns) & set(obs_columns))
if len(duplicate_columns) > 0:
raise KeyError(
"Labels file may not contain column names which overlap " f"with h5ad obs columns {duplicate_columns}"
)
# labels must have same count as obs annotations
shape = self.get_shape()
if labels_df.shape[0] != shape[0]:
raise ValueError("Labels file must have same number of rows as data file.")
# This will convert a float column that contains integer data into an integer type.
# This case can occur when a user makes a copy of a category that originally contained integer data.
# The client always copies array data to floats, therefore the copy will contain floats instead of integers.
# float data is not allowed as a categorical type.
if any([np.issubdtype(coltype.type, np.floating) for coltype in labels_df.dtypes]):
labels_df = labels_df.convert_dtypes()
for col, dtype in zip(labels_df, labels_df.dtypes):
if isinstance(dtype, pd.Int32Dtype):
labels_df[col] = labels_df[col].astype("int32")
if isinstance(dtype, pd.Int64Dtype):
labels_df[col] = labels_df[col].astype("int64")
if any([np.issubdtype(coltype.type, np.floating) for coltype in labels_df.dtypes]):
raise ValueError("Columns may not have floating point types")
return labels_df
def data_frame_to_fbs_matrix(self, filter, axis):
"""
Retrieves data 'X' and returns in a flatbuffer Matrix.
:param filter: filter: dictionary with filter params
:param axis: string obs or var
:return: flatbuffer Matrix
Caveats:
* currently only supports access on VAR axis
* currently only supports filtering on VAR axis
"""
if axis != Axis.VAR:
raise ValueError("Only VAR dimension access is supported")
try:
obs_selector, var_selector = self._filter_to_mask(filter)
except (KeyError, IndexError, TypeError, AttributeError, DatasetAccessError):
raise FilterError("Error parsing filter")
if obs_selector is not None:
raise FilterError("filtering on obs unsupported")
num_columns = self.get_shape()[1] if var_selector is None else np.count_nonzero(var_selector)
if self.server_config.exceeds_limit("column_request_max", num_columns):
raise ExceedsLimitError("Requested dataframe columns exceed column request limit")
X = self.get_X_array(obs_selector, var_selector)
col_idx = np.nonzero([] if var_selector is None else var_selector)[0]
return encode_matrix_fbs(X, col_idx=col_idx, row_idx=None)
def diffexp_topN(self, obsFilterA, obsFilterB, top_n=None):
"""
Computes the top N differentially expressed variables between two observation sets. If mode
is "TOP_N", then stats for the top N
dataframes
contain a subset of variables, then statistics for all variables will be returned, otherwise
only the top N vars will be returned.
:param obsFilterA: filter: dictionary with filter params for first set of observations
:param obsFilterB: filter: dictionary with filter params for second set of observations
:param top_n: Limit results to top N (Top var mode only)
:return: top N genes and corresponding stats
"""
if Axis.VAR in obsFilterA or Axis.VAR in obsFilterB:
raise FilterError("Observation filters may not contain variable conditions")
try:
shape = self.get_shape()
obs_mask_A = self._axis_filter_to_mask(Axis.OBS, obsFilterA["obs"], shape[0])
obs_mask_B = self._axis_filter_to_mask(Axis.OBS, obsFilterB["obs"], shape[0])
except (KeyError, IndexError):
raise FilterError("Error parsing filter")
if top_n is None:
top_n = self.dataset_config.diffexp__top_n
if self.server_config.exceeds_limit(
"diffexp_cellcount_max", np.count_nonzero(obs_mask_A) + np.count_nonzero(obs_mask_B)
):
raise ExceedsLimitError("Diffexp request exceeds max cell count limit")
result = self.compute_diffexp_ttest(
maskA=obs_mask_A, maskB=obs_mask_B, top_n=top_n, lfc_cutoff=self.dataset_config.diffexp__lfc_cutoff
)
try:
return jsonify_numpy(result)
except ValueError:
raise JSONEncodingValueError("Error encoding differential expression to JSON")
@abstractmethod
def compute_diffexp_ttest(self, maskA, maskB, top_n, lfc_cutoff):
pass
@staticmethod
def normalize_embedding(embedding):
"""Normalize embedding layout to meet client assumptions.
Embedding is an ndarray, shape (n_obs, n)., where n is normally 2
"""
# scale isotropically
try:
min = np.nanmin(embedding, axis=0)
max = np.nanmax(embedding, axis=0)
except RuntimeError:
# indicates entire array was NaN, which should propagate
min = np.NaN
max = np.NaN
scale = np.amax(max - min)
normalized_layout = (embedding - min) / scale
# translate to center on both axis
translate = 0.5 - ((max - min) / scale / 2)
normalized_layout = normalized_layout + translate
normalized_layout = normalized_layout.astype(dtype=np.float32)
return normalized_layout
def layout_to_fbs_matrix(self, fields):
"""
return specified embeddings as a flatbuffer, using the cellxgene matrix fbs encoding.
* returns only first two dimensions, with name {ename}_0 and {ename}_1,
where {ename} is the embedding name.
* client assumes each will be individually centered & scaled (isotropically)
to a [0, 1] range.
* does not support filtering
"""
embeddings = self.get_embedding_names() if fields is None or len(fields) == 0 else fields
layout_data = []
with ServerTiming.time("layout.query"):
for ename in embeddings:
embedding = self.get_embedding_array(ename, 2)
normalized_layout = DataAdaptor.normalize_embedding(embedding)
layout_data.append(pd.DataFrame(normalized_layout, columns=[f"{ename}_0", f"{ename}_1"]))
with ServerTiming.time("layout.encode"):
if layout_data:
df = pd.concat(layout_data, axis=1, copy=False)
else:
df = pd.DataFrame()
fbs = encode_matrix_fbs(df, col_idx=df.columns, row_idx=None)
return fbs
def get_last_mod_time(self):
try:
lastmod = self.get_data_locator().lastmodtime()
except RuntimeError:
lastmod = None
return lastmod
def summarize_var(self, method, filter, query_hash):
if method != "mean":
raise UnsupportedSummaryMethod("Unknown gene set summary method.")
obs_selector, var_selector = self._filter_to_mask(filter)
if obs_selector is not None:
raise FilterError("filtering on obs unsupported")
# if no filter, just return zeros. We don't have a use case
# for summarizing the entire X without a filter, and it would
# potentially be quite compute / memory intensive.
if var_selector is None or np.count_nonzero(var_selector) == 0:
mean = np.zeros((self.get_shape()[0], 1), dtype=np.float32)
else:
X = self.get_X_array(obs_selector, var_selector)
if sparse.issparse(X):
mean = X.mean(axis=1).A
else:
mean = X.mean(axis=1, keepdims=True)
col_idx = pd.Index([query_hash])
return encode_matrix_fbs(mean, col_idx=col_idx, row_idx=None)
@@ -0,0 +1,288 @@
from enum import Enum
import threading
import time
from backend.common.utils.data_locator import DataLocator
from backend.common.errors import DatasetAccessError
from contextlib import contextmanager
from http import HTTPStatus
from backend.czi_hosted.data_common.rwlock import RWLock
class MatrixDataCacheItem(object):
"""This class provides access and caching for a dataset. The first time a dataset is accessed, it is
opened and cached. Later accesses use the cached version. It may also be deleted by the
MatrixDataCacheManager to make room for another dataset. While a dataset is actively being used
(during the lifetime of a api request), a reader lock is locked. During that time, the dataset cannot
be removed."""
def __init__(self, loader):
self.loader = loader
self.data_adaptor = None
self.data_lock = RWLock()
def acquire_existing(self):
"""If the data_adaptor exists, take a read lock and return it, else return None"""
self.data_lock.r_acquire()
if self.data_adaptor:
return self.data_adaptor
self.data_lock.r_release()
return None
def acquire_and_open(self, app_config, dataset_config=None):
"""returns the data_adaptor if cached. opens the data_adaptor if not.
In either case, the a reader lock is taken. Must call release when
the data_adaptor is no longer needed"""
self.data_lock.r_acquire()
if self.data_adaptor:
return self.data_adaptor
self.data_lock.r_release()
self.data_lock.w_acquire()
# the data may have been loaded while waiting on the lock
if not self.data_adaptor:
try:
self.loader.pre_load_validation()
self.data_adaptor = self.loader.open(app_config, dataset_config)
except Exception as e:
# necessary to hold the reader lock after an exception, since
# the release will occur when the context exits.
self.data_lock.w_demote()
raise DatasetAccessError(str(e))
# demote the write lock to a read lock.
self.data_lock.w_demote()
return self.data_adaptor
def release(self):
"""Release the reader lock"""
self.data_lock.r_release()
def delete(self):
"""Clear resources used by this dataset"""
with self.data_lock.w_locked():
if self.data_adaptor:
self.data_adaptor.cleanup()
self.data_adaptor = None
def attempt_delete(self):
"""Delete, but only if the write lock can be immediately locked. Return True if the delete happened"""
if self.data_lock.w_acquire_non_blocking():
if self.data_adaptor:
try:
self.data_adaptor.cleanup()
self.data_adaptor = None
except Exception:
# catch all exceptions to ensure the lock is released
pass
self.data_lock.w_release()
return True
else:
return False
class MatrixDataCacheInfo(object):
def __init__(self, cache_item, timestamp):
# The MatrixDataCacheItem in the cache
self.cache_item = cache_item
# The last time the cache_item was accessed
self.last_access = timestamp
# The number of times the cache_item was accessed (used for testing)
self.num_access = 1
class MatrixDataCacheManager(object):
"""A class to manage the cached datasets. This is intended to be used as a context manager
for handling api requests. When the context is created, the data_adator is either loaded or
retrieved from a cache. In either case, the reader lock is taken during this time, and release
when the context ends. This class currently implements a simple least recently used cache,
which can delete a dataset from the cache to make room for a new one.
This is the intended usage pattern:
m = MatrixDataCacheManager(max_cached=..., timelimmit_s = ...)
with m.data_adaptor(location, app_config) as data_adaptor:
# use the data_adaptor for some operation
"""
# FIXME: If the number of active datasets exceeds the max_cached, then each request could
# lead to a dataset being deleted and a new only being opened: the cache will get thrashed.
# In this case, we may need to send back a 503 (Server Unavailable), or some other error message.
# NOTE: If the actual dataset is changed. E.g. a new set of datafiles replaces an existing set,
# then the cache will not react to this, however once the cache time limit is reached, the dataset
# will automatically be refreshed.
def __init__(self, max_cached, timelimit_s=None):
# key is tuple(url_dataroot, location), value is a MatrixDataCacheInfo
self.datasets = {}
# lock to protect the datasets
self.lock = threading.Lock()
# The number of datasets to cache. When max_cached is reached, the least recently used
# cache is replaced with the newly requested one.
# TODO: This is very simple. This can be improved by taking into account how much space is actually
# taken by each dataset, instead of arbitrarily picking a max datasets to cache.
self.max_cached = max_cached
# items are automatically removed from the cache once this time limit is reached
self.timelimit_s = timelimit_s
@contextmanager
def data_adaptor(self, url_dataroot, location, app_config):
# create a loader for to this location if it does not already exist
delete_adaptor = None
data_adaptor = None
cache_item = None
key = (url_dataroot, location)
with self.lock:
self.evict_old_datasets()
info = self.datasets.get(key)
if info is not None:
info.last_access = time.time()
info.num_access += 1
self.datasets[key] = info
data_adaptor = info.cache_item.acquire_existing()
cache_item = info.cache_item
if data_adaptor is None:
while True:
if len(self.datasets) < self.max_cached:
break
items = list(self.datasets.items())
items = sorted(items, key=lambda x: x[1].last_access)
# close the least recently used loader
oldest = items[0]
oldest_cache = oldest[1].cache_item
oldest_key = oldest[0]
del self.datasets[oldest_key]
delete_adaptor = oldest_cache
loader = MatrixDataLoader(location, app_config=app_config)
cache_item = MatrixDataCacheItem(loader)
item = MatrixDataCacheInfo(cache_item, time.time())
self.datasets[key] = item
try:
assert cache_item
if delete_adaptor:
delete_adaptor.delete()
if data_adaptor is None:
dataset_config = app_config.get_dataset_config(url_dataroot)
data_adaptor = cache_item.acquire_and_open(app_config, dataset_config)
yield data_adaptor
except DatasetAccessError:
cache_item.release()
with self.lock:
del self.datasets[key]
cache_item.delete()
cache_item = None
raise
finally:
if cache_item:
cache_item.release()
def evict_old_datasets(self):
# must be called with the lock held
if self.timelimit_s is None:
return
now = time.time()
to_del = []
for key, info in self.datasets.items():
if (now - info.last_access) > self.timelimit_s:
# remove the data_cache when if it has been in the cache too long
to_del.append((key, info))
for key, info in to_del:
# try and get the write_lock for the dataset.
# if this returns false, it means the dataset is being used, and should
# not be removed.
if info.cache_item.attempt_delete():
del self.datasets[key]
class MatrixDataType(Enum):
H5AD = "h5ad"
CXG = "cxg"
UNKNOWN = "unknown"
class MatrixDataLoader(object):
def __init__(self, location, matrix_data_type=None, app_config=None):
""" location can be a string or DataLocator """
region_name = None if app_config is None else app_config.server_config.data_locator__s3__region_name
self.location = DataLocator(location, region_name=region_name)
if not self.location.exists():
raise DatasetAccessError("Dataset does not exist.", HTTPStatus.NOT_FOUND)
# matrix_data_type is an enum value of type MatrixDataType
self.matrix_data_type = matrix_data_type
# matrix_type is a DataAdaptor type, which corresponds to the matrix_data_type
self.matrix_type = None
if matrix_data_type is None:
self.matrix_data_type = self.__matrix_data_type()
if not self.__matrix_data_type_allowed(app_config):
raise DatasetAccessError("Dataset does not have an allowed type.")
if self.matrix_data_type == MatrixDataType.H5AD:
from backend.czi_hosted.data_anndata.anndata_adaptor import AnndataAdaptor
self.matrix_type = AnndataAdaptor
elif self.matrix_data_type == MatrixDataType.CXG:
from backend.czi_hosted.data_cxg.cxg_adaptor import CxgAdaptor
self.matrix_type = CxgAdaptor
def __matrix_data_type(self):
if self.location.path.endswith(".h5ad"):
return MatrixDataType.H5AD
elif ".cxg" in self.location.path:
return MatrixDataType.CXG
else:
return MatrixDataType.UNKNOWN
def __matrix_data_type_allowed(self, app_config):
if self.matrix_data_type == MatrixDataType.UNKNOWN:
return False
if not app_config:
return True
if not app_config.is_multi_dataset():
return True
if len(app_config.server_config.multi_dataset__allowed_matrix_types) == 0:
return True
for val in app_config.server_config.multi_dataset__allowed_matrix_types:
try:
if self.matrix_data_type == MatrixDataType(val):
return True
except ValueError:
# Check case where multi_dataset_allowed_matrix_type does not have a
# valid MatrixDataType value. TODO: Add a feature to check
# the AppConfig for errors on startup
return False
return False
def pre_load_validation(self):
if self.matrix_data_type == MatrixDataType.UNKNOWN:
raise DatasetAccessError("Dataset does not have a recognized type: .h5ad or .cxg")
self.matrix_type.pre_load_validation(self.location)
def file_size(self):
return self.matrix_type.file_size(self.location)
def open(self, app_config, dataset_config=None):
# create and return a DataAdaptor object
return self.matrix_type.open(self.location, app_config, dataset_config)
+135
View File
@@ -0,0 +1,135 @@
# -*- coding: utf-8 -*-
""" rwlock.py
A class to implement read-write locks on top of the standard threading
library.
This is implemented with two mutexes (threading.Lock instances) as per this
wikipedia pseudocode:
https://en.wikipedia.org/wiki/Readers%E2%80%93writer_lock#Using_two_mutexes
Code written by Tyler Neylon at Unbox Research.
This file is public domain.
Modified to add a w_demote function to convert a writer lock to a reader lock
"""
# _______________________________________________________________________
# Imports
from contextlib import contextmanager
from threading import Lock
# _______________________________________________________________________
# Class
class RWLock(object):
""" RWLock class; this is meant to allow an object to be read from by
multiple threads, but only written to by a single thread at a time. See:
https://en.wikipedia.org/wiki/Readers%E2%80%93writer_lock
Usage:
from rwlock import RWLock
my_obj_rwlock = RWLock()
# When reading from my_obj:
with my_obj_rwlock.r_locked():
do_read_only_things_with(my_obj)
# When writing to my_obj:
with my_obj_rwlock.w_locked():
mutate(my_obj)
"""
def __init__(self):
self.w_lock = Lock()
self.num_r_lock = Lock()
self.num_r = 0
# The d_lock is needed to handle the demotion case,
# so that the writer can become a reader without releasing the w_lock.
# the d_lock is held by the writer, and prevents any other thread from taking the
# num_r_lock during that time, which means the writer thread is able to take the
# num_r_lock to update the num_r.
self.d_lock = Lock()
# ___________________________________________________________________
# Reading methods.
def r_acquire(self):
self.d_lock.acquire()
self.num_r_lock.acquire()
self.num_r += 1
if self.num_r == 1:
self.w_lock.acquire()
self.num_r_lock.release()
self.d_lock.release()
def r_release(self):
assert self.num_r > 0
self.num_r_lock.acquire()
self.num_r -= 1
if self.num_r == 0:
self.w_lock.release()
self.num_r_lock.release()
@contextmanager
def r_locked(self):
""" This method is designed to be used via the `with` statement. """
try:
self.r_acquire()
yield
finally:
self.r_release()
# ___________________________________________________________________
# Writing methods.
def w_acquire(self):
self.d_lock.acquire()
self.w_lock.acquire()
def w_acquire_non_blocking(self):
# if d_lock and w_lock can be acquired without blocking, acquire and return True,
# else immediately return False.
if self.d_lock.acquire(blocking=False):
if self.w_lock.acquire(blocking=False):
return True
else:
self.d_lock.release()
return False
def w_release(self):
self.w_lock.release()
self.d_lock.release()
def w_demote(self):
"""demote a writer lock to a reader lock"""
# the d_lock is already held from w_acquire.
# releasing the d_lock at the end of this function allows multiple readers.
# incrementing num_r makes this thread one of those readers.
self.num_r_lock.acquire()
self.num_r += 1
self.num_r_lock.release()
self.d_lock.release()
@contextmanager
def w_locked(self):
""" This method is designed to be used via the `with` statement. """
try:
self.w_acquire()
yield
finally:
self.w_release()

Some files were not shown because too many files have changed in this diff Show More