Skip to content

Commit

Permalink
Fixed initial (untested) implementation of multi-database wrapper
Browse files Browse the repository at this point in the history
  • Loading branch information
bartvanb committed Dec 7, 2022
1 parent d3a6c7c commit da8da7d
Show file tree
Hide file tree
Showing 2 changed files with 14 additions and 7 deletions.
14 changes: 10 additions & 4 deletions vantage6-client/vantage6/tools/docker_wrapper.py
Expand Up @@ -10,6 +10,7 @@
import io
from abc import ABC, abstractmethod
import pandas
import json

from vantage6.tools.dispatch_rpc import dispatch_rpc
from vantage6.tools.util import info
Expand Down Expand Up @@ -40,6 +41,11 @@ def parquet_wrapper(module: str):
wrapper.wrap_algorithm(module)


def multidb_wrapper(module: str):
wrapper = MultiDBWrapper()
wrapper.wrap_algorithm(module)


class WrapperBase(ABC):

def wrap_algorithm(self, module):
Expand Down Expand Up @@ -152,11 +158,11 @@ def load_data(database_uri, input_data):
class MultiDBWrapper(WrapperBase):
@staticmethod
def load_data(database_uri, input_data):
db_env_vars = os.environ.get("ALL_DATABASE_ENVVARS")
db_labels = json.loads(os.environ.get("DB_LABELS"))
databases = {}
for db_label in db_env_vars:
databases[db_label] = os.environ.get(db_label)
info(f"{databases}")
for db_label in db_labels:
db_env_var = f'{db_label.upper()}_DATABASE_URI'
databases[db_label] = os.environ.get(db_env_var)
return databases


Expand Down
7 changes: 4 additions & 3 deletions vantage6-node/vantage6/node/docker/task_manager.py
Expand Up @@ -4,6 +4,7 @@
import os
import pickle
import docker.errors
import json

from enum import Enum
from typing import Dict, List, Union
Expand Down Expand Up @@ -404,15 +405,15 @@ def _setup_environment_vars(self, algorithm_env: Dict = {},
# Only prepend the data_folder is it is a file-based database
# This allows algorithms to access multiple data-sources at the
# same time
db_env_vars = []
db_labels = []
for label in self.databases:
db = self.databases[label]
var_name = f'{label.upper()}_DATABASE_URI'
environment_variables[var_name] = \
f"{self.data_folder}/{os.path.basename(db['uri'])}" \
if db['is_file'] else db['uri']
db_env_vars.append(var_name)
environment_variables['ALL_DATABASE_ENVVARS'] = db_env_vars
db_labels.append(label)
environment_variables['DB_LABELS'] = json.dumps(db_labels)

# Support legacy algorithms
# TODO remove in v4+
Expand Down

0 comments on commit da8da7d

Please sign in to comment.