diff --git a/include/multipass/ssh/plain_ssh_process.h b/include/multipass/ssh/plain_ssh_process.h index 145e01a446..fa840d6e48 100644 --- a/include/multipass/ssh/plain_ssh_process.h +++ b/include/multipass/ssh/plain_ssh_process.h @@ -36,7 +36,7 @@ class PlainSftpSession; class PlainSSHProcess : public SSHProcess { public: - PlainSSHProcess(ssh_session_struct& raw_session, + PlainSSHProcess(ssh_session_struct* raw_session, // non-null const std::string& cmd, std::unique_lock session_lock); diff --git a/src/ssh/plain_ssh_process.cpp b/src/ssh/plain_ssh_process.cpp index 27929cfbb0..f5ac55961b 100644 --- a/src/ssh/plain_ssh_process.cpp +++ b/src/ssh/plain_ssh_process.cpp @@ -81,6 +81,8 @@ void mp::PlainSSHProcess::ChannelDeleter::operator()(ssh_channel_struct* channel mp::PlainSSHProcess::ChannelUPtr mp::PlainSSHProcess::make_channel(ssh_session raw_session, const std::string& cmd) { + assert(raw_session && "precondition - need an actual session"); + if (!MP_LIBSSH.ssh_is_connected(raw_session)) throw mp::SSHException(fmt::format( "unable to create a channel for remote process: '{}', the SSH session is not connected", @@ -99,14 +101,14 @@ mp::PlainSSHProcess::ChannelUPtr mp::PlainSSHProcess::make_channel(ssh_session r return channel; } -mp::PlainSSHProcess::PlainSSHProcess(ssh_session_struct& raw_session, +mp::PlainSSHProcess::PlainSSHProcess(ssh_session_struct* raw_session, const std::string& cmd, std::unique_lock session_lock) : session_lock{std::move( session_lock)}, // this is held until the exit code is requested or this is destroyed - raw_session{&raw_session}, + raw_session{raw_session}, cmd{cmd}, - channel{make_channel(this->raw_session, cmd)}, + channel{make_channel(raw_session, cmd)}, exit_result{} { assert(this->session_lock.owns_lock()); diff --git a/src/ssh/plain_ssh_session.cpp b/src/ssh/plain_ssh_session.cpp index d0b032d7b0..2a0a99c32e 100644 --- a/src/ssh/plain_ssh_session.cpp +++ b/src/ssh/plain_ssh_session.cpp @@ -155,7 +155,7 @@ std::unique_ptr mp::PlainSSHSession::exec_plain(const std:: auto lvl = whisper ? mpl::Level::trace : mpl::Level::debug; mpl::log(lvl, category, "Executing '{}'", cmd); - return std::make_unique(*raw_session.get(), cmd, std::move(lock)); + return std::make_unique(raw_session.get(), cmd, std::move(lock)); } std::unique_ptr mp::PlainSSHSession::make_sftp_session( diff --git a/src/sshfs_mount/sftp_server.cpp b/src/sshfs_mount/sftp_server.cpp index 613b18de92..1bf0e4fc4a 100644 --- a/src/sshfs_mount/sftp_server.cpp +++ b/src/sshfs_mount/sftp_server.cpp @@ -25,6 +25,7 @@ #include #include #include +#include #include #include @@ -275,28 +276,34 @@ int reverse_id_for(const mp::id_mappings& id_maps, const int id, const int defau : found->first; } -constexpr bool follows_symlinks(uint8_t type) +constexpr bool follows_symlinks(mp::SftpMessageType type) { switch (type) { - case SSH_FXP_OPEN: - case SSH_FXP_OPENDIR: - case SSH_FXP_REALPATH: - case SSH_FXP_SETSTAT: - case SSH_FXP_STAT: + case mp::SftpMessageType::open: + case mp::SftpMessageType::opendir: + case mp::SftpMessageType::realpath: + case mp::SftpMessageType::setstat: + case mp::SftpMessageType::stat: return true; - case SSH_FXP_LSTAT: - case SSH_FXP_READLINK: - case SSH_FXP_REMOVE: - case SSH_FXP_RMDIR: - case SSH_FXP_MKDIR: - case SSH_FXP_RENAME: - case SSH_FXP_SYMLINK: + case mp::SftpMessageType::lstat: + case mp::SftpMessageType::readlink: + case mp::SftpMessageType::remove: + case mp::SftpMessageType::rmdir: + case mp::SftpMessageType::mkdir: + case mp::SftpMessageType::rename: + case mp::SftpMessageType::symlink: return false; default: return false; // Fail-safe default } } + +// TODO@sftp dump once handlers take SftpMessage +mp::SftpMessageType type_of(sftp_client_message msg) +{ + return static_cast(MP_LIBSSH.sftp_client_message_get_type(msg)); +} } // namespace mp::SftpServer::SftpServer(std::unique_ptr&& session, @@ -339,15 +346,15 @@ sftp_attributes_struct mp::SftpServer::attr_from(const QFileInfo& file_info) attr.permissions = to_unix_permissions(file_info.permissions()); attr.atime = file_info.lastRead().toUTC().toMSecsSinceEpoch() / 1000; attr.mtime = file_info.lastModified().toUTC().toMSecsSinceEpoch() / 1000; - attr.flags = SSH_FILEXFER_ATTR_SIZE | SSH_FILEXFER_ATTR_UIDGID | SSH_FILEXFER_ATTR_PERMISSIONS | - SSH_FILEXFER_ATTR_ACMODTIME; + attr.flags = SftpAttrFlags::size | SftpAttrFlags::uidgid | SftpAttrFlags::permissions | + SftpAttrFlags::acmodtime; if (file_info.isSymLink()) - attr.permissions |= SSH_S_IFLNK | 0777; + attr.permissions |= SftpFileMode::symlink | 0777; else if (file_info.isDir()) - attr.permissions |= SSH_S_IFDIR; + attr.permissions |= SftpFileMode::directory; else if (file_info.isFile()) - attr.permissions |= SSH_S_IFREG; + attr.permissions |= SftpFileMode::regular; return attr; } @@ -465,7 +472,7 @@ fs::path mp::SftpServer::get_absolute_path(const char* path) const std::optional mp::SftpServer::get_validated_path(sftp_client_message msg) const { - bool follows{follows_symlinks(MP_LIBSSH.sftp_client_message_get_type(msg))}; + bool follows{follows_symlinks(type_of(msg))}; const auto path = get_absolute_path(MP_LIBSSH.sftp_client_message_get_filename(msg)); if (!validate_path(path, follows)) { @@ -499,60 +506,60 @@ std::string mp::SftpServer::host_to_guest_path(const fs::path& host_path) const void mp::SftpServer::process_message(sftp_client_message msg) { int ret = 0; - const auto type = MP_LIBSSH.sftp_client_message_get_type(msg); + const auto type = type_of(msg); switch (type) { - case SFTP_REALPATH: + case SftpMessageType::realpath: ret = handle_realpath(msg); break; - case SFTP_OPENDIR: + case SftpMessageType::opendir: ret = handle_opendir(msg); break; - case SFTP_MKDIR: + case SftpMessageType::mkdir: ret = handle_mkdir(msg); break; - case SFTP_RMDIR: + case SftpMessageType::rmdir: ret = handle_rmdir(msg); break; - case SFTP_LSTAT: - case SFTP_STAT: - ret = handle_stat(msg, type == SFTP_STAT); + case SftpMessageType::lstat: + case SftpMessageType::stat: + ret = handle_stat(msg, type == SftpMessageType::stat); break; - case SFTP_FSTAT: + case SftpMessageType::fstat: ret = handle_fstat(msg); break; - case SFTP_READDIR: + case SftpMessageType::readdir: ret = handle_readdir(msg); break; - case SFTP_CLOSE: + case SftpMessageType::close: ret = handle_close(msg); break; - case SFTP_OPEN: + case SftpMessageType::open: ret = handle_open(msg); break; - case SFTP_READ: + case SftpMessageType::read: ret = handle_read(msg); break; - case SFTP_WRITE: + case SftpMessageType::write: ret = handle_write(msg); break; - case SFTP_RENAME: + case SftpMessageType::rename: ret = handle_rename(msg); break; - case SFTP_REMOVE: + case SftpMessageType::remove: ret = handle_remove(msg); break; - case SFTP_SETSTAT: - case SFTP_FSETSTAT: + case SftpMessageType::setstat: + case SftpMessageType::fsetstat: ret = handle_setstat(msg); break; - case SFTP_READLINK: + case SftpMessageType::readlink: ret = handle_readlink(msg); break; - case SFTP_SYMLINK: + case SftpMessageType::symlink: ret = handle_symlink(msg); break; - case SFTP_EXTENDED: + case SftpMessageType::extended: ret = handle_extended(msg); break; default: @@ -657,7 +664,7 @@ int mp::SftpServer::handle_fstat(sftp_client_message msg) const auto& [path, _] = *handle; - if (!validate_path(path, follows_symlinks(MP_LIBSSH.sftp_client_message_get_type(msg)))) + if (!validate_path(path, follows_symlinks(type_of(msg)))) { mpl::trace(category, "{}: cannot validate target path \'{}\' against source \'{}\'", @@ -791,29 +798,23 @@ int mp::SftpServer::handle_open(sftp_client_message msg) int mode = 0; const auto flags = MP_LIBSSH.sftp_client_message_get_flags(msg); - if (flags & SSH_FXF_READ) + if ((flags & SftpOpenFlags::read) && (flags & SftpOpenFlags::write)) + mode |= O_RDWR; + else if (flags & SftpOpenFlags::read) mode |= O_RDONLY; - - if (flags & SSH_FXF_WRITE) + else if (flags & SftpOpenFlags::write) mode |= O_WRONLY; - if ((flags & SSH_FXF_READ) && (flags & SSH_FXF_WRITE)) - { - mode &= ~O_RDONLY; - mode &= ~O_WRONLY; - mode |= O_RDWR; - } - - if (flags & SSH_FXF_APPEND) + if (flags & SftpOpenFlags::append) mode |= O_APPEND; - if (flags & SSH_FXF_TRUNC) + if (flags & SftpOpenFlags::trunc) mode |= O_TRUNC; - if (flags & SSH_FXF_CREAT) + if (flags & SftpOpenFlags::creat) mode |= O_CREAT; - if (flags & SSH_FXF_EXCL) + if (flags & SftpOpenFlags::excl) mode |= O_EXCL; auto named_fd = MP_FILEOPS.open_fd(*filename, mode, msg->attr ? msg->attr->permissions : 0); @@ -1142,7 +1143,7 @@ int mp::SftpServer::handle_setstat(sftp_client_message msg) { fs::path filename; - if (MP_LIBSSH.sftp_client_message_get_type(msg) == SFTP_FSETSTAT) + if (type_of(msg) == SftpMessageType::fsetstat) { const auto handle = get_handle(msg); if (handle == nullptr) @@ -1183,7 +1184,7 @@ int mp::SftpServer::handle_setstat(sftp_client_message msg) return reply_perm_denied(msg); } - if (msg->attr->flags & SSH_FILEXFER_ATTR_SIZE) + if (msg->attr->flags & SftpAttrFlags::size) { QFile file{filename}; if (!MP_FILEOPS.resize(file, msg->attr->size)) @@ -1193,7 +1194,7 @@ int mp::SftpServer::handle_setstat(sftp_client_message msg) } } - if (msg->attr->flags & SSH_FILEXFER_ATTR_PERMISSIONS) + if (msg->attr->flags & SftpAttrFlags::permissions) { if (!MP_PLATFORM.set_permissions(filename, static_cast(msg->attr->permissions))) { @@ -1205,7 +1206,7 @@ int mp::SftpServer::handle_setstat(sftp_client_message msg) } } - if (msg->attr->flags & SSH_FILEXFER_ATTR_ACMODTIME) + if (msg->attr->flags & SftpAttrFlags::acmodtime) { if (MP_PLATFORM.utime(filename.string().c_str(), msg->attr->atime, msg->attr->mtime) < 0) { @@ -1217,7 +1218,7 @@ int mp::SftpServer::handle_setstat(sftp_client_message msg) } } - if (msg->attr->flags & SSH_FILEXFER_ATTR_UIDGID) + if (msg->attr->flags & SftpAttrFlags::uidgid) { if (!has_reverse_uid_mapping_for(msg->attr->uid) && !has_reverse_gid_mapping_for(msg->attr->gid)) @@ -1391,7 +1392,7 @@ int mp::SftpServer::handle_extended(sftp_client_message msg) return reply_perm_denied(msg); const auto new_name = get_absolute_path(MP_LIBSSH.sftp_client_message_get_data(msg)); - if (!validate_path(new_name, follows_symlinks(MP_LIBSSH.sftp_client_message_get_type(msg)))) + if (!validate_path(new_name, follows_symlinks(type_of(msg)))) { mpl::trace(category, "{}: cannot validate path \'{}\' against source \'{}\'", diff --git a/tests/unit/CMakeLists.txt b/tests/unit/CMakeLists.txt index af28042f38..5e38c10c52 100644 --- a/tests/unit/CMakeLists.txt +++ b/tests/unit/CMakeLists.txt @@ -111,8 +111,8 @@ add_executable(multipass_cpp_tests test_permission_utils.cpp test_persistent_settings_handler.cpp test_plain_sftp_session.cpp + test_plain_ssh_process.cpp test_plain_ssh_session.cpp - test_plain_ssh_session_mocked_libssh.cpp test_private_pass_provider.cpp test_qemu_img_utils.cpp test_qemuimg_process_spec.cpp diff --git a/tests/unit/test_plain_sftp_session.cpp b/tests/unit/test_plain_sftp_session.cpp index 542491e6c0..617517a286 100644 --- a/tests/unit/test_plain_sftp_session.cpp +++ b/tests/unit/test_plain_sftp_session.cpp @@ -103,7 +103,7 @@ struct TestPlainSftpSession : public Test TEST_F(TestPlainSftpSession, makeSftpSessionRunsSshfsCommand) { - sshfs_exit_code = 1; // TODO@sftp mock success path instead + sshfs_exit_code = 1; auto session = make_ssh_session(); EXPECT_CALL(mock_libssh, ssh_channel_request_exec(fake_channel, StrEq("sshfs -o slave"))) diff --git a/tests/unit/test_plain_ssh_process.cpp b/tests/unit/test_plain_ssh_process.cpp new file mode 100644 index 0000000000..fb50f1ecba --- /dev/null +++ b/tests/unit/test_plain_ssh_process.cpp @@ -0,0 +1,82 @@ +/* + * Copyright (C) Canonical, Ltd. + * + * This program is free software; you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation; version 3. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with this program. If not, see . + * + */ + +#include "common.h" +#include "mock_libssh.h" + +#include +#include + +#include + +namespace mp = multipass; +namespace mpt = multipass::test; +using namespace testing; + +namespace +{ +struct TestPlainSSHProcess : public Test +{ + TestPlainSSHProcess() + { + ON_CALL(mock_libssh, ssh_is_connected).WillByDefault(Return(1)); + ON_CALL(mock_libssh, ssh_channel_new).WillByDefault(Return(fake_channel)); + } + + mp::PlainSSHProcess make_ssh_process(const std::string& cmd = "cmd") + { + return mp::PlainSSHProcess{fake_session, cmd, std::unique_lock{mutex}}; + } + + mpt::MockLibssh::GuardedMock guarded_mock = mpt::MockLibssh::inject(); + mpt::MockLibssh& mock_libssh = *guarded_mock.first; + + constexpr static auto bad_addr = 0xdeadbeefdeadbeefull; // should reliably segfault on 32/64-bit + ssh_session fake_session = reinterpret_cast(bad_addr); + ssh_channel fake_channel = reinterpret_cast(bad_addr); + + std::mutex mutex; +}; +} // namespace + +TEST_F(TestPlainSSHProcess, execThrowsOnADeadSession) +{ + EXPECT_CALL(mock_libssh, ssh_is_connected(fake_session)).WillOnce(Return(0)); + + MP_EXPECT_THROW_THAT(make_ssh_process(), + mp::SSHException, + mpt::match_what(HasSubstr("not connected"))); +} + +TEST_F(TestPlainSSHProcess, execThrowsWhenUnableToOpenAChannelSession) +{ + constexpr auto err = "mocked error"; + EXPECT_CALL(mock_libssh, ssh_channel_open_session).WillOnce(Return(SSH_ERROR)); + EXPECT_CALL(mock_libssh, ssh_get_error(fake_session)).WillOnce(Return(err)); + + MP_EXPECT_THROW_THAT(make_ssh_process(), mp::SSHException, mpt::match_what(HasSubstr(err))); +} + +TEST_F(TestPlainSSHProcess, execThrowsWhenUnableToRequestChannelExec) +{ + constexpr auto err = "mocked error"; + ON_CALL(mock_libssh, ssh_channel_open_session).WillByDefault(Return(SSH_OK)); + EXPECT_CALL(mock_libssh, ssh_channel_request_exec).WillOnce(Return(SSH_ERROR)); + EXPECT_CALL(mock_libssh, ssh_get_error(fake_session)).WillOnce(Return(err)); + + MP_EXPECT_THROW_THAT(make_ssh_process(), mp::SSHException, mpt::match_what(HasSubstr(err))); +} diff --git a/tests/unit/test_plain_ssh_session.cpp b/tests/unit/test_plain_ssh_session.cpp index 73687428a9..70e7de6fbc 100644 --- a/tests/unit/test_plain_ssh_session.cpp +++ b/tests/unit/test_plain_ssh_session.cpp @@ -16,133 +16,175 @@ */ #include "common.h" +#include "mock_libssh.h" #include "mock_platform.h" -#include "mock_ssh.h" #include "stub_ssh_key_provider.h" +#include #include +#include #include +#include +#include + namespace mp = multipass; namespace mpt = multipass::test; using namespace testing; namespace { +static_assert(std::is_final_v, "required to prevent chopping on move"); +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(std::is_move_constructible_v); +static_assert(std::is_move_assignable_v); +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); + struct TestPlainSSHSession : public Test { - mp::PlainSSHSession make_ssh_session() + TestPlainSSHSession() + { + ON_CALL(mock_libssh, ssh_new()).WillByDefault(Return(fake_session)); + ON_CALL(mock_libssh, ssh_options_set).WillByDefault(Return(SSH_OK)); + ON_CALL(mock_libssh, ssh_connect).WillByDefault(Return(SSH_OK)); + ON_CALL(mock_libssh, ssh_userauth_publickey).WillByDefault(Return(SSH_AUTH_SUCCESS)); + ON_CALL(mock_libssh, ssh_is_connected).WillByDefault(Return(1)); + ON_CALL(mock_libssh, ssh_channel_new).WillByDefault(Return(fake_channel)); + ON_CALL(mock_libssh, ssh_channel_open_session).WillByDefault(Return(SSH_OK)); + ON_CALL(mock_libssh, ssh_channel_request_exec).WillByDefault(Return(SSH_OK)); + ON_CALL(mock_libssh, ssh_get_fd).WillByDefault(Return(-1)); // no socket to shutdown + ON_CALL(mock_libssh, ssh_get_error).WillByDefault(Return("mocked error")); + } + + mp::PlainSSHSession make_ssh_session() const { - return mp::PlainSSHSession("theanswertoeverything", 42, "ubuntu", key_provider); + return mp::PlainSSHSession{"host", 42, "ubuntu", key_provider}; } - mp::test::StubSSHKeyProvider key_provider; + mpt::MockLibssh::GuardedMock guarded_mock = mpt::MockLibssh::inject(); + mpt::MockLibssh& mock_libssh = *guarded_mock.first; + + constexpr static auto bad_addr = 0xdeadbeefdeadbeefull; // should reliably segfault on 32/64-bit + constexpr static auto bad_addr_too = 0xbadadd4f0ccac1adull; // idem + ssh_session fake_session = reinterpret_cast(bad_addr); + ssh_channel fake_channel = reinterpret_cast(bad_addr); + + mpt::StubSSHKeyProvider key_provider; }; } // namespace TEST_F(TestPlainSSHSession, throwsWhenUnableToAllocateSession) { - REPLACE(ssh_new, []() { return nullptr; }); - EXPECT_THROW(make_ssh_session(), std::runtime_error); + EXPECT_CALL(mock_libssh, ssh_new()).WillOnce(Return(nullptr)); + EXPECT_THROW(make_ssh_session(), mp::SSHException); } TEST_F(TestPlainSSHSession, throwsWhenUnableToSetOption) { - REPLACE(ssh_options_set, [](auto...) { return SSH_ERROR; }); - EXPECT_THROW(make_ssh_session(), std::runtime_error); + constexpr auto err = "mocked error"; + EXPECT_CALL(mock_libssh, ssh_options_set).WillOnce(Return(SSH_ERROR)); + EXPECT_CALL(mock_libssh, ssh_get_error(fake_session)).WillOnce(Return(err)); + + MP_EXPECT_THROW_THAT(make_ssh_session(), mp::SSHException, mpt::match_what(HasSubstr(err))); } TEST_F(TestPlainSSHSession, throwsWhenUnableToConnect) { - REPLACE(ssh_connect, [](auto...) { return SSH_ERROR; }); - EXPECT_THROW(make_ssh_session(), std::runtime_error); + constexpr auto err = "mocked error"; + EXPECT_CALL(mock_libssh, ssh_connect).WillOnce(Return(SSH_ERROR)); + EXPECT_CALL(mock_libssh, ssh_get_error(fake_session)).WillOnce(Return(err)); + + MP_EXPECT_THROW_THAT(make_ssh_session(), mp::SSHException, mpt::match_what(HasSubstr(err))); } TEST_F(TestPlainSSHSession, throwsWhenUnableToAuth) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_ERROR; }); - EXPECT_THROW(make_ssh_session(), std::runtime_error); + constexpr auto err = "mocked error"; + EXPECT_CALL(mock_libssh, ssh_userauth_publickey).WillOnce(Return(SSH_AUTH_ERROR)); + EXPECT_CALL(mock_libssh, ssh_get_error(fake_session)).WillOnce(Return(err)); + + MP_EXPECT_THROW_THAT(make_ssh_session(), mp::SSHException, mpt::match_what(HasSubstr(err))); } -TEST_F(TestPlainSSHSession, execThrowsOnADeadSession) +TEST_F(TestPlainSSHSession, execPlainReturnsConcreteProcessRunningGivenCommand) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session = make_ssh_session(); + auto session = make_ssh_session(); + EXPECT_CALL(mock_libssh, ssh_channel_request_exec(fake_channel, StrEq("ls -la"))) + .WillOnce(Return(SSH_OK)); + + std::unique_ptr proc = session.exec_plain("ls -la"); - REPLACE(ssh_is_connected, [](auto...) { return false; }); - EXPECT_THROW(static_cast(session.exec("dummy")), std::runtime_error); + ASSERT_THAT(proc, NotNull()); + EXPECT_EQ(proc->get_cmd(), "ls -la"); } -TEST_F(TestPlainSSHSession, execThrowsIfSshIsDead) +TEST_F(TestPlainSSHSession, execPlainThrowsOnDisconnectedSession) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session = make_ssh_session(); + auto session = make_ssh_session(); + EXPECT_CALL(mock_libssh, ssh_is_connected(fake_session)).WillOnce(Return(0)); - REPLACE(ssh_is_connected, [](auto...) { return false; }); - EXPECT_THROW(static_cast(session.exec("dummy")), std::runtime_error); + MP_EXPECT_THROW_THAT(static_cast(session.exec_plain("cmd")), + mp::SSHException, + mpt::match_what(HasSubstr("not connected"))); } -TEST_F(TestPlainSSHSession, execThrowsWhenUnableToOpenAChannelSession) +TEST_F(TestPlainSSHSession, execProducesPlainProcess) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session = make_ssh_session(); + auto session = make_ssh_session(); + + auto proc = session.exec("true"); - REPLACE(ssh_is_connected, [](auto...) { return true; }); - REPLACE(ssh_channel_open_session, [](auto...) { return SSH_ERROR; }); - EXPECT_THROW(static_cast(session.exec("dummy")), std::runtime_error); + ASSERT_THAT(proc, NotNull()); + EXPECT_THAT(dynamic_cast(proc.get()), NotNull()); } -TEST_F(TestPlainSSHSession, execThrowsWhenUnableToRequestChannelExec) +TEST_F(TestPlainSSHSession, moveConstructionLeavesSourceMoved) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session = make_ssh_session(); - - REPLACE(ssh_is_connected, [](auto...) { return true; }); - REPLACE(ssh_channel_open_session, [](auto...) { return SSH_OK; }); - REPLACE(ssh_channel_request_exec, [](auto...) { return SSH_ERROR; }); - EXPECT_THROW(static_cast(session.exec("dummy")), std::runtime_error); + auto session1 = make_ssh_session(); + EXPECT_FALSE(session1.is_moved()); + + auto session2 = std::move(session1); + + EXPECT_TRUE(session1.is_moved()); + EXPECT_FALSE(session2.is_moved()); } -TEST_F(TestPlainSSHSession, execSucceeds) +TEST_F(TestPlainSSHSession, moveAssignmentTransfersUnderlyingSession) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session = make_ssh_session(); + auto other_session = reinterpret_cast(bad_addr_too); + EXPECT_CALL(mock_libssh, ssh_new()) + .WillOnce(Return(fake_session)) + .WillOnce(Return(other_session)); + + auto session1 = make_ssh_session(); + auto session2 = make_ssh_session(); - REPLACE(ssh_is_connected, [](auto...) { return true; }); - REPLACE(ssh_channel_open_session, [](auto...) { return SSH_OK; }); - REPLACE(ssh_channel_request_exec, [](auto...) { return SSH_OK; }); + session1 = std::move(session2); + + EXPECT_TRUE(session2.is_moved()); + EXPECT_FALSE(session1.is_moved()); - EXPECT_NO_THROW(static_cast(session.exec("dummy"))); + EXPECT_CALL(mock_libssh, ssh_channel_new(other_session)).WillOnce(Return(fake_channel)); + ASSERT_THAT(session1.exec_plain("cmd"), NotNull()); } -TEST_F(TestPlainSSHSession, moveAssigns) +TEST_F(TestPlainSSHSession, movedSessionReleasesOnce) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - mp::PlainSSHSession session1 = make_ssh_session(); - mp::PlainSSHSession session2 = make_ssh_session(); - ssh_session ssh_session2 = session2; + EXPECT_CALL(mock_libssh, ssh_free(fake_session)).Times(1); + EXPECT_CALL(mock_libssh, ssh_disconnect(fake_session)).Times(1); - session1 = std::move(session2); - EXPECT_EQ(ssh_session{session1}, ssh_session2); - EXPECT_EQ(ssh_session{session2}, nullptr); + auto session1 = make_ssh_session(); + auto session2 = std::move(session1); } TEST_F(TestPlainSSHSession, forceShutdownCallsShutdownSocketWhenFdIsValid) { constexpr socket_t fake_fd = 5; + auto session = make_ssh_session(); - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - REPLACE(ssh_get_fd, [](auto...) -> socket_t { return fake_fd; }); - - mp::PlainSSHSession session = make_ssh_session(); + ON_CALL(mock_libssh, ssh_get_fd(fake_session)).WillByDefault(Return(fake_fd)); auto [mock_platform, guard] = mpt::MockPlatform::inject(); EXPECT_CALL(*mock_platform, shutdown_socket(Field(&mp::Socket::fd, fake_fd))); @@ -151,11 +193,7 @@ TEST_F(TestPlainSSHSession, forceShutdownCallsShutdownSocketWhenFdIsValid) TEST_F(TestPlainSSHSession, forceShutdownSkipsShutdownSocketWhenNoFd) { - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - REPLACE(ssh_get_fd, [](auto...) { return (socket_t)-1; }); - - mp::PlainSSHSession session = make_ssh_session(); + auto session = make_ssh_session(); auto [mock_platform, guard] = mpt::MockPlatform::inject(); EXPECT_CALL(*mock_platform, shutdown_socket).Times(0); @@ -165,14 +203,10 @@ TEST_F(TestPlainSSHSession, forceShutdownSkipsShutdownSocketWhenNoFd) TEST_F(TestPlainSSHSession, dtorCallsShutdownSocket) { constexpr socket_t fake_fd = 12; - REPLACE(ssh_connect, [](auto...) { return SSH_OK; }); - REPLACE(ssh_userauth_publickey, [](auto...) { return SSH_AUTH_SUCCESS; }); - REPLACE(ssh_get_fd, [](auto...) -> socket_t { return fake_fd; }); + EXPECT_CALL(mock_libssh, ssh_get_fd(fake_session)).WillOnce(Return(fake_fd)); - { - auto [mock_platform, guard] = mpt::MockPlatform::inject(); - EXPECT_CALL(*mock_platform, shutdown_socket(Field(&mp::Socket::fd, fake_fd))); + auto [mock_platform, guard] = mpt::MockPlatform::inject(); + EXPECT_CALL(*mock_platform, shutdown_socket(Field(&mp::Socket::fd, fake_fd))); - mp::PlainSSHSession session = make_ssh_session(); - } + auto session = make_ssh_session(); } diff --git a/tests/unit/test_plain_ssh_session_mocked_libssh.cpp b/tests/unit/test_plain_ssh_session_mocked_libssh.cpp deleted file mode 100644 index 0d4af963d3..0000000000 --- a/tests/unit/test_plain_ssh_session_mocked_libssh.cpp +++ /dev/null @@ -1,146 +0,0 @@ -/* - * Copyright (C) Canonical, Ltd. - * - * This program is free software; you can redistribute it and/or modify - * it under the terms of the GNU General Public License as published by - * the Free Software Foundation; version 3. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU General Public License for more details. - * - * You should have received a copy of the GNU General Public License - * along with this program. If not, see . - * - */ - -#include "common.h" -#include "mock_libssh.h" -#include "stub_ssh_key_provider.h" - -#include -#include -#include - -#include -#include - -namespace mp = multipass; -namespace mpt = multipass::test; -using namespace testing; - -namespace -{ -static_assert(std::is_final_v, "required to prevent chopping on move"); -static_assert(!std::is_copy_constructible_v); -static_assert(!std::is_copy_assignable_v); -static_assert(std::is_move_constructible_v); -static_assert(std::is_move_assignable_v); -static_assert(!std::is_copy_constructible_v); -static_assert(!std::is_copy_assignable_v); - -// TODO@sftp transfer premock-based session tests to this scheme; then, rename this file/suite -struct TestPlainSSHSessionMockedLibssh : public Test -{ - TestPlainSSHSessionMockedLibssh() - { - ON_CALL(mock_libssh, ssh_new()).WillByDefault(Return(fake_session)); - ON_CALL(mock_libssh, ssh_options_set).WillByDefault(Return(SSH_OK)); - ON_CALL(mock_libssh, ssh_connect).WillByDefault(Return(SSH_OK)); - ON_CALL(mock_libssh, ssh_userauth_publickey).WillByDefault(Return(SSH_AUTH_SUCCESS)); - ON_CALL(mock_libssh, ssh_is_connected).WillByDefault(Return(1)); - ON_CALL(mock_libssh, ssh_channel_new).WillByDefault(Return(fake_channel)); - ON_CALL(mock_libssh, ssh_channel_open_session).WillByDefault(Return(SSH_OK)); - ON_CALL(mock_libssh, ssh_channel_request_exec).WillByDefault(Return(SSH_OK)); - ON_CALL(mock_libssh, ssh_get_fd).WillByDefault(Return(-1)); // no socket to shutdown - ON_CALL(mock_libssh, ssh_get_error).WillByDefault(Return("mocked error")); - } - - mp::PlainSSHSession make_ssh_session() const - { - return mp::PlainSSHSession{"host", 42, "ubuntu", key_provider}; - } - - mpt::MockLibssh::GuardedMock guarded_mock = mpt::MockLibssh::inject(); - mpt::MockLibssh& mock_libssh = *guarded_mock.first; - - constexpr static auto bad_addr = 0xdeadbeefdeadbeefull; // should reliably segfault on 32/64-bit - constexpr static auto bad_addr_too = 0xbadadd4f0ccac1adull; // idem - ssh_session fake_session = reinterpret_cast(bad_addr); - ssh_channel fake_channel = reinterpret_cast(bad_addr); - - mpt::StubSSHKeyProvider key_provider; -}; -} // namespace - -TEST_F(TestPlainSSHSessionMockedLibssh, execPlainReturnsConcreteProcessRunningGivenCommand) -{ - auto session = make_ssh_session(); - EXPECT_CALL(mock_libssh, ssh_channel_request_exec(fake_channel, StrEq("ls -la"))) - .WillOnce(Return(SSH_OK)); - - std::unique_ptr proc = session.exec_plain("ls -la"); - - ASSERT_THAT(proc, NotNull()); - EXPECT_EQ(proc->get_cmd(), "ls -la"); -} - -TEST_F(TestPlainSSHSessionMockedLibssh, execPlainThrowsOnDisconnectedSession) -{ - auto session = make_ssh_session(); - EXPECT_CALL(mock_libssh, ssh_is_connected(fake_session)).WillOnce(Return(0)); - - MP_EXPECT_THROW_THAT(static_cast(session.exec_plain("cmd")), - mp::SSHException, - mpt::match_what(HasSubstr("not connected"))); -} - -TEST_F(TestPlainSSHSessionMockedLibssh, execProducesPlainProcess) -{ - auto session = make_ssh_session(); - - auto proc = session.exec("true"); - - ASSERT_THAT(proc, NotNull()); - EXPECT_THAT(dynamic_cast(proc.get()), NotNull()); -} - -TEST_F(TestPlainSSHSessionMockedLibssh, moveConstructionLeavesSourceMoved) -{ - auto session1 = make_ssh_session(); - EXPECT_FALSE(session1.is_moved()); - - auto session2 = std::move(session1); - - EXPECT_TRUE(session1.is_moved()); - EXPECT_FALSE(session2.is_moved()); -} - -TEST_F(TestPlainSSHSessionMockedLibssh, moveAssignmentTransfersUnderlyingSession) -{ - auto other_session = reinterpret_cast(bad_addr_too); - EXPECT_CALL(mock_libssh, ssh_new()) - .WillOnce(Return(fake_session)) - .WillOnce(Return(other_session)); - - auto session1 = make_ssh_session(); - auto session2 = make_ssh_session(); - - session1 = std::move(session2); - - EXPECT_TRUE(session2.is_moved()); - EXPECT_FALSE(session1.is_moved()); - - EXPECT_CALL(mock_libssh, ssh_channel_new(other_session)).WillOnce(Return(fake_channel)); - ASSERT_THAT(session1.exec_plain("cmd"), NotNull()); -} - -TEST_F(TestPlainSSHSessionMockedLibssh, movedSessionReleasesOnce) -{ - EXPECT_CALL(mock_libssh, ssh_free(fake_session)).Times(1); - EXPECT_CALL(mock_libssh, ssh_disconnect(fake_session)).Times(1); - - auto session1 = make_ssh_session(); - auto session2 = std::move(session1); -}