From e825a9a7c2f564871035e567b5813259f79bb6e3 Mon Sep 17 00:00:00 2001 From: Matthew Printz Date: Wed, 22 Jul 2026 17:36:00 -0600 Subject: [PATCH 1/4] Add SecretManager service and initial design --- src/beaker_notebook/app/base.py | 23 +++- src/beaker_notebook/services/auth/__init__.py | 7 +- .../services/kernel/manager.py | 5 + .../services/secrets/__init__.py | 0 .../services/secrets/manager.py | 107 +++++++++++++++ .../services/secrets/policies.py | 129 ++++++++++++++++++ src/beaker_notebook/services/secrets/types.py | 93 +++++++++++++ .../services/secrets/validations.py | 13 ++ 8 files changed, 375 insertions(+), 2 deletions(-) create mode 100644 src/beaker_notebook/services/secrets/__init__.py create mode 100644 src/beaker_notebook/services/secrets/manager.py create mode 100644 src/beaker_notebook/services/secrets/policies.py create mode 100644 src/beaker_notebook/services/secrets/types.py create mode 100644 src/beaker_notebook/services/secrets/validations.py diff --git a/src/beaker_notebook/app/base.py b/src/beaker_notebook/app/base.py index 93dfaed8..4026df2e 100644 --- a/src/beaker_notebook/app/base.py +++ b/src/beaker_notebook/app/base.py @@ -1,10 +1,11 @@ +import asyncio import getpass import inspect import logging import os import re import urllib.parse -from typing import ClassVar +from typing import ClassVar, TYPE_CHECKING import traitlets from traitlets.traitlets import Unicode @@ -24,6 +25,9 @@ from beaker_notebook.lib.utils import import_dotted_class from beaker_notebook.app.handlers import register_handlers +if TYPE_CHECKING: + from beaker_notebook.services.secrets.manager import BeakerSecretsManager + logger = logging.getLogger("beaker_server") @@ -69,6 +73,12 @@ class BaseBeakerApp(ServerApp): help="Path pointing to where user directories should be stored. Defaults to 'root_dir' if not set.", config=True, ) + secrets_manager: "BeakerSecretsManager" = traitlets.Instance( + f"beaker_notebook.services.secrets.manager.BeakerSecretsManager", + help="Beaker Secrets Manager singleton instance", + config=True, + ) + kernel_spec_include_local = traitlets.Bool(True, help="Include local kernel specs", config=True) kernel_spec_managers = traitlets.Dict(help="Kernel specification managers indexed by extension name", config=True) @@ -112,6 +122,13 @@ def _default_authorizer_class(self): from beaker_notebook.services.auth.notebook import NotebookAuthorizer return NotebookAuthorizer + @traitlets.default("secrets_manager") + def _default_secrets_manager(self): + from beaker_notebook.services.secrets.manager import BeakerSecretsManager + from beaker_notebook.services.secrets.types import UserEnvironmentSecret, SystemEnvironmentSecret + secrets_manager = BeakerSecretsManager() + return secrets_manager + @traitlets.default("config_file_name") def _default_config_file_name(self): if self.app_slug: @@ -176,6 +193,10 @@ def initialize(self, argv = None, find_extensions = False, new_httpserver = True super().initialize(argv, find_extensions, new_httpserver, starter_extension, **kwargs) + ioloop = asyncio.get_event_loop() + system_secrets = ioloop.run_until_complete(self.secrets_manager.collect_system_secrets(self)) + self.secrets_manager.add_secrets(system_secrets) + self.config["KernelProvisionerFactory"].setdefault("default_provisioner_name", "beaker-local-provisioner") if config.jupyter_token: self.config["IdentityProvider"].setdefault("token", config.jupyter_token) diff --git a/src/beaker_notebook/services/auth/__init__.py b/src/beaker_notebook/services/auth/__init__.py index aa69ee36..f4b4760f 100644 --- a/src/beaker_notebook/services/auth/__init__.py +++ b/src/beaker_notebook/services/auth/__init__.py @@ -6,7 +6,7 @@ import os from dataclasses import dataclass, field from functools import lru_cache, update_wrapper, wraps -from typing import Optional +from typing import Optional, TYPE_CHECKING from jupyter_server.auth.authorizer import Authorizer from jupyter_server.auth.identity import IdentityProvider, User @@ -15,6 +15,9 @@ from jupyter_server.services.config.manager import ConfigManager +if TYPE_CHECKING: + from beaker_notebook.services.secrets.types import BaseSecret + current_user = contextvars.ContextVar("current_user", default=None) current_request = contextvars.ContextVar("current_request", default=None) @@ -134,6 +137,7 @@ class BeakerAuthorizer(Authorizer): class BeakerUser(User): home_dir: Optional[str] = field(default=None) config: Optional[dict] = field(default=None) + secrets: "list[BaseSecret]" = field(default_factory=list) def __post_init__(self): """Initialize home directory if not provided. @@ -146,6 +150,7 @@ def __post_init__(self): if self.config is None: # TODO: Fetch config from somewhere self.config = {} + # TODO: Populate user secrets return super().__post_init__() @staticmethod diff --git a/src/beaker_notebook/services/kernel/manager.py b/src/beaker_notebook/services/kernel/manager.py index aa69d528..2a675328 100644 --- a/src/beaker_notebook/services/kernel/manager.py +++ b/src/beaker_notebook/services/kernel/manager.py @@ -10,6 +10,7 @@ from beaker_notebook.lib.app import BeakerApp from beaker_notebook.lib.config import config +from beaker_notebook.services.auth import BeakerUser, current_user if TYPE_CHECKING: from beaker_notebook.app.base import BaseBeakerApp @@ -124,6 +125,8 @@ async def _async_pre_start_kernel(self, **kw): if beaker_session and not self.beaker_session: self.beaker_session = beaker_session + beaker_user: BeakerUser = current_user.get() + cmd, kw = await super()._async_pre_start_kernel(**kw) env = kw.pop("env", {}) @@ -134,9 +137,11 @@ async def _async_pre_start_kernel(self, **kw): home_dir = os.path.expanduser(f"~{kernel_user}") kw["cwd"] = home_dir env["HOME"] = home_dir + env = await self.app.secrets_manager.sanitize_kernel_environment_vars(env=env) else: kernel_user = self.app.subkernel_user home_dir = kw.get("cwd") + env = await self.app.secrets_manager.sanitize_subkernel_envionment_vars(user=beaker_user, env=env) user_info = pwd.getpwnam(kernel_user) home_dir = os.path.expanduser(f"~{kernel_user}") diff --git a/src/beaker_notebook/services/secrets/__init__.py b/src/beaker_notebook/services/secrets/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/beaker_notebook/services/secrets/manager.py b/src/beaker_notebook/services/secrets/manager.py new file mode 100644 index 00000000..4a9505e9 --- /dev/null +++ b/src/beaker_notebook/services/secrets/manager.py @@ -0,0 +1,107 @@ +import copy +import inspect +import os +from dataclasses import dataclass, is_dataclass, asdict +from typing import TYPE_CHECKING, Any, Collection, Literal, TypeAlias + +import traitlets +from traitlets import Type, default +from traitlets.config.configurable import LoggingConfigurable +from traitlets.utils.importstring import import_item + +from beaker_notebook.services.secrets.policies import PolicyTypes, BasePolicy, Allow, Redact, Remove +from beaker_notebook.services.secrets.types import BaseSecret, EnvironmentSecret, UserEnvironmentSecret, SystemEnvironmentSecret, is_env_secret, is_system_env_secret, is_user_env_secret + + +if TYPE_CHECKING: + from beaker_notebook.app.base import BaseBeakerApp + from beaker_notebook.lib.context import BeakerContext + from beaker_notebook.services.auth import BeakerUser + + +class BeakerSecretsManager(LoggingConfigurable): + parent: "BaseBeakerApp" + + _secrets: list[BaseSecret] + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._secrets = [] + + def add_secret(self, secret: BaseSecret): + if secret in self._secrets: + return + self._secrets.append(secret) + + def add_secrets(self, secrets: Collection[BaseSecret]): + for secret in secrets: + self.add_secret(secret) + + async def sanitize_kernel_environment_vars(self, env: dict) -> dict[str, str]: + result = copy.copy(env) + for secret in self._secrets: + policy = secret.subkernel_environment_policy + if is_env_secret(secret) and secret.name in result: + if isinstance(policy, Allow): + continue + elif isinstance(policy, Remove): + result.pop(secret.name) + else: + result[secret.name] = await policy.sanitize(secret, result[secret.name]) + return result + + async def sanitize_subkernel_envionment_vars(self, user: "BeakerUser", env: dict) -> dict[str, str]: + result = copy.copy(env) + if user and user.secrets: + self.add_secrets(user.secrets) + for secret in self._secrets: + policy = secret.subkernel_environment_policy + if is_env_secret(secret) and secret.name in result: + if isinstance(policy, Allow): + continue + elif isinstance(policy, Remove): + result.pop(secret.name) + else: + result[secret.name] = await policy.sanitize(secret, result[secret.name]) + if is_user_env_secret(secret) and secret.name not in result: + # Add the user secret to the environment if it doesn't exist TODO: Determine if this is the right thing + result[secret.name] = secret.get_value() + return result + + + def _secret_env_vars(self) -> list[str]: + sensitive_env_keys = set() + anywhere = ( + "secret", + "private", + "password", + "passwd", + "creds", + "credentials", + "token", + "auth", + "passphrase", + ) + suffix = ( + "_key", + "_key_id", + ) + for key in os.environ.keys(): + lkey = key.lower() + if any(( + any(substr in lkey for substr in anywhere), + lkey.endswith(suffix), + )): + sensitive_env_keys.add(key) + return sorted(sensitive_env_keys) + + async def collect_system_secrets(self, app: "BaseBeakerApp") -> list[BaseSecret]: + system_secrets = [] + system_secrets.extend( + SystemEnvironmentSecret(name=env_name) for env_name in self._secret_env_vars() + ) + # TODO: Extract secrets from app config + # TODO: Extract secrets from beaker config + return system_secrets + + diff --git a/src/beaker_notebook/services/secrets/policies.py b/src/beaker_notebook/services/secrets/policies.py new file mode 100644 index 00000000..64fcb2c6 --- /dev/null +++ b/src/beaker_notebook/services/secrets/policies.py @@ -0,0 +1,129 @@ +import asyncio +import copy +from dataclasses import dataclass, field, is_dataclass, asdict +from typing import TYPE_CHECKING, Any, Awaitable, Callable, ClassVar, Collection, Literal, TypeAlias + +from beaker_notebook.services.secrets.validations import Validation, not_in + +if TYPE_CHECKING: + from .types import BaseSecret + +PolicyTypes: TypeAlias = Literal["redact", "remove", "last4", "allow"] + +@dataclass +class BasePolicy: + type: PolicyTypes + validations: Collection[Validation] = field(default_factory=lambda: [not_in,]) + + async def replacement(self, secret_str: str): + raise NotImplementedError(f"Policy '{BasePolicy}' does not define a sanitize function") + + + async def _sanitize_str(self, secret: "BaseSecret", content: str) -> str: + secret_str = secret.get_value() + replacement = await self.replacement(secret_str) + sanitized = content.replace(secret_str, replacement) + + await asyncio.gather(*[ + validation_func(secret_str, sanitized) for validation_func in self.validations + ]) + + return sanitized + + async def _sanitize_list(self, secret: "BaseSecret", content: list) -> list: + # Walk a list, sanitizing each item + for idx, value in enumerate(content): + match value: + case str(): + content[idx] = await self._sanitize_str(secret, value) + case dict(): + content[idx] = await self._sanitize_dict(secret, value) + case list(): + content[idx] = await self._sanitize_list(secret, value) + case bytes(): + content[idx] = (await self._sanitize_str(secret, value.decode())).encode() + case _: + pass + return content + + async def _sanitize_dict(self, secret: "BaseSecret", content: dict) -> dict: + # Walk dict, looking for strings and do a sanitization on each string + for key, value in content.items(): + match value: + case str(): + content[key] = await self._sanitize_str(secret, value) + case dict(): + content[key] = await self._sanitize_dict(secret, value) + case list(): + content[key] = await self._sanitize_list(secret, value) + case bytes(): + content[key] = (await self._sanitize_str(secret, value.decode())).encode() + case _: + pass + + async def sanitize(self, secret: "BaseSecret", content: str|dict): + target = copy.deepcopy(content) + match target: + case str() | bytes(): + return await self._sanitize_str(secret, target) + case dict(): + return await self._sanitize_dict(secret, target) + case list(): + return await self._sanitize_list(secret, target) + case _: + return content + +@dataclass +class Allow(BasePolicy): + """ + This policy does not modify the content in any way. + Only for use in high-trust situations or debugging + """ + + type: PolicyTypes = "allow" + + # TODO: Add warning on initialization about this being insecure? + + async def sanitize(self, secret, content): + return content + + +@dataclass +class Redact(BasePolicy): + type: PolicyTypes = "redact" + + async def replacement(self, secret_str): + return "#" * len(secret_str) + + +@dataclass +class Remove(BasePolicy): + type: PolicyTypes = "remove" + + async def _sanitize_list(self, secret, content): + secret_value = secret.get_value() + while secret_value in content: + content.remove(secret_value) + return content + + async def _sanitize_dict(self, secret, content): + name = getattr(secret, name, None) + if name and name in content: + content.pop(name) + return content + + async def replacement(self, secret_str): + return "" + + +@dataclass +class Last4(BasePolicy): + type: PolicyTypes = "last4" + + async def replacement(self, secret_str): + if len(secret_str) < 6: + # Returning last four chars could is too much information for such a short secret. + return "####" + return f"###{secret_str[:-4]}" + + diff --git a/src/beaker_notebook/services/secrets/types.py b/src/beaker_notebook/services/secrets/types.py new file mode 100644 index 00000000..f106cdcf --- /dev/null +++ b/src/beaker_notebook/services/secrets/types.py @@ -0,0 +1,93 @@ +import os + +from dataclasses import dataclass, field, is_dataclass, asdict +from typing import TYPE_CHECKING, Any, ClassVar, Literal, Optional, TypeAlias, Type, TypeIs + +from beaker_notebook.services.auth import BeakerUser +from beaker_notebook.services.secrets.policies import BasePolicy, Allow, Redact, Remove + +if TYPE_CHECKING: + pass + + +PolicyRef: TypeAlias = BasePolicy | Type[BasePolicy] + + +@dataclass(kw_only=True) +class BaseSecret: + type: ClassVar[str] = "base-secret" + + subkernel_message_policy: PolicyRef + ui_message_policy: PolicyRef + agent_message_policy: PolicyRef + beaker_kernel_environment_policy: PolicyRef + subkernel_environment_policy: PolicyRef + + def __post_init__(self, *args, **kwargs): + for name in self.__dataclass_fields__: + value = getattr(self, name) + if isinstance(value, type) and issubclass(value, BasePolicy): + setattr(self, name, value()) + + def get_value(self) -> Optional[str]: + raise NotImplementedError() + + +@dataclass(kw_only=True) +class EnvironmentSecret(BaseSecret): + type = "env-secret" + agent_message_policy: PolicyRef = Redact + subkernel_message_policy: PolicyRef = Redact + ui_message_policy: PolicyRef = Redact + + name: str + + def get_value(self): + return os.environ.get(self.name) + + +@dataclass(kw_only=True) +class SystemEnvironmentSecret(EnvironmentSecret): + subkernel_message_policy: PolicyRef = Redact + ui_message_policy: PolicyRef = Redact + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Remove + + def get_value(self): + import os + return os.environ.get(self.name) + + +@dataclass(kw_only=True) +class UserEnvironmentSecret(EnvironmentSecret): + subkernel_message_policy: PolicyRef = Allow + ui_message_policy: PolicyRef = Allow + beaker_kernel_environment_policy: PolicyRef = Remove + subkernel_environment_policy: PolicyRef = Allow + + user: Optional[BeakerUser] = None + + def get_value(self): + return "USER ENV VALUE" + + +@dataclass(kw_only=True) +class SkillSecret(BaseSecret): + type: str = "skill-secret" + skill_name: str + name: str + default_value: Optional[str] + + def get_value(self): + return os.environ.get(self.name) + + +def is_env_secret(secret: BaseSecret) -> TypeIs[EnvironmentSecret]: + return isinstance(secret, EnvironmentSecret) + +def is_system_env_secret(secret: BaseSecret) -> TypeIs[SystemEnvironmentSecret]: + return isinstance(secret, SystemEnvironmentSecret) + +def is_user_env_secret(secret: BaseSecret) -> TypeIs[UserEnvironmentSecret]: + return isinstance(secret, UserEnvironmentSecret) + diff --git a/src/beaker_notebook/services/secrets/validations.py b/src/beaker_notebook/services/secrets/validations.py new file mode 100644 index 00000000..4fb76b05 --- /dev/null +++ b/src/beaker_notebook/services/secrets/validations.py @@ -0,0 +1,13 @@ +from typing import TYPE_CHECKING, Any, Awaitable, Callable, ClassVar, Collection, Literal, TypeAlias + + +Validation: TypeAlias = Callable[[str, str], Awaitable[None]] + + +class SecretValidationError(Exception): + pass + + +async def not_in(secret_str: str, content: str) -> None: + if secret_str in content: + raise SecretValidationError() From e87f80de32d787ccc80bc0b922368fc7c6c63945 Mon Sep 17 00:00:00 2001 From: Matthew Printz Date: Thu, 23 Jul 2026 13:58:03 -0600 Subject: [PATCH 2/4] Implementation of App Trait secrets --- .../services/secrets/manager.py | 151 ++++++++++++++---- src/beaker_notebook/services/secrets/types.py | 51 +++++- .../services/storage/notebook.py | 2 +- 3 files changed, 169 insertions(+), 35 deletions(-) diff --git a/src/beaker_notebook/services/secrets/manager.py b/src/beaker_notebook/services/secrets/manager.py index 4a9505e9..bd874a1e 100644 --- a/src/beaker_notebook/services/secrets/manager.py +++ b/src/beaker_notebook/services/secrets/manager.py @@ -5,12 +5,15 @@ from typing import TYPE_CHECKING, Any, Collection, Literal, TypeAlias import traitlets -from traitlets import Type, default -from traitlets.config.configurable import LoggingConfigurable +from traitlets import Type, default, HasTraits +from traitlets.config.configurable import Configurable, LoggingConfigurable from traitlets.utils.importstring import import_item from beaker_notebook.services.secrets.policies import PolicyTypes, BasePolicy, Allow, Redact, Remove -from beaker_notebook.services.secrets.types import BaseSecret, EnvironmentSecret, UserEnvironmentSecret, SystemEnvironmentSecret, is_env_secret, is_system_env_secret, is_user_env_secret +from beaker_notebook.services.secrets.types import ( + BaseSecret, UserEnvironmentSecret, SystemEnvironmentSecret, AppTraitSecret, + is_env_secret, is_system_env_secret, is_user_env_secret +) if TYPE_CHECKING: @@ -19,10 +22,93 @@ from beaker_notebook.services.auth import BeakerUser +def index_configurables(root): + from traitlets.config import SingletonConfigurable, LoggingConfigurable, Configurable + seen, index, queue = set(), {}, [root] + while queue: + obj = queue.pop() + if id(obj) in seen: + continue + seen.add(id(obj)) + for klass in type(obj).__mro__: + if klass in (LoggingConfigurable, Configurable, SingletonConfigurable): + break + index.setdefault(klass.__name__, obj) + for _, value in inspect.getmembers(obj, lambda member: isinstance(obj, Configurable)): + if isinstance(value, HasTraits): + queue.append(value) + elif isinstance(value, (list, tuple, dict)): + queue.extend(v for v in (value.values() if isinstance(value, dict) else value) + if isinstance(v, HasTraits)) + return index + + class BeakerSecretsManager(LoggingConfigurable): parent: "BaseBeakerApp" - _secrets: list[BaseSecret] + _index: dict[str, Configurable] + + app_trait_secrets: list[str] = traitlets.List( + trait=traitlets.Unicode, + config=True, + default_value=[ + "Application.cookie_secret", + "IdentityProvider.token", + "NotebookNotary.secret", + "GatewayClient.auth_token", + ] + ) + extra_app_trait_secrets: list[str] = traitlets.List( + trait=traitlets.Unicode, + config=True, + default_value=[], + ) + + extra_user_env_vars: list[str] = traitlets.List( + trait=traitlets.Unicode, + config=True, + default_value=[], + ) + extra_system_env_vars: list[str] = traitlets.List( + trait=traitlets.Unicode, + config=True, + default_value=[], + ) + + env_key_detection_anywhere: list[str] = traitlets.Tuple( + trait=traitlets.Unicode, + config=True, + default_value=( + "secret", + "private", + "password", + "passwd", + "creds", + "credentials", + "token", + "auth", + "passphrase", + ), + ) + extra_env_key_detection_anywhere: list[str] = traitlets.Tuple( + trait=traitlets.Unicode, + config=True, + default_value=tuple(), + ) + env_key_detection_suffix: list[str] = traitlets.Tuple( + trait=traitlets.Unicode, + config=True, + default_value=( + "_key", + "_key_id", + ), + ) + extra_env_key_detection_suffix: list[str] = traitlets.Tuple( + trait=traitlets.Unicode, + config=True, + default_value=tuple(), + ) + def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -37,6 +123,10 @@ def add_secrets(self, secrets: Collection[BaseSecret]): for secret in secrets: self.add_secret(secret) + @property + def secrets(self) -> list[BaseSecret]: + return self._secrets + async def sanitize_kernel_environment_vars(self, env: dict) -> dict[str, str]: result = copy.copy(env) for secret in self._secrets: @@ -71,37 +161,40 @@ async def sanitize_subkernel_envionment_vars(self, user: "BeakerUser", env: dict def _secret_env_vars(self) -> list[str]: sensitive_env_keys = set() - anywhere = ( - "secret", - "private", - "password", - "passwd", - "creds", - "credentials", - "token", - "auth", - "passphrase", - ) - suffix = ( - "_key", - "_key_id", - ) + # Ensure all keys are lowercase for matching + suffix_keys = tuple(map(str.lower, self.env_key_detection_suffix + self.extra_env_key_detection_suffix)) + anywhere_keys = tuple(map(str.lower, self.env_key_detection_anywhere + self.extra_env_key_detection_anywhere)) + # Check each key for match for key in os.environ.keys(): lkey = key.lower() - if any(( - any(substr in lkey for substr in anywhere), - lkey.endswith(suffix), - )): + if ( + lkey.endswith(suffix_keys) or + any(substr in lkey for substr in anywhere_keys) + ): sensitive_env_keys.add(key) return sorted(sensitive_env_keys) + async def collect_system_secrets(self, app: "BaseBeakerApp") -> list[BaseSecret]: system_secrets = [] - system_secrets.extend( - SystemEnvironmentSecret(name=env_name) for env_name in self._secret_env_vars() - ) - # TODO: Extract secrets from app config - # TODO: Extract secrets from beaker config - return system_secrets + # Environment secrets from configuration + discovered_system_env_secrets = [SystemEnvironmentSecret(name=env_name) for env_name in self._secret_env_vars()] + extra_system_envs = [SystemEnvironmentSecret(name=env_name) for env_name in self.extra_system_env_vars] + extra_user_envs = [UserEnvironmentSecret(name=env_name) for env_name in self.extra_user_env_vars] + configurable_index = index_configurables(app) + app_trait_secrets = [] + for config_string in self.app_trait_secrets + self.extra_app_trait_secrets: + configurable_name, trait_name = config_string.split(".", maxsplit=1) + configurable = configurable_index[configurable_name] + app_trait_secrets.append(AppTraitSecret(configurable=configurable, trait_name=trait_name)) + + # TODO: Extract secrets from beaker config + + # TODO: Clean this up + system_secrets.extend(discovered_system_env_secrets) + system_secrets.extend(extra_system_envs) + system_secrets.extend(extra_user_envs) + system_secrets.extend(app_trait_secrets) + return system_secrets diff --git a/src/beaker_notebook/services/secrets/types.py b/src/beaker_notebook/services/secrets/types.py index f106cdcf..39dc3a4d 100644 --- a/src/beaker_notebook/services/secrets/types.py +++ b/src/beaker_notebook/services/secrets/types.py @@ -1,8 +1,10 @@ import os - -from dataclasses import dataclass, field, is_dataclass, asdict +import weakref +from dataclasses import dataclass, InitVar, field, is_dataclass, asdict from typing import TYPE_CHECKING, Any, ClassVar, Literal, Optional, TypeAlias, Type, TypeIs +from traitlets.config import Application, Configurable + from beaker_notebook.services.auth import BeakerUser from beaker_notebook.services.secrets.policies import BasePolicy, Allow, Redact, Remove @@ -24,9 +26,9 @@ class BaseSecret: subkernel_environment_policy: PolicyRef def __post_init__(self, *args, **kwargs): - for name in self.__dataclass_fields__: - value = getattr(self, name) - if isinstance(value, type) and issubclass(value, BasePolicy): + for name, field_info in self.__dataclass_fields__.items(): + if field_info.type == PolicyRef: + value = getattr(self, name) setattr(self, name, value()) def get_value(self) -> Optional[str]: @@ -82,6 +84,43 @@ def get_value(self): return os.environ.get(self.name) +@dataclass(kw_only=True) +class AppTraitSecret(BaseSecret): + subkernel_message_policy: PolicyRef = Remove + ui_message_policy: PolicyRef = Redact + agent_message_policy: PolicyRef = Redact + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Remove + + configurable: InitVar["Optional[Configurable]"] = field(default=None, repr=False, compare=False) + _configurable_ref: "weakref.ref[Configurable] | None" = field(default=None, init=False, repr=False, compare=False) + trait_name: str + + def __post_init__(self, configurable=None): + if configurable is not None: + self._configurable_ref = weakref.ref(configurable) + super().__post_init__() + if configurable: + # We only need to validate when we have something to validate against + self._validate(configurable) + + def _validate(self, configurable: Configurable): + from traitlets import Unicode, Bytes + + traits = configurable.trait_names() + if self.trait_name not in traits: + raise ValueError(f"{self.__class__.__name__}: Class {configurable.__class__.__name__} does not contain a trait named '{self.trait_name}'") + trait = configurable.traits().get(self.trait_name) + if not isinstance(trait, (Unicode, Bytes)): + raise ValueError(f"{self.__class__.__name__} targeted traits should point to a string (Unicode or Bytes), not {trait.__class__.__name__}") + + def get_value(self): + configurable = self._configurable_ref and self._configurable_ref() + if configurable is None: + raise ValueError(f"<{self.__class__.__name__}: name='{self.name}'> is not available.") + return getattr(configurable, self.trait_name, None) + + def is_env_secret(secret: BaseSecret) -> TypeIs[EnvironmentSecret]: return isinstance(secret, EnvironmentSecret) @@ -91,3 +130,5 @@ def is_system_env_secret(secret: BaseSecret) -> TypeIs[SystemEnvironmentSecret]: def is_user_env_secret(secret: BaseSecret) -> TypeIs[UserEnvironmentSecret]: return isinstance(secret, UserEnvironmentSecret) +def is_app_trait_secret(secret: BaseSecret) -> TypeIs[AppTraitSecret]: + return isinstance(secret, AppTraitSecret) \ No newline at end of file diff --git a/src/beaker_notebook/services/storage/notebook.py b/src/beaker_notebook/services/storage/notebook.py index 78d7ebd5..7cfecff5 100644 --- a/src/beaker_notebook/services/storage/notebook.py +++ b/src/beaker_notebook/services/storage/notebook.py @@ -89,7 +89,7 @@ async def delete_snapshot(self, session_id: str) -> None: class FileNotebookManager(BaseNotebookManager): contents_manager_class = traitlets.Type( - default_value=None, + default_value="jupyter_server.services.contents.filemanager.AsyncFileContentsManager", klass=ContentsManager, allow_none=True, config=True, From 7681ac443d02eb7d07158532e5b8d9bb844869b5 Mon Sep 17 00:00:00 2001 From: Matthew Printz Date: Mon, 27 Jul 2026 15:56:29 -0600 Subject: [PATCH 3/4] Finish implementation of secrets, policies, and add initial scrubbing - Finishes defining initial secrets types and policies and related logic. - Adds scrubbing for jupyter message over the wire. - Adds backend, beaker-kernel-only api endpoint for loading secrets to kernel process - Full test suite --- src/beaker_notebook/app/base.py | 3 +- src/beaker_notebook/kernel.py | 51 ++++- .../lib/jupyter_kernel_proxy.py | 19 +- src/beaker_notebook/lib/secrets/__init__.py | 0 src/beaker_notebook/lib/secrets/policies.py | 196 ++++++++++++++++++ .../lib/secrets/secret_types.py | 169 +++++++++++++++ .../{services => lib}/secrets/validations.py | 0 src/beaker_notebook/services/auth/__init__.py | 2 +- .../services/secrets/app_secrets.py | 65 ++++++ .../services/secrets/handlers.py | 40 ++++ .../services/secrets/manager.py | 55 +++-- .../services/secrets/policies.py | 129 ------------ src/beaker_notebook/services/secrets/types.py | 134 ------------ .../services/storage/notebook.py | 17 +- tests/secrets/__init__.py | 0 tests/secrets/test_message_scrubbing.py | 110 ++++++++++ tests/secrets/test_policies.py | 187 +++++++++++++++++ tests/secrets/util.py | 133 ++++++++++++ tests/test_generate_config.py | 53 +++++ 19 files changed, 1063 insertions(+), 300 deletions(-) create mode 100644 src/beaker_notebook/lib/secrets/__init__.py create mode 100644 src/beaker_notebook/lib/secrets/policies.py create mode 100644 src/beaker_notebook/lib/secrets/secret_types.py rename src/beaker_notebook/{services => lib}/secrets/validations.py (100%) create mode 100644 src/beaker_notebook/services/secrets/app_secrets.py create mode 100644 src/beaker_notebook/services/secrets/handlers.py delete mode 100644 src/beaker_notebook/services/secrets/policies.py delete mode 100644 src/beaker_notebook/services/secrets/types.py create mode 100644 tests/secrets/__init__.py create mode 100644 tests/secrets/test_message_scrubbing.py create mode 100644 tests/secrets/test_policies.py create mode 100644 tests/secrets/util.py create mode 100644 tests/test_generate_config.py diff --git a/src/beaker_notebook/app/base.py b/src/beaker_notebook/app/base.py index 4026df2e..b7616d50 100644 --- a/src/beaker_notebook/app/base.py +++ b/src/beaker_notebook/app/base.py @@ -125,8 +125,7 @@ def _default_authorizer_class(self): @traitlets.default("secrets_manager") def _default_secrets_manager(self): from beaker_notebook.services.secrets.manager import BeakerSecretsManager - from beaker_notebook.services.secrets.types import UserEnvironmentSecret, SystemEnvironmentSecret - secrets_manager = BeakerSecretsManager() + secrets_manager = BeakerSecretsManager(parent=self) return secrets_manager @traitlets.default("config_file_name") diff --git a/src/beaker_notebook/kernel.py b/src/beaker_notebook/kernel.py index 9304ac07..5f28280a 100644 --- a/src/beaker_notebook/kernel.py +++ b/src/beaker_notebook/kernel.py @@ -1,5 +1,4 @@ import asyncio -import contextvars import copy import inspect import json @@ -18,6 +17,7 @@ from beaker_notebook.lib.config import reset_config, config from beaker_notebook.lib.context import BeakerContext, autodiscover_contexts +from beaker_notebook.lib.secrets.secret_types import BaseSecret, BeakerConfigProviderSecret, SkillSecret, UserEnvironmentSecret from beaker_notebook.lib.subkernel import BeakerSubkernel from beaker_notebook.lib.jupyter_kernel_proxy import InterceptionFilter, JupyterMessage, KernelProxyManager from beaker_notebook.lib.utils import (message_handler, LogMessageEncoder, magic, @@ -63,6 +63,7 @@ class BeakerKernel(KernelProxyManager): magic_commands: dict[str, callable] ready: asyncio.Future running_actions: dict[str, Awaitable] + secrets: list[BaseSecret] def __init__(self, session_config, kernel_id=None, connection_file=None): self.session_config = session_config @@ -76,6 +77,7 @@ def __init__(self, session_config, kernel_id=None, connection_file=None): self.internal_executions = set() self.subkernel_execution_tracking = {} self.running_actions = {} + self.secrets = [] context_args = session_config.get("context", {}) super().__init__(session_config, session_id=(self.beaker_session or self.kernel_id)) self.register_magic_commands() @@ -145,6 +147,10 @@ async def start_default_context(self, default_context=None, default_context_payl if not default_context: default_context = "default" default_context_payload = {} + + # Delay fetching secrets until context is established + self.secrets = self.get_session_secrets() + await self.set_context(default_context, default_context_payload, **optional_args) def add_base_intercepts(self): @@ -172,6 +178,7 @@ def add_base_intercepts(self): self.server.intercept_message("control", "shutdown_request", self.shutdown) self.server.intercept_message("shell", "notebook_state_response", self.notebook_state_response) self.server.intercept_message("shell", "beaker_session_info_request", self.beaker_session_info) + self.server.intercept_message(None, None, self.redact_secrets) def register_magic_commands(self): for _, method in inspect.getmembers(self, lambda member: inspect.ismethod(member) and hasattr(member, "_magic_prefix")): @@ -224,6 +231,29 @@ def get_session_attachments(self, current_attachment_ids: list[str] | None = Non item["current"] = item.get("id") in current_ids return attachments + def get_session_secrets(self, *args, **kwargs) -> list: + """Fetch session secrets to scrub""" + session_id = self.beaker_session or self.session_id + url = url_path_join( + self.jupyter_server, + "/beaker/secrets/", + urllib.parse.quote(str(session_id), safe=""), + ) + response = requests.get( + url, + headers={"X-AUTH-BEAKER": self.api_auth()}, + timeout=10, + ) + if response.status_code >= 400: + raise ValueError( + f"Unable to load session secrets (status {response.status_code}): {response.text}" + ) + secrets_json = response.json() + # Reify secrets, skipping ones for which the sent value is None + secrets = [BaseSecret.from_dict(secret_dict) for secret_dict in secrets_json if secret_dict.get("_value", None) is not None] + return secrets + + def clear_session_attachments(self) -> None: """Delete every temporary attachment owned by the current notebook session.""" session_id = self.beaker_session or self.session_id @@ -902,6 +932,25 @@ async def notebook_state_response(self, server, target_stream, data): return None setattr(notebook_state_response, 'result', None) + async def redact_secrets(self, server, target_stream, data): + destination = server.get_destination(target_stream) + policy_attr = { + "client": "ui_message_policy", + "subkernel": "subkernel_message_policy", + }.get(destination) + if policy_attr is None: + return data + message = JupyterMessage.parse(data) + # Scrub the decoded payload fields. Routing headers (header/parent_header) + # and binary buffers are left intact so a false-positive substring match + # can't corrupt correlation ids or buffer bytes. + for secret in self.secrets: + policy = getattr(secret, policy_attr) + for field_name in ("metadata", "content"): + sanitized = await policy.sanitize(secret=secret, content=getattr(message, field_name)) + message = message._replace(**{field_name: sanitized}) + return message.parts + async def request_notebook_state(self, parent_message=None): msg_id = str(uuid.uuid4()) self.send_response( diff --git a/src/beaker_notebook/lib/jupyter_kernel_proxy.py b/src/beaker_notebook/lib/jupyter_kernel_proxy.py index ebd4b2c1..35b882f1 100644 --- a/src/beaker_notebook/lib/jupyter_kernel_proxy.py +++ b/src/beaker_notebook/lib/jupyter_kernel_proxy.py @@ -15,6 +15,7 @@ import uuid from collections import OrderedDict, namedtuple from operator import attrgetter +from typing import Callable, Awaitable, TypeAlias, Literal import six import zmq @@ -209,20 +210,23 @@ class ProxyKernelClient(AbstractProxyKernel): def __init__(self, config, role="client", zmq_context=zmq.Context.instance(), session_id=None): super(ProxyKernelClient, self).__init__(config, role, zmq_context, session_id=session_id) +FilterType: TypeAlias = "tuple[str|None, str|None, Callable[[ProxyKernelServer, zmqstream.ZMQStream, list], Awaitable[list[JupyterMessage]|None]]]" -InterceptionFilter = namedtuple( +InterceptionFilter: FilterType = namedtuple( "InterceptionFilter", ("stream_type", "msg_type", "callback") ) - class ProxyKernelServer(AbstractProxyKernel): def __init__(self, config, role="server", zmq_context=zmq.Context.instance(), session_id=None): self.manager = None super(ProxyKernelServer, self).__init__(config, role, zmq_context, session_id=session_id) - self.filters = [] + self.filters: list[FilterType] = [] self.session_id = session_id self.proxy_target = None + def get_destination(self, stream: zmqstream.ZMQStream) -> Literal["client", "subkernel"]: + return "client" if stream in self.streams else "subkernel" + def _proxy_to( self, other_stream, socktype=None, validate_using=None, resign_using=None ): @@ -246,10 +250,9 @@ async def handler(data): msg.identities = [] if self.session_id and self.session_id not in msg.identities: msg.identities.append(msg.parent_header.get("session")) + # None on stream_type or msg_type matches any value of either for stream_type, msg_type, callback in self.filters: - if stream_type == socktype and msg_type == msg.header.get( - "msg_type" - ): + if ((stream_type is None or stream_type == socktype) and (msg_type is None or msg_type == msg.header.get("msg_type"))): new_data = await callback(self, other_stream, data) if new_data is None: return @@ -280,9 +283,9 @@ def set_proxy_target(self, proxy_client): def intercept_message(self, stream_type=None, msg_type=None, callback=None): if stream_type in KERNEL_SOCKETS_NAMES: stream_type = KERNEL_SOCKETS[KERNEL_SOCKETS_NAMES.index(stream_type)] - if stream_type not in KERNEL_SOCKETS: + if stream_type not in KERNEL_SOCKETS and stream_type is not None: raise ValueError( - "stream_type should be one of " + ", ".join(KERNEL_SOCKETS_NAMES) + f'stream_type should be one of {", ".join(KERNEL_SOCKETS_NAMES)} or None (to match all)' ) if not callable(callback): raise ValueError("callback must be callable") diff --git a/src/beaker_notebook/lib/secrets/__init__.py b/src/beaker_notebook/lib/secrets/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/beaker_notebook/lib/secrets/policies.py b/src/beaker_notebook/lib/secrets/policies.py new file mode 100644 index 00000000..e34048d9 --- /dev/null +++ b/src/beaker_notebook/lib/secrets/policies.py @@ -0,0 +1,196 @@ +import asyncio +import copy +import re +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, MutableSequence, MutableMapping, Literal, TypeAlias, TypeVar + +from beaker_notebook.lib.secrets.validations import Validation, not_in + + +if TYPE_CHECKING: + from beaker_notebook.lib.secrets.secret_types import BaseSecret + + +PolicyTypes: TypeAlias = Literal["redact", "remove", "last4", "allow"] + +MappedContent = TypeVar("MappedContent", bound=MutableMapping) +SequencedContent = TypeVar("SequencedContent", bound=MutableSequence) + +key_regex = re.compile(r'__([^:]+):([0-9#]+)__') +def next_key(key: str, mapping: MutableMapping): + match = key_regex.search(key) + if not match: + return '?' + cls_name = match.group(1) + count = len([k for k in mapping.keys() if (m := key_regex.search(k)) and m.group(1) == cls_name]) + num = count + while (new_key := key_regex.sub(lambda m: f"__{m.group(1)}:{num}__", key)) in mapping: + num += 1 + return new_key + +@dataclass +class BasePolicy: + type: PolicyTypes + validations: MutableSequence[Validation] = field(default_factory=lambda: [not_in,]) + + async def replacement(self, secret_str: str): + raise NotImplementedError(f"Policy '{BasePolicy}' does not define a sanitize function") + + def _sanitize_func_for(self, target: Any): + match target: + case str(): + return self._sanitize_str + case {**mapping}: + return self._sanitize_mapping + case [*list]: + return self._sanitize_list + case bytes(): + return self._sanitize_bytes + case _: + return None + + async def _sanitize_bytes(self, secret: "BaseSecret", content: bytes) -> bytes: + secret_str = secret.value + if secret_str is None: + return content + secret_bytes = secret_str.encode() + replacement: str = await self.replacement(secret_str) + replacement_bytes = replacement.encode() + sanitized = content.replace(secret_bytes, replacement_bytes) + + await asyncio.gather(*[ + validation_func(secret_bytes, sanitized) for validation_func in self.validations + ]) + + return sanitized + + + async def _sanitize_str(self, secret: "BaseSecret", content: str) -> str: + secret_str = secret.value + if secret_str is None: + return content + replacement = await self.replacement(secret_str) + sanitized = content.replace(secret_str, replacement) + + await asyncio.gather(*[ + validation_func(secret_str, sanitized) for validation_func in self.validations + ]) + + return sanitized + + async def _sanitize_list(self, secret: "BaseSecret", content: SequencedContent) -> SequencedContent: + # Walk a list, sanitizing each item + for idx, value in enumerate(content): + sanitize_func = self._sanitize_func_for(value) + if sanitize_func is None: + continue + content[idx] = await sanitize_func(secret, value) + return content + + async def _sanitize_mapping(self, secret: "BaseSecret", content: MappedContent) -> MappedContent: + # Walk dict, looking for strings and do a sanitization on each string + to_remap: list[tuple[str, str]] = [] + key_policy = MappingKey() + for key, value in content.items(): + key_sanitize_func = key_policy._sanitize_func_for(key) + value_sanitize_func = self._sanitize_func_for(value) + if key_sanitize_func: + new_key = await key_sanitize_func(secret, key) + if new_key != key: + to_remap.append((key, new_key)) + if value_sanitize_func: + content[key] = await value_sanitize_func(secret, value) + + # Remap after walking to avoid mutating keys during iteration + for key, new_key in to_remap: + new_key = next_key(new_key, content) + content[new_key] = content.pop(key) + return content + + async def sanitize(self, secret: "BaseSecret", content: str|dict|list|bytes): + sanitize_func = self._sanitize_func_for(content) + if sanitize_func is None: + return content + target = copy.deepcopy(content) + return await sanitize_func(secret, target) + +@dataclass +class Allow(BasePolicy): + """ + This policy does not modify the content in any way. + Only for use in high-trust situations or debugging + """ + + type: PolicyTypes = "allow" + + # TODO: Add warning on initialization about this being insecure? + + async def sanitize(self, secret, content): + return content + + +@dataclass +class Redact(BasePolicy): + type: PolicyTypes = "redact" + + async def replacement(self, secret_str): + return "#" * len(secret_str) + + +@dataclass +class Remove(BasePolicy): + type: PolicyTypes = "remove" + + async def replacement(self, secret_str): + return "" + + +@dataclass +class Last4(BasePolicy): + type: PolicyTypes = "last4" + + async def replacement(self, secret_str): + secret_len = len(secret_str) + if secret_len < 6: + # Returning last four chars could is too much information for such a short secret. + return "####" + hash = "#" * (secret_len - 4) + return f"{hash}{secret_str[-4:]}" + + +@dataclass +class MappingKey(BasePolicy): + type: PolicyTypes = "mapping-key" + + def _secret_value_str(self, secret: "BaseSecret"): + secret_str = secret.type.upper() + return f"__{secret_str}:#__" + + async def _sanitize_bytes(self, secret: "BaseSecret", content: bytes) -> bytes: + secret_str = secret.value + if secret_str is None: + return content + secret_bytes = secret_str.encode() + replacement: str = self._secret_value_str(secret) + replacement_bytes = replacement.encode() + sanitized = content.replace(secret_bytes, replacement_bytes) + + await asyncio.gather(*[ + validation_func(secret_bytes, sanitized) for validation_func in self.validations + ]) + + return sanitized + + async def _sanitize_str(self, secret: "BaseSecret", content: str) -> str: + secret_str = secret.value + if secret_str is None: + return content + replacement = self._secret_value_str(secret) + sanitized = content.replace(secret_str, replacement) + + await asyncio.gather(*[ + validation_func(secret_str, sanitized) for validation_func in self.validations + ]) + + return sanitized + diff --git a/src/beaker_notebook/lib/secrets/secret_types.py b/src/beaker_notebook/lib/secrets/secret_types.py new file mode 100644 index 00000000..e0a08d0d --- /dev/null +++ b/src/beaker_notebook/lib/secrets/secret_types.py @@ -0,0 +1,169 @@ +import os +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, ClassVar, Optional, TypeAlias, Type, TypeGuard, get_origin + +# from beaker_notebook.services.auth import BeakerUser +from beaker_notebook.lib.secrets.policies import BasePolicy, Allow, Redact, Remove +from beaker_notebook.lib.utils import to_import_string, import_dotted_class + +if TYPE_CHECKING: + pass + + +PolicyRef: TypeAlias = BasePolicy | Type[BasePolicy] + + +@dataclass(kw_only=True) +class BaseSecret: + type: ClassVar[str] = "base-secret" + _value: Optional[str] = field(init=False, repr=False, hash=False, compare=False) + + subkernel_message_policy: PolicyRef + ui_message_policy: PolicyRef + agent_message_policy: PolicyRef + beaker_kernel_environment_policy: PolicyRef + subkernel_environment_policy: PolicyRef + + def __post_init__(self, *args, **kwargs): + # Allows defining policies as a class or an instance + # If the policy is a class, instantiate it as part of init + for name, field_info in self.__dataclass_fields__.items(): + if field_info.type == PolicyRef: + value = getattr(self, name) + setattr(self, name, value()) + self._value = None + + def to_dict(self, with_value: bool = False): + output = { + "cls": to_import_string(self) + } + for key, field_def in self.__dataclass_fields__.items(): + if get_origin(field_def.type) == ClassVar: + continue + if key == "_value": + continue + field = getattr(self, key) + match key, field: + case _, BasePolicy(): + field = { + "import_str": to_import_string(field) + } + case _, val if callable(val): + continue + case _, _: + pass + output[key] = field + if with_value: + try: + output["_value"] = self.get_value() + except Exception: + pass + return output + + @classmethod + def from_dict(cls, data: dict): + result_cls = data.pop("cls") + result_cls = import_dotted_class(result_cls) + secret_value = data.pop("_value", None) + + for key, value in data.items(): + if isinstance(value, dict) and (import_str := value.get("import_str")): + data[key] = import_dotted_class(import_str) + + secret = result_cls(**data) + secret._value = secret_value + return secret + + @property + def value(self) -> Optional[str]: + if self._value is not None: + return self._value + else: + return self.get_value() + + def get_value(self) -> Optional[str]: + raise NotImplementedError() + + +@dataclass(kw_only=True) +class EnvironmentSecret(BaseSecret): + type = "env-secret" + agent_message_policy: PolicyRef = Redact + subkernel_message_policy: PolicyRef = Redact + ui_message_policy: PolicyRef = Redact + + name: str + + def get_value(self): + return os.environ.get(self.name) + + +@dataclass(kw_only=True) +class SystemEnvironmentSecret(EnvironmentSecret): + type = "system-env-secret" + subkernel_message_policy: PolicyRef = Redact + ui_message_policy: PolicyRef = Redact + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Remove + + def get_value(self): + import os + return os.environ.get(self.name) + + +@dataclass(kw_only=True) +class UserEnvironmentSecret(EnvironmentSecret): + type = "user-env-secret" + subkernel_message_policy: PolicyRef = Allow + ui_message_policy: PolicyRef = Allow + beaker_kernel_environment_policy: PolicyRef = Remove + subkernel_environment_policy: PolicyRef = Allow + + # user: Optional[BeakerUser] = None + + def get_value(self): + return "USER ENV VALUE" + + +@dataclass(kw_only=True) +class SkillSecret(BaseSecret): + type: str = "skill-secret" + skill_name: str + name: str + default_value: Optional[str] + + def get_value(self): + return os.environ.get(self.name) + + +@dataclass(kw_only=True) +class BeakerConfigProviderSecret(BaseSecret): + OVERRIDE_KEY: ClassVar[str] = "__OVERRIDE__" + + type: str = "config-provider-secret" + subkernel_message_policy: PolicyRef = Remove + ui_message_policy: PolicyRef = Redact + agent_message_policy: PolicyRef = Redact + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Remove + + provider_name: str + + def get_value(self): + from beaker_notebook.lib.config import config as beaker_config + if self.provider_name == self.OVERRIDE_KEY: + return beaker_config.llm_service_token + else: + return beaker_config.provider.get(self.provider_name, {}).get("api_key", None) + + +def is_env_secret(secret: BaseSecret) -> TypeGuard[EnvironmentSecret]: + return isinstance(secret, EnvironmentSecret) + + +def is_system_env_secret(secret: BaseSecret) -> TypeGuard[SystemEnvironmentSecret]: + return isinstance(secret, SystemEnvironmentSecret) + + +def is_user_env_secret(secret: BaseSecret) -> TypeGuard[UserEnvironmentSecret]: + return isinstance(secret, UserEnvironmentSecret) diff --git a/src/beaker_notebook/services/secrets/validations.py b/src/beaker_notebook/lib/secrets/validations.py similarity index 100% rename from src/beaker_notebook/services/secrets/validations.py rename to src/beaker_notebook/lib/secrets/validations.py diff --git a/src/beaker_notebook/services/auth/__init__.py b/src/beaker_notebook/services/auth/__init__.py index f4b4760f..6b82d87a 100644 --- a/src/beaker_notebook/services/auth/__init__.py +++ b/src/beaker_notebook/services/auth/__init__.py @@ -16,7 +16,7 @@ from jupyter_server.services.config.manager import ConfigManager if TYPE_CHECKING: - from beaker_notebook.services.secrets.types import BaseSecret + from beaker_notebook.services.secrets.app_secrets import BaseSecret current_user = contextvars.ContextVar("current_user", default=None) current_request = contextvars.ContextVar("current_request", default=None) diff --git a/src/beaker_notebook/services/secrets/app_secrets.py b/src/beaker_notebook/services/secrets/app_secrets.py new file mode 100644 index 00000000..7eae51a6 --- /dev/null +++ b/src/beaker_notebook/services/secrets/app_secrets.py @@ -0,0 +1,65 @@ +import inspect +import weakref +from dataclasses import dataclass +from typing import TYPE_CHECKING, ClassVar, TypeGuard + +from traitlets import HasTraits +from traitlets.config import Application, Configurable, SingletonConfigurable, LoggingConfigurable + +from beaker_notebook.services.auth import BeakerUser +from beaker_notebook.lib.secrets.policies import Allow, Redact, Remove +from beaker_notebook.lib.secrets.secret_types import BaseSecret, PolicyRef + +if TYPE_CHECKING: + pass + + +@dataclass(kw_only=True) +class AppTraitSecret(BaseSecret): + type: str = "app-trait-secret" + subkernel_message_policy: PolicyRef = Remove + ui_message_policy: PolicyRef = Redact + agent_message_policy: PolicyRef = Redact + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Remove + + _index: ClassVar[dict[str, Configurable]] = {} + config_str: str + + def _update_index(self): + # Fetch the global singleton of the app instance + app = Application.instance() + seen, queue = set(), [app] + while queue: + obj = queue.pop() + if id(obj) in seen: + continue + seen.add(id(obj)) + for klass in type(obj).__mro__: + if klass in (LoggingConfigurable, Configurable, SingletonConfigurable): + break + self._index.setdefault(klass.__name__, obj) + for _, value in inspect.getmembers_static(obj, lambda member: isinstance(member, Configurable)): + if isinstance(value, HasTraits): + queue.append(value) + elif isinstance(value, (list, tuple, dict)): + queue.extend(v for v in (value.values() if isinstance(value, dict) else value) + if isinstance(v, HasTraits)) + + def __post_init__(self, configurable=None): + if configurable is not None: + self._configurable_ref = weakref.ref(configurable) + super().__post_init__() + + def get_value(self): + configurable_name, trait_name = self.config_str.split(".", maxsplit=1) + if configurable_name not in self._index: + self._update_index() + configurable = self._index.get(configurable_name, None) + if configurable is None: + raise ValueError(f"<{self.__class__.__name__}: name='{self.config_str}'> is not available as '{configurable_name}' cannot be located.") + return getattr(configurable, trait_name, None) + + +def is_app_trait_secret(secret: BaseSecret) -> TypeGuard[AppTraitSecret]: + return isinstance(secret, AppTraitSecret) \ No newline at end of file diff --git a/src/beaker_notebook/services/secrets/handlers.py b/src/beaker_notebook/services/secrets/handlers.py new file mode 100644 index 00000000..f5860221 --- /dev/null +++ b/src/beaker_notebook/services/secrets/handlers.py @@ -0,0 +1,40 @@ +import typing + +from beaker_notebook.lib.secrets.secret_types import BaseSecret +from beaker_notebook.services.secrets.app_secrets import AppTraitSecret +from beaker_notebook.services import ServiceApi, ServiceApiHandler, HTTPError + +if typing.TYPE_CHECKING: + from .manager import BeakerSecretsManager + + +class SecretsApi(ServiceApi): + prefix = r"secrets" + + class ContextInfo(ServiceApiHandler): + pattern = r"(?P[\w_-]+)?" + + @staticmethod + def _valid_secret(secret: BaseSecret) -> bool: + # TODO: Stub to be filled out once boundaries are more clear + return True + + @property + def secrets_manager(self) -> "BeakerSecretsManager": + secrets_manager = getattr(self.serverapp, "secrets_manager", None) + if secrets_manager is None: + raise HTTPError(404, "Secrets manager not found") + return secrets_manager + + async def get(self, session=None): + from dataclasses import asdict + user = self.current_user + + secrets = await self.secrets_manager.get_kernel_secrets_for_user(user) + # First we filter to only secrets that might apply, then we filter those to secrets that have a + # replaceable value. E.g. get rid of empty strings, Nones, etc. + secret_dicts = [secret.to_dict(with_value=True) for secret in secrets if self._valid_secret(secret)] + filtered_output = [secret_dict for secret_dict in secret_dicts if secret_dict.get("_value", None)] + + self.write(filtered_output) + diff --git a/src/beaker_notebook/services/secrets/manager.py b/src/beaker_notebook/services/secrets/manager.py index bd874a1e..1668d7f5 100644 --- a/src/beaker_notebook/services/secrets/manager.py +++ b/src/beaker_notebook/services/secrets/manager.py @@ -1,19 +1,21 @@ import copy import inspect import os +from itertools import chain from dataclasses import dataclass, is_dataclass, asdict -from typing import TYPE_CHECKING, Any, Collection, Literal, TypeAlias +from typing import TYPE_CHECKING, Any, Collection, Literal, TypeAlias, Optional import traitlets from traitlets import Type, default, HasTraits from traitlets.config.configurable import Configurable, LoggingConfigurable from traitlets.utils.importstring import import_item -from beaker_notebook.services.secrets.policies import PolicyTypes, BasePolicy, Allow, Redact, Remove -from beaker_notebook.services.secrets.types import ( - BaseSecret, UserEnvironmentSecret, SystemEnvironmentSecret, AppTraitSecret, +from beaker_notebook.lib.secrets.policies import PolicyTypes, BasePolicy, Allow, Redact, Remove +from beaker_notebook.lib.secrets.secret_types import ( + BaseSecret, UserEnvironmentSecret, SystemEnvironmentSecret, BeakerConfigProviderSecret, is_env_secret, is_system_env_secret, is_user_env_secret ) +from beaker_notebook.services.secrets.app_secrets import AppTraitSecret if TYPE_CHECKING: @@ -46,7 +48,6 @@ def index_configurables(root): class BeakerSecretsManager(LoggingConfigurable): parent: "BaseBeakerApp" _secrets: list[BaseSecret] - _index: dict[str, Configurable] app_trait_secrets: list[str] = traitlets.List( trait=traitlets.Unicode, @@ -54,7 +55,6 @@ class BeakerSecretsManager(LoggingConfigurable): default_value=[ "Application.cookie_secret", "IdentityProvider.token", - "NotebookNotary.secret", "GatewayClient.auth_token", ] ) @@ -113,6 +113,9 @@ class BeakerSecretsManager(LoggingConfigurable): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._secrets = [] + from .handlers import SecretsApi + # TODO: Defining handlers here is clunky, There should be a more elegant way. + self.parent.handlers.extend(SecretsApi.handlers) def add_secret(self, secret: BaseSecret): if secret in self._secrets: @@ -155,7 +158,7 @@ async def sanitize_subkernel_envionment_vars(self, user: "BeakerUser", env: dict result[secret.name] = await policy.sanitize(secret, result[secret.name]) if is_user_env_secret(secret) and secret.name not in result: # Add the user secret to the environment if it doesn't exist TODO: Determine if this is the right thing - result[secret.name] = secret.get_value() + result[secret.name] = secret.value return result @@ -176,25 +179,37 @@ def _secret_env_vars(self) -> list[str]: async def collect_system_secrets(self, app: "BaseBeakerApp") -> list[BaseSecret]: - system_secrets = [] + from beaker_notebook.lib.config import config as beaker_config # Environment secrets from configuration discovered_system_env_secrets = [SystemEnvironmentSecret(name=env_name) for env_name in self._secret_env_vars()] extra_system_envs = [SystemEnvironmentSecret(name=env_name) for env_name in self.extra_system_env_vars] extra_user_envs = [UserEnvironmentSecret(name=env_name) for env_name in self.extra_user_env_vars] - configurable_index = index_configurables(app) - app_trait_secrets = [] - for config_string in self.app_trait_secrets + self.extra_app_trait_secrets: - configurable_name, trait_name = config_string.split(".", maxsplit=1) - configurable = configurable_index[configurable_name] - app_trait_secrets.append(AppTraitSecret(configurable=configurable, trait_name=trait_name)) + # Secrets from server app traits/configuration + app_trait_secrets = [AppTraitSecret(config_str=config_st) for config_st in self.app_trait_secrets + self.extra_app_trait_secrets] - # TODO: Extract secrets from beaker config + # LLM secrets from Beaker Config + beaker_config_secrets = [ + BeakerConfigProviderSecret(provider_name=provider_name) + for provider_name, provider_config in beaker_config.providers.items() + if provider_config["api_key"] + ] + if beaker_config.llm_service_token: + beaker_config_secrets.append(BeakerConfigProviderSecret(provider_name=BeakerConfigProviderSecret.OVERRIDE_KEY)) + + system_secrets = list(chain( + discovered_system_env_secrets, + extra_system_envs, + extra_user_envs, + app_trait_secrets, + beaker_config_secrets, + )) - # TODO: Clean this up - system_secrets.extend(discovered_system_env_secrets) - system_secrets.extend(extra_system_envs) - system_secrets.extend(extra_user_envs) - system_secrets.extend(app_trait_secrets) return system_secrets + + + async def get_kernel_secrets_for_user(self, user: "Optional[BeakerUser]"): + # TODO: filter secrets and add in user's secrets if any + # set to push = { secrets owned by this session's scope } ∩ { secrets with any non-Allow message policy }. + return self.secrets diff --git a/src/beaker_notebook/services/secrets/policies.py b/src/beaker_notebook/services/secrets/policies.py deleted file mode 100644 index 64fcb2c6..00000000 --- a/src/beaker_notebook/services/secrets/policies.py +++ /dev/null @@ -1,129 +0,0 @@ -import asyncio -import copy -from dataclasses import dataclass, field, is_dataclass, asdict -from typing import TYPE_CHECKING, Any, Awaitable, Callable, ClassVar, Collection, Literal, TypeAlias - -from beaker_notebook.services.secrets.validations import Validation, not_in - -if TYPE_CHECKING: - from .types import BaseSecret - -PolicyTypes: TypeAlias = Literal["redact", "remove", "last4", "allow"] - -@dataclass -class BasePolicy: - type: PolicyTypes - validations: Collection[Validation] = field(default_factory=lambda: [not_in,]) - - async def replacement(self, secret_str: str): - raise NotImplementedError(f"Policy '{BasePolicy}' does not define a sanitize function") - - - async def _sanitize_str(self, secret: "BaseSecret", content: str) -> str: - secret_str = secret.get_value() - replacement = await self.replacement(secret_str) - sanitized = content.replace(secret_str, replacement) - - await asyncio.gather(*[ - validation_func(secret_str, sanitized) for validation_func in self.validations - ]) - - return sanitized - - async def _sanitize_list(self, secret: "BaseSecret", content: list) -> list: - # Walk a list, sanitizing each item - for idx, value in enumerate(content): - match value: - case str(): - content[idx] = await self._sanitize_str(secret, value) - case dict(): - content[idx] = await self._sanitize_dict(secret, value) - case list(): - content[idx] = await self._sanitize_list(secret, value) - case bytes(): - content[idx] = (await self._sanitize_str(secret, value.decode())).encode() - case _: - pass - return content - - async def _sanitize_dict(self, secret: "BaseSecret", content: dict) -> dict: - # Walk dict, looking for strings and do a sanitization on each string - for key, value in content.items(): - match value: - case str(): - content[key] = await self._sanitize_str(secret, value) - case dict(): - content[key] = await self._sanitize_dict(secret, value) - case list(): - content[key] = await self._sanitize_list(secret, value) - case bytes(): - content[key] = (await self._sanitize_str(secret, value.decode())).encode() - case _: - pass - - async def sanitize(self, secret: "BaseSecret", content: str|dict): - target = copy.deepcopy(content) - match target: - case str() | bytes(): - return await self._sanitize_str(secret, target) - case dict(): - return await self._sanitize_dict(secret, target) - case list(): - return await self._sanitize_list(secret, target) - case _: - return content - -@dataclass -class Allow(BasePolicy): - """ - This policy does not modify the content in any way. - Only for use in high-trust situations or debugging - """ - - type: PolicyTypes = "allow" - - # TODO: Add warning on initialization about this being insecure? - - async def sanitize(self, secret, content): - return content - - -@dataclass -class Redact(BasePolicy): - type: PolicyTypes = "redact" - - async def replacement(self, secret_str): - return "#" * len(secret_str) - - -@dataclass -class Remove(BasePolicy): - type: PolicyTypes = "remove" - - async def _sanitize_list(self, secret, content): - secret_value = secret.get_value() - while secret_value in content: - content.remove(secret_value) - return content - - async def _sanitize_dict(self, secret, content): - name = getattr(secret, name, None) - if name and name in content: - content.pop(name) - return content - - async def replacement(self, secret_str): - return "" - - -@dataclass -class Last4(BasePolicy): - type: PolicyTypes = "last4" - - async def replacement(self, secret_str): - if len(secret_str) < 6: - # Returning last four chars could is too much information for such a short secret. - return "####" - return f"###{secret_str[:-4]}" - - diff --git a/src/beaker_notebook/services/secrets/types.py b/src/beaker_notebook/services/secrets/types.py deleted file mode 100644 index 39dc3a4d..00000000 --- a/src/beaker_notebook/services/secrets/types.py +++ /dev/null @@ -1,134 +0,0 @@ -import os -import weakref -from dataclasses import dataclass, InitVar, field, is_dataclass, asdict -from typing import TYPE_CHECKING, Any, ClassVar, Literal, Optional, TypeAlias, Type, TypeIs - -from traitlets.config import Application, Configurable - -from beaker_notebook.services.auth import BeakerUser -from beaker_notebook.services.secrets.policies import BasePolicy, Allow, Redact, Remove - -if TYPE_CHECKING: - pass - - -PolicyRef: TypeAlias = BasePolicy | Type[BasePolicy] - - -@dataclass(kw_only=True) -class BaseSecret: - type: ClassVar[str] = "base-secret" - - subkernel_message_policy: PolicyRef - ui_message_policy: PolicyRef - agent_message_policy: PolicyRef - beaker_kernel_environment_policy: PolicyRef - subkernel_environment_policy: PolicyRef - - def __post_init__(self, *args, **kwargs): - for name, field_info in self.__dataclass_fields__.items(): - if field_info.type == PolicyRef: - value = getattr(self, name) - setattr(self, name, value()) - - def get_value(self) -> Optional[str]: - raise NotImplementedError() - - -@dataclass(kw_only=True) -class EnvironmentSecret(BaseSecret): - type = "env-secret" - agent_message_policy: PolicyRef = Redact - subkernel_message_policy: PolicyRef = Redact - ui_message_policy: PolicyRef = Redact - - name: str - - def get_value(self): - return os.environ.get(self.name) - - -@dataclass(kw_only=True) -class SystemEnvironmentSecret(EnvironmentSecret): - subkernel_message_policy: PolicyRef = Redact - ui_message_policy: PolicyRef = Redact - beaker_kernel_environment_policy: PolicyRef = Allow - subkernel_environment_policy: PolicyRef = Remove - - def get_value(self): - import os - return os.environ.get(self.name) - - -@dataclass(kw_only=True) -class UserEnvironmentSecret(EnvironmentSecret): - subkernel_message_policy: PolicyRef = Allow - ui_message_policy: PolicyRef = Allow - beaker_kernel_environment_policy: PolicyRef = Remove - subkernel_environment_policy: PolicyRef = Allow - - user: Optional[BeakerUser] = None - - def get_value(self): - return "USER ENV VALUE" - - -@dataclass(kw_only=True) -class SkillSecret(BaseSecret): - type: str = "skill-secret" - skill_name: str - name: str - default_value: Optional[str] - - def get_value(self): - return os.environ.get(self.name) - - -@dataclass(kw_only=True) -class AppTraitSecret(BaseSecret): - subkernel_message_policy: PolicyRef = Remove - ui_message_policy: PolicyRef = Redact - agent_message_policy: PolicyRef = Redact - beaker_kernel_environment_policy: PolicyRef = Allow - subkernel_environment_policy: PolicyRef = Remove - - configurable: InitVar["Optional[Configurable]"] = field(default=None, repr=False, compare=False) - _configurable_ref: "weakref.ref[Configurable] | None" = field(default=None, init=False, repr=False, compare=False) - trait_name: str - - def __post_init__(self, configurable=None): - if configurable is not None: - self._configurable_ref = weakref.ref(configurable) - super().__post_init__() - if configurable: - # We only need to validate when we have something to validate against - self._validate(configurable) - - def _validate(self, configurable: Configurable): - from traitlets import Unicode, Bytes - - traits = configurable.trait_names() - if self.trait_name not in traits: - raise ValueError(f"{self.__class__.__name__}: Class {configurable.__class__.__name__} does not contain a trait named '{self.trait_name}'") - trait = configurable.traits().get(self.trait_name) - if not isinstance(trait, (Unicode, Bytes)): - raise ValueError(f"{self.__class__.__name__} targeted traits should point to a string (Unicode or Bytes), not {trait.__class__.__name__}") - - def get_value(self): - configurable = self._configurable_ref and self._configurable_ref() - if configurable is None: - raise ValueError(f"<{self.__class__.__name__}: name='{self.name}'> is not available.") - return getattr(configurable, self.trait_name, None) - - -def is_env_secret(secret: BaseSecret) -> TypeIs[EnvironmentSecret]: - return isinstance(secret, EnvironmentSecret) - -def is_system_env_secret(secret: BaseSecret) -> TypeIs[SystemEnvironmentSecret]: - return isinstance(secret, SystemEnvironmentSecret) - -def is_user_env_secret(secret: BaseSecret) -> TypeIs[UserEnvironmentSecret]: - return isinstance(secret, UserEnvironmentSecret) - -def is_app_trait_secret(secret: BaseSecret) -> TypeIs[AppTraitSecret]: - return isinstance(secret, AppTraitSecret) \ No newline at end of file diff --git a/src/beaker_notebook/services/storage/notebook.py b/src/beaker_notebook/services/storage/notebook.py index 7cfecff5..3cb4412d 100644 --- a/src/beaker_notebook/services/storage/notebook.py +++ b/src/beaker_notebook/services/storage/notebook.py @@ -89,7 +89,6 @@ async def delete_snapshot(self, session_id: str) -> None: class FileNotebookManager(BaseNotebookManager): contents_manager_class = traitlets.Type( - default_value="jupyter_server.services.contents.filemanager.AsyncFileContentsManager", klass=ContentsManager, allow_none=True, config=True, @@ -104,14 +103,22 @@ class FileNotebookManager(BaseNotebookManager): config=True, ) + @traitlets.default("contents_manager_class") + def _default_contents_manager_class(self): + metadata = self.traits()["contents_manager_class"].metadata + metadata["is_default_val"] = True + return AsyncFileContentsManager + @traitlets.default("contents_manager") def _default_contents_manager(self): - if self.contents_manager_class not in (traitlets.Undefined, None, ""): - return self.contents_manager_class(parent=self, **self.contents_manager_params) - if getattr(self.parent, "contents_manager", None): + # If the conents_manager_class trait has not been set by a user or configuration, + # default to using the parent classes' content manager instance. + getattr(self, "contents_manager_class") + metadata = self.traits()["contents_manager_class"].metadata + if metadata.get("is_default_val", False) and getattr(self.parent, "contents_manager", None): return self.parent.contents_manager else: - return AsyncFileContentsManager(parent=self.parent) + return self.contents_manager_class(parent=self.parent) async def get_notebook_info(self, notebook_id: str) -> NotebookInfo: """Retrieve notebook metadata for a given session ID. diff --git a/tests/secrets/__init__.py b/tests/secrets/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/secrets/test_message_scrubbing.py b/tests/secrets/test_message_scrubbing.py new file mode 100644 index 00000000..5ae41b6a --- /dev/null +++ b/tests/secrets/test_message_scrubbing.py @@ -0,0 +1,110 @@ +""" +Tests for secret scrubbing over the wire (multipart Jupyter messages). + +These drive the real ``BeakerKernel.redact_secrets`` flow (via ``run_redact_secrets``): +parse the multipart message, sanitize the decoded payload fields, reserialize. +They assert the security invariants a scrubbed message must uphold: + +- The secret value never appears anywhere on the wire after scrubbing +- The message remains a structurally valid, re-parseable Jupyter message +- Per-channel policy selection is honored by destination (client vs subkernel) +- Secrets containing JSON-escaped characters (quotes/backslashes/non-ascii) + are still scrubbed from the decoded content +- Binary buffers are passed through untouched and never raise on decode +""" + +from beaker_notebook.lib.jupyter_kernel_proxy import JupyterMessage +from beaker_notebook.lib.secrets.policies import Allow, Redact, Remove + +from tests.secrets.util import ( + make_secret, + run_redact_secrets, + signed_parts, + wire_bytes, + decoded_strings, +) + + +# -- redact over the wire -- + + +async def test_redact_scrubs_secret_and_keeps_message_valid(): + secret = make_secret("s3cr3t", subkernel=Redact) + parts = signed_parts({"data": {"text/plain": "k=s3cr3t"}}) + + out = await run_redact_secrets([secret], "subkernel", parts) + + assert b"s3cr3t" not in wire_bytes(out) + # Message is still structurally valid and can be re-parsed / re-signed. + reparsed = JupyterMessage.parse(out) + assert reparsed.content["data"]["text/plain"] == "k=######" + + +# -- remove over the wire -- + + +async def test_remove_scrubs_secret_from_message(): + secret = make_secret("s3cr3t", subkernel=Remove) + parts = signed_parts({"data": {"text/plain": "k=s3cr3t"}}) + + out = await run_redact_secrets([secret], "subkernel", parts) + + assert b"s3cr3t" not in wire_bytes(out) + + +async def test_remove_uses_injected_value_not_get_value(): + """On the kernel side a secret carries an injected value and its server-bound + get_value() may be unavailable; Remove must scrub via the injected value.""" + secret = make_secret("s3cr3t", subkernel=Remove, break_get_value=True) + parts = signed_parts({"data": {"text/plain": "k=s3cr3t"}}) + + out = await run_redact_secrets([secret], "subkernel", parts) + + assert b"s3cr3t" not in wire_bytes(out) + + +# -- per-channel policy selection -- + + +async def test_channel_policies_are_independent(): + # Allow to the UI, Redact to the subkernel. + secret = make_secret("s3cr3t", ui=Allow, subkernel=Redact) + parts = signed_parts({"data": {"text/plain": "k=s3cr3t"}}) + + ui_out = await run_redact_secrets([secret], "client", parts) + sub_out = await run_redact_secrets([secret], "subkernel", parts) + + assert b"s3cr3t" in wire_bytes(ui_out) + assert b"s3cr3t" not in wire_bytes(sub_out) + + +# -- secrets with JSON-escaped characters -- + + +async def test_scrubs_secret_with_json_special_characters(): + secret_value = 'p@ss"w\\ordé' # contains a quote, a backslash, and non-ascii + secret = make_secret(secret_value, subkernel=Redact) + parts = signed_parts({"data": {"text/plain": f"token={secret_value}"}}) + + out = await run_redact_secrets([secret], "subkernel", parts) + + # Must remain a valid message (a naive byte-replace can corrupt the JSON)... + reparsed = JupyterMessage.parse(out) + assert reparsed.header["msg_type"] == "execute_result" + # ...and the secret must not survive in any decoded string field. + leaked = [s for s in decoded_strings(out) if secret_value in s] + assert not leaked, f"secret leaked in decoded content: {leaked}" + + +# -- binary buffers -- + + +async def test_binary_buffers_are_preserved_and_do_not_crash(): + secret = make_secret("s3cr3t", subkernel=Redact) + blob = b"\x89PNG\r\n\x1a\n\xff\xfe\x00binary" # not valid utf-8 + parts = signed_parts({"data": {"text/plain": "k=s3cr3t"}}, buffers=(blob,)) + + out = await run_redact_secrets([secret], "subkernel", parts) + + assert blob in out # buffer survives untouched + assert b"s3cr3t" not in wire_bytes(out) diff --git a/tests/secrets/test_policies.py b/tests/secrets/test_policies.py new file mode 100644 index 00000000..7159de3d --- /dev/null +++ b/tests/secrets/test_policies.py @@ -0,0 +1,187 @@ +""" +Tests for beaker_notebook.lib.secrets policy engine and secret (de)serialization. + +Covers the *correct* behavior of the sanitization policies independent of the +wire/message plumbing (see test_secrets_message_scrubbing.py for that path): + +- Allow leaves content untouched +- Redact replaces the secret with equal-length masking across str/bytes/list/dict +- Redact preserves surrounding structure (nested dicts/lists are not dropped) +- Remove eliminates the secret value from str / bytes-frame list / dict content +- Last4 reveals at most the trailing four characters, never the prefix +- not_in validation raises iff the secret survives +- BaseSecret.value falls back to get_value() when no explicit value was injected +- to_dict / from_dict round-trips the class, per-channel policies, and value + +Several assertions encode behavior the current implementation does not yet +satisfy (Remove over a frame list, nested-dict preservation, Last4 masking, +the .value fallback); those are deliberate and describe the target behavior. +""" + +import json + +import pytest + +from beaker_notebook.lib.secrets.policies import Allow, Redact, Remove, Last4 +from beaker_notebook.lib.secrets.secret_types import BaseSecret, SystemEnvironmentSecret +from beaker_notebook.lib.secrets.validations import not_in, SecretValidationError + +from tests.secrets.util import make_secret + + +# -- Allow -- + + +async def test_allow_leaves_string_untouched(): + secret = make_secret("s3cr3t") + assert await Allow().sanitize(secret, "token=s3cr3t;") == "token=s3cr3t;" + + +async def test_allow_leaves_container_untouched(): + secret = make_secret("s3cr3t") + payload = {"a": ["s3cr3t"], "b": {"c": "s3cr3t"}} + assert await Allow().sanitize(secret, payload) == payload + + +# -- Redact -- + + +async def test_redact_string_masks_equal_length(): + secret = make_secret("s3cr3t") + out = await Redact().sanitize(secret, "token=s3cr3t;") + assert out == "token=######;" + assert "s3cr3t" not in out + + +async def test_redact_bytes(): + secret = make_secret("s3cr3t") + out = await Redact().sanitize(secret, b"key=s3cr3t") + assert out == b"key=######" + + +async def test_redact_nested_dict_preserves_structure(): + secret = make_secret("s3cr3t") + out = await Redact().sanitize( + secret, {"outer": {"inner": "x s3cr3t"}, "keep": "ok"} + ) + # The dict (and its nested dict) must survive as a dict, not be dropped/nulled. + assert isinstance(out, dict) + assert out["keep"] == "ok" + assert isinstance(out["outer"], dict) + assert "s3cr3t" not in json.dumps(out) + + +async def test_redact_list_preserves_elements(): + secret = make_secret("s3cr3t") + out = await Redact().sanitize(secret, ["x s3cr3t", {"k": "s3cr3t"}]) + assert out[0] == "x ######" + assert out[1] == {"k": "######"} + + +# -- Remove -- + + +async def test_remove_string_eliminates_secret(): + secret = make_secret("s3cr3t") + out = await Remove().sanitize(secret, "a s3cr3t b") + assert "s3cr3t" not in out + + +async def test_remove_bytes_frame_list_eliminates_secret(): + """Remove applied to a multipart-style list of byte frames (the wire case) + must eliminate the secret from every frame.""" + secret = make_secret("s3cr3t") + frames = [b'{"code":"key=s3cr3t"}', b""] + out = await Remove().sanitize(secret, frames) + assert all(b"s3cr3t" not in frame for frame in out) + + +async def test_remove_dict_eliminates_secret_value(): + secret = make_secret("s3cr3t") + out = await Remove().sanitize(secret, {"code": "key=s3cr3t"}) + assert isinstance(out, dict) + assert "s3cr3t" not in json.dumps(out) + + +# -- Last4 -- + + +async def test_last4_reveals_only_trailing_four(): + secret = make_secret("abcdefghij") # length 10 + out = await Last4().sanitize(secret, "abcdefghij") + assert "abcdefghij" not in out + # Only the final four characters may be revealed; the prefix must not leak. + assert out.endswith("ghij") + assert "abcdef" not in out + assert out == "######ghij" + + +async def test_last4_short_secret_fully_masked(): + secret = make_secret("abcd") # length < 6 + out = await Last4().sanitize(secret, "abcd") + assert out == "####" + + +# -- validations -- + + +async def test_not_in_raises_when_secret_survives(): + with pytest.raises(SecretValidationError): + await not_in("s3cr3t", "leftover s3cr3t here") + + +async def test_not_in_passes_when_absent(): + await not_in("s3cr3t", "nothing sensitive here") + + +# -- value property -- + + +async def test_value_falls_back_to_get_value(): + """A secret constructed without an injected value must resolve via get_value().""" + secret = make_secret("resolved-val", set_value=False) + assert secret.value == "resolved-val" + + +# -- serialization round-trip -- + + +def test_to_dict_from_dict_round_trip(monkeypatch): + monkeypatch.setenv("MY_TOKEN", "abc123") + secret = SystemEnvironmentSecret(name="MY_TOKEN") + + data = secret.to_dict(with_value=True) + assert data["_value"] == "abc123" + + restored = BaseSecret.from_dict(dict(data)) + assert isinstance(restored, SystemEnvironmentSecret) + assert restored.name == "MY_TOKEN" + assert restored.value == "abc123" + # Per-channel policies survive as instantiated policy objects. + assert isinstance(restored.subkernel_message_policy, Redact) + assert isinstance(restored.ui_message_policy, Redact) + + +# -- misc + +async def test_dict_key_scrubbed(): + secret1 = make_secret("s3cr3t") + secret2 = make_secret("0th3r") + payload = { + "token": "s3cr3t;", + "s3cr3t": "abc", + "my-s3cr3t": "def", + "0th3r": "ghi", + "my-0th3r": "jkl", + } + out1 = await Remove().sanitize(secret1, payload) + out2 = await Remove().sanitize(secret2, out1) + + assert "s3cr3t" not in out2 + assert out2 == { + "token": ";", + "__FAKE-SECRET:0__": "abc", + "my-__FAKE-SECRET:1__": "def", + "__FAKE-SECRET:2__": "ghi", + "my-__FAKE-SECRET:3__": "jkl", + } diff --git a/tests/secrets/util.py b/tests/secrets/util.py new file mode 100644 index 00000000..bf5c0aeb --- /dev/null +++ b/tests/secrets/util.py @@ -0,0 +1,133 @@ +"""Shared fixtures and helpers for the beaker_notebook secrets test suite. + +Provides a fully-controllable concrete secret (``FakeSecret`` / ``make_secret``) +and helpers for building and inspecting signed multipart Jupyter messages. +""" + +import json +from dataclasses import dataclass +from types import SimpleNamespace + +from beaker_notebook.lib.jupyter_kernel_proxy import JupyterMessage +from beaker_notebook.lib.secrets.policies import Allow +from beaker_notebook.lib.secrets.secret_types import BaseSecret, PolicyRef + + +# -- secrets -- + + +@dataclass(kw_only=True) +class FakeSecret(BaseSecret): + """A concrete secret with fully-controllable policies and value for tests.""" + type = "fake-secret" + + subkernel_message_policy: PolicyRef = Allow + ui_message_policy: PolicyRef = Allow + agent_message_policy: PolicyRef = Allow + beaker_kernel_environment_policy: PolicyRef = Allow + subkernel_environment_policy: PolicyRef = Allow + + name: str = "SECRET_KEY" + source_value: str = "s3cr3t-value" + break_get_value: bool = False + + def get_value(self): + if self.break_get_value: + # Simulates a server-bound secret (app trait, config provider) whose + # resolution cannot run in the kernel runtime -- only the injected + # _value is usable there. + raise RuntimeError("resolution unavailable in kernel runtime") + return self.source_value + + +def make_secret( + value="s3cr3t", + *, + subkernel=Allow, + ui=Allow, + agent=Allow, + name="SECRET_KEY", + set_value=True, + break_get_value=False, +) -> FakeSecret: + secret = FakeSecret( + subkernel_message_policy=subkernel, + ui_message_policy=ui, + agent_message_policy=agent, + beaker_kernel_environment_policy=Allow, + subkernel_environment_policy=Allow, + name=name, + source_value=value, + break_get_value=break_get_value, + ) + if set_value: + # Mirror the kernel-side path where the resolved value is injected + # directly (BaseSecret.from_dict) rather than looked up locally. + secret._value = value + return secret + + +# -- multipart messages -- + + +def signed_parts(content, *, key=b"testkey", identities=(b"id",), buffers=()): + """Build a signed multipart message (list of byte frames) carrying `content`.""" + header = {"msg_id": "m1", "msg_type": "execute_result", "session": "sess"} + parent = {"msg_id": "p1"} + metadata = {} + raw = ( + list(identities) + + [ + JupyterMessage.DELIMITER, + b"placeholder-signature", + json.dumps(header).encode(), + json.dumps(parent).encode(), + json.dumps(metadata).encode(), + json.dumps(content).encode(), + ] + + list(buffers) + ) + return JupyterMessage.parse(raw).sign_using(key).parts + + +async def run_redact_secrets(secrets, destination, parts): + """Drive the real ``BeakerKernel.redact_secrets`` flow with lightweight fakes. + + Exercises the actual scrubbing path (parse -> sanitize decoded fields -> + reserialize) rather than calling a policy on raw frames directly, so tests + reflect what the kernel really does on the wire. ``destination`` is what the + proxy's ``get_destination`` would return ("client" or "subkernel"). + """ + from beaker_notebook.kernel import BeakerKernel + + fake_self = SimpleNamespace(secrets=list(secrets)) + fake_server = SimpleNamespace(get_destination=lambda _stream: destination) + return await BeakerKernel.redact_secrets(fake_self, fake_server, target_stream=None, data=parts) + + +def wire_bytes(parts) -> bytes: + return b"".join(parts) + + +def _iter_strings(obj): + if isinstance(obj, str): + yield obj + elif isinstance(obj, dict): + for key, value in obj.items(): + yield from _iter_strings(key) + yield from _iter_strings(value) + elif isinstance(obj, (list, tuple)): + for value in obj: + yield from _iter_strings(value) + + +def decoded_strings(parts) -> list[str]: + """Every decoded (unescaped) string leaf in the message fields. + + Checking membership against these -- not a re-serialized JSON blob, which + would re-escape the secret and hide a leak -- is what catches a secret that + survived scrubbing inside a JSON-escaped payload. + """ + msg = JupyterMessage.parse(parts) + fields = [msg.header, msg.parent_header, msg.metadata, msg.content] + return list(_iter_strings(fields)) diff --git a/tests/test_generate_config.py b/tests/test_generate_config.py new file mode 100644 index 00000000..4bd93674 --- /dev/null +++ b/tests/test_generate_config.py @@ -0,0 +1,53 @@ +""" +Tests for the ``beaker server generate-config`` CLI command. + +These exercise the command end-to-end via click's ``CliRunner``, verifying that +a config file is actually written to disk with the expected name and content. +""" + +from pathlib import Path + +from click.testing import CliRunner + +from beaker_notebook.cli.server import server + + +def test_generate_config_default_writes_config_file(): + """With no arguments, a ``beaker_config.py`` file is written.""" + runner = CliRunner() + with runner.isolated_filesystem(): + result = runner.invoke(server, ["generate-config"]) + + assert result.exit_code == 0, result.output + + config_path = Path("beaker_config.py") + assert config_path.exists() + + content = config_path.read_text(encoding="utf-8") + assert "Beaker Notebook Service Configuration File" in content + assert "c = get_config()" in content + + +def test_generate_config_custom_file_option(): + """The ``--file`` option controls the output path.""" + runner = CliRunner() + with runner.isolated_filesystem(): + result = runner.invoke(server, ["generate-config", "--file", "custom_config.py"]) + + assert result.exit_code == 0, result.output + + assert Path("custom_config.py").exists() + assert not Path("beaker_config.py").exists() + + +def test_generate_config_for_server_type(): + """Passing a server type names the file after that app's slug.""" + runner = CliRunner() + with runner.isolated_filesystem(): + result = runner.invoke(server, ["generate-config", "server"]) + + assert result.exit_code == 0, result.output + + config_path = Path("beaker_server_config.py") + assert config_path.exists() + assert "c = get_config()" in config_path.read_text(encoding="utf-8") From 1ea7d2a461167d643c113868d8ad26de5fd40a2d Mon Sep 17 00:00:00 2001 From: Matthew Printz Date: Tue, 28 Jul 2026 17:46:04 -0600 Subject: [PATCH 4/4] Fix policy mixup --- src/beaker_notebook/services/secrets/manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/beaker_notebook/services/secrets/manager.py b/src/beaker_notebook/services/secrets/manager.py index 1668d7f5..1098198c 100644 --- a/src/beaker_notebook/services/secrets/manager.py +++ b/src/beaker_notebook/services/secrets/manager.py @@ -133,7 +133,7 @@ def secrets(self) -> list[BaseSecret]: async def sanitize_kernel_environment_vars(self, env: dict) -> dict[str, str]: result = copy.copy(env) for secret in self._secrets: - policy = secret.subkernel_environment_policy + policy = secret.beaker_kernel_environment_policy if is_env_secret(secret) and secret.name in result: if isinstance(policy, Allow): continue