diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/BasicHtmlWebResponseObject.Common.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/BasicHtmlWebResponseObject.Common.cs index 75775428bc..9bd76f9941 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/BasicHtmlWebResponseObject.Common.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/BasicHtmlWebResponseObject.Common.cs @@ -3,6 +3,7 @@ #nullable enable +using System; using System.Collections.Generic; using System.Diagnostics; using System.Diagnostics.CodeAnalysis; @@ -26,8 +27,9 @@ namespace Microsoft.PowerShell.Commands /// Initializes a new instance of the class. /// /// The response. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// Cancellation token. - public BasicHtmlWebResponseObject(HttpResponseMessage response, CancellationToken cancellationToken) : this(response, null, cancellationToken) { } + public BasicHtmlWebResponseObject(HttpResponseMessage response, TimeSpan perReadTimeout, CancellationToken cancellationToken) : this(response, null, perReadTimeout, cancellationToken) { } /// /// Initializes a new instance of the class @@ -35,8 +37,9 @@ namespace Microsoft.PowerShell.Commands /// /// The response. /// The content stream associated with the response. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// Cancellation token. - public BasicHtmlWebResponseObject(HttpResponseMessage response, Stream? contentStream, CancellationToken cancellationToken) : base(response, contentStream, cancellationToken) + public BasicHtmlWebResponseObject(HttpResponseMessage response, Stream? contentStream, TimeSpan perReadTimeout, CancellationToken cancellationToken) : base(response, contentStream, perReadTimeout, cancellationToken) { InitializeContent(cancellationToken); InitializeRawContent(response); @@ -157,7 +160,7 @@ namespace Microsoft.PowerShell.Commands // Fill the Content buffer string? characterSet = WebResponseHelper.GetCharacterSet(BaseResponse); - Content = StreamHelper.DecodeStream(RawContentStream, characterSet, out Encoding encoding, cancellationToken); + Content = StreamHelper.DecodeStream(RawContentStream, characterSet, out Encoding encoding, perReadTimeout, cancellationToken); Encoding = encoding; } else diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/InvokeRestMethodCommand.Common.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/InvokeRestMethodCommand.Common.cs index df1ef75052..ca78ff370b 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/InvokeRestMethodCommand.Common.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/InvokeRestMethodCommand.Common.cs @@ -80,11 +80,12 @@ namespace Microsoft.PowerShell.Commands ArgumentNullException.ThrowIfNull(response); ArgumentNullException.ThrowIfNull(_cancelToken); + TimeSpan perReadTimeout = ConvertTimeoutSecondsToTimeSpan(OperationTimeoutSeconds); Stream baseResponseStream = StreamHelper.GetResponseStream(response, _cancelToken.Token); if (ShouldWriteToPipeline) { - using BufferingStreamReader responseStream = new(baseResponseStream, _cancelToken.Token); + using BufferingStreamReader responseStream = new(baseResponseStream, perReadTimeout, _cancelToken.Token); // First see if it is an RSS / ATOM feed, in which case we can // stream it - unless the user has overridden it with a return type of "XML" @@ -96,8 +97,7 @@ namespace Microsoft.PowerShell.Commands { // Try to get the response encoding from the ContentType header. string? characterSet = WebResponseHelper.GetCharacterSet(response); - - string str = StreamHelper.DecodeStream(responseStream, characterSet, out Encoding encoding, _cancelToken.Token); + string str = StreamHelper.DecodeStream(responseStream, characterSet, out Encoding encoding, perReadTimeout, _cancelToken.Token); string encodingVerboseName; try @@ -139,12 +139,12 @@ namespace Microsoft.PowerShell.Commands } } else if (ShouldSaveToOutFile) - { + { string outFilePath = WebResponseHelper.GetOutFilePath(response, _qualifiedOutFile); WriteVerbose(string.Create(System.Globalization.CultureInfo.InvariantCulture, $"File Name: {Path.GetFileName(_qualifiedOutFile)}")); - StreamHelper.SaveStreamToFile(baseResponseStream, outFilePath, this, response.Content.Headers.ContentLength.GetValueOrDefault(), _cancelToken.Token); + StreamHelper.SaveStreamToFile(baseResponseStream, outFilePath, this, response.Content.Headers.ContentLength.GetValueOrDefault(), perReadTimeout, _cancelToken.Token); } if (!string.IsNullOrEmpty(StatusCodeVariable)) @@ -349,18 +349,20 @@ namespace Microsoft.PowerShell.Commands internal class BufferingStreamReader : Stream { - internal BufferingStreamReader(Stream baseStream, CancellationToken cancellationToken) + internal BufferingStreamReader(Stream baseStream, TimeSpan perReadTimeout, CancellationToken cancellationToken) { _baseStream = baseStream; _streamBuffer = new MemoryStream(); _length = long.MaxValue; _copyBuffer = new byte[4096]; + _perReadTimeout = perReadTimeout; _cancellationToken = cancellationToken; } private readonly Stream _baseStream; private readonly MemoryStream _streamBuffer; private readonly byte[] _copyBuffer; + private readonly TimeSpan _perReadTimeout; private readonly CancellationToken _cancellationToken; public override bool CanRead => true; @@ -395,7 +397,7 @@ namespace Microsoft.PowerShell.Commands // If we don't have enough data to fill this from memory, cache more. // We try to read 4096 bytes from base stream every time, so at most we // may cache 4095 bytes more than what is required by the Read operation. - int bytesRead = _baseStream.ReadAsync(_copyBuffer, 0, _copyBuffer.Length, _cancellationToken).GetAwaiter().GetResult(); + int bytesRead = _baseStream.ReadAsync(_copyBuffer.AsMemory(), _perReadTimeout, _cancellationToken).GetAwaiter().GetResult(); if (_streamBuffer.Position < _streamBuffer.Length) { diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebRequestPSCmdlet.Common.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebRequestPSCmdlet.Common.cs index f7b6f3e270..ba7d66bb9a 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebRequestPSCmdlet.Common.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebRequestPSCmdlet.Common.cs @@ -266,11 +266,25 @@ namespace Microsoft.PowerShell.Commands public virtual SwitchParameter DisableKeepAlive { get; set; } /// - /// Gets or sets the TimeOut property. + /// Gets or sets the ConnectionTimeoutSeconds property. /// + /// + /// This property applies to sending the request and receiving the response headers only. + /// + [Alias("TimeoutSec")] [Parameter] [ValidateRange(0, int.MaxValue)] - public virtual int TimeoutSec { get; set; } + public virtual int ConnectionTimeoutSeconds { get; set; } + + /// + /// Gets or sets the OperationTimeoutSeconds property. + /// + /// + /// This property applies to each read operation when receiving the response body. + /// + [Parameter] + [ValidateRange(0, int.MaxValue)] + public virtual int OperationTimeoutSeconds { get; set; } /// /// Gets or sets the Headers property. @@ -570,7 +584,7 @@ namespace Microsoft.PowerShell.Commands string respVerboseMsg = contentLength is null ? string.Format(CultureInfo.CurrentCulture, WebCmdletStrings.WebResponseNoSizeVerboseMsg, response.Version, contentType) : string.Format(CultureInfo.CurrentCulture, WebCmdletStrings.WebResponseVerboseMsg, response.Version, contentLength, contentType); - + WriteVerbose(respVerboseMsg); bool _isSuccess = response.IsSuccessStatusCode; @@ -621,12 +635,19 @@ namespace Microsoft.PowerShell.Commands string detailMsg = string.Empty; try { - string error = StreamHelper.GetResponseString(response, _cancelToken.Token); + // We can't use ReadAsStringAsync because it doesn't have per read timeouts + TimeSpan perReadTimeout = ConvertTimeoutSecondsToTimeSpan(OperationTimeoutSeconds); + string characterSet = WebResponseHelper.GetCharacterSet(response); + var responseStream = StreamHelper.GetResponseStream(response, _cancelToken.Token); + int initialCapacity = (int)Math.Min(contentLength ?? StreamHelper.DefaultReadBuffer, StreamHelper.DefaultReadBuffer); + var bufferedStream = new WebResponseContentMemoryStream(responseStream, initialCapacity, this, contentLength, perReadTimeout, _cancelToken.Token); + string error = StreamHelper.DecodeStream(bufferedStream, characterSet, out Encoding encoding, perReadTimeout, _cancelToken.Token); detailMsg = FormatErrorMessage(error, contentType); } - catch + catch (Exception ex) { // Catch all + er.ErrorDetails = new ErrorDetails(ex.ToString()); } if (!string.IsNullOrEmpty(detailMsg)) @@ -656,6 +677,11 @@ namespace Microsoft.PowerShell.Commands WriteError(er); } } + catch (TimeoutException ex) + { + ErrorRecord er = new(ex, "OperationTimeoutReached", ErrorCategory.OperationTimeout, null); + ThrowTerminatingError(er); + } catch (HttpRequestException ex) { ErrorRecord er = new(ex, "WebCmdletWebResponseException", ErrorCategory.InvalidOperation, request); @@ -666,7 +692,7 @@ namespace Microsoft.PowerShell.Commands ThrowTerminatingError(er); } - finally + finally { _cancelToken?.Dispose(); _cancelToken = null; @@ -970,7 +996,7 @@ namespace Microsoft.PowerShell.Commands } else { - webProxy.UseDefaultCredentials = ProxyUseDefaultCredentials; + webProxy.UseDefaultCredentials = ProxyUseDefaultCredentials; } // We don't want to update the WebSession unless the proxies are different @@ -1020,7 +1046,7 @@ namespace Microsoft.PowerShell.Commands WebSession.RetryIntervalInSeconds = RetryIntervalSec; } - WebSession.TimeoutSec = TimeoutSec; + WebSession.ConnectionTimeout = ConvertTimeoutSecondsToTimeSpan(ConnectionTimeoutSeconds); } internal virtual HttpClient GetHttpClient(bool handleRedirect) @@ -1263,8 +1289,24 @@ namespace Microsoft.PowerShell.Commands Uri currentUri = currentRequest.RequestUri; _cancelToken = new CancellationTokenSource(); - response = client.SendAsync(currentRequest, HttpCompletionOption.ResponseHeadersRead, _cancelToken.Token).GetAwaiter().GetResult(); + try + { + response = client.SendAsync(currentRequest, HttpCompletionOption.ResponseHeadersRead, _cancelToken.Token).GetAwaiter().GetResult(); + } + catch (TaskCanceledException ex) + { + if (ex.InnerException is TimeoutException) + { + // HTTP Request timed out + ErrorRecord er = new(ex, "ConnectionTimeoutReached", ErrorCategory.OperationTimeout, null); + ThrowTerminatingError(er); + } + else + { + throw; + } + } if (handleRedirect && _maximumRedirection is not 0 && IsRedirectCode(response.StatusCode) @@ -1388,6 +1430,9 @@ namespace Microsoft.PowerShell.Commands #endregion Virtual Methods #region Helper Methods + + internal static TimeSpan ConvertTimeoutSecondsToTimeSpan(int timeout) => timeout > 0 ? TimeSpan.FromSeconds(timeout) : Timeout.InfiniteTimeSpan; + private Uri PrepareUri(Uri uri) { uri = CheckProtocol(uri); diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebResponseObject.Common.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebResponseObject.Common.cs index 5b89a2352f..81bd4f13c6 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebResponseObject.Common.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/Common/WebResponseObject.Common.cs @@ -72,14 +72,24 @@ namespace Microsoft.PowerShell.Commands #endregion Properties + #region Protected Fields + + /// + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. + /// + protected TimeSpan perReadTimeout; + + #endregion Protected Fields + #region Constructors /// /// Initializes a new instance of the class. /// /// The Http response. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// The cancellation token. - public WebResponseObject(HttpResponseMessage response, CancellationToken cancellationToken) : this(response, null, cancellationToken) { } + public WebResponseObject(HttpResponseMessage response, TimeSpan perReadTimeout, CancellationToken cancellationToken) : this(response, null, perReadTimeout, cancellationToken) { } /// /// Initializes a new instance of the class @@ -87,9 +97,11 @@ namespace Microsoft.PowerShell.Commands /// /// Http response. /// The http content stream. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// The cancellation token. - public WebResponseObject(HttpResponseMessage response, Stream? contentStream, CancellationToken cancellationToken) + public WebResponseObject(HttpResponseMessage response, Stream? contentStream, TimeSpan perReadTimeout, CancellationToken cancellationToken) { + this.perReadTimeout = perReadTimeout; SetResponse(response, contentStream, cancellationToken); InitializeContent(); InitializeRawContent(response); @@ -149,7 +161,7 @@ namespace Microsoft.PowerShell.Commands } int initialCapacity = (int)Math.Min(contentLength, StreamHelper.DefaultReadBuffer); - RawContentStream = new WebResponseContentMemoryStream(st, initialCapacity, cmdlet: null, response.Content.Headers.ContentLength.GetValueOrDefault(), cancellationToken); + RawContentStream = new WebResponseContentMemoryStream(st, initialCapacity, cmdlet: null, response.Content.Headers.ContentLength.GetValueOrDefault(), perReadTimeout, cancellationToken); } // Set the position of the content stream to the beginning diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/CoreCLR/InvokeWebRequestCommand.CoreClr.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/CoreCLR/InvokeWebRequestCommand.CoreClr.cs index 1e66157ef0..026bbe866e 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/CoreCLR/InvokeWebRequestCommand.CoreClr.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/CoreCLR/InvokeWebRequestCommand.CoreClr.cs @@ -7,6 +7,7 @@ using System; using System.IO; using System.Management.Automation; using System.Net.Http; +using System.Threading; namespace Microsoft.PowerShell.Commands { @@ -35,7 +36,7 @@ namespace Microsoft.PowerShell.Commands internal override void ProcessResponse(HttpResponseMessage response) { ArgumentNullException.ThrowIfNull(response); - + TimeSpan perReadTimeout = ConvertTimeoutSecondsToTimeSpan(OperationTimeoutSeconds); Stream responseStream = StreamHelper.GetResponseStream(response, _cancelToken.Token); if (ShouldWriteToPipeline) { @@ -45,8 +46,9 @@ namespace Microsoft.PowerShell.Commands StreamHelper.ChunkSize, this, response.Content.Headers.ContentLength.GetValueOrDefault(), + perReadTimeout, _cancelToken.Token); - WebResponseObject ro = WebResponseHelper.IsText(response) ? new BasicHtmlWebResponseObject(response, responseStream, _cancelToken.Token) : new WebResponseObject(response, responseStream, _cancelToken.Token); + WebResponseObject ro = WebResponseHelper.IsText(response) ? new BasicHtmlWebResponseObject(response, responseStream, perReadTimeout, _cancelToken.Token) : new WebResponseObject(response, responseStream, perReadTimeout, _cancelToken.Token); ro.RelationLink = _relationLink; WriteObject(ro); @@ -63,7 +65,7 @@ namespace Microsoft.PowerShell.Commands WriteVerbose(string.Create(System.Globalization.CultureInfo.InvariantCulture, $"File Name: {Path.GetFileName(_qualifiedOutFile)}")); - StreamHelper.SaveStreamToFile(responseStream, outFilePath, this, response.Content.Headers.ContentLength.GetValueOrDefault(), _cancelToken.Token); + StreamHelper.SaveStreamToFile(responseStream, outFilePath, this, response.Content.Headers.ContentLength.GetValueOrDefault(), perReadTimeout, _cancelToken.Token); } } diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/StreamHelper.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/StreamHelper.cs index 6d217cdc90..d03debf3fc 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/StreamHelper.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/StreamHelper.cs @@ -4,6 +4,7 @@ #nullable enable using System; +using System.Buffers; using System.IO; using System.Management.Automation; using System.Management.Automation.Internal; @@ -29,6 +30,7 @@ namespace Microsoft.PowerShell.Commands private readonly Stream _originalStreamToProxy; private readonly Cmdlet? _ownerCmdlet; private readonly CancellationToken _cancellationToken; + private readonly TimeSpan _perReadTimeout; private bool _isInitialized = false; #endregion Data @@ -41,13 +43,15 @@ namespace Microsoft.PowerShell.Commands /// Presize the memory stream. /// Owner cmdlet if any. /// Expected download size in Bytes. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// Cancellation token. - internal WebResponseContentMemoryStream(Stream stream, int initialCapacity, Cmdlet? cmdlet, long? contentLength, CancellationToken cancellationToken) : base(initialCapacity) + internal WebResponseContentMemoryStream(Stream stream, int initialCapacity, Cmdlet? cmdlet, long? contentLength, TimeSpan perReadTimeout, CancellationToken cancellationToken) : base(initialCapacity) { this._contentLength = contentLength; _originalStreamToProxy = stream; _ownerCmdlet = cmdlet; _cancellationToken = cancellationToken; + _perReadTimeout = perReadTimeout; } #endregion Constructors @@ -228,7 +232,7 @@ namespace Microsoft.PowerShell.Commands } } - read = _originalStreamToProxy.ReadAsync(buffer, 0, buffer.Length, cancellationToken).GetAwaiter().GetResult(); + read = _originalStreamToProxy.ReadAsync(buffer.AsMemory(), _perReadTimeout, cancellationToken).GetAwaiter().GetResult(); if (read > 0) { @@ -255,6 +259,84 @@ namespace Microsoft.PowerShell.Commands } } + internal static class StreamTimeoutExtensions + { + internal static async Task ReadAsync(this Stream stream, Memory buffer, TimeSpan readTimeout, CancellationToken cancellationToken) + { + if (readTimeout == Timeout.InfiniteTimeSpan) + { + return await stream.ReadAsync(buffer, cancellationToken); + } + + using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + try + { + cts.CancelAfter(readTimeout); + return await stream.ReadAsync(buffer, cts.Token).ConfigureAwait(false); + } + catch (TaskCanceledException ex) + { + if (cts.IsCancellationRequested) + { + throw new TimeoutException($"The request was canceled due to the configured OperationTimeout of {readTimeout.TotalSeconds} seconds elapsing", ex); + } + else + { + throw; + } + } + } + + internal static async Task CopyToAsync(this Stream source, Stream destination, TimeSpan perReadTimeout, CancellationToken cancellationToken) + { + if (perReadTimeout == Timeout.InfiniteTimeSpan) + { + // No timeout - use fast path + await source.CopyToAsync(destination, cancellationToken); + return; + } + + byte[] buffer = ArrayPool.Shared.Rent(StreamHelper.ChunkSize); + CancellationTokenSource cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + try + { + while (true) + { + if (!cts.TryReset()) + { + cts.Dispose(); + cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + } + + cts.CancelAfter(perReadTimeout); + int bytesRead = await source.ReadAsync(buffer, cts.Token).ConfigureAwait(false); + if (bytesRead == 0) + { + break; + } + + await destination.WriteAsync(buffer.AsMemory(0, bytesRead), cancellationToken).ConfigureAwait(false); + } + } + catch (TaskCanceledException ex) + { + if (cts.IsCancellationRequested) + { + throw new TimeoutException($"The request was canceled due to the configured OperationTimeout of {perReadTimeout.TotalSeconds} seconds elapsing", ex); + } + else + { + throw; + } + } + finally + { + cts.Dispose(); + ArrayPool.Shared.Return(buffer); + } + } + } + internal static class StreamHelper { #region Constants @@ -270,11 +352,11 @@ namespace Microsoft.PowerShell.Commands #region Static Methods - internal static void WriteToStream(Stream input, Stream output, PSCmdlet cmdlet, long? contentLength, CancellationToken cancellationToken) + internal static void WriteToStream(Stream input, Stream output, PSCmdlet cmdlet, long? contentLength, TimeSpan perReadTimeout, CancellationToken cancellationToken) { ArgumentNullException.ThrowIfNull(cmdlet); - Task copyTask = input.CopyToAsync(output, cancellationToken); + Task copyTask = input.CopyToAsync(output, perReadTimeout, cancellationToken); bool wroteProgress = false; ProgressRecord record = new( @@ -328,16 +410,17 @@ namespace Microsoft.PowerShell.Commands /// Output file name. /// Current cmdlet (Invoke-WebRequest or Invoke-RestMethod). /// Expected download size in Bytes. + /// Time permitted between reads or Timeout.InfiniteTimeSpan for no timeout. /// CancellationToken to track the cmdlet cancellation. - internal static void SaveStreamToFile(Stream stream, string filePath, PSCmdlet cmdlet, long? contentLength, CancellationToken cancellationToken) + internal static void SaveStreamToFile(Stream stream, string filePath, PSCmdlet cmdlet, long? contentLength, TimeSpan perReadTimeout, CancellationToken cancellationToken) { // If the web cmdlet should resume, append the file instead of overwriting. FileMode fileMode = cmdlet is WebRequestPSCmdlet webCmdlet && webCmdlet.ShouldResume ? FileMode.Append : FileMode.Create; using FileStream output = new(filePath, fileMode, FileAccess.Write, FileShare.Read); - WriteToStream(stream, output, cmdlet, contentLength, cancellationToken); + WriteToStream(stream, output, cmdlet, contentLength, perReadTimeout, cancellationToken); } - private static string StreamToString(Stream stream, Encoding encoding, CancellationToken cancellationToken) + private static string StreamToString(Stream stream, Encoding encoding, TimeSpan perReadTimeout, CancellationToken cancellationToken) { StringBuilder result = new(capacity: ChunkSize); Decoder decoder = encoding.GetDecoder(); @@ -348,51 +431,59 @@ namespace Microsoft.PowerShell.Commands useBufferSize = encoding.GetMaxCharCount(10); } - char[] chars = new char[useBufferSize]; - byte[] bytes = new byte[useBufferSize * 4]; - int bytesRead = 0; - do + char[] chars = ArrayPool.Shared.Rent(useBufferSize); + byte[] bytes = ArrayPool.Shared.Rent(useBufferSize * 4); + try { - // Read at most the number of bytes that will fit in the input buffer. The - // return value is the actual number of bytes read, or zero if no bytes remain. - bytesRead = stream.ReadAsync(bytes, 0, useBufferSize * 4, cancellationToken).GetAwaiter().GetResult(); - - bool completed = false; - int byteIndex = 0; - - while (!completed) + int bytesRead = 0; + do { - // If this is the last input data, flush the decoder's internal buffer and state. - bool flush = bytesRead is 0; - decoder.Convert(bytes, byteIndex, bytesRead - byteIndex, chars, 0, useBufferSize, flush, out int bytesUsed, out int charsUsed, out completed); + // Read at most the number of bytes that will fit in the input buffer. The + // return value is the actual number of bytes read, or zero if no bytes remain. + bytesRead = stream.ReadAsync(bytes.AsMemory(), perReadTimeout, cancellationToken).GetAwaiter().GetResult(); - // The conversion produced the number of characters indicated by charsUsed. Write that number - // of characters to our result buffer - result.Append(chars, 0, charsUsed); + bool completed = false; + int byteIndex = 0; - // Increment byteIndex to the next block of bytes in the input buffer, if any, to convert. - byteIndex += bytesUsed; - - // The behavior of decoder.Convert changed start .NET 3.1-preview2. - // The change was made in https://github.com/dotnet/coreclr/pull/27229 - // The recommendation from .NET team is to not check for 'completed' if 'flush' is false. - // Break out of the loop if all bytes have been read. - if (!flush && bytesRead == byteIndex) + while (!completed) { - break; + // If this is the last input data, flush the decoder's internal buffer and state. + bool flush = bytesRead is 0; + decoder.Convert(bytes, byteIndex, bytesRead - byteIndex, chars, 0, useBufferSize, flush, out int bytesUsed, out int charsUsed, out completed); + + // The conversion produced the number of characters indicated by charsUsed. Write that number + // of characters to our result buffer + result.Append(chars, 0, charsUsed); + + // Increment byteIndex to the next block of bytes in the input buffer, if any, to convert. + byteIndex += bytesUsed; + + // The behavior of decoder.Convert changed start .NET 3.1-preview2. + // The change was made in https://github.com/dotnet/coreclr/pull/27229 + // The recommendation from .NET team is to not check for 'completed' if 'flush' is false. + // Break out of the loop if all bytes have been read. + if (!flush && bytesRead == byteIndex) + { + break; + } } } - } - while (bytesRead != 0); + while (bytesRead != 0); - return result.ToString(); + return result.ToString(); + } + finally + { + ArrayPool.Shared.Return(chars); + ArrayPool.Shared.Return(bytes); + } } - internal static string DecodeStream(Stream stream, string? characterSet, out Encoding encoding, CancellationToken cancellationToken) + internal static string DecodeStream(Stream stream, string? characterSet, out Encoding encoding, TimeSpan perReadTimeout, CancellationToken cancellationToken) { bool isDefaultEncoding = !TryGetEncoding(characterSet, out encoding); - string content = StreamToString(stream, encoding, cancellationToken); + string content = StreamToString(stream, encoding, perReadTimeout, cancellationToken); if (isDefaultEncoding) { // We only look within the first 1k characters as the meta element and @@ -415,7 +506,7 @@ namespace Microsoft.PowerShell.Commands if (TryGetEncoding(characterSet, out Encoding localEncoding)) { stream.Seek(0, SeekOrigin.Begin); - content = StreamToString(stream, localEncoding, cancellationToken); + content = StreamToString(stream, localEncoding, perReadTimeout, cancellationToken); encoding = localEncoding; } } @@ -459,8 +550,6 @@ namespace Microsoft.PowerShell.Commands return encoding.GetBytes(str); } - internal static string GetResponseString(HttpResponseMessage response, CancellationToken cancellationToken) => response.Content.ReadAsStringAsync(cancellationToken).GetAwaiter().GetResult(); - internal static Stream GetResponseStream(HttpResponseMessage response, CancellationToken cancellationToken) => response.Content.ReadAsStreamAsync(cancellationToken).GetAwaiter().GetResult(); #endregion Static Methods diff --git a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/WebRequestSession.cs b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/WebRequestSession.cs index d4a9a5cc48..a55a7dde38 100644 --- a/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/WebRequestSession.cs +++ b/src/Microsoft.PowerShell.Commands.Utility/commands/utility/WebCmdlet/WebRequestSession.cs @@ -32,7 +32,7 @@ namespace Microsoft.PowerShell.Commands private bool _skipCertificateCheck; private bool _noProxy; private bool _disposed; - private int _timeoutSec; + private TimeSpan _connectionTimeout; /// /// Contains true if an existing HttpClient had to be disposed and recreated since the WebSession was last used. @@ -142,7 +142,7 @@ namespace Microsoft.PowerShell.Commands internal bool SkipCertificateCheck { set => SetStructVar(ref _skipCertificateCheck, value); } - internal int TimeoutSec { set => SetStructVar(ref _timeoutSec, value); } + internal TimeSpan ConnectionTimeout { set => SetStructVar(ref _connectionTimeout, value); } internal bool NoProxy { @@ -240,7 +240,7 @@ namespace Microsoft.PowerShell.Commands // Check timeout setting (in seconds instead of milliseconds as in HttpWebRequest) return new HttpClient(handler) { - Timeout = _timeoutSec is 0 ? TimeSpan.FromMilliseconds(Timeout.Infinite) : TimeSpan.FromSeconds(_timeoutSec) + Timeout = _connectionTimeout }; } diff --git a/test/powershell/Modules/Microsoft.PowerShell.Utility/WebCmdlets.Tests.ps1 b/test/powershell/Modules/Microsoft.PowerShell.Utility/WebCmdlets.Tests.ps1 index da7a395ff2..21b28949e7 100644 --- a/test/powershell/Modules/Microsoft.PowerShell.Utility/WebCmdlets.Tests.ps1 +++ b/test/powershell/Modules/Microsoft.PowerShell.Utility/WebCmdlets.Tests.ps1 @@ -241,10 +241,10 @@ function ExecuteRequestWithCustomUserAgent { try { $Params = @{ - Uri = $Uri - TimeoutSec = 5 - UserAgent = $UserAgent - SkipHeaderValidation = $SkipHeaderValidation.IsPresent + Uri = $Uri + ConnectionTimeoutSeconds = 5 + UserAgent = $UserAgent + SkipHeaderValidation = $SkipHeaderValidation.IsPresent } if ($Cmdlet -eq 'Invoke-WebRequest') { $result.Output = Invoke-WebRequest @Params @@ -608,12 +608,20 @@ Describe "Invoke-WebRequest tests" -Tags "Feature", "RequireAdminOnWindows" { $Result.Output.Content | Should -Match '测试123' } - It "Invoke-WebRequest validate timeout option" { + It "Invoke-WebRequest validate ConnectionTimeoutSeconds option" { + $uri = Get-WebListenerUrl -Test 'Delay' -TestValue '5' + $command = "Invoke-WebRequest -Uri '$uri' -ConnectionTimeoutSeconds 2" + + $result = ExecuteWebCommand -command $command + $result.Error.FullyQualifiedErrorId | Should -Be "ConnectionTimeoutReached,Microsoft.PowerShell.Commands.InvokeWebRequestCommand" + } + + It "Invoke-WebRequest validate TimeoutSec alias" { $uri = Get-WebListenerUrl -Test 'Delay' -TestValue '5' $command = "Invoke-WebRequest -Uri '$uri' -TimeoutSec 2" $result = ExecuteWebCommand -command $command - $result.Error.FullyQualifiedErrorId | Should -Be "System.Threading.Tasks.TaskCanceledException,Microsoft.PowerShell.Commands.InvokeWebRequestCommand" + $result.Error.FullyQualifiedErrorId | Should -Be "ConnectionTimeoutReached,Microsoft.PowerShell.Commands.InvokeWebRequestCommand" } It "Validate Invoke-WebRequest error with -Proxy and -NoProxy option" { @@ -2650,12 +2658,20 @@ Describe "Invoke-RestMethod tests" -Tags "Feature", "RequireAdminOnWindows" { $Result.Output | Should -Match '测试123' } - It "Invoke-RestMethod validate timeout option" { + It "Invoke-RestMethod validate ConnectionTimeoutSeconds option" { + $uri = Get-WebListenerUrl -Test 'Delay' -TestValue '5' + $command = "Invoke-RestMethod -Uri '$uri' -ConnectionTimeoutSeconds 2" + + $result = ExecuteWebCommand -command $command + $result.Error.FullyQualifiedErrorId | Should -Be "ConnectionTimeoutReached,Microsoft.PowerShell.Commands.InvokeRestMethodCommand" + } + + It "Invoke-RestMethod validate TimeoutSec alias" { $uri = Get-WebListenerUrl -Test 'Delay' -TestValue '5' $command = "Invoke-RestMethod -Uri '$uri' -TimeoutSec 2" $result = ExecuteWebCommand -command $command - $result.Error.FullyQualifiedErrorId | Should -Be "System.Threading.Tasks.TaskCanceledException,Microsoft.PowerShell.Commands.InvokeRestMethodCommand" + $result.Error.FullyQualifiedErrorId | Should -Be "ConnectionTimeoutReached,Microsoft.PowerShell.Commands.InvokeRestMethodCommand" } It "Validate Invoke-RestMethod error with -Proxy and -NoProxy option" { @@ -4373,7 +4389,12 @@ Describe 'Invoke-WebRequest and Invoke-RestMethod support Cancellation through C RunWithCancellation -Uri $uri } - It 'Invoke-WebRequest: Defalate Compression CTRL-C Cancels request after request headers' { + It 'Invoke-WebRequest: Gzip Compression CTRL-C Cancels request after request headers with Content-Length' { + $uri = Get-WebListenerUrl -Test StallGzip -TestValue '30/application%2fjson' -Query @{ contentLength = $true } + RunWithCancellation -Uri $uri + } + + It 'Invoke-WebRequest: Deflate Compression CTRL-C Cancels request after request headers' { $uri = Get-WebListenerUrl -Test StallDeflate -TestValue '30/application%2fjson' RunWithCancellation -Uri $uri } @@ -4388,7 +4409,7 @@ Describe 'Invoke-WebRequest and Invoke-RestMethod support Cancellation through C RunWithCancellation -Uri $uri -Arguments '-SkipCertificateCheck' } - It 'Invoke-WebRequest: HTTPS with Defalte compression CTRL-C Cancels request after request headers' { + It 'Invoke-WebRequest: HTTPS with Deflate compression CTRL-C Cancels request after request headers' { $uri = Get-WebListenerUrl -Https -Test StallDeflate -TestValue '30/application%2fjson' RunWithCancellation -Uri $uri -Arguments '-SkipCertificateCheck' } @@ -4443,3 +4464,82 @@ Describe 'Invoke-WebRequest and Invoke-RestMethod support Cancellation through C RunWithCancellation -Command 'Invoke-RestMethod' -Uri $uri } } + +Describe 'Invoke-WebRequest and Invoke-RestMethod support OperationTimeoutSeconds' -Tags "CI", "RequireAdminOnWindows" { + BeforeAll { + $oldProgress = $ProgressPreference + $ProgressPreference = 'SilentlyContinue' + $WebListener = Start-WebListener + } + + AfterAll { + $ProgressPreference = $oldProgress + } + + function RunWithNetworkTimeout { + param( + [ValidateSet('Invoke-WebRequest', 'Invoke-RestMethod')] + [string]$Command = 'Invoke-WebRequest', + [string]$Arguments = '', + [uri]$Uri, + [int]$OperationTimeoutSeconds, + [switch]$WillTimeout + ) + + $invoke = "$Command -Uri `"$Uri`" $Arguments" + if ($PSBoundParameters.ContainsKey('OperationTimeoutSeconds')) { + $invoke = "$invoke -OperationTimeoutSeconds $OperationTimeoutSeconds" + } + + $result = ExecuteWebCommand -command $invoke + if ($WillTimeout) { + $result.Error | Should -Not -BeNullOrEmpty + $fqErrorClass = if ($Command -eq 'Invoke-WebRequest') { 'InvokeWebRequestCommand'} else { 'InvokeRestMethodCommand'} + $result.Error.FullyQualifiedErrorId | Should -Be "OperationTimeoutReached,Microsoft.PowerShell.Commands.$fqErrorClass" + $result.Output | Should -BeNullOrEmpty + } else { + $result.Error | Should -BeNullOrEmpty + $result.Output | Should -Not -BeNullOrEmpty + } + } + + It 'Invoke-WebRequest: OperationTimeoutSeconds does not cancel if stalls shorter than timeout but download takes longer than timeout' { + $uri = Get-WebListenerUrl -Test Stall -TestValue '2' -Query @{ chunks = 5 } + RunWithNetworkTimeout -Uri $uri -OperationTimeoutSeconds 4 + } + + It 'Invoke-WebRequest: OperationTimeoutSeconds cancels if stall lasts longer than OperationTimeoutSeconds value' { + $uri = Get-WebListenerUrl -Test Stall -TestValue 30 + RunWithNetworkTimeout -Uri $uri -OperationTimeoutSeconds 3 -WillTimeout + } + + It 'Invoke-WebRequest: OperationTimeoutSeconds cancels if stall lasts longer than OperationTimeoutSeconds value for HTTPS/gzip compression' { + $uri = Get-WebListenerUrl -Https -Test StallGzip -TestValue 30 + RunWithNetworkTimeout -Uri $uri -OperationTimeoutSeconds 3 -WillTimeout -Arguments '-SkipCertificateCheck' + } + + It 'Invoke-RestMethod: OperationTimeoutSeconds does not cancel if stalls shorter than timeout but download takes longer than timeout' { + $uri = Get-WebListenerUrl -Test Stall -TestValue '2' -Query @{ chunks = 5 } + RunWithNetworkTimeout -Command Invoke-RestMethod -Uri $uri -OperationTimeoutSeconds 4 + } + + It 'Invoke-RestMethod: OperationTimeoutSeconds cancels if stall lasts longer than OperationTimeoutSeconds value' { + $uri = Get-WebListenerUrl -Test Stall -TestValue 30 + RunWithNetworkTimeout -Command Invoke-RestMethod -Uri $uri -OperationTimeoutSeconds 2 -WillTimeout + } + + It 'Invoke-RestMethod: OperationTimeoutSeconds cancels when doing XML atom processing' { + $uri = Get-WebListenerUrl -Test Stall -TestValue '30/application%2fxml' + RunWithNetworkTimeout -Command Invoke-RestMethod -Uri $uri -OperationTimeoutSeconds 2 -WillTimeout + } + + It 'Invoke-RestMethod: OperationTimeoutSeconds cancels when doing JSON processing' { + $uri = Get-WebListenerUrl -Test Stall -TestValue '30/application%2fjson' + RunWithNetworkTimeout -Command Invoke-RestMethod -Uri $uri -OperationTimeoutSeconds 2 -WillTimeout + } + + It 'Invoke-RestMethod: OperationTimeoutSeconds cancels when doing XML atom processing for HTTPS/gzip compression' { + $uri = Get-WebListenerUrl -Https -Test StallGzip -TestValue 30/application%2fXML + RunWithNetworkTimeout -Command Invoke-RestMethod -Uri $uri -OperationTimeoutSeconds 2 -WillTimeout -Arguments '-SkipCertificateCheck' + } +} diff --git a/test/tools/WebListener/Controllers/DelayController.cs b/test/tools/WebListener/Controllers/DelayController.cs index 5ab081d580..5fdd9051d3 100644 --- a/test/tools/WebListener/Controllers/DelayController.cs +++ b/test/tools/WebListener/Controllers/DelayController.cs @@ -54,33 +54,33 @@ namespace mvc.Controllers return getController.Index(); } - public async Task Stall(int seconds, string contentType, CancellationToken cancellationToken) + public async Task Stall(int seconds, string contentType, int chunks, bool contentLength, CancellationToken cancellationToken) { - await WriteStallResponse(seconds, contentType, null, null, cancellationToken); + await WriteStallResponse(seconds, contentType, chunks, contentLength, null, null, cancellationToken); } - public async Task StallBrotli(int seconds, string contentType, CancellationToken cancellationToken) + public async Task StallBrotli(int seconds, string contentType, int chunks, bool contentLength, CancellationToken cancellationToken) { using var memStream = new MemoryStream(); using var compressedStream = new BrotliStream(memStream, CompressionLevel.Fastest); Response.Headers.ContentEncoding = "br"; - await WriteStallResponse(seconds, contentType, compressedStream, memStream, cancellationToken); + await WriteStallResponse(seconds, contentType, chunks, contentLength, compressedStream, memStream, cancellationToken); } - public async Task StallDeflate(int seconds, string contentType, CancellationToken cancellationToken) + public async Task StallDeflate(int seconds, string contentType, int chunks, bool contentLength, CancellationToken cancellationToken) { using var memStream = new MemoryStream(); using var compressedStream = new DeflateStream(memStream, CompressionLevel.Fastest); Response.Headers.ContentEncoding = "deflate"; - await WriteStallResponse(seconds, contentType, compressedStream, memStream, cancellationToken); + await WriteStallResponse(seconds, contentType, chunks, contentLength, compressedStream, memStream, cancellationToken); } - public async Task StallGZip(int seconds, string contentType, CancellationToken cancellationToken) + public async Task StallGZip(int seconds, string contentType, int chunks, bool contentLength, CancellationToken cancellationToken) { using var memStream = new MemoryStream(); using var compressedStream = new GZipStream(memStream, CompressionLevel.Fastest); Response.Headers.ContentEncoding = "gzip"; - await WriteStallResponse(seconds, contentType, compressedStream, memStream, cancellationToken); + await WriteStallResponse(seconds, contentType, chunks, contentLength, compressedStream, memStream, cancellationToken); } public IActionResult Error() @@ -88,7 +88,7 @@ namespace mvc.Controllers return View(new ErrorViewModel { RequestId = Activity.Current?.Id ?? HttpContext.TraceIdentifier }); } - private async Task WriteStallResponse(int seconds, string contentType, Stream stream, MemoryStream memStream, CancellationToken cancellationToken) + private async Task WriteStallResponse(int seconds, string contentType, int chunks, bool contentLength, Stream stream, MemoryStream memStream, CancellationToken cancellationToken) { if (string.IsNullOrWhiteSpace(contentType)) { @@ -124,20 +124,41 @@ namespace mvc.Controllers stream.Close(); response = memStream.ToArray(); } - int midPoint = response.Length / 2; - - // Start writing approx half the content, including headers and then delay before writing the rest. - await Response.Body.WriteAsync(response, 0, midPoint, cancellationToken); - await Response.Body.FlushAsync(cancellationToken); - - if (seconds > 0) + if (chunks < 2) { - int milliseconds = seconds * 1000; - await Task.Delay(milliseconds); + chunks = 2; + } + if (chunks > response.Length) + { + throw new InvalidDataException($"Response message is not big enough to break into {chunks} chunks. (Size {response.Length} bytes)."); } - await Response.Body.WriteAsync(response, midPoint, response.Length - midPoint, cancellationToken); - await Response.Body.FlushAsync(cancellationToken); + if (contentLength) + { + Response.ContentLength = response.Length; + } + int chunkSize = response.Length / chunks; + int currentPos = 0; + + // Write each of the content chunks followed by a delay + // The last segment makes up the remainder of the content if + // it doesn't divide neatly into the required chunks + for (int i = 0; i < chunks; i++) + { + if (i == chunks - 1) + { + chunkSize = response.Length - currentPos; + seconds = 0; + } + await Response.Body.WriteAsync(response, currentPos, chunkSize, cancellationToken); + await Response.Body.FlushAsync(cancellationToken); + currentPos += chunkSize; + if (seconds > 0) + { + int milliseconds = seconds * 1000; + await Task.Delay(milliseconds); + } + } } } }