| 1 | /*++ |
| 2 | |
| 3 | Copyright (c) Microsoft. All rights reserved. |
| 4 | |
| 5 | Module Name: |
| 6 | |
| 7 | hvsocket.cpp |
| 8 | |
| 9 | Abstract: |
| 10 | |
| 11 | This file contains hvsocket helper function definitions. |
| 12 | |
| 13 | --*/ |
| 14 | |
| 15 | #include "precomp.h" |
| 16 | #include <mutex> |
| 17 | #include "socket.hpp" |
| 18 | #include "hvsocket.hpp" |
| 19 | #pragma hdrstop |
| 20 | |
| 21 | namespace { |
| 22 | void InitializeSocketAddress(_In_ const GUID& VmId, _In_ unsigned long Port, _Out_ PSOCKADDR_HV Address) |
| 23 | { |
| 24 | RtlZeroMemory(Address, sizeof(*Address)); |
| 25 | Address->Family = AF_HYPERV; |
| 26 | Address->VmId = VmId; |
| 27 | Address->ServiceId = HV_GUID_VSOCK_TEMPLATE; |
| 28 | Address->ServiceId.Data1 = Port; |
| 29 | } |
| 30 | |
| 31 | void InitializeWildcardSocketAddress(_Out_ PSOCKADDR_HV Address) |
| 32 | { |
| 33 | RtlZeroMemory(Address, sizeof(*Address)); |
| 34 | Address->Family = AF_HYPERV; |
| 35 | Address->VmId = HV_GUID_WILDCARD; |
| 36 | Address->ServiceId = HV_GUID_WILDCARD; |
| 37 | } |
| 38 | } // namespace |
| 39 | |
| 40 | wil::unique_socket wsl::windows::common::hvsocket::Connect( |
| 41 | _In_ const GUID& VmId, _In_ unsigned long Port, _In_opt_ HANDLE ExitHandle, _In_opt_ ULONG Timeout, _In_ const std::source_location& Location) |
| 42 | { |
| 43 | OVERLAPPED Overlapped{}; |
| 44 | const wil::unique_event OverlappedEvent(wil::EventOptions::ManualReset); |
| 45 | Overlapped.hEvent = OverlappedEvent.get(); |
| 46 | |
| 47 | auto Socket = Create(); |
| 48 | |
| 49 | static constexpr GUID ConnectExGuid = WSAID_CONNECTEX; |
| 50 | LPFN_CONNECTEX ConnectFn{}; |
| 51 | DWORD BytesReturned; |
| 52 | const auto Result = WSAIoctl( |
| 53 | Socket.get(), |
| 54 | SIO_GET_EXTENSION_FUNCTION_POINTER, |
| 55 | const_cast<GUID*>(&ConnectExGuid), |
| 56 | sizeof(ConnectExGuid), |
| 57 | &ConnectFn, |
| 58 | sizeof(ConnectFn), |
| 59 | &BytesReturned, |
| 60 | &Overlapped, |
| 61 | nullptr); |
| 62 | |
| 63 | if (Result != 0) |
| 64 | { |
| 65 | socket::GetResult(Socket.get(), Overlapped, INFINITE, ExitHandle, Location); |
| 66 | } |
| 67 | |
| 68 | THROW_LAST_ERROR_IF_MSG( |
| 69 | setsockopt(Socket.get(), HV_PROTOCOL_RAW, HVSOCKET_CONNECT_TIMEOUT, reinterpret_cast<char*>(&Timeout), sizeof(Timeout)) == SOCKET_ERROR, |
| 70 | "Timeout: %lu", |
| 71 | Timeout); |
| 72 | |
| 73 | SOCKADDR_HV Addr; |
| 74 | InitializeWildcardSocketAddress(&Addr); |
| 75 | THROW_LAST_ERROR_IF(bind(Socket.get(), reinterpret_cast<sockaddr*>(&Addr), sizeof(Addr)) == SOCKET_ERROR); |
| 76 | InitializeSocketAddress(VmId, Port, &Addr); |
| 77 | OverlappedEvent.ResetEvent(); |
| 78 | const BOOL Success = ConnectFn(Socket.get(), reinterpret_cast<sockaddr*>(&Addr), sizeof(Addr), nullptr, 0, nullptr, &Overlapped); |
| 79 | if (Success == FALSE) |
| 80 | { |
| 81 | socket::GetResult(Socket.get(), Overlapped, INFINITE, ExitHandle, Location); |
| 82 | } |
| 83 | |
| 84 | // Mark the socket as connected (required to call shutdown() later). |
| 85 | THROW_LAST_ERROR_IF(setsockopt(Socket.get(), SOL_SOCKET, SO_UPDATE_CONNECT_CONTEXT, nullptr, 0) == SOCKET_ERROR); |
| 86 | |
| 87 | return Socket; |
| 88 | } |
| 89 | |
| 90 | wil::unique_socket wsl::windows::common::hvsocket::Create() |
| 91 | { |
| 92 | wil::unique_socket Socket(WSASocket(AF_HYPERV, SOCK_STREAM, HV_PROTOCOL_RAW, nullptr, 0, WSA_FLAG_OVERLAPPED)); |
| 93 | THROW_LAST_ERROR_IF(!Socket); |
| 94 | |
| 95 | ULONG Enable = 1; |
| 96 | THROW_LAST_ERROR_IF( |
| 97 | setsockopt(Socket.get(), HV_PROTOCOL_RAW, HVSOCKET_CONNECTED_SUSPEND, reinterpret_cast<char*>(&Enable), sizeof(Enable)) == SOCKET_ERROR); |
| 98 | |
| 99 | return Socket; |
| 100 | } |
| 101 | |
| 102 | wil::unique_socket wsl::windows::common::hvsocket::Listen(_In_ const GUID& VmId, _In_ unsigned long Port, _In_ int Backlog) |
| 103 | { |
| 104 | SOCKADDR_HV Addr; |
| 105 | InitializeSocketAddress(VmId, Port, &Addr); |
| 106 | auto Socket = Create(); |
| 107 | THROW_LAST_ERROR_IF(bind(Socket.get(), reinterpret_cast<sockaddr*>(&Addr), sizeof(Addr)) == SOCKET_ERROR); |
| 108 | THROW_LAST_ERROR_IF(listen(Socket.get(), Backlog) == SOCKET_ERROR); |
| 109 | return Socket; |
| 110 | } |