master
cpp 203 lines 6.92 KB
Raw
1 /*++
2
3 Copyright (c) Microsoft. All rights reserved.
4
5 Module Name:
6
7 socket.cpp
8
9 Abstract:
10
11 This file contains socket helper function definitions.
12
13 --*/
14
15 #include "precomp.h"
16 #include <mutex>
17 #include "socket.hpp"
18 #pragma hdrstop
19
20 void wsl::windows::common::socket::SetAcceptContext(_In_ SOCKET AcceptedSocket, _In_ SOCKET ListenSocket, _In_ const std::source_location& Location)
21 {
22 // Set the accept context to mark the socket as connected.
23 THROW_LAST_ERROR_IF_MSG(
24 setsockopt(AcceptedSocket, SOL_SOCKET, SO_UPDATE_ACCEPT_CONTEXT, reinterpret_cast<const char*>(&ListenSocket), sizeof(ListenSocket)) == SOCKET_ERROR,
25 "From: %hs",
26 std::format("{}", Location).c_str());
27 }
28
29 std::optional<wil::unique_socket> wsl::windows::common::socket::CancellableAccept(
30 _In_ SOCKET ListenSocket, _In_ DWORD Timeout, _In_opt_ HANDLE ExitHandle, _In_ const std::source_location& Location)
31 {
32 io::MultiHandleWait io;
33
34 std::optional<wil::unique_socket> accepted;
35
36 io.AddHandle(
37 std::make_unique<io::AcceptHandle>(
38 ListenSocket, true, [&accepted](wil::unique_socket&& socket) { accepted = std::move(socket); }),
39 io::MultiHandleWait::CancelOnCompleted);
40
41 if (ExitHandle != nullptr)
42 {
43 io.AddHandle(std::make_unique<io::EventHandle>(ExitHandle), io::MultiHandleWait::CancelOnCompleted);
44 }
45
46 std::optional<std::chrono::milliseconds> timeout;
47 if (Timeout != INFINITE)
48 {
49 timeout = std::chrono::milliseconds(Timeout);
50 }
51
52 try
53 {
54
55 io.Run(timeout);
56 }
57 catch (...)
58 {
59 auto hr = wil::ResultFromCaughtException();
60 THROW_HR_MSG(hr, "Failed to accept socket. From: %hs", std::format("{}", Location).c_str());
61 }
62
63 return accepted;
64 }
65
66 std::pair<DWORD, DWORD> wsl::windows::common::socket::GetResult(
67 _In_ SOCKET Socket, _In_ OVERLAPPED& Overlapped, _In_ DWORD Timeout, _In_ HANDLE ExitHandle, _In_ const std::source_location& Location)
68 {
69 const int error = WSAGetLastError();
70 THROW_HR_IF(HRESULT_FROM_WIN32(error), error != WSA_IO_PENDING);
71
72 std::vector<HANDLE> waitObjects{};
73 waitObjects.push_back(Overlapped.hEvent);
74 if (ARGUMENT_PRESENT(ExitHandle))
75 {
76 waitObjects.push_back(ExitHandle);
77 }
78
79 DWORD bytesProcessed;
80 DWORD flagsReturned;
81 auto cancelFunction = wil::scope_exit_log(WI_DIAGNOSTICS_INFO, [&] {
82 CancelIoEx(reinterpret_cast<HANDLE>(Socket), &Overlapped);
83 WSAGetOverlappedResult(Socket, &Overlapped, &bytesProcessed, TRUE, &flagsReturned);
84 });
85
86 const DWORD waitStatus = WaitForMultipleObjects(gsl::narrow_cast<DWORD>(waitObjects.size()), waitObjects.data(), FALSE, Timeout);
87 if (waitObjects.size() > 1 && waitStatus == WAIT_OBJECT_0 + 1)
88 {
89 return {0, 0};
90 }
91
92 THROW_HR_IF_MSG(HCS_E_CONNECTION_TIMEOUT, (waitStatus != WAIT_OBJECT_0), "From: %hs", std::format("{}", Location).c_str());
93
94 cancelFunction.release();
95 const bool result = WSAGetOverlappedResult(Socket, &Overlapped, &bytesProcessed, FALSE, &flagsReturned);
96 if (!result)
97 {
98 const auto lastError = WSAGetLastError();
99 if (lastError != WSAECONNABORTED || (ExitHandle != nullptr && WaitForSingleObject(ExitHandle, 0) == WAIT_TIMEOUT))
100 {
101 THROW_WIN32(lastError);
102 }
103 else
104 {
105 return {0, 0};
106 }
107 }
108 return {bytesProcessed, flagsReturned};
109 }
110
111 int wsl::windows::common::socket::Receive(
112 _In_ SOCKET Socket, _In_ gsl::span<gsl::byte> Buffer, _In_opt_ HANDLE ExitHandle, _In_ DWORD Flags, _In_ DWORD Timeout, _In_ const std::source_location& Location)
113 {
114 const int BytesRead = ReceiveNoThrow(Socket, Buffer, ExitHandle, Flags, Timeout, Location);
115 THROW_LAST_ERROR_IF(BytesRead == SOCKET_ERROR);
116
117 return BytesRead;
118 }
119
120 int wsl::windows::common::socket::ReceiveNoThrow(
121 _In_ SOCKET Socket, _In_ gsl::span<gsl::byte> Buffer, _In_opt_ HANDLE ExitHandle, _In_ DWORD Flags, _In_ DWORD Timeout, _In_ const std::source_location& Location)
122 {
123 OVERLAPPED Overlapped{};
124 const wil::unique_event OverlappedEvent(wil::EventOptions::ManualReset);
125 WSABUF VectorBuffer = {gsl::narrow_cast<ULONG>(Buffer.size()), reinterpret_cast<CHAR*>(Buffer.data())};
126 Overlapped.hEvent = OverlappedEvent.get();
127 DWORD BytesReturned{};
128 if (WSARecv(Socket, &VectorBuffer, 1, &BytesReturned, &Flags, &Overlapped, nullptr) != 0)
129 {
130 try
131 {
132 BytesReturned = SOCKET_ERROR;
133 auto [innerBytes, Flags] = GetResult(Socket, Overlapped, Timeout, ExitHandle, Location);
134 BytesReturned = innerBytes;
135 }
136 catch (...)
137 {
138 LOG_CAUGHT_EXCEPTION();
139 // Receive will call GetLastError to look for the error code
140 SetLastError(wil::ResultFromCaughtException());
141 }
142 }
143
144 return BytesReturned;
145 }
146
147 std::vector<gsl::byte> wsl::windows::common::socket::Receive(
148 _In_ SOCKET Socket, _In_opt_ HANDLE ExitHandle, _In_ DWORD Timeout, _In_ const std::source_location& Location)
149 {
150 Receive(Socket, {}, ExitHandle, MSG_PEEK, Timeout, Location);
151
152 ULONG Size = 0;
153 THROW_LAST_ERROR_IF(ioctlsocket(Socket, FIONREAD, &Size) == SOCKET_ERROR);
154
155 std::vector<gsl::byte> Buffer(Size);
156 WI_VERIFY(Receive(Socket, gsl::make_span(Buffer), ExitHandle, MSG_WAITALL, Timeout, Location) == static_cast<int>(Size));
157
158 return Buffer;
159 }
160
161 int wsl::windows::common::socket::Send(
162 _In_ SOCKET Socket, _In_ gsl::span<const gsl::byte> Buffer, _In_opt_ HANDLE ExitHandle, _In_ const std::source_location& Location)
163 {
164 const wil::unique_event OverlappedEvent(wil::EventOptions::ManualReset);
165 OVERLAPPED Overlapped{};
166 Overlapped.hEvent = OverlappedEvent.get();
167
168 DWORD Offset = 0;
169 while (Offset < Buffer.size())
170 {
171 OverlappedEvent.ResetEvent();
172
173 WSABUF VectorBuffer = {
174 gsl::narrow_cast<ULONG>(Buffer.size() - Offset), const_cast<CHAR*>(reinterpret_cast<const CHAR*>(Buffer.data() + Offset))};
175
176 DWORD BytesWritten{};
177 if (WSASend(Socket, &VectorBuffer, 1, &BytesWritten, 0, &Overlapped, nullptr) != 0)
178 {
179 // If WSASend returns non-zero, expect WSA_IO_PENDING.
180 if (auto error = WSAGetLastError(); error != WSA_IO_PENDING)
181 {
182 THROW_WIN32_MSG(error, "WSASend failed. From: %hs", std::format("{}", Location).c_str());
183 }
184
185 DWORD Flags;
186 std::tie(BytesWritten, Flags) = GetResult(Socket, Overlapped, INFINITE, ExitHandle, Location);
187 if (BytesWritten == 0)
188 {
189 THROW_WIN32_MSG(ERROR_CONNECTION_ABORTED, "Socket closed during WSASend(). From: %hs", std::format("{}", Location).c_str());
190 }
191 }
192
193 Offset += BytesWritten;
194 if (Offset < Buffer.size())
195 {
196 WSL_LOG("PartialSocketWrite", TraceLoggingValue(Buffer.size(), "MessageSize"), TraceLoggingValue(Offset, "Offset"));
197 }
198 }
199
200 WI_ASSERT(Offset == gsl::narrow_cast<DWORD>(Buffer.size()));
201
202 return Offset;
203 }