diff --git a/src/linux/init/GnsPortTracker.cpp b/src/linux/init/GnsPortTracker.cpp index fb1383f77..73d0be841 100644 --- a/src/linux/init/GnsPortTracker.cpp +++ b/src/linux/init/GnsPortTracker.cpp @@ -375,7 +375,7 @@ std::optional GnsPortTracker::GetCallInfo( uint64_t CallId, pid_t Pid, int Arch, int SysCallNumber, const gsl::span& Arguments) { auto ParseSocket = [&](int Socket, size_t AddressPtr, size_t AddressLength) -> std::optional { - if (AddressLength < sizeof(sockaddr)) + if (AddressLength < sizeof(sockaddr) || AddressLength > sizeof(sockaddr_storage)) { return {{{}, {}, CallId}}; // Invalid sockaddr. Let it go through. } diff --git a/src/linux/init/main.cpp b/src/linux/init/main.cpp index b95fee1fa..8188813b6 100644 --- a/src/linux/init/main.cpp +++ b/src/linux/init/main.cpp @@ -2170,7 +2170,9 @@ Return Value: try { - wil::unique_fd InitFd{open(Target, (O_CREAT | O_WRONLY | O_TRUNC), 0755)}; + THROW_LAST_ERROR_IF(unlink(Target) < 0 && errno != ENOENT); + + wil::unique_fd InitFd{open(Target, (O_CREAT | O_EXCL | O_WRONLY), 0755)}; THROW_LAST_ERROR_IF(!InitFd); THROW_LAST_ERROR_IF(mount(LX_INIT_PATH, Target, nullptr, (MS_RDONLY | MS_BIND), nullptr) < 0); diff --git a/src/windows/common/WslClient.cpp b/src/windows/common/WslClient.cpp index 3c120c8ce..8c5ef60bc 100644 --- a/src/windows/common/WslClient.cpp +++ b/src/windows/common/WslClient.cpp @@ -223,8 +223,9 @@ int ExportDistribution(_In_ std::wstring_view commandLine) ArgumentParser parser(std::wstring{commandLine}, WSL_BINARY_NAME); std::filesystem::path filePath; LPCWSTR name{}; + int tarFormatSet = 0; - auto parseFormat = [&flags](LPCWSTR Value) { + auto parseFormat = [&flags, &tarFormatSet](LPCWSTR Value) { if (Value == nullptr) { return -1; @@ -242,7 +243,11 @@ int ExportDistribution(_In_ std::wstring_view commandLine) { WI_SetFlag(flags, LXSS_EXPORT_DISTRO_FLAGS_VHD); } - else if (!wsl::shared::string::IsEqual(L"tar", Value)) + else if (wsl::shared::string::IsEqual(L"tar", Value)) + { + tarFormatSet = 1; + } + else { THROW_HR(E_INVALIDARG); } @@ -256,9 +261,8 @@ int ExportDistribution(_In_ std::wstring_view commandLine) parser.AddArgument(parseFormat, WSL_EXPORT_ARG_FORMAT_OPTION); parser.Parse(); - THROW_HR_IF( - WSL_E_INVALID_USAGE, - filePath.empty() || (WI_IsFlagSet(flags, LXSS_EXPORT_DISTRO_FLAGS_GZIP) && WI_IsFlagSet(flags, LXSS_EXPORT_DISTRO_FLAGS_VHD))); + constexpr ULONG c_exportFormatFlags = LXSS_EXPORT_DISTRO_FLAGS_VHD | LXSS_EXPORT_DISTRO_FLAGS_GZIP | LXSS_EXPORT_DISTRO_FLAGS_XZIP; + THROW_HR_IF(WSL_E_INVALID_USAGE, filePath.empty() || std::popcount(flags & c_exportFormatFlags) + tarFormatSet > 1); // Determine if the target is stdout, or an on-disk file. wil::unique_hfile file; @@ -926,12 +930,10 @@ int Manage(_In_ std::wstring_view commandLine) else if (defaultUser) { auto wslExe = wil::GetModuleFileNameW(wil::GetModuleInstanceHandle()); - - auto commandLine = std::format( - L"\"{}\" {} -u root /usr/bin/id -u -- '{}'", - wslExe, - wsl::shared::string::GuidToString(distroGuid), - defaultUser.value()); + const auto distroGuidString = wsl::shared::string::GuidToString(distroGuid); + const std::array arguments{ + wslExe, distroGuidString, WSL_USER_ARG, L"root", WSL_EXEC_ARG, L"/usr/bin/id", L"-u", L"--", defaultUser.value()}; + const auto commandLine = wil::ArgvToCommandLine(arguments); wsl::windows::common::SubProcess process{wslExe.c_str(), commandLine.c_str()}; diff --git a/src/windows/common/precomp.h b/src/windows/common/precomp.h index c7b53ee54..a903ca27b 100644 --- a/src/windows/common/precomp.h +++ b/src/windows/common/precomp.h @@ -97,6 +97,7 @@ Module Name: #include #include #include +#include // Socket APIs #include diff --git a/src/windows/service/exe/WslCoreGuestNetworkService.cpp b/src/windows/service/exe/WslCoreGuestNetworkService.cpp index 5fae67f0a..b79e11254 100644 --- a/src/windows/service/exe/WslCoreGuestNetworkService.cpp +++ b/src/windows/service/exe/WslCoreGuestNetworkService.cpp @@ -92,6 +92,7 @@ void wsl::core::networking::GuestNetworkService::CreateGuestNetworkService( TraceLoggingHResult(result, "result"), TraceLoggingValue(error.is_valid() ? error.get() : L"null", "errorString")); THROW_IF_FAILED_MSG(result, "%ls", error.get()); + m_id = VmId; m_guestNetworkServiceCallback = windows::common::hcs::RegisterGuestNetworkServiceCallback(m_service, Callback, CallbackContext); SetGuestNetworkServiceState(hns::GuestNetworkServiceState::Bootstrapping); diff --git a/src/windows/service/exe/WslCoreNetworkEndpoint.h b/src/windows/service/exe/WslCoreNetworkEndpoint.h index 7dfdf57e3..7db7c3aec 100644 --- a/src/windows/service/exe/WslCoreNetworkEndpoint.h +++ b/src/windows/service/exe/WslCoreNetworkEndpoint.h @@ -3,6 +3,7 @@ #pragma once #include #include +#include #include #include "WslCoreNetworkEndpointSettings.h" @@ -16,15 +17,29 @@ struct NetworkEndpoint ~NetworkEndpoint() noexcept { - if (Endpoint) + DeleteEndpoint(); + } + + NetworkEndpoint(NetworkEndpoint&&) = default; + NetworkEndpoint& operator=(NetworkEndpoint&& source) noexcept + { + if (this != &source) { - wil::unique_cotaskmem_string error; - LOG_IF_FAILED_MSG(::HcnDeleteEndpoint(EndpointId, &error), "error message: %ls", error.get()); + DeleteEndpoint(); + StateTracking.reset(); + + Network = std::move(source.Network); + NetworkId = source.NetworkId; + EndpointId = source.EndpointId; + InterfaceGuid = source.InterfaceGuid; + InterfaceLuid = source.InterfaceLuid; + Endpoint = std::move(source.Endpoint); + StateTracking = std::move(source.StateTracking); } + + return *this; } - NetworkEndpoint(NetworkEndpoint&&) = default; - NetworkEndpoint& operator=(NetworkEndpoint&& source) = default; NetworkEndpoint(const NetworkEndpoint&) = delete; NetworkEndpoint& operator=(const NetworkEndpoint&) = delete; @@ -36,6 +51,16 @@ struct NetworkEndpoint windows::common::hcs::unique_hcn_endpoint Endpoint{}; std::optional StateTracking; + void DeleteEndpoint() noexcept + { + if (Endpoint) + { + wil::unique_cotaskmem_string error; + LOG_IF_FAILED_MSG(::HcnDeleteEndpoint(EndpointId, &error), "error message: %ls", error.get()); + Endpoint.reset(); + } + } + void TraceLoggingRundown() const { if (Network) diff --git a/test/windows/UnitTests.cpp b/test/windows/UnitTests.cpp index 035323b75..e3c846180 100644 --- a/test/windows/UnitTests.cpp +++ b/test/windows/UnitTests.cpp @@ -169,6 +169,9 @@ class UnitTests VERIFY_ARE_EQUAL(out, L"This operation is only supported by WSL2.\r\nError code: Wsl/Service/WSL_E_WSL2_NEEDED\r\n"); VERIFY_ARE_EQUAL(err, L""); } + + VerifyInvalidUsage(std::format(L"--export {} {} --format tar.gz --format tar.xz", LXSS_DISTRO_NAME_TEST_L, tarPath)); + VerifyInvalidUsage(std::format(L"--export {} {} --format tar.xz --vhd", LXSS_DISTRO_NAME_TEST_L, tarPath)); } WSL2_TEST_METHOD(SystemdSafeMode) @@ -4365,6 +4368,22 @@ localhostForwarding=true VERIFY_ARE_EQUAL( out, L"There is no distribution with the supplied name.\r\nError code: Wsl/Service/WSL_E_DISTRO_NOT_FOUND\r\n"); + + constexpr auto injectionMarker = L"/tmp/wsl-manage-default-user-injection"; + LxsstuLaunchWsl(std::format(L"-u root -e /usr/bin/rm -f {}", injectionMarker)); + auto cleanupInjectionMarker = wil::scope_exit_log(WI_DIAGNOSTICS_INFO, [injectionMarker]() { + LxsstuLaunchWsl(std::format(L"-u root -e /usr/bin/rm -f {}", injectionMarker)); + }); + + const auto injectionUsername = std::format(L"' || touch {} || '", injectionMarker); + const std::array injectionArguments{ + WSL_MANAGE_ARG, LXSS_DISTRO_NAME_TEST_L, WSL_MANAGE_ARG_SET_DEFAULT_USER_OPTION_LONG, injectionUsername}; + const auto injectionCommand = wil::ArgvToCommandLine(injectionArguments, wil::ArgvToCommandLineFlags::FirstArgumentIsNotPath); + auto injectionCommandLine = LxssGenerateWslCommandLine(injectionCommand.c_str()); + const auto injectionExitCode = LxsstuRunCommand(injectionCommandLine.data()); + + VERIFY_ARE_EQUAL(LxsstuLaunchWsl(std::format(L"-u root -e /usr/bin/test ! -e {}", injectionMarker)), 0L); + VERIFY_ARE_EQUAL(injectionExitCode, 1L); } TEST_METHOD(PostDistroRegistrationSettingsOOBE)