master
h 128 lines 3.34 KB
Raw
1 /*++
2
3 Copyright (c) Microsoft. All rights reserved.
4
5 Module Name:
6
7 socketshared.h
8
9 Abstract:
10
11 This file contains shared socket helper functions.
12
13 --*/
14
15 #pragma once
16 #include <cassert>
17
18 namespace wsl::shared::socket {
19
20 #if defined(_MSC_VER)
21 inline gsl::span<gsl::byte> RecvMessage(SOCKET Socket, std::vector<gsl::byte>& Buffer, std::optional<HANDLE> ExitHandle = {}, DWORD Timeout = INFINITE)
22 #elif defined(__GNUC__)
23 inline gsl::span<gsl::byte> RecvMessage(int Socket, std::vector<gsl::byte>& Buffer, const timeval* Timeout = nullptr)
24 #endif
25 try
26 {
27 auto MessageSize = sizeof(MESSAGE_HEADER);
28 if (Buffer.size() < MessageSize)
29 {
30 Buffer.resize(MessageSize);
31 }
32
33 auto Message = gsl::make_span(Buffer.data(), MessageSize);
34 #if defined(_MSC_VER)
35 auto BytesRead = wsl::windows::common::socket::Receive(Socket, Message, ExitHandle.value_or(nullptr), MSG_WAITALL, Timeout);
36 #elif defined(__GNUC__)
37 // 'Timeout' is not implemented on Linux.
38 assert(Timeout == nullptr);
39
40 auto BytesRead = TEMP_FAILURE_RETRY(recv(Socket, Message.data(), Message.size(), MSG_WAITALL));
41 THROW_LAST_ERROR_IF(BytesRead < 0);
42 #endif
43 if (BytesRead == 0)
44 {
45 return {};
46 }
47 else if (BytesRead < MessageSize)
48 {
49 #if defined(_MSC_VER)
50 THROW_HR(E_UNEXPECTED);
51 #elif defined(__GNUC__)
52 THROW_UNEXPECTED();
53 #endif
54 }
55
56 // Grow the message buffer if needed and read the rest of the message.
57 MessageSize = gslhelpers::get_struct<MESSAGE_HEADER>(Message)->MessageSize;
58 if (MessageSize < sizeof(MESSAGE_HEADER))
59 {
60 #if defined(_MSC_VER)
61 THROW_HR_MSG(E_UNEXPECTED, "Unexpected message size: %llu", MessageSize);
62 #elif defined(__GNUC__)
63 THROW_UNEXPECTED();
64 #endif
65 }
66
67 if (MessageSize > 16 * 1024 * 1024) // 16 MiB
68 {
69 #if defined(_MSC_VER)
70 THROW_HR_MSG(E_UNEXPECTED, "Message size too large: %llu", MessageSize);
71 #elif defined(__GNUC__)
72 THROW_UNEXPECTED();
73 #endif
74 }
75
76 if (Buffer.size() < MessageSize)
77 {
78 Buffer.resize(MessageSize);
79 }
80
81 Message = gsl::make_span(Buffer.data(), MessageSize).subspan(sizeof(MESSAGE_HEADER));
82 while (Message.size() > 0)
83 {
84 #if defined(_MSC_VER)
85 BytesRead = wsl::windows::common::socket::Receive(Socket, Message, ExitHandle.value_or(nullptr), 0);
86 #elif defined(__GNUC__)
87 BytesRead = TEMP_FAILURE_RETRY(recv(Socket, Message.data(), Message.size(), 0));
88 THROW_LAST_ERROR_IF(BytesRead < 0);
89 #endif
90 if (BytesRead <= 0)
91 {
92 const auto* Header = reinterpret_cast<const MESSAGE_HEADER*>(Buffer.data());
93
94 #if defined(_MSC_VER)
95
96 LOG_HR_MSG(
97 E_UNEXPECTED,
98 "Socket closed while reading message. Size: %u, type: %i, id: %u",
99 Header->MessageSize,
100 Header->MessageType,
101 Header->TransactionId);
102
103 #elif defined(__GNUC__)
104
105 LOG_ERROR(
106 "Socket closed while reading message. Size: {}, type: {}, id: {}",
107 Header->MessageSize,
108 Header->MessageType,
109 Header->TransactionId);
110
111 #endif
112
113 return {};
114 }
115
116 Message = Message.subspan(BytesRead);
117 }
118
119 return gsl::make_span(Buffer.data(), MessageSize);
120 }
121 catch (...)
122 {
123 LOG_CAUGHT_EXCEPTION();
124 errno = wil::ResultFromCaughtException();
125 return {};
126 }
127
128 } // namespace wsl::shared::socket