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
4 changes: 2 additions & 2 deletions sdcm/remote/docker_cmd_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from io import BytesIO
from typing import TYPE_CHECKING
from shlex import quote
from pathlib import Path
from pathlib import Path, PurePosixPath
import shutil

from invoke.runners import Result
Expand Down Expand Up @@ -215,7 +215,7 @@ def send_files(
self.run(f"rm -rf {quote(dst)}", ignore_status=True, verbose=False)

tar_stream = self._create_tar_stream(src, dst)
dst_path = Path(dst)
dst_path = PurePosixPath(dst)
extraction_dir = dst if dst.endswith("/") or not dst_path.suffix else str(dst_path.parent)
self.run(f"mkdir -p {quote(extraction_dir)}", ignore_status=True, verbose=False)
container.put_archive(path=extraction_dir, data=tar_stream)
Expand Down
5 changes: 5 additions & 0 deletions sdcm/remote/remote_file.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,11 @@ def remote_file(
LOGGER.debug("New content of `%s':\n%s", remote_path, content)

remote_tempfile = remoter.run("mktemp").stdout.strip()
if not remote_tempfile:
raise RuntimeError(
f"'mktemp' on {remoter.hostname} returned an empty path; "
f"cannot update '{remote_path}' without a valid remote temporary file"
)
remote_tempfile_move_cmd = shell_script_cmd(f"""\
cat '{remote_tempfile}' > '{remote_path}'
rm '{remote_tempfile}'
Expand Down
47 changes: 47 additions & 0 deletions sdcm/utils/docker_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
#
# Copyright (c) 2020 ScyllaDB

import io
import os
import re
import logging
Expand All @@ -22,10 +23,13 @@
import itertools

import docker
from docker.api.client import APIClient
from docker.errors import DockerException, NotFound, ImageNotFound, NullResource, BuildError
from docker.models.images import Image
from docker.models.containers import Container
from docker.utils import socket as docker_socket
from docker.utils.json_stream import json_stream
from docker.utils.socket import consume_socket_output, demux_adaptor, frames_iter

from sdcm.remote import LOCALRUNNER
from sdcm.remote.base import CommandRunner
Expand Down Expand Up @@ -71,8 +75,51 @@ class ContainerAlreadyRegistered(DockerException):
pass


# docker-py frame reader must accept an io.BufferedReader from BufferedStreamAPIClient.
# Save the original docker_socket.read before patching it, so this wrapper can call the
# original implementation directly and avoid recursion.
_unpatched_socket_read = docker_socket.read


def _read_through_buffer(sock, n=4096):
if isinstance(sock, io.BufferedReader):
return sock.read(n)
return _unpatched_socket_read(sock, n)


docker_socket.read = _read_through_buffer


class BufferedStreamAPIClient(APIClient):
"""API client that reads hijacked streams from the buffered HTTP reader.

This avoids losing the first exec/attach frame when it arrives together
with the '101 UPGRADED' headers (SCT-952, docker/docker-py#3332).
"""

def _read_from_socket(self, response, stream, tty=True, demux=False):
buffered_reader = response.raw._fp.fp
if not isinstance(buffered_reader, io.BufferedReader):
return super()._read_from_socket(response, stream, tty=tty, demux=demux)

self._raise_for_status(response)
frames = frames_iter(buffered_reader, tty)
frames = (demux_adaptor(*frame) for frame in frames) if demux else (data for _, data in frames)

if stream:
return frames

try:
return consume_socket_output(frames, demux=demux)
finally:
response.close()


# TODO: remove this wrapper when migrated to Docker Python module completely.
class DockerClient(docker.DockerClient):
def __init__(self, *args, **kwargs):
self.api = BufferedStreamAPIClient(*args, **kwargs)

def __call__(self, cmd, timeout=10):
deprecation("consider to use Docker Python module instead of using Docker CLI commands")

Expand Down
6 changes: 4 additions & 2 deletions unit_tests/unit/test_docker_cmd_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,11 +165,13 @@ def test_send_files(self, mock_retrying):
mock_create_tar.return_value = mock_tar_stream
mock_run.return_value = MagicMock(ok=True)

result = self.runner.send_files(temp_file.name, "/dest/path")
result = self.runner.send_files(temp_file.name, "/etc/scylla/scylla.yaml")

assert result
mock_container.put_archive.assert_called_once()
mock_create_tar.assert_called_once_with(temp_file.name, "/dest/path")
_, put_archive_kwargs = mock_container.put_archive.call_args
assert put_archive_kwargs["path"] == "/etc/scylla"
mock_create_tar.assert_called_once_with(temp_file.name, "/etc/scylla/scylla.yaml")

@patch("sdcm.remote.docker_cmd_runner.retrying")
def test_receive_files(self, mock_retrying):
Expand Down
18 changes: 18 additions & 0 deletions unit_tests/unit/test_remoter.py
Original file line number Diff line number Diff line change
Expand Up @@ -545,3 +545,21 @@ def test_remote_file_preserve_readonly(self):
assert remoter.rf_dst.startswith("/tmp/sct")
assert remoter.rf_dst.endswith(os.path.basename(some_file))
assert remoter.sf_data is None

def test_remote_file_empty_mktemp_raises(self):
"""If 'mktemp' comes back with an empty result, remote_file() must raise clearly
instead of proceeding with dst="" into a 300s retry loop that always fails."""

class _EmptyMktempRunner(self.remoter_cls):
def run(self, cmd, *_, **__):
if cmd == "mktemp":
return Result(stdout="")
return super().run(cmd, *_, **__)

remoter = _EmptyMktempRunner()
some_file = "/some/path/some.file"
with pytest.raises(RuntimeError, match="mktemp.*empty"):
with remote_file(
remoter=remoter, remote_path=some_file, preserve_ownership=False, preserve_permissions=False
) as fobj:
fobj.write("test data")
53 changes: 53 additions & 0 deletions unit_tests/unit/test_utils_docker.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,14 +14,20 @@

from __future__ import absolute_import

import contextlib
import http.client
import io
import os
import socket
import pytest
from types import SimpleNamespace
from unittest.mock import Mock, patch, mock_open, sentinel
from collections import namedtuple

from sdcm.utils.docker_utils import (
_Name,
ContainerManager,
DockerClient,
DockerException,
NotFound,
ImageNotFound,
Expand All @@ -30,6 +36,20 @@
ContainerAlreadyRegistered,
)

# Docker sends 101 upgrade headers, then multiplexed frames on the same connection.
# http.client reads through an io.BufferedReader, so frame bytes can already be in that buffer.
HIJACK_HEADERS = (
b"HTTP/1.1 101 UPGRADED\r\n"
b"Content-Type: application/vnd.docker.multiplexed-stream\r\n"
b"Connection: Upgrade\r\n"
b"Upgrade: tcp\r\n"
b"Api-Version: 1.49\r\n"
b"\r\n"
)
# Frame header: stream id 1, three padding bytes, then the payload length as a big-endian
# uint32. The declared length has to match the payload, 0x10 for these 16 bytes.
MKTEMP_FRAME = b"\x01\x00\x00\x00\x00\x00\x00\x10" + b"/tmp/tmp.abcdef\n"

build_args = {}


Expand Down Expand Up @@ -728,3 +748,36 @@ def stub_list(*args, **kwargs):

# Registered not labeled container not destroyed
r_c2.remove.assert_not_called()


@contextlib.contextmanager
def hijacked_response(frames: bytes):
"""Yield a response where headers and frames arrive in the same recv()."""
server, client = socket.socketpair()
try:
server.sendall(HIJACK_HEADERS + frames)
# signal EOF after the last frame
server.shutdown(socket.SHUT_WR)
http_response = http.client.HTTPResponse(client)
http_response.begin()
assert isinstance(http_response.fp, io.BufferedReader), "http.client must buffer the stream"

yield SimpleNamespace(
raw=SimpleNamespace(_fp=http_response),
raise_for_status=lambda: None,
close=lambda: None,
)
finally:
client.close()
server.close()


def test_frame_buffered_with_the_headers_is_not_lost():
"""Ensure a frame buffered with 101 headers is still read."""
docker_client = DockerClient(base_url="unix:///var/run/docker.sock", version="1.49")

with hijacked_response(MKTEMP_FRAME) as response:
stdout, stderr = docker_client.api._read_from_socket(response, stream=False, tty=False, demux=True)

assert stdout == b"/tmp/tmp.abcdef\n"
assert stderr is None
Loading