Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 21 additions & 1 deletion src/beaker_notebook/app/base.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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")

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
51 changes: 50 additions & 1 deletion src/beaker_notebook/kernel.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import asyncio
import contextvars
import copy
import inspect
import json
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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()
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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")):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
19 changes: 11 additions & 8 deletions src/beaker_notebook/lib/jupyter_kernel_proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
):
Expand All @@ -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
Expand Down Expand Up @@ -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")
Expand Down
Empty file.
Loading