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);
+ }
+ }
}
}
}