Team Ai
Datasetpublic

MegaBites-AI/Windows-powershell

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes372downloads
test_RemoteHyperV.cs807 linesDownload Raw Back to csharp
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