MegaBites-AI/Windows-powershell
0372
1// Copyright (c) Microsoft Corporation.2// Licensed under the MIT License.3 4using System;5using System.Collections.Generic;6using System.Management.Automation.Language;7using System.Management.Automation.Subsystem;8using System.Management.Automation.Subsystem.Prediction;9using System.Threading;10using System.Net;11using System.Net.Sockets;12using System.Text;13using System.Reflection;14using System.Threading.Tasks;15using Xunit;16using Xunit.Abstractions;17 18namespace PSTests.Sequential19{20 public class RemoteHyperVTests21 {22 private static ITestOutputHelper _output;23 24 public RemoteHyperVTests(ITestOutputHelper output)25 {26 if (!System.Management.Automation.Platform.IsWindows)27 {28 throw new SkipException("RemoteHyperVTests are only supported on Windows.");29 }30 31 _output = output;32 }33 34 // Helper method to connect with retries35 private static void ConnectWithRetry(Socket client, IPAddress address, int port, ITestOutputHelper output, int maxRetries = 10)36 {37 int retryDelayMs = 500;38 int attempt = 0;39 bool connected = false;40 while (attempt < maxRetries && !connected)41 {42 try43 {44 client.Connect(address, port);45 connected = true;46 }47 catch (SocketException)48 {49 attempt++;50 if (attempt < maxRetries)51 {52 output?.WriteLine($"Connect attempt {attempt} failed, retrying in {retryDelayMs}ms...");53 Thread.Sleep(retryDelayMs);54 retryDelayMs *= 2;55 }56 else57 {58 output?.WriteLine($"Failed to connect after {maxRetries} attempts. This is most likely an intermittent failure due to environmental issues.");59 throw;60 }61 }62 }63 }64 65 private static void SendResponse(string name, Socket client, Queue<(byte[] bytes, int delayMs)> serverResponses)66 {67 if (serverResponses.Count > 0)68 {69 _output.WriteLine($"Mock {name} ----------------------------------------------------");70 var respTuple = serverResponses.Dequeue();71 var resp = respTuple.bytes;72 73 if (respTuple.delayMs > 0)74 {75 _output.WriteLine($"Mock {name} - delaying response by {respTuple.delayMs} ms");76 Thread.Sleep(respTuple.delayMs);77 }78 if (resp.Length > 0) {79 client.Send(resp, resp.Length, SocketFlags.None);80 _output.WriteLine($"Mock {name} - sent response: " + Encoding.ASCII.GetString(resp));81 }82 }83 }84 85 private static void StartHandshakeServer(86 string name,87 int port,88 IEnumerable<(string message, Encoding encoding)> expectedClientSends,89 IEnumerable<(string message, Encoding encoding)> serverResponses,90 bool verifyConnectionClosed,91 CancellationToken cancellationToken,92 bool sendFirst = false)93 {94 IEnumerable<(string message, Encoding encoding, int delayMs)> serverResponsesWithDelay = new List<(string message, Encoding encoding, int delayMs)>();95 foreach (var item in serverResponses)96 {97 ((List<(string message, Encoding encoding, int delayMs)>)serverResponsesWithDelay).Add((item.message, item.encoding, 1));98 }99 StartHandshakeServer(name, port, expectedClientSends, serverResponsesWithDelay, verifyConnectionClosed, cancellationToken, sendFirst);100 }101 102 private static void StartHandshakeServer(103 string name,104 int port,105 IEnumerable<(string message, Encoding encoding)> expectedClientSends,106 IEnumerable<(string message, Encoding encoding, int delayMs)> serverResponses,107 bool verifyConnectionClosed,108 CancellationToken cancellationToken,109 bool sendFirst = false)110 {111 var expectedMessages = new Queue<(string message, byte[] bytes, Encoding encoding)>();112 foreach (var item in expectedClientSends)113 {114 var itemBytes = item.encoding.GetBytes(item.message);115 expectedMessages.Enqueue((message: item.message, bytes: itemBytes, encoding: item.encoding));116 }117 118 var serverResponseBytes = new Queue<(byte[] bytes, int delayMs)>();119 foreach (var item in serverResponses)120 {121 (byte[] bytes, int delayMs) queueItem = (item.encoding.GetBytes(item.message), item.delayMs);122 serverResponseBytes.Enqueue(queueItem);123 }124 125 _output.WriteLine($"Mock {name} - starting listener on port {port} with {expectedMessages.Count} expected messages and {serverResponseBytes.Count} responses.");126 StartHandshakeServerImplementation(name, port, expectedMessages, serverResponseBytes, verifyConnectionClosed, cancellationToken, sendFirst);127 }128 129 private static void StartHandshakeServerImplementation(130 string name,131 int port,132 Queue<(string message, byte[] bytes, Encoding encoding)> expectedClientSends,133 Queue<(byte[] bytes, int delayMs)> serverResponses,134 bool verifyConnectionClosed,135 CancellationToken cancellationToken,136 bool sendFirst = false)137 {138 DateTime startTime = DateTime.UtcNow;139 var buffer = new byte[1024];140 var listener = new TcpListener(IPAddress.Loopback, port);141 listener.Start();142 try143 {144 using (var client = listener.AcceptSocket())145 {146 if (sendFirst)147 {148 // Send the first message from the serverResponses queue149 SendResponse(name, client, serverResponses);150 }151 152 while (expectedClientSends.Count > 0)153 {154 _output.WriteLine($"Mock {name} - time elapsed: {(DateTime.UtcNow - startTime).TotalMilliseconds} milliseconds");155 client.ReceiveTimeout = 2 * 1000; // 2 seconds timeout for receiving data156 cancellationToken.ThrowIfCancellationRequested();157 var expectedMessage = expectedClientSends.Dequeue();158 _output.WriteLine($"Mock {name} - remaining expected messages: {expectedClientSends.Count}");159 var expected = expectedMessage.bytes;160 Array.Clear(buffer, 0, buffer.Length);161 int received = client.Receive(buffer);162 // Optionally validate received data matches expected163 string expectedString = expectedMessage.message;164 string bufferString = expectedMessage.encoding.GetString(buffer, 0, received);165 string alternativeEncodedString = string.Empty;166 if (expectedMessage.encoding == Encoding.Unicode)167 {168 alternativeEncodedString = Encoding.UTF8.GetString(buffer, 0, received);169 }170 else if (expectedMessage.encoding == Encoding.UTF8)171 {172 alternativeEncodedString = Encoding.Unicode.GetString(buffer, 0, received);173 }174 175 if (received != expected.Length)176 {177 string errorMessage = $"Mock {name} - Expected {expected.Length} bytes, but received {received} bytes: `{bufferString}`(alt encoding: `{alternativeEncodedString}`); expected: {expectedString}";178 _output.WriteLine(errorMessage);179 throw new Exception(errorMessage);180 }181 if (!string.Equals(bufferString, expectedString, StringComparison.OrdinalIgnoreCase))182 {183 string errorMessage = $"Mock {name} - Expected `{expectedString}`; length {expected.Length}, but received; length {received}; `{bufferString}`(alt encoding: `{alternativeEncodedString}`) instead.";184 _output.WriteLine(errorMessage);185 throw new Exception(errorMessage);186 }187 _output.WriteLine($"Mock {name} - received expected message: " + expectedString);188 SendResponse(name, client, serverResponses);189 }190 191 if (verifyConnectionClosed)192 {193 _output.WriteLine($"Mock {name} - verifying client connection is closed.");194 // Wait for the client to close the connection synchronously (no timeout)195 try196 {197 while (true)198 {199 int bytesRead = client.Receive(buffer, SocketFlags.None);200 if (bytesRead == 0)201 {202 break;203 }204 205 // If we receive any data, log and throw (assume UTF8 encoding)206 string unexpectedData = Encoding.UTF8.GetString(buffer, 0, bytesRead);207 _output.WriteLine($"Mock {name} - received unexpected data after handshake: {unexpectedData}");208 throw new Exception($"Mock {name} - received unexpected data after handshake: {unexpectedData}");209 }210 _output.WriteLine($"Mock {name} - client closed the connection.");211 }212 catch (SocketException ex)213 {214 _output.WriteLine($"Mock {name} - socket exception while waiting for client close: {ex.Message} {ex.GetType().FullName}");215 }216 catch (ObjectDisposedException)217 {218 _output.WriteLine($"Mock {name} - socket already closed.");219 // Socket already closed220 }221 }222 }223 224 _output.WriteLine($"Mock {name} - on port {port} completed successfully.");225 }226 catch (Exception ex)227 {228 _output.WriteLine($"Mock {name} - Exception: {ex.Message} {ex.GetType().FullName}");229 _output.WriteLine(ex.StackTrace);230 throw;231 }232 finally233 {234 _output.WriteLine($"Mock {name} - remaining expected messages: {expectedClientSends.Count}");235 _output.WriteLine($"Mock {name} - stopping listener on port {port}.");236 listener.Stop();237 }238 }239 240 // Helper function to create a random 4-character ASCII response241 private static string CreateRandomAsciiResponse()242 {243 var rand = new Random();244 // Randomly return either "PASS" or "FAIL"245 return rand.Next(0, 2) == 0 ? "PASS" : "FAIL";246 }247 248 // Helper method to create test data249 private static (List<(string, Encoding)> expectedClientSends, List<(string, Encoding)> serverResponses) CreateHandshakeTestData(NetworkCredential cred)250 {251 var expectedClientSends = new List<(string message, Encoding encoding)>252 {253 (message: cred.Domain, encoding: Encoding.Unicode),254 (message: cred.UserName, encoding: Encoding.Unicode),255 (message: "NONEMPTYPW", encoding: Encoding.ASCII),256 (message: cred.Password, encoding: Encoding.Unicode)257 };258 259 var serverResponses = new List<(string message, Encoding encoding)>260 {261 (message: CreateRandomAsciiResponse(), encoding: Encoding.ASCII), // Response to domain262 (message: CreateRandomAsciiResponse(), encoding: Encoding.ASCII), // Response to username263 (message: CreateRandomAsciiResponse(), encoding: Encoding.ASCII) // Response to non-empty password264 };265 266 return (expectedClientSends, serverResponses);267 }268 269 private static List<(string message, Encoding encoding)> CreateVersionNegotiationClientSends()270 {271 return new List<(string message, Encoding encoding)>272 {273 (message: "VERSION", encoding: Encoding.UTF8),274 (message: "VERSION_2", encoding: Encoding.UTF8),275 };276 }277 278 private static List<(string, Encoding)> CreateV2Sends(NetworkCredential cred, string configurationName)279 {280 var sends = CreateVersionNegotiationClientSends();281 var password = cred.Password;282 var emptyPassword = string.IsNullOrEmpty(password);283 284 sends.AddRange(new List<(string message, Encoding encoding)>285 {286 (message: cred.Domain, encoding: Encoding.Unicode),287 (message: cred.UserName, encoding: Encoding.Unicode)288 });289 290 if (!emptyPassword)291 {292 sends.AddRange(new List<(string message, Encoding encoding)>293 {294 (message: "NONEMPTYPW", encoding: Encoding.UTF8),295 (message: cred.Password, encoding: Encoding.Unicode)296 });297 }298 else299 {300 sends.Add((message: "EMPTYPW", encoding: Encoding.UTF8)); // Empty password and we don't expect a response301 }302 303 if (!string.IsNullOrEmpty(configurationName))304 {305 sends.Add((message: "NONEMPTYCF", encoding: Encoding.UTF8));306 sends.Add((message: configurationName, encoding: Encoding.Unicode)); // Configuration string and we don't expect a response307 }308 else309 {310 sends.Add((message: "EMPTYCF", encoding: Encoding.UTF8)); // Configuration string and we don't expect a response311 }312 313 sends.Add((message: "PASS", encoding: Encoding.ASCII)); // Response to TOKEN314 315 return sends;316 }317 318 private static List<(string, Encoding)> CreateV2Responses(string version = "VERSION_2", bool emptyConfig = false, string token = "FakeToken0+/=", bool emptyPassword = false)319 {320 var responses = new List<(string message, Encoding encoding)>321 {322 (message: version, encoding: Encoding.ASCII), // Response to VERSION323 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2324 (message: "PASS", encoding: Encoding.ASCII), // Response to domain325 (message: "PASS", encoding: Encoding.ASCII), // Response to username326 };327 328 if (!emptyPassword)329 {330 responses.Add((message: "PASS", encoding: Encoding.ASCII)); // Response to non-empty password331 }332 333 responses.Add((message: "CONF", encoding: Encoding.ASCII)); // Response to configuration334 335 if (!emptyConfig)336 {337 responses.Add((message: "PASS", encoding: Encoding.ASCII)); // Response to non-empty configuration338 }339 responses.Add((message: "TOKEN " + token, encoding: Encoding.ASCII)); // Response to with a token than uses each class of character in base 64 encoding340 341 return responses;342 }343 344 // Helper method to create test data345 private static (List<(string, Encoding)> expectedClientSends, List<(string, Encoding)> serverResponses)346 CreateHandshakeTestDataV2(NetworkCredential cred, string version, string configurationName, string token)347 {348 bool emptyConfig = string.IsNullOrEmpty(configurationName);349 bool emptyPassword = string.IsNullOrEmpty(cred.Password);350 return (CreateV2Sends(cred, configurationName), CreateV2Responses(version, emptyConfig, token, emptyPassword));351 }352 353 // Helper method to create test data354 private static (List<(string, Encoding)> expectedClientSends, List<(string, Encoding)> serverResponses) CreateHandshakeTestDataForFallback(NetworkCredential cred)355 {356 var expectedClientSends = new List<(string message, Encoding encoding)>357 {358 (message: "VERSION", encoding: Encoding.UTF8),359 (message: @"?<PSDirectVMLegacy>", encoding: Encoding.Unicode),360 (message: "EMPTYPW", encoding: Encoding.UTF8), // Response to domain361 (message: "FAIL", encoding: Encoding.UTF8), // Response to domain362 };363 364 List<(string message, Encoding encoding)> serverResponses = new List<(string message, Encoding encoding)>365 {366 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION but v1 server expects domain so it says "PASS"367 (message: "PASS", encoding: Encoding.ASCII), // Response to username368 (message: "FAIL", encoding: Encoding.ASCII) // Response to EMPTYPW369 };370 371 return (expectedClientSends, serverResponses);372 }373 374 // Helper to create a password with at least one non-ASCII Unicode character375 public static string CreateRandomUnicodePassword(string prefix)376 {377 var rand = new Random();378 var asciiPart = new char[6 + prefix.Length];379 // Copy prefix into asciiPart380 Array.Copy(prefix.ToCharArray(), 0, asciiPart, 0, prefix.Length);381 for (int i = prefix.Length; i < asciiPart.Length; i++)382 {383 asciiPart[i] = (char)rand.Next(33, 127); // ASCII printable384 }385 // Add a random Unicode character outside ASCII range (e.g., U+0100 to U+017F)386 char unicodeChar = (char)rand.Next(0x0100, 0x017F);387 // Insert the unicode character at a random position388 int insertPos = rand.Next(0, asciiPart.Length + 1);389 var passwordChars = new List<char>(asciiPart);390 passwordChars.Insert(insertPos, unicodeChar);391 return new string(passwordChars.ToArray());392 }393 394 public static NetworkCredential CreateTestCredential()395 {396 return new NetworkCredential(CreateRandomUnicodePassword("username"), CreateRandomUnicodePassword("password"), CreateRandomUnicodePassword("domain"));397 }398 399 [SkippableFact]400 public async Task PerformCredentialAndConfigurationHandshake_V1_Pass()401 {402 // Arrange403 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);404 var cred = CreateTestCredential();405 string configurationName = CreateRandomUnicodePassword("config");406 407 var (expectedClientSends, serverResponses) = CreateHandshakeTestData(cred);408 expectedClientSends.Add(("PASS", Encoding.ASCII));409 serverResponses.Add(("PASS", Encoding.ASCII));410 411 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));412 var serverTask = Task.Run(() => StartHandshakeServer("Broker", port, expectedClientSends, serverResponses, verifyConnectionClosed: false, cts.Token), cts.Token);413 414 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))415 {416 ConnectWithRetry(client, IPAddress.Loopback, port, _output);417 var exchangeResult = System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.ExchangeCredentialsAndConfiguration(cred, configurationName, client, true);418 var result = exchangeResult.success;419 _output.WriteLine($"Exchange result: {result}, Token: {exchangeResult.authenticationToken}");420 System.Threading.Thread.Sleep(100); // Allow time for server to process421 Assert.True(result, $"Expected Exchange to pass");422 }423 424 await serverTask;425 }426 427 [SkippableTheory]428 [InlineData("VERSION_2", "configurationname1", "FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0+/==")] // a fake base64 token about 512 bits long (double the size when this was spec'ed)429 [InlineData("VERSION_10", null, "FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0+/=")] // a fake base64 token about 256 bits Long (the size when this was spec'ed)430 public async Task PerformCredentialAndConfigurationHandshake_V2_Pass(string versionResponse, string configurationName, string token)431 {432 // Arrange433 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);434 var cred = CreateTestCredential();435 436 var (expectedClientSends, serverResponses) = CreateHandshakeTestDataV2(cred, versionResponse, configurationName, token);437 438 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));439 var serverTask = Task.Run(() => StartHandshakeServer("Broker", port, expectedClientSends, serverResponses, verifyConnectionClosed: true, cts.Token), cts.Token);440 441 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))442 {443 client.Connect(IPAddress.Loopback, port);444 var exchangeResult = System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.ExchangeCredentialsAndConfiguration(cred, configurationName, client, false);445 var result = exchangeResult.success;446 System.Threading.Thread.Sleep(100); // Allow time for server to process447 Assert.True(result, $"Expected Exchange to pass for version response '{versionResponse}'");448 Assert.Equal(token, exchangeResult.authenticationToken);449 }450 451 await serverTask;452 }453 454 [SkippableFact]455 public async Task PerformCredentialAndConfigurationHandshake_V1_Fallback()456 {457 // Arrange458 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);459 var cred = CreateTestCredential();460 string configurationName = CreateRandomUnicodePassword("config");461 462 var (expectedClientSends, serverResponses) = CreateHandshakeTestDataForFallback(cred);463 464 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));465 var serverTask = Task.Run(() => StartHandshakeServer("Broker", port, expectedClientSends, serverResponses, verifyConnectionClosed: false, cts.Token), cts.Token);466 467 bool isFallback = false;468 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))469 {470 _output.WriteLine("Starting handshake with V2 protocol.");471 client.Connect(IPAddress.Loopback, port);472 var exchangeResult = System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.ExchangeCredentialsAndConfiguration(cred, configurationName, client, false);473 isFallback = !exchangeResult.success;474 475 System.Threading.Thread.Sleep(100); // Allow time for server to process476 _output.WriteLine("Handshake indicated fallback to V1.");477 Assert.True(isFallback, "Expected fallback to V1.");478 }479 _output.WriteLine("Handshake completed successfully with fallback to V1.");480 481 await serverTask;482 }483 484 [SkippableFact]485 public async Task PerformCredentialAndConfigurationHandshake_V2_InvalidResponse()486 {487 // Arrange488 int port = 51000 + (int)(DateTime.Now.Ticks % 10000);489 var cred = CreateTestCredential();490 491 var (expectedClientSends, serverResponses) = CreateHandshakeTestData(cred);492 //expectedClientSends.Add("FAI1");493 serverResponses.Add(("FAI1", Encoding.ASCII));494 495 using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(30));496 497 //cts.Token.Register(() => throw new OperationCanceledException("Test timed out."));498 499 var serverTask = Task.Run(() => StartHandshakeServer("Broker", port, expectedClientSends, serverResponses, verifyConnectionClosed: false, cts.Token), cts.Token);500 501 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))502 {503 _output.WriteLine("connecting on port " + port);504 ConnectWithRetry(client, IPAddress.Loopback, port, _output);505 506 var ex = Record.Exception(() => System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.ExchangeCredentialsAndConfiguration(cred, "config", client, true));507 508 try509 {510 await serverTask;511 }512 catch (AggregateException exAgg)513 {514 Assert.Null(exAgg.Flatten().InnerExceptions[1].Message);515 }516 cts.Token.ThrowIfCancellationRequested();517 518 Assert.NotNull(ex);519 Assert.NotNull(ex.Message);520 Assert.Contains("Hyper-V Broker sent an invalid Credential response", ex.Message);521 }522 }523 524 [SkippableFact]525 public async Task PerformCredentialAndConfigurationHandshake_V1_Fail()526 {527 // Arrange528 int port = 51000 + (int)(DateTime.Now.Ticks % 10000);529 var cred = CreateTestCredential();530 531 var (expectedClientSends, serverResponses) = CreateHandshakeTestData(cred);532 expectedClientSends.Add(("FAIL", Encoding.ASCII));533 serverResponses.Add(("FAIL", Encoding.ASCII));534 535 using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(15));536 537 // This scenario does not close the connection in a timely manner, so we set verifyConnectionClosed to false538 var serverTask = Task.Run(() => StartHandshakeServer("Broker", port, expectedClientSends, serverResponses, verifyConnectionClosed: false, cts.Token), cts.Token);539 540 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))541 {542 client.Connect(IPAddress.Loopback, port);543 544 var ex = Record.Exception(() => System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.ExchangeCredentialsAndConfiguration(cred, "config", client, true));545 546 try547 {548 await serverTask;549 }550 catch (AggregateException exAgg)551 {552 Assert.Null(exAgg.Flatten().InnerExceptions[1].Message);553 }554 555 cts.Token.ThrowIfCancellationRequested();556 557 Assert.NotNull(ex);558 Assert.NotNull(ex.Message);559 Assert.Contains("The credential is invalid.", ex.Message);560 }561 }562 563 [SkippableTheory]564 [InlineData("VERSION_2", "FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0+/==")] // a fake base64 token about 512 bits long (double the size when this was spec'ed)565 [InlineData("VERSION_10", "FakeTokenaaaaaaaaaAAAAAAAAAAAAAAAAAAAAAA0+/=")] // a fake base64 token about 256 bits Long (the size when this was spec'ed)566 public async Task PerformTransportVersionAndTokenExchange_Pass(string version, string token)567 {568 // Arrange569 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);570 var cred = CreateTestCredential();571 572 var expectedClientSends = CreateVersionNegotiationClientSends();573 expectedClientSends.Add((message: "TOKEN " + token, encoding: Encoding.ASCII));574 575 var serverResponses = new List<(string message, Encoding encoding)>{576 (message: version, encoding: Encoding.ASCII), // Response to VERSION577 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2578 (message: "PASS", encoding: Encoding.ASCII) // Response to token579 };580 581 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));582 var serverTask = Task.Run(() => StartHandshakeServer("Server", port, expectedClientSends, serverResponses, verifyConnectionClosed: true, cts.Token), cts.Token);583 584 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))585 {586 ConnectWithRetry(client, IPAddress.Loopback, port, _output);587 System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.PerformTransportVersionAndTokenExchange(client, token);588 System.Threading.Thread.Sleep(100); // Allow time for server to process589 }590 591 await serverTask;592 }593 594 [SkippableTheory]595 [InlineData(1, true)]596 [InlineData(2, true)]597 [InlineData(0, false)]598 [InlineData(null, false)]599 [System.Runtime.Versioning.SupportedOSPlatform("windows")]600 public void IsRequirePsDirectAuthenticationEnabled(int? regValue, bool expected)601 {602 const string testKeyPath = @"SOFTWARE\Microsoft\TestRequirePsDirectAuthentication";603 const string valueName = "RequirePsDirectAuthentication";604 if (!System.Management.Automation.Platform.IsWindows)605 {606 throw new SkipException("RemoteHyperVTests are only supported on Windows.");607 }608 609 // Clean up any previous test key610 var regHive = Microsoft.Win32.RegistryHive.CurrentUser;611 var baseKey = Microsoft.Win32.RegistryKey.OpenBaseKey(regHive, Microsoft.Win32.RegistryView.Registry64);612 baseKey.DeleteSubKeyTree(testKeyPath, false);613 614 bool? result = null;615 616 // Create the test key617 using (var key = baseKey.CreateSubKey(testKeyPath))618 {619 if (regValue.HasValue)620 {621 key.SetValue(valueName, regValue.Value, Microsoft.Win32.RegistryValueKind.DWord);622 }623 else624 {625 // Ensure the value does not exist626 key.DeleteValue(valueName, false);627 }628 629 result = System.Management.Automation.Remoting.RemoteSessionHyperVSocketClient.IsRequirePsDirectAuthenticationEnabled(testKeyPath, regHive);630 }631 632 Assert.True(result.HasValue, "IsRequirePsDirectAuthenticationEnabled should return a value.");633 Assert.True(expected == result.Value,634 $"Expected IsRequirePsDirectAuthenticationEnabled to return {expected} when registry value is {(regValue.HasValue ? regValue.ToString() : "not set")}.");635 636 return;637 }638 639 [SkippableTheory]640 [InlineData("testToken", "testToken")]641 [InlineData("testToken\0", "testToken")]642 public async Task ValidatePassesWhenTokensMatch(string token, string expectedToken)643 {644 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);645 646 var expectedClientSends = new List<(string message, Encoding encoding)>{647 (message: "VERSION", encoding: Encoding.ASCII), // Response to VERSION648 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION_2649 (message: $"TOKEN {token}", encoding: Encoding.ASCII)650 };651 652 var serverResponses = new List<(string message, Encoding encoding)>{653 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION_2654 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2655 (message: "PASS", encoding: Encoding.ASCII) // Response to token656 };657 658 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));659 var serverTask = Task.Run(() => StartHandshakeServer("Client", port, serverResponses, expectedClientSends, verifyConnectionClosed: true, cts.Token, sendFirst: true), cts.Token);660 661 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))662 {663 ConnectWithRetry(client, IPAddress.Loopback, port, _output);664 System.Management.Automation.Remoting.RemoteSessionHyperVSocketServer.ValidateToken(client, expectedToken, DateTimeOffset.UtcNow, 1);665 System.Threading.Thread.Sleep(100); // Allow time for server to process666 }667 668 await serverTask;669 }670 671 [SkippableTheory]672 [InlineData(5500, "A connection attempt failed because the connected party did not properly respond after a period of time, or established connection failed because connected host has failed to respond.", "SocketException")] // test the socket timeout673 [InlineData(3200, "canceled", "System.OperationCanceledException")] // test the cancellation token674 [InlineData(10, "", "")]675 public async Task ValidateTokenTimeoutFails(int timeoutMs, string expectedMessage, string expectedExceptionType = "SocketException")676 {677 string token = "testToken";678 string expectedToken = token;679 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);680 681 var expectedClientSends = new List<(string message, Encoding encoding, int delayMs)>{682 (message: "VERSION", encoding: Encoding.ASCII, delayMs: timeoutMs), // Response to VERSION683 (message: "VERSION_2", encoding: Encoding.ASCII, delayMs: timeoutMs), // Response to VERSION_2684 (message: $"TOKEN {token}", encoding: Encoding.ASCII, delayMs: 1)685 };686 687 var serverResponses = new List<(string message, Encoding encoding)>{688 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION_2689 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2690 (message: "PASS", encoding: Encoding.ASCII) // Response to token691 };692 693 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));694 var serverTask = Task.Run(() => StartHandshakeServer("Client", port, serverResponses, expectedClientSends, verifyConnectionClosed: true, cts.Token, sendFirst: true), cts.Token);695 696 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))697 {698 ConnectWithRetry(client, IPAddress.Loopback, port, _output);699 if (expectedMessage.Length > 0)700 {701 var exception = Record.Exception(702 () => System.Management.Automation.Remoting.RemoteSessionHyperVSocketServer.ValidateToken(client, expectedToken, DateTimeOffset.UtcNow, 5)); // set the timeout to 5 seconds or 5000 ms703 Assert.NotNull(exception);704 string exceptionType = exception.GetType().FullName;705 _output.WriteLine($"Caught exception of type {exceptionType} with message: {exception.Message}");706 Assert.Contains(expectedExceptionType, exceptionType, StringComparison.OrdinalIgnoreCase);707 Assert.Contains(expectedMessage, exception.Message, StringComparison.OrdinalIgnoreCase);708 }709 else710 {711 System.Management.Automation.Remoting.RemoteSessionHyperVSocketServer.ValidateToken(client, expectedToken, DateTimeOffset.UtcNow, 5);712 }713 System.Threading.Thread.Sleep(100); // Allow time for server to process714 }715 716 if (expectedMessage.Length == 0)717 {718 await serverTask;719 }720 }721 722 [SkippableFact]723 public async Task ValidateTokenTimeoutDoesAffectSession()724 {725 string token = "testToken";726 string expectedToken = token;727 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);728 729 var expectedClientSends = new List<(string message, Encoding encoding, int delayMs)>{730 (message: "VERSION", encoding: Encoding.ASCII, delayMs: 1), // Response to VERSION731 (message: "VERSION_2", encoding: Encoding.ASCII, delayMs: 1), // Response to VERSION_2732 (message: $"TOKEN {token}", encoding: Encoding.ASCII, delayMs: 1),733 (message: string.Empty, encoding: Encoding.ASCII, delayMs: 99), // Send some data after the handshake734 (message: string.Empty, encoding: Encoding.ASCII, delayMs: 100), // Send some data after the handshake735 (message: string.Empty, encoding: Encoding.ASCII, delayMs: 101), // Send some data after the handshake736 (message: string.Empty, encoding: Encoding.ASCII, delayMs: 102), // Send some data after the handshake737 (message: string.Empty, encoding: Encoding.ASCII, delayMs: 103) // Send some data after the handshake738 };739 740 var serverResponses = new List<(string message, Encoding encoding)>{741 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION_2742 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2743 (message: "PASS", encoding: Encoding.ASCII), // Response to token744 (message: "PSRP-Message0", encoding: Encoding.ASCII), // Indicate server is ready to receive data745 (message: "PSRP-Message1", encoding: Encoding.ASCII), // Indicate server is ready to receive data746 (message: "PSRP-Message2", encoding: Encoding.ASCII), // Indicate server is ready to receive data747 (message: "PSRP-Message3", encoding: Encoding.ASCII), // Indicate server is ready to receive data748 (message: "PSRP-Message4", encoding: Encoding.ASCII) //749 750 };751 752 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(2));753 var serverTask = Task.Run(() => StartHandshakeServer("Client", port, serverResponses, expectedClientSends, verifyConnectionClosed: false, cts.Token, sendFirst: true), cts.Token);754 755 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))756 {757 ConnectWithRetry(client, IPAddress.Loopback, port, _output);758 System.Management.Automation.Remoting.RemoteSessionHyperVSocketServer.ValidateToken(client, expectedToken, DateTimeOffset.UtcNow, 5);759 for (int i = 0; i < 5; i++)760 {761 System.Threading.Thread.Sleep(1500);762 client.Send(Encoding.ASCII.GetBytes($"PSRP-Message{i}")); // Send some data after the handshake763 }764 }765 766 await serverTask;767 }768 769 [SkippableTheory]770 [InlineData("abc", "xyz")]771 [InlineData("abc", "abcdef")]772 [InlineData("abcdef", "abc")]773 [InlineData("abc\0def", "abc")]774 public async Task ValidateFailsWhenTokensMismatch(string token, string expectedToken)775 {776 int port = 50000 + (int)(DateTime.Now.Ticks % 10000);777 778 var expectedClientSends = new List<(string message, Encoding encoding)>{779 (message: "VERSION", encoding: Encoding.ASCII), // Initial request780 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION_2781 (message: $"TOKEN {token}", encoding: Encoding.ASCII)782 };783 784 var serverResponses = new List<(string message, Encoding encoding)>{785 (message: "VERSION_2", encoding: Encoding.ASCII), // Response to VERSION786 (message: "PASS", encoding: Encoding.ASCII), // Response to VERSION_2787 (message: "FAIL", encoding: Encoding.ASCII) // Response to token788 };789 790 using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(1));791 var serverTask = Task.Run(() => StartHandshakeServer("Client", port, serverResponses, expectedClientSends, verifyConnectionClosed: true, cts.Token, sendFirst: true), cts.Token);792 793 using (var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp))794 {795 ConnectWithRetry(client, IPAddress.Loopback, port, _output);796 DateTimeOffset tokenCreationTime = DateTimeOffset.UtcNow; // Token created 10 minutes ago797 var exception = Assert.Throws<System.Management.Automation.Remoting.PSDirectException>(798 () => System.Management.Automation.Remoting.RemoteSessionHyperVSocketServer.ValidateToken(client, expectedToken, tokenCreationTime, 5));799 System.Threading.Thread.Sleep(100); // Allow time for server to process800 Assert.Contains("The credential is invalid.", exception.Message);801 }802 803 await serverTask;804 }805 }806}807 