master
cpp 138 lines 4.48 KB
Raw
1 /*++
2
3 Copyright (c) Microsoft. All rights reserved.
4
5 Module Name:
6
7 PortRelayHandle.cpp
8
9 Abstract:
10
11 Contains the implementation of the PortRelayAcceptHandle class.
12
13 --*/
14
15 #include "PortRelayHandle.h"
16 #include "IORelay.h"
17 #include "hvsocket.hpp"
18 #include "socket.hpp"
19 #include "wslutil.h"
20 #include "lxinitshared.h"
21 #include <gslhelpers.h>
22 #include <mswsock.h>
23 #include <thread>
24
25 using namespace wsl::windows::service::wslc;
26 using namespace wsl::windows::common;
27
28 PortRelayAcceptHandle::PortRelayAcceptHandle(
29 wil::unique_socket&& ListenSocket, const GUID& VmId, uint32_t RelayPort, uint32_t LinuxPort, int Family, IORelay& IoRelay) :
30 ListenSocket(std::move(ListenSocket)), VmId(VmId), RelayPort(RelayPort), LinuxPort(LinuxPort), Family(Family), IoRelay(IoRelay)
31 {
32 Overlapped.hEvent = Event.get();
33 }
34
35 PortRelayAcceptHandle::~PortRelayAcceptHandle()
36 {
37 if (State == io::IOHandleStatus::Pending)
38 {
39 LOG_IF_WIN32_BOOL_FALSE(CancelIoEx(reinterpret_cast<HANDLE>(ListenSocket.get()), &Overlapped));
40
41 DWORD bytesProcessed{};
42 DWORD flagsReturned{};
43 if (!WSAGetOverlappedResult(ListenSocket.get(), &Overlapped, &bytesProcessed, TRUE, &flagsReturned))
44 {
45 auto error = GetLastError();
46 LOG_LAST_ERROR_IF(error != ERROR_CONNECTION_ABORTED && error != ERROR_OPERATION_ABORTED);
47 }
48 }
49 }
50
51 void PortRelayAcceptHandle::Schedule()
52 {
53 WI_ASSERT(State == io::IOHandleStatus::Standby);
54
55 // Create a new socket for accepting
56 AcceptedSocket.reset(WSASocket(Family, SOCK_STREAM, IPPROTO_TCP, nullptr, 0, WSA_FLAG_OVERLAPPED));
57 THROW_LAST_ERROR_IF(!AcceptedSocket);
58
59 memset(AcceptBuffer, 0, sizeof(AcceptBuffer));
60 DWORD bytesReturned{};
61 if (AcceptEx(ListenSocket.get(), AcceptedSocket.get(), AcceptBuffer, 0, sizeof(SOCKADDR_STORAGE), sizeof(SOCKADDR_STORAGE), &bytesReturned, &Overlapped))
62 {
63 // Accept completed immediately
64 State = io::IOHandleStatus::Completed;
65 }
66 else
67 {
68 auto error = WSAGetLastError();
69 THROW_HR_IF_MSG(HRESULT_FROM_WIN32(error), error != ERROR_IO_PENDING, "Handle: 0x%p", reinterpret_cast<void*>(ListenSocket.get()));
70
71 State = io::IOHandleStatus::Pending;
72 }
73 }
74
75 void PortRelayAcceptHandle::Collect()
76 {
77 WI_ASSERT(State == io::IOHandleStatus::Pending || State == io::IOHandleStatus::Completed);
78
79 if (State == io::IOHandleStatus::Pending)
80 {
81 DWORD bytesReceived{};
82 DWORD flagsReturned{};
83 THROW_IF_WIN32_BOOL_FALSE(WSAGetOverlappedResult(ListenSocket.get(), &Overlapped, &bytesReceived, false, &flagsReturned));
84 }
85
86 // Set the accept context to mark the socket as connected.
87 socket::SetAcceptContext(AcceptedSocket.get(), ListenSocket.get());
88
89 // Launch a relay for this accepted connection
90 LaunchRelay(std::move(AcceptedSocket));
91
92 // Go back to standby to accept the next connection
93 State = io::IOHandleStatus::Standby;
94 }
95
96 HANDLE PortRelayAcceptHandle::GetHandle() const
97 {
98 return Event.get();
99 }
100
101 void PortRelayAcceptHandle::LaunchRelay(wil::unique_socket&& AcceptedSocket)
102 {
103 WSL_LOG(
104 "StartPortRelay",
105 TraceLoggingValue(LinuxPort, "LinuxPort"),
106 TraceLoggingValue(Family, "Family"),
107 TraceLoggingValue(AcceptedSocket.get(), "Socket"));
108
109 // Launch relay in a dedicated thread
110 std::thread relayThread{
111 [Socket = std::move(AcceptedSocket), VmId = VmId, LinuxPort = LinuxPort, RelayPort = RelayPort, Family = Family]() mutable {
112 try
113 {
114 wslutil::SetThreadDescription(L"Port relay");
115
116 // Connect to the HvSocket
117 auto hvSocket = hvsocket::Connect(VmId, RelayPort);
118
119 // Send relay start message
120 LX_INIT_START_SOCKET_RELAY message{};
121 message.Header.MessageType = LxInitMessageStartSocketRelay;
122 message.Header.MessageSize = sizeof(message);
123 message.Family = (Family == AF_INET) ? LX_AF_INET : LX_AF_INET6;
124 message.Port = LinuxPort;
125 message.BufferSize = 0x20000; // LOCALHOST_RELAY_BUFFER_SIZE
126
127 socket::Send(hvSocket.get(), gslhelpers::struct_as_bytes(message));
128
129 // Relay data between the two sockets
130 relay::SocketRelay(Socket.get(), hvSocket.get(), message.BufferSize);
131
132 WSL_LOG("StopPortRelay", TraceLoggingValue(LinuxPort, "LinuxPort"), TraceLoggingValue(Socket.get(), "Socket"));
133 }
134 CATCH_LOG();
135 }};
136
137 relayThread.detach();
138 }