master
cpp 110 lines 3.46 KB
Raw
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 }