mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Fix typing errors in vault
This commit is contained in:
@@ -182,6 +182,7 @@ class MockAppConfig(GalaxyDataTestConfig, CommonConfigurationMixin):
|
||||
self.enable_tool_shed_check = False
|
||||
self.monitor_thread_join_timeout = 1
|
||||
self.integrated_tool_panel_config = None
|
||||
self.vault_config_file = None
|
||||
|
||||
@property
|
||||
def config_dict(self):
|
||||
|
||||
@@ -2356,4 +2356,4 @@ galaxy:
|
||||
# Vault config file
|
||||
# The value of this option will be resolved with respect to
|
||||
# <config_dir>.
|
||||
vault_config_file: vault_conf.yml
|
||||
vault_config_file: vault_conf.yml
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import abc
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from abc import ABC
|
||||
from typing import Optional
|
||||
|
||||
import yaml
|
||||
from cryptography.fernet import Fernet, MultiFernet
|
||||
@@ -28,9 +29,20 @@ class UnknownVaultTypeException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Vault(ABC):
|
||||
class Vault(abc.ABC):
|
||||
|
||||
def read_secret(self, key: str) -> str:
|
||||
@abc.abstractmethod
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def write_secret(self, key: str, value: str) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class NullVault(Vault):
|
||||
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
raise UnknownVaultTypeException("No vault configured. Make sure the vault_config_file setting is defined in galaxy.yml")
|
||||
|
||||
def write_secret(self, key: str, value: str) -> None:
|
||||
@@ -46,7 +58,7 @@ class HashicorpVault(Vault):
|
||||
self.vault_token = config.get('vault_token')
|
||||
self.client = hvac.Client(url=self.vault_address, token=self.vault_token)
|
||||
|
||||
def read_secret(self, key: str) -> str:
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
try:
|
||||
response = self.client.secrets.kv.read_secret_version(path=key)
|
||||
return response['data']['data'].get('value')
|
||||
@@ -76,7 +88,7 @@ class DatabaseVault(Vault):
|
||||
self.sa_session.merge(vault_entry)
|
||||
self.sa_session.flush()
|
||||
|
||||
def read_secret(self, key: str) -> str:
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
key_obj = self.sa_session.query(model.Vault).filter_by(key=key).first()
|
||||
if key_obj:
|
||||
f = self._get_multi_fernet()
|
||||
@@ -101,7 +113,7 @@ class CustosVault(Vault):
|
||||
self.b64_encoded_custos_token = custos_util.get_token(custos_settings=self.custos_settings)
|
||||
self.client = ResourceSecretManagementClient(self.custos_settings)
|
||||
|
||||
def read_secret(self, key: str) -> str:
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
try:
|
||||
response = self.client.get_KV_credential(token=self.b64_encoded_custos_token,
|
||||
client_id=self.custos_settings.CUSTOS_CLIENT_ID,
|
||||
@@ -127,7 +139,7 @@ class UserVaultWrapper(Vault):
|
||||
self.vault = vault
|
||||
self.user = user
|
||||
|
||||
def read_secret(self, key: str) -> str:
|
||||
def read_secret(self, key: str) -> Optional[str]:
|
||||
return self.vault.read_secret(f"user/{self.user.id}/{key}")
|
||||
|
||||
def write_secret(self, key: str, value: str) -> None:
|
||||
@@ -137,14 +149,14 @@ class UserVaultWrapper(Vault):
|
||||
class VaultFactory(object):
|
||||
|
||||
@staticmethod
|
||||
def load_vault_config(vault_conf_yml: str) -> dict:
|
||||
def load_vault_config(vault_conf_yml: str) -> Optional[dict]:
|
||||
if os.path.exists(vault_conf_yml):
|
||||
with open(vault_conf_yml) as f:
|
||||
return yaml.safe_load(f)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def from_vault_type(app, vault_type: str, cfg: dict) -> Vault:
|
||||
def from_vault_type(app, vault_type: Optional[str], cfg: dict) -> Vault:
|
||||
if vault_type == "hashicorp":
|
||||
return HashicorpVault(cfg)
|
||||
elif vault_type == "database":
|
||||
@@ -158,5 +170,6 @@ class VaultFactory(object):
|
||||
def from_app(app) -> Vault:
|
||||
vault_config = VaultFactory.load_vault_config(app.config.vault_config_file)
|
||||
if vault_config:
|
||||
return VaultFactory.from_vault_type(app, vault_config.get('type'), vault_config)
|
||||
return VaultFactory.from_vault_type(app, vault_config.get('type', None), vault_config)
|
||||
log.warning("No vault configured. We recommend defining the vault_config_file setting in galaxy.yml")
|
||||
return NullVault()
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
from typing import cast, Any
|
||||
|
||||
from requests import (
|
||||
get,
|
||||
@@ -21,7 +22,7 @@ class VaultApiTestCase(ApiTestCase):
|
||||
def test_extra_prefs_vault_storage(self):
|
||||
user = self._setup_user(TEST_USER_EMAIL)
|
||||
url = self.__url("information/inputs", user)
|
||||
app = self._test_driver.app
|
||||
app = cast(Any, self._test_driver.app if self._test_driver else None)
|
||||
|
||||
# create some initial data
|
||||
put(url, data=json.dumps({
|
||||
|
||||
@@ -7,10 +7,13 @@ from abc import ABC
|
||||
from cryptography.fernet import InvalidToken
|
||||
|
||||
from galaxy.app_unittest_utils.galaxy_mock import MockApp, MockAppConfig
|
||||
from galaxy.security import vault
|
||||
from galaxy.security.vault import NullVault, Vault, VaultFactory
|
||||
|
||||
|
||||
class VaultTestBase(ABC):
|
||||
class VaultTestBase(ABC, unittest.TestCase):
|
||||
|
||||
def __init__(self):
|
||||
self.vault = NullVault() # type: Vault
|
||||
|
||||
def test_read_write_secret(self):
|
||||
self.vault.write_secret("my/test/secret", "hello world")
|
||||
@@ -30,7 +33,7 @@ class TestHashicorpVault(VaultTestBase, unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
config = MockAppConfig(vault_config_file=VAULT_CONF_HASHICORP)
|
||||
app = MockApp(config=config)
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
|
||||
|
||||
VAULT_CONF_DATABASE = os.path.join(os.path.dirname(__file__), "fixtures/vault_conf_database.yaml")
|
||||
@@ -43,28 +46,28 @@ class TestDatabaseVault(VaultTestBase, unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
config = MockAppConfig(vault_config_file=VAULT_CONF_DATABASE)
|
||||
app = MockApp(config=config)
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
|
||||
def test_rotate_keys(self):
|
||||
config = MockAppConfig(vault_config_file=VAULT_CONF_DATABASE)
|
||||
app = MockApp(config=config)
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
self.vault.write_secret("my/rotated/secret", "hello rotated")
|
||||
|
||||
# should succeed after rotation
|
||||
app.config.vault_config_file = VAULT_CONF_DATABASE_ROTATED
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
self.assertEqual(self.vault.read_secret("my/rotated/secret"), "hello rotated")
|
||||
|
||||
def test_wrong_keys(self):
|
||||
config = MockAppConfig(vault_config_file=VAULT_CONF_DATABASE)
|
||||
app = MockApp(config=config)
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
self.vault.write_secret("my/incorrect/secret", "hello incorrect")
|
||||
|
||||
# should fail because decryption keys are the wrong
|
||||
app.config.vault_config_file = VAULT_CONF_DATABASE_INVALID
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
with self.assertRaises(InvalidToken):
|
||||
self.vault.read_secret("my/incorrect/secret")
|
||||
|
||||
@@ -82,7 +85,7 @@ class TestCustosVault(VaultTestBase, unittest.TestCase):
|
||||
self.vault_temp_conf = tempconf.name
|
||||
config = MockAppConfig(vault_config_file=self.vault_temp_conf)
|
||||
app = MockApp(config=config)
|
||||
self.vault = vault.VaultFactory.from_app(app)
|
||||
self.vault = VaultFactory.from_app(app)
|
||||
|
||||
def tearDown(self) -> None:
|
||||
os.remove(self.vault_temp_conf)
|
||||
|
||||
Reference in New Issue
Block a user