master
cpp 417 lines 17 KB
Raw
1 // Copyright (C) Microsoft Corporation. All rights reserved.
2
3 #include <LxssDynamicFunction.h>
4 #include "precomp.h"
5 #include "DnsResolver.h"
6
7 using wsl::core::networking::DnsResolver;
8
9 static constexpr auto c_dnsModuleName = L"dnsapi.dll";
10
11 std::optional<LxssDynamicFunction<decltype(DnsQueryRaw)>> DnsResolver::s_dnsQueryRaw;
12 std::optional<LxssDynamicFunction<decltype(DnsCancelQueryRaw)>> DnsResolver::s_dnsCancelQueryRaw;
13 std::optional<LxssDynamicFunction<decltype(DnsQueryRawResultFree)>> DnsResolver::s_dnsQueryRawResultFree;
14
15 HRESULT DnsResolver::LoadDnsResolverMethods() noexcept
16 {
17 static wil::shared_hmodule dnsModule;
18 static DWORD loadError = ERROR_SUCCESS;
19 static std::once_flag dnsLoadFlag;
20
21 // Load DNS dll only once
22 std::call_once(dnsLoadFlag, [&]() {
23 dnsModule.reset(LoadLibraryEx(c_dnsModuleName, nullptr, LOAD_LIBRARY_SEARCH_SYSTEM32));
24 if (!dnsModule)
25 {
26 loadError = GetLastError();
27 }
28 });
29
30 RETURN_IF_WIN32_ERROR_MSG(loadError, "LoadLibraryEx %ls", c_dnsModuleName);
31
32 // Initialize dynamic functions for the DNS tunneling Windows APIs.
33 // using the non-throwing instance of LxssDynamicFunction as to not end up in the Error telemetry
34 LxssDynamicFunction<decltype(DnsQueryRaw)> local_dnsQueryRaw{DynamicFunctionErrorLogs::None};
35 RETURN_IF_FAILED_EXPECTED(local_dnsQueryRaw.load(dnsModule, "DnsQueryRaw"));
36 LxssDynamicFunction<decltype(DnsCancelQueryRaw)> local_dnsCancelQueryRaw{DynamicFunctionErrorLogs::None};
37 RETURN_IF_FAILED_EXPECTED(local_dnsCancelQueryRaw.load(dnsModule, "DnsCancelQueryRaw"));
38 LxssDynamicFunction<decltype(DnsQueryRawResultFree)> local_dnsQueryRawResultFree{DynamicFunctionErrorLogs::None};
39 RETURN_IF_FAILED_EXPECTED(local_dnsQueryRawResultFree.load(dnsModule, "DnsQueryRawResultFree"));
40
41 // Make a dummy call to the DNS APIs to verify if they are working. The APIs are going to be present
42 // on older Windows versions, where they can be turned on/off. If turned off, the APIs
43 // will be unusable and will return ERROR_CALL_NOT_IMPLEMENTED.
44 if (local_dnsQueryRaw(nullptr, nullptr) == ERROR_CALL_NOT_IMPLEMENTED)
45 {
46 RETURN_IF_WIN32_ERROR_EXPECTED(ERROR_CALL_NOT_IMPLEMENTED);
47 }
48
49 s_dnsQueryRaw.emplace(std::move(local_dnsQueryRaw));
50 s_dnsCancelQueryRaw.emplace(std::move(local_dnsCancelQueryRaw));
51 s_dnsQueryRawResultFree.emplace(std::move(local_dnsQueryRawResultFree));
52 return S_OK;
53 }
54
55 DnsResolver::DnsResolver(wil::unique_socket&& dnsHvsocket, DnsResolverFlags flags) :
56 m_dnsChannel(
57 std::move(dnsHvsocket),
58 [this](const gsl::span<gsl::byte> dnsBuffer, const LX_GNS_DNS_CLIENT_IDENTIFIER& dnsClientIdentifier) {
59 ProcessDnsRequest(dnsBuffer, dnsClientIdentifier);
60 }),
61 m_flags(flags)
62 {
63 // Initialize as signaled, as there are no requests yet
64 m_allRequestsFinished.SetEvent();
65
66 // Read external interface constraint regkey
67 const auto lxssKey = windows::common::registry::OpenLxssMachineKey(KEY_READ);
68 m_externalInterfaceConstraintName =
69 windows::common::registry::ReadString(lxssKey.get(), nullptr, c_interfaceConstraintKey, L"");
70
71 if (!m_externalInterfaceConstraintName.empty())
72 {
73 ResolveExternalInterfaceConstraintIndex();
74
75 WSL_LOG(
76 "DnsResolver::DnsResolver",
77 TraceLoggingValue(m_externalInterfaceConstraintName.c_str(), "m_externalInterfaceConstraintName"),
78 TraceLoggingValue(m_externalInterfaceConstraintIndex, "m_externalInterfaceConstraintIndex"));
79
80 // Register for interface change notifications. Notifications are used to determine if the external interface constraint setting is applicable.
81 THROW_IF_WIN32_ERROR(NotifyIpInterfaceChange(AF_UNSPEC, &DnsResolver::InterfaceChangeCallback, this, FALSE, &m_interfaceNotificationHandle));
82 }
83 }
84
85 DnsResolver::~DnsResolver() noexcept
86 {
87 Stop();
88 }
89
90 void DnsResolver::GenerateTelemetry() noexcept
91 try
92 {
93 // Find the 3 most common DNS API failures
94 uint32_t mostCommonDnsStatusError = 0;
95 uint32_t mostCommonDnsStatusErrorCount = 0;
96 uint32_t secondCommonDnsStatusError = 0;
97 uint32_t secondCommonDnsStatusErrorCount = 0;
98 uint32_t thirdCommonDnsStatusError = 0;
99 uint32_t thirdCommonDnsStatusErrorCount = 0;
100
101 std::vector<std::pair<uint32_t, uint32_t>> failures(m_dnsApiFailures.size());
102 std::copy(m_dnsApiFailures.begin(), m_dnsApiFailures.end(), failures.begin());
103
104 // Sort in descending order based on failure count
105 std::sort(failures.begin(), failures.end(), [](const auto& lhs, const auto& rhs) { return lhs.second > rhs.second; });
106
107 if (failures.size() >= 1)
108 {
109 mostCommonDnsStatusError = failures[0].first;
110 mostCommonDnsStatusErrorCount = failures[0].second;
111 }
112 if (failures.size() >= 2)
113 {
114 secondCommonDnsStatusError = failures[1].first;
115 secondCommonDnsStatusErrorCount = failures[1].second;
116 }
117 if (failures.size() >= 3)
118 {
119 thirdCommonDnsStatusError = failures[2].first;
120 thirdCommonDnsStatusErrorCount = failures[2].second;
121 }
122
123 // Add telemetry with DNS tunneling statistics, before shutting down
124 WSL_LOG(
125 "DnsTunnelingStatistics",
126 TraceLoggingValue(m_totalUdpQueries.load(), "totalUdpQueries"),
127 TraceLoggingValue(m_successfulUdpQueries.load(), "successfulUdpQueries"),
128 TraceLoggingValue(m_totalTcpQueries.load(), "totalTcpQueries"),
129 TraceLoggingValue(m_successfulTcpQueries.load(), "successfulTcpQueries"),
130 TraceLoggingValue(m_queriesWithNullResult.load(), "queriesWithNullResult"),
131 TraceLoggingValue(m_failedDnsQueryRawCalls.load(), "FailedDnsQueryRawCalls"),
132 TraceLoggingValue(m_dnsApiFailures.size(), "totalDnsStatusErrorInstances"),
133 TraceLoggingValue(mostCommonDnsStatusError, "mostCommonDnsStatusError"),
134 TraceLoggingValue(mostCommonDnsStatusErrorCount, "mostCommonDnsStatusErrorCount"),
135 TraceLoggingValue(secondCommonDnsStatusError, "secondCommonDnsStatusError"),
136 TraceLoggingValue(secondCommonDnsStatusErrorCount, "secondCommonDnsStatusErrorCount"),
137 TraceLoggingValue(thirdCommonDnsStatusError, "thirdCommonDnsStatusError"),
138 TraceLoggingValue(thirdCommonDnsStatusErrorCount, "thirdCommonDnsStatusErrorCount"));
139 }
140 CATCH_LOG()
141
142 void DnsResolver::Stop() noexcept
143 try
144 {
145 WSL_LOG("DnsResolver::Stop");
146
147 // Scoped m_dnsLock
148 {
149 const std::lock_guard lock(m_dnsLock);
150
151 m_stopped = true;
152
153 // Cancel existing requests. Cancel is complete when DnsQueryRawCallback is
154 // invoked with status == ERROR_CANCELLED
155 // N.B. Cancelling can end up calling the DnsQueryRawCallback directly on this same thread. i.e., while this
156 // lock is held. Which is fine because m_dnsLock is a recursive mutex.
157 // N.B. Cancelling a query will synchronously remove the query from m_dnsRequests, which invalidates iterators.
158
159 std::vector<DNS_QUERY_RAW_CANCEL*> cancelHandles;
160 cancelHandles.reserve(m_dnsRequests.size());
161
162 for (auto& [_, context] : m_dnsRequests)
163 {
164 cancelHandles.emplace_back(&context->m_cancelHandle);
165 }
166
167 for (const auto e : cancelHandles)
168 {
169 LOG_IF_WIN32_ERROR(s_dnsCancelQueryRaw.value()(e));
170 }
171 }
172
173 // Wait for all requests to complete. At this point no new requests can be started since the object is stopped.
174 // We are only waiting for existing requests to finish.
175 m_allRequestsFinished.wait();
176
177 // Stop the response queue first as it can make calls in m_dnsChannel
178 m_dnsResponseQueue.cancel();
179
180 m_dnsChannel.Stop();
181
182 // Stop interface change notifications
183 m_interfaceNotificationHandle.reset();
184
185 GenerateTelemetry();
186 }
187 CATCH_LOG()
188
189 void DnsResolver::ProcessDnsRequest(const gsl::span<gsl::byte> dnsBuffer, const LX_GNS_DNS_CLIENT_IDENTIFIER& dnsClientIdentifier) noexcept
190 try
191 {
192 const std::lock_guard lock(m_dnsLock);
193 if (m_stopped)
194 {
195 return;
196 }
197
198 WSL_LOG_DEBUG(
199 "DnsResolver::ProcessDnsRequest - received new DNS request",
200 TraceLoggingValue(dnsBuffer.size(), "DNS buffer size"),
201 TraceLoggingValue(dnsClientIdentifier.Protocol == IPPROTO_UDP ? "UDP" : "TCP", "Protocol"),
202 TraceLoggingValue(dnsClientIdentifier.DnsClientId, "DNS client id"),
203 TraceLoggingValue(!m_externalInterfaceConstraintName.empty(), "Is ExternalInterfaceConstraint configured"),
204 TraceLoggingValue(m_externalInterfaceConstraintIndex, "m_externalInterfaceConstraintIndex"));
205
206 // If the external interface constraint is configured but it is *not* present/up, WSL should be net-blind, so we avoid making DNS requests.
207 if (!m_externalInterfaceConstraintName.empty() && m_externalInterfaceConstraintIndex == 0)
208 {
209 return;
210 }
211
212 dnsClientIdentifier.Protocol == IPPROTO_UDP ? m_totalUdpQueries++ : m_totalTcpQueries++;
213
214 // Get next request id. If value reaches UINT_MAX + 1 it will be automatically reset to 0
215 const auto requestId = m_currentRequestId++;
216
217 // Create the DNS request context
218 auto context = std::make_unique<DnsResolver::DnsQueryContext>(
219 requestId, dnsClientIdentifier, [this](_Inout_ DnsResolver::DnsQueryContext* context, _Inout_opt_ DNS_QUERY_RAW_RESULT* queryResults) {
220 HandleDnsQueryCompletion(context, queryResults);
221 });
222
223 auto [it, _] = m_dnsRequests.emplace(requestId, std::move(context));
224 const auto localContext = it->second.get();
225
226 auto removeContextOnError = wil::scope_exit([&] { WI_VERIFY(m_dnsRequests.erase(requestId) == 1); });
227
228 // Fill DNS request structure
229 DNS_QUERY_RAW_REQUEST request{};
230
231 request.version = DNS_QUERY_RAW_REQUEST_VERSION1;
232 request.resultsVersion = DNS_QUERY_RAW_RESULTS_VERSION1;
233 request.dnsQueryRawSize = static_cast<ULONG>(dnsBuffer.size());
234 request.dnsQueryRaw = (PBYTE)dnsBuffer.data();
235 request.protocol = (dnsClientIdentifier.Protocol == IPPROTO_TCP) ? DNS_PROTOCOL_TCP : DNS_PROTOCOL_UDP;
236 request.queryCompletionCallback = DnsResolver::DnsQueryRawCallback;
237 request.queryContext = localContext;
238 // Only unicast UDP & TCP queries are tunneled. Pass this flag to tell Windows DNS client to *not* resolve using multicast.
239 request.queryOptions |= DNS_QUERY_NO_MULTICAST;
240
241 // In a DNS request from Linux there might be DNS records that Windows DNS client does not know how to parse.
242 // By default in this case Windows will fail the request. When the flag is enabled, Windows will extract the
243 // question from the DNS request and attempt to resolve it, ignoring the unknown records.
244 if (WI_IsFlagSet(m_flags, DnsResolverFlags::BestEffortDnsParsing))
245 {
246 request.queryRawOptions |= DNS_QUERY_RAW_OPTION_BEST_EFFORT_PARSE;
247 }
248
249 // If the external interface constraint is configured and present on the host, only send DNS requests on that interface.
250 if (m_externalInterfaceConstraintIndex != 0)
251 {
252 request.interfaceIndex = m_externalInterfaceConstraintIndex;
253 }
254
255 // Start the DNS request
256 // N.B. All DNS requests will bypass the Windows DNS cache
257 const auto result = s_dnsQueryRaw.value()(&request, &localContext->m_cancelHandle);
258 if (result != DNS_REQUEST_PENDING)
259 {
260 m_failedDnsQueryRawCalls++;
261
262 WSL_LOG(
263 "ProcessDnsRequestFailed",
264 TraceLoggingValue(requestId, "requestId"),
265 TraceLoggingValue(result, "result"),
266 TraceLoggingValue("DnsQueryRaw", "executionStep"));
267 return;
268 }
269
270 removeContextOnError.release();
271
272 m_allRequestsFinished.ResetEvent();
273 }
274 CATCH_LOG()
275
276 void DnsResolver::HandleDnsQueryCompletion(_Inout_ DnsResolver::DnsQueryContext* queryContext, _Inout_opt_ DNS_QUERY_RAW_RESULT* queryResults) noexcept
277 try
278 {
279 // Always free the query result structure
280 const auto freeQueryResults = wil::scope_exit([&] {
281 if (queryResults != nullptr)
282 {
283 s_dnsQueryRawResultFree.value()(queryResults);
284 }
285 });
286
287 const std::lock_guard lock(m_dnsLock);
288
289 if (queryResults != nullptr)
290 {
291 WSL_LOG(
292 "DnsResolver::HandleDnsQueryCompletion",
293 TraceLoggingValue(queryContext->m_id, "queryContext->m_id"),
294 TraceLoggingValue(queryResults->queryStatus, "queryResults->queryStatus"),
295 TraceLoggingValue(queryResults->queryRawResponse != nullptr, "validResponse"));
296
297 // Note: The response may be valid even if queryResults->queryStatus is not 0, for example when the DNS server returns a negative response.
298 if (queryResults->queryRawResponse != nullptr)
299 {
300 queryContext->m_dnsClientIdentifier.Protocol == IPPROTO_UDP ? m_successfulUdpQueries++ : m_successfulTcpQueries++;
301 }
302 // the Windows DNS API returned failure
303 else
304 {
305 if (m_dnsApiFailures.find(queryResults->queryStatus) == m_dnsApiFailures.end())
306 {
307 m_dnsApiFailures[queryResults->queryStatus] = 1;
308 }
309 else
310 {
311 m_dnsApiFailures[queryResults->queryStatus]++;
312 }
313 }
314 }
315 else
316 {
317 WSL_LOG(
318 "DnsResolver::HandleDnsQueryCompletion - received a NULL queryResults",
319 TraceLoggingValue(queryContext->m_id, "queryContext->m_id"));
320 m_queriesWithNullResult++;
321 }
322
323 if (!m_stopped && queryResults != nullptr && queryResults->queryRawResponse != nullptr)
324 {
325 // Copy DNS response buffer
326 std::vector<gsl::byte> dnsResponse(queryResults->queryRawResponseSize);
327 CopyMemory(dnsResponse.data(), queryResults->queryRawResponse, queryResults->queryRawResponseSize);
328
329 WSL_LOG_DEBUG(
330 "DnsResolver::HandleDnsQueryCompletion - received new DNS response",
331 TraceLoggingValue(dnsResponse.size(), "DNS buffer size"),
332 TraceLoggingValue(queryContext->m_dnsClientIdentifier.Protocol == IPPROTO_UDP ? "UDP" : "TCP", "Protocol"),
333 TraceLoggingValue(queryContext->m_dnsClientIdentifier.DnsClientId, "DNS client id"));
334
335 // Schedule the DNS response to be sent to Linux
336 m_dnsResponseQueue.submit([this, dnsResponse = std::move(dnsResponse), dnsClientIdentifier = queryContext->m_dnsClientIdentifier]() mutable {
337 m_dnsChannel.SendDnsMessage(gsl::make_span(dnsResponse), dnsClientIdentifier);
338 });
339 }
340
341 // Stop tracking this DNS request and delete the request context
342 WI_VERIFY(m_dnsRequests.erase(queryContext->m_id) == 1);
343
344 // Set event if all tracked requests have finished
345 if (m_dnsRequests.empty())
346 {
347 m_allRequestsFinished.SetEvent();
348 }
349 }
350 CATCH_LOG()
351
352 void DnsResolver::ResolveExternalInterfaceConstraintIndex() noexcept
353 try
354 {
355 const std::lock_guard lock(m_dnsLock);
356 if (m_stopped)
357 {
358 return;
359 }
360
361 if (m_externalInterfaceConstraintName.empty())
362 {
363 return;
364 }
365
366 NET_LUID interfaceLuid{};
367 ULONG interfaceIndex = 0;
368
369 // Update the interface index on every exit path.
370 // The calls below to convert interface name to index will fail if the interface does not exist anymore,
371 // in which case we still need to reset the interface index to its default value of 0.
372 const auto setInterfaceIndex = wil::scope_exit([&] {
373 if (interfaceIndex != m_externalInterfaceConstraintIndex)
374 {
375 WSL_LOG(
376 "DnsResolver::ResolveExternalInterfaceConstraintIndex - setting m_externalInterfaceConstraintIndex to new value",
377 TraceLoggingValue(m_externalInterfaceConstraintIndex, "old interface index"),
378 TraceLoggingValue(interfaceIndex, "new interface index"));
379
380 m_externalInterfaceConstraintIndex = interfaceIndex;
381 }
382 });
383
384 // If external interface constraint is configured, query to see if it's present on the host.
385 auto errorCode = ConvertInterfaceAliasToLuid(m_externalInterfaceConstraintName.c_str(), &interfaceLuid);
386 if (FAILED_WIN32_LOG(errorCode))
387 {
388 return;
389 }
390
391 errorCode = ConvertInterfaceLuidToIndex(&interfaceLuid, reinterpret_cast<PNET_IFINDEX>(&interfaceIndex));
392 if (FAILED_WIN32_LOG(errorCode))
393 {
394 return;
395 }
396 }
397 CATCH_LOG()
398
399 VOID CALLBACK DnsResolver::DnsQueryRawCallback(_In_ VOID* queryContext, _Inout_opt_ DNS_QUERY_RAW_RESULT* queryResults) noexcept
400 try
401 {
402 assert(queryContext != nullptr);
403
404 const auto context = static_cast<DnsQueryContext*>(queryContext);
405
406 // Call into DnsResolver parent object to process the query result
407 context->m_handleQueryCompletion(context, queryResults);
408 }
409 CATCH_LOG()
410
411 VOID CALLBACK DnsResolver::InterfaceChangeCallback(_In_ PVOID context, PMIB_IPINTERFACE_ROW, MIB_NOTIFICATION_TYPE) noexcept
412 try
413 {
414 const auto dnsResolver = static_cast<DnsResolver*>(context);
415 dnsResolver->ResolveExternalInterfaceConstraintIndex();
416 }
417 CATCH_LOG()