diff --git a/src/beaker_notebook/app/base.py b/src/beaker_notebook/app/base.py index 93dfaed8..b7616d50 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,12 @@ 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 + secrets_manager = BeakerSecretsManager(parent=self) + return secrets_manager + @traitlets.default("config_file_name") def _default_config_file_name(self): if self.app_slug: @@ -176,6 +192,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/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/lib/secrets/validations.py b/src/beaker_notebook/lib/secrets/validations.py new file mode 100644 index 00000000..4fb76b05 --- /dev/null +++ b/src/beaker_notebook/lib/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() diff --git a/src/beaker_notebook/services/auth/__init__.py b/src/beaker_notebook/services/auth/__init__.py index aa69ee36..6b82d87a 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.app_secrets 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/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 new file mode 100644 index 00000000..1098198c --- /dev/null +++ b/src/beaker_notebook/services/secrets/manager.py @@ -0,0 +1,215 @@ +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, 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.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: + from beaker_notebook.app.base import BaseBeakerApp + from beaker_notebook.lib.context import BeakerContext + 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] + + app_trait_secrets: list[str] = traitlets.List( + trait=traitlets.Unicode, + config=True, + default_value=[ + "Application.cookie_secret", + "IdentityProvider.token", + "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) + 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: + return + self._secrets.append(secret) + + 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: + policy = secret.beaker_kernel_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.value + return result + + + def _secret_env_vars(self) -> list[str]: + sensitive_env_keys = set() + # 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 ( + 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]: + 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] + + # 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] + + # 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, + )) + + 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/storage/notebook.py b/src/beaker_notebook/services/storage/notebook.py index 78d7ebd5..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=None, 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")