diff --git a/sdk/core/Azure.Core/api/Azure.Core.net461.cs b/sdk/core/Azure.Core/api/Azure.Core.net461.cs index 4766740cf92a..3e2d895e4135 100644 --- a/sdk/core/Azure.Core/api/Azure.Core.net461.cs +++ b/sdk/core/Azure.Core/api/Azure.Core.net461.cs @@ -194,6 +194,7 @@ public partial class RequestContext public RequestContext() { } public System.Threading.CancellationToken CancellationToken { get { throw null; } set { } } public Azure.ErrorOptions ErrorOptions { get { throw null; } set { } } + public void AddPolicy(Azure.Core.Pipeline.HttpPipelinePolicy policy, Azure.Core.HttpPipelinePosition position) { } public static implicit operator Azure.RequestContext (Azure.ErrorOptions options) { throw null; } } public partial class RequestFailedException : System.Exception, System.Runtime.Serialization.ISerializable @@ -764,6 +765,7 @@ public HttpPipeline(Azure.Core.Pipeline.HttpPipelineTransport transport, Azure.C public static System.IDisposable CreateClientRequestIdScope(string? clientRequestId) { throw null; } public static System.IDisposable CreateHttpMessagePropertiesScope(System.Collections.Generic.IDictionary messageProperties) { throw null; } public Azure.Core.HttpMessage CreateMessage() { throw null; } + public Azure.Core.HttpMessage CreateMessage(Azure.RequestContext context) { throw null; } public Azure.Core.Request CreateRequest() { throw null; } public void Send(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { } public System.Threading.Tasks.ValueTask SendAsync(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { throw null; } diff --git a/sdk/core/Azure.Core/api/Azure.Core.net5.0.cs b/sdk/core/Azure.Core/api/Azure.Core.net5.0.cs index 48642f2d38ab..113ff6909e9e 100644 --- a/sdk/core/Azure.Core/api/Azure.Core.net5.0.cs +++ b/sdk/core/Azure.Core/api/Azure.Core.net5.0.cs @@ -194,6 +194,7 @@ public partial class RequestContext public RequestContext() { } public System.Threading.CancellationToken CancellationToken { get { throw null; } set { } } public Azure.ErrorOptions ErrorOptions { get { throw null; } set { } } + public void AddPolicy(Azure.Core.Pipeline.HttpPipelinePolicy policy, Azure.Core.HttpPipelinePosition position) { } public static implicit operator Azure.RequestContext (Azure.ErrorOptions options) { throw null; } } public partial class RequestFailedException : System.Exception, System.Runtime.Serialization.ISerializable @@ -764,6 +765,7 @@ public HttpPipeline(Azure.Core.Pipeline.HttpPipelineTransport transport, Azure.C public static System.IDisposable CreateClientRequestIdScope(string? clientRequestId) { throw null; } public static System.IDisposable CreateHttpMessagePropertiesScope(System.Collections.Generic.IDictionary messageProperties) { throw null; } public Azure.Core.HttpMessage CreateMessage() { throw null; } + public Azure.Core.HttpMessage CreateMessage(Azure.RequestContext context) { throw null; } public Azure.Core.Request CreateRequest() { throw null; } public void Send(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { } public System.Threading.Tasks.ValueTask SendAsync(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { throw null; } diff --git a/sdk/core/Azure.Core/api/Azure.Core.netcoreapp2.1.cs b/sdk/core/Azure.Core/api/Azure.Core.netcoreapp2.1.cs index 4766740cf92a..3e2d895e4135 100644 --- a/sdk/core/Azure.Core/api/Azure.Core.netcoreapp2.1.cs +++ b/sdk/core/Azure.Core/api/Azure.Core.netcoreapp2.1.cs @@ -194,6 +194,7 @@ public partial class RequestContext public RequestContext() { } public System.Threading.CancellationToken CancellationToken { get { throw null; } set { } } public Azure.ErrorOptions ErrorOptions { get { throw null; } set { } } + public void AddPolicy(Azure.Core.Pipeline.HttpPipelinePolicy policy, Azure.Core.HttpPipelinePosition position) { } public static implicit operator Azure.RequestContext (Azure.ErrorOptions options) { throw null; } } public partial class RequestFailedException : System.Exception, System.Runtime.Serialization.ISerializable @@ -764,6 +765,7 @@ public HttpPipeline(Azure.Core.Pipeline.HttpPipelineTransport transport, Azure.C public static System.IDisposable CreateClientRequestIdScope(string? clientRequestId) { throw null; } public static System.IDisposable CreateHttpMessagePropertiesScope(System.Collections.Generic.IDictionary messageProperties) { throw null; } public Azure.Core.HttpMessage CreateMessage() { throw null; } + public Azure.Core.HttpMessage CreateMessage(Azure.RequestContext context) { throw null; } public Azure.Core.Request CreateRequest() { throw null; } public void Send(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { } public System.Threading.Tasks.ValueTask SendAsync(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { throw null; } diff --git a/sdk/core/Azure.Core/api/Azure.Core.netstandard2.0.cs b/sdk/core/Azure.Core/api/Azure.Core.netstandard2.0.cs index 4766740cf92a..3e2d895e4135 100644 --- a/sdk/core/Azure.Core/api/Azure.Core.netstandard2.0.cs +++ b/sdk/core/Azure.Core/api/Azure.Core.netstandard2.0.cs @@ -194,6 +194,7 @@ public partial class RequestContext public RequestContext() { } public System.Threading.CancellationToken CancellationToken { get { throw null; } set { } } public Azure.ErrorOptions ErrorOptions { get { throw null; } set { } } + public void AddPolicy(Azure.Core.Pipeline.HttpPipelinePolicy policy, Azure.Core.HttpPipelinePosition position) { } public static implicit operator Azure.RequestContext (Azure.ErrorOptions options) { throw null; } } public partial class RequestFailedException : System.Exception, System.Runtime.Serialization.ISerializable @@ -764,6 +765,7 @@ public HttpPipeline(Azure.Core.Pipeline.HttpPipelineTransport transport, Azure.C public static System.IDisposable CreateClientRequestIdScope(string? clientRequestId) { throw null; } public static System.IDisposable CreateHttpMessagePropertiesScope(System.Collections.Generic.IDictionary messageProperties) { throw null; } public Azure.Core.HttpMessage CreateMessage() { throw null; } + public Azure.Core.HttpMessage CreateMessage(Azure.RequestContext context) { throw null; } public Azure.Core.Request CreateRequest() { throw null; } public void Send(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { } public System.Threading.Tasks.ValueTask SendAsync(Azure.Core.HttpMessage message, System.Threading.CancellationToken cancellationToken) { throw null; } diff --git a/sdk/core/Azure.Core/src/HttpMessage.cs b/sdk/core/Azure.Core/src/HttpMessage.cs index 6a16bb58e694..c387c028089e 100644 --- a/sdk/core/Azure.Core/src/HttpMessage.cs +++ b/sdk/core/Azure.Core/src/HttpMessage.cs @@ -4,7 +4,6 @@ using System; using System.Collections.Generic; using System.IO; -using System.Net.Http; using System.Threading; using Azure.Core.Pipeline; @@ -13,7 +12,7 @@ namespace Azure.Core /// /// Represents a context flowing through the . /// - public sealed class HttpMessage: IDisposable + public sealed class HttpMessage : IDisposable { private Dictionary? _properties; @@ -81,6 +80,19 @@ public Response Response /// public TimeSpan? NetworkTimeout { get; set; } + internal void AddPolicies(RequestContext context) + { + if (context == null || context.Policies == null || context.Policies.Count == 0) + { + return; + } + + Policies ??= new(context.Policies.Count); + Policies.AddRange(context.Policies); + } + + internal List<(HttpPipelinePosition Position, HttpPipelinePolicy Policy)>? Policies { get; set; } + /// /// Gets a property that modifies the pipeline behavior. Please refer to individual policies documentation on what properties it supports. /// @@ -115,7 +127,7 @@ public void SetProperty(string name, object value) { case ResponseShouldNotBeUsedStream responseContent: return responseContent.Original; - case Stream stream : + case Stream stream: _response.ContentStream = new ResponseShouldNotBeUsedStream(_response.ContentStream); return stream; default: @@ -132,7 +144,7 @@ public void Dispose() _response?.Dispose(); } - private class ResponseShouldNotBeUsedStream: Stream + private class ResponseShouldNotBeUsedStream : Stream { public Stream Original { get; } diff --git a/sdk/core/Azure.Core/src/Pipeline/DisposableHttpPipeline.cs b/sdk/core/Azure.Core/src/Pipeline/DisposableHttpPipeline.cs index df7e0d3ac8d7..a5652f19c590 100644 --- a/sdk/core/Azure.Core/src/Pipeline/DisposableHttpPipeline.cs +++ b/sdk/core/Azure.Core/src/Pipeline/DisposableHttpPipeline.cs @@ -14,11 +14,14 @@ public sealed class DisposableHttpPipeline : HttpPipeline, IDisposable /// Creates a new instance of with the provided transport, policies and response classifier. /// /// The to use for sending the requests. + /// + /// /// Policies to be invoked as part of the pipeline in order. /// The response classifier to be used in invocations. - internal DisposableHttpPipeline(HttpPipelineTransport transport, HttpPipelinePolicy[]? policies = null, ResponseClassifier? responseClassifier = null) - : base(transport, policies, responseClassifier) - { } + internal DisposableHttpPipeline(HttpPipelineTransport transport, int perCallIndex, int perRetryIndex, HttpPipelinePolicy[]? policies = null, ResponseClassifier? responseClassifier = null) + : base(transport, perCallIndex, perRetryIndex, policies, responseClassifier) + { + } /// /// Calls Dispose on the underlying . diff --git a/sdk/core/Azure.Core/src/Pipeline/HttpPipeline.cs b/sdk/core/Azure.Core/src/Pipeline/HttpPipeline.cs index defcaaa9e006..817454226122 100644 --- a/sdk/core/Azure.Core/src/Pipeline/HttpPipeline.cs +++ b/sdk/core/Azure.Core/src/Pipeline/HttpPipeline.cs @@ -2,6 +2,7 @@ // Licensed under the MIT License. using System; +using System.Buffers; using System.Collections.Generic; using System.Threading; using System.Threading.Tasks; @@ -19,6 +20,25 @@ public class HttpPipeline private readonly ReadOnlyMemory _pipeline; + /// + /// Indicates whether or not the pipeline was created using its internal constructor. + /// If it was, we know the indices where we can add per-request policies at positions + /// and . + /// + private readonly bool _internallyConstructed; + + /// + /// The pipeline index where policies will be added, + /// if any are specified using . + /// + private readonly int _perCallIndex; + + /// + /// The pipeline index where policies will be added, + /// if any are specified using . + /// + private readonly int _perRetryIndex; + /// /// Creates a new instance of with the provided transport, policies and response classifier. /// @@ -39,6 +59,13 @@ public HttpPipeline(HttpPipelineTransport transport, HttpPipelinePolicy[]? polic _pipeline = all; } + internal HttpPipeline(HttpPipelineTransport transport, int perCallIndex, int perRetryIndex, HttpPipelinePolicy[]? policies = null, ResponseClassifier? responseClassifier = null) : this(transport, policies, responseClassifier) + { + _perCallIndex = perCallIndex; + _perRetryIndex = perRetryIndex; + _internallyConstructed = true; + } + /// /// Creates a new instance. /// @@ -55,6 +82,18 @@ public HttpMessage CreateMessage() return new HttpMessage(CreateRequest(), ResponseClassifier); } + /// + /// Creates a new instance. + /// + /// Context specifying the message options. + /// The message. + public HttpMessage CreateMessage(RequestContext context) + { + var message = CreateMessage(); + message.AddPolicies(context); + return message; + } + /// /// The instance used in this pipeline invocations. /// @@ -70,7 +109,28 @@ public ValueTask SendAsync(HttpMessage message, CancellationToken cancellationTo { message.CancellationToken = cancellationToken; AddHttpMessageProperties(message); - return _pipeline.Span[0].ProcessAsync(message, _pipeline.Slice(1)); + + if (message.Policies == null || message.Policies.Count == 0) + { + return _pipeline.Span[0].ProcessAsync(message, _pipeline.Slice(1)); + } + + return SendAsync(message); + } + + private async ValueTask SendAsync(HttpMessage message) + { + var length = _pipeline.Length + message.Policies!.Count; + var policies = ArrayPool.Shared.Rent(length); + try + { + var pipeline = CreateRequestPipeline(policies, message.Policies); + await pipeline.Span[0].ProcessAsync(message, pipeline.Slice(1)).ConfigureAwait(false); + } + finally + { + ArrayPool.Shared.Return(policies); + } } /// @@ -82,8 +142,27 @@ public void Send(HttpMessage message, CancellationToken cancellationToken) { message.CancellationToken = cancellationToken; AddHttpMessageProperties(message); - _pipeline.Span[0].Process(message, _pipeline.Slice(1)); + + if (message.Policies == null || message.Policies.Count == 0) + { + _pipeline.Span[0].Process(message, _pipeline.Slice(1)); + } + else + { + var length = _pipeline.Length + message.Policies.Count; + var policies = ArrayPool.Shared.Rent(length); + try + { + var pipeline = CreateRequestPipeline(policies, message.Policies); + pipeline.Span[0].Process(message, pipeline.Slice(1)); + } + finally + { + ArrayPool.Shared.Return(policies); + } + } } + /// /// Invokes the pipeline asynchronously with the provided request. /// @@ -144,6 +223,60 @@ public static IDisposable CreateHttpMessagePropertiesScope(IDictionary CreateRequestPipeline(HttpPipelinePolicy[] policies, List<(HttpPipelinePosition Position, HttpPipelinePolicy Policy)> customPolicies) + { + if (!_internallyConstructed) + { + throw new InvalidOperationException("Cannot send messages with per-request policies if the pipeline wasn't constructed with HttpPipelineBuilder."); + } + + // Copy over client policies and splice in custom policies at designated indices + var pipeline = _pipeline.Span; + int transportIndex = pipeline.Length - 1; + + pipeline.Slice(0, _perCallIndex).CopyTo(policies); + + int index = _perCallIndex; + int count = AddCustomPolicies(customPolicies, policies, HttpPipelinePosition.PerCall, index); + + index += count; + count = _perRetryIndex - _perCallIndex; + pipeline.Slice(_perCallIndex, count).CopyTo(policies.AsSpan(index, count)); + + index += count; + count = AddCustomPolicies(customPolicies, policies, HttpPipelinePosition.PerRetry, index); + + index += count; + count = transportIndex - _perRetryIndex; + pipeline.Slice(_perRetryIndex, count).CopyTo(policies.AsSpan(index, count)); + + index += count; + count = AddCustomPolicies(customPolicies, policies, HttpPipelinePosition.BeforeTransport, index); + + index += count; + policies[index] = pipeline[transportIndex]; + + return new ReadOnlyMemory(policies, 0, index + 1); + } + + private static int AddCustomPolicies(List<(HttpPipelinePosition Position, HttpPipelinePolicy Policy)> source, HttpPipelinePolicy[] target, HttpPipelinePosition position, int start) + { + int count = 0; + if (source != null) + { + foreach (var policy in source) + { + if (policy.Position == position) + { + target[start + count] = policy.Policy; + count++; + } + } + } + + return count; + } + private static void AddHttpMessageProperties(HttpMessage message) { if (CurrentHttpMessagePropertiesScope.Value != null) diff --git a/sdk/core/Azure.Core/src/Pipeline/HttpPipelineBuilder.cs b/sdk/core/Azure.Core/src/Pipeline/HttpPipelineBuilder.cs index 6bb73d23993c..8cb7371ae000 100644 --- a/sdk/core/Azure.Core/src/Pipeline/HttpPipelineBuilder.cs +++ b/sdk/core/Azure.Core/src/Pipeline/HttpPipelineBuilder.cs @@ -53,6 +53,8 @@ public static HttpPipeline Build( /// A new instance of public static HttpPipeline Build(ClientOptions options, HttpPipelinePolicy[] perCallPolicies, HttpPipelinePolicy[] perRetryPolicies, ResponseClassifier responseClassifier, HttpPipelineTransportOptions? defaultTransportOptions) { + int perCallIndex; + int perRetryIndex; if (perCallPolicies == null) { throw new ArgumentNullException(nameof(perCallPolicies)); @@ -94,6 +96,9 @@ void AddCustomerPolicies(HttpPipelinePosition position) AddCustomerPolicies(HttpPipelinePosition.PerCall); + policies.RemoveAll(static policy => policy == null); + perCallIndex = policies.Count; + policies.Add(ClientRequestIdPolicy.Shared); if (diagnostics.IsTelemetryEnabled) @@ -110,6 +115,9 @@ void AddCustomerPolicies(HttpPipelinePosition position) AddCustomerPolicies(HttpPipelinePosition.PerRetry); + policies.RemoveAll(static policy => policy == null); + perRetryIndex = policies.Count; + if (diagnostics.IsLoggingEnabled) { string assemblyName = options.GetType().Assembly!.GetName().Name!; @@ -122,7 +130,6 @@ void AddCustomerPolicies(HttpPipelinePosition position) policies.Add(new RequestActivityPolicy(isDistributedTracingEnabled, ClientDiagnostics.GetResourceProviderNamespace(options.GetType().Assembly), sanitizer)); AddCustomerPolicies(HttpPipelinePosition.BeforeTransport); - policies.RemoveAll(static policy => policy == null); // Override the provided Transport with the provided transport options if the transport has not been set after default construction and options are not null. @@ -141,12 +148,16 @@ void AddCustomerPolicies(HttpPipelinePosition position) { transport = HttpPipelineTransport.Create(defaultTransportOptions); return new DisposableHttpPipeline(transport, + perCallIndex, + perRetryIndex, policies.ToArray(), responseClassifier); } } return new HttpPipeline(transport, + perCallIndex, + perRetryIndex, policies.ToArray(), responseClassifier); } diff --git a/sdk/core/Azure.Core/src/Pipeline/HttpPipelineTransport.cs b/sdk/core/Azure.Core/src/Pipeline/HttpPipelineTransport.cs index d1c5d020e39b..61257d8e0f49 100644 --- a/sdk/core/Azure.Core/src/Pipeline/HttpPipelineTransport.cs +++ b/sdk/core/Azure.Core/src/Pipeline/HttpPipelineTransport.cs @@ -25,7 +25,7 @@ public abstract class HttpPipelineTransport /// /// Creates a new transport specific instance of . This should not be called directly, or - /// should be used instead. + /// should be used instead. /// /// public abstract Request CreateRequest(); diff --git a/sdk/core/Azure.Core/src/Request.cs b/sdk/core/Azure.Core/src/Request.cs index 7c9f13d99420..011af269df7b 100644 --- a/sdk/core/Azure.Core/src/Request.cs +++ b/sdk/core/Azure.Core/src/Request.cs @@ -10,7 +10,7 @@ namespace Azure.Core { /// - /// Represents an HTTP request. Use or to create an instance. + /// Represents an HTTP request. Use or to create an instance. /// #pragma warning disable AZC0012 // Avoid single word type names public abstract class Request : IDisposable diff --git a/sdk/core/Azure.Core/src/RequestContext.cs b/sdk/core/Azure.Core/src/RequestContext.cs index e3322fed7109..b569d9bac49d 100644 --- a/sdk/core/Azure.Core/src/RequestContext.cs +++ b/sdk/core/Azure.Core/src/RequestContext.cs @@ -2,6 +2,7 @@ // Licensed under the MIT License. using System; +using System.Collections.Generic; using System.Threading; using Azure.Core; using Azure.Core.Pipeline; @@ -13,6 +14,8 @@ namespace Azure /// public class RequestContext { + internal List<(HttpPipelinePosition Position, HttpPipelinePolicy Policy)>? Policies { get; private set; } + /// /// Initializes a new instance of the class. /// @@ -35,5 +38,19 @@ public RequestContext() /// Controls under what conditions the operation raises an exception if the underlying response indicates a failure. /// public ErrorOptions ErrorOptions { get; set; } = ErrorOptions.Default; + + /// + /// Adds an into the pipeline for the duration of this request. + /// The position of policy in the pipeline is controlled by parameter. + /// If you want the policy to execute once per client request use + /// otherwise use to run the policy for every retry. + /// + /// The instance to be added to the pipeline. + /// The position of the policy in the pipeline. + public void AddPolicy(HttpPipelinePolicy policy, HttpPipelinePosition position) + { + Policies ??= new(); + Policies.Add((position, policy)); + } } } diff --git a/sdk/core/Azure.Core/tests/HttpPipelineBuilderTest.cs b/sdk/core/Azure.Core/tests/HttpPipelineBuilderTest.cs index 272e7c141b71..9a61131b9aea 100644 --- a/sdk/core/Azure.Core/tests/HttpPipelineBuilderTest.cs +++ b/sdk/core/Azure.Core/tests/HttpPipelineBuilderTest.cs @@ -234,6 +234,19 @@ public void SetTransportOptions([Values(true, false)] bool isCustomTransportSet) } } + [Test] + public void CanPassNullPolicies([Values(true, false)] bool isCustomTransportSet) + { + var pipeline = HttpPipelineBuilder.Build( + new TestOptions(), + new HttpPipelinePolicy[] { null }, + new HttpPipelinePolicy[] { null }, + null); + + var message = pipeline.CreateMessage(); + pipeline.SendAsync(message, message.CancellationToken); + } + private class TestOptions : ClientOptions { public TestOptions() diff --git a/sdk/core/Azure.Core/tests/HttpPipelineTests.cs b/sdk/core/Azure.Core/tests/HttpPipelineTests.cs index 7582b41237b2..3f2a682c5395 100644 --- a/sdk/core/Azure.Core/tests/HttpPipelineTests.cs +++ b/sdk/core/Azure.Core/tests/HttpPipelineTests.cs @@ -1,6 +1,8 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +using System; +using System.Linq; using System.Threading.Tasks; using Azure.Core.Pipeline; using Azure.Core.TestFramework; @@ -29,5 +31,215 @@ public async Task DoesntDisposeRequestInSendRequestAsync() private class TestOptions : ClientOptions { } + + [Test] + public async Task CanAddPolicy_PerCall() + { + var mockTransport = new MockTransport(new MockResponse(200)); + var options = new TestOptions() + { + Transport = mockTransport, + }; + var pipeline = HttpPipelineBuilder.Build(options); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("PerCallHeader", "Value"), HttpPipelinePosition.PerCall); + + var message = pipeline.CreateMessage(context); + await pipeline.SendAsync(message, message.CancellationToken); + + Request request = mockTransport.Requests[0]; + Assert.IsTrue(request.Headers.TryGetValues("PerCallHeader", out var values)); + Assert.AreEqual(1, values.Count()); + Assert.AreEqual("Value", values.ElementAt(0)); + } + + [Test] + public async Task CanAddPolicy_PerRetry() + { + var retryResponse = new MockResponse(408); // Request Timeout + var mockTransport = new MockTransport(retryResponse, retryResponse, new MockResponse(200)); + var options = new TestOptions() + { + Transport = mockTransport, + }; + + var pipeline = HttpPipelineBuilder.Build(options); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("PerRetryHeader", "Value"), HttpPipelinePosition.PerRetry); + + var message = pipeline.CreateMessage(context); + await pipeline.SendAsync(message, message.CancellationToken); + + Request request = mockTransport.Requests[0]; + Assert.IsTrue(request.Headers.TryGetValues("PerRetryHeader", out var values)); + Assert.AreEqual(3, values.Count()); + Assert.AreEqual("Value", values.ElementAt(0)); + Assert.AreEqual("Value", values.ElementAt(1)); + Assert.AreEqual("Value", values.ElementAt(2)); + } + + [Test] + public async Task CanAddPolicy_BeforeTransport() + { + var retryResponse = new MockResponse(408); // Request Timeout + + // retry twice + var mockTransport = new MockTransport(retryResponse, retryResponse, new MockResponse(200)); + var options = new TestOptions() + { + Transport = mockTransport, + }; + + var pipeline = HttpPipelineBuilder.Build(options); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("BeforeTransportHeader", "Value"), HttpPipelinePosition.BeforeTransport); + + var message = pipeline.CreateMessage(context); + await pipeline.SendAsync(message, message.CancellationToken); + + Request request = mockTransport.Requests[0]; + + Assert.IsTrue(request.Headers.TryGetValues("BeforeTransportHeader", out var values)); + Assert.AreEqual(3, values.Count()); + Assert.AreEqual("Value", values.ElementAt(0)); + Assert.AreEqual("Value", values.ElementAt(1)); + Assert.AreEqual("Value", values.ElementAt(2)); + } + + [Test] + public async Task CanAddRequestPolicies_AllPositions() + { + var retryResponse = new MockResponse(408); // Request Timeout + + // retry twice -- this will add the header three times. + var mockTransport = new MockTransport(retryResponse, retryResponse, new MockResponse(200)); + var options = new TestOptions() + { + Transport = mockTransport, + }; + + var pipeline = HttpPipelineBuilder.Build(options); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("PerCallHeader1", "PerCall1"), HttpPipelinePosition.PerCall); + context.AddPolicy(new AddHeaderPolicy("PerCallHeader2", "PerCall2"), HttpPipelinePosition.PerCall); + context.AddPolicy(new AddHeaderPolicy("PerRetryHeader", "PerRetry"), HttpPipelinePosition.PerRetry); + context.AddPolicy(new AddHeaderPolicy("BeforeTransportHeader", "BeforeTransport"), HttpPipelinePosition.BeforeTransport); + + var message = pipeline.CreateMessage(context); + await pipeline.SendAsync(message, message.CancellationToken); + + Request request = mockTransport.Requests[0]; + + Assert.IsTrue(request.Headers.TryGetValues("PerCallHeader1", out var perCall1Values)); + Assert.AreEqual(1, perCall1Values.Count()); + Assert.AreEqual("PerCall1", perCall1Values.ElementAt(0)); + + Assert.IsTrue(request.Headers.TryGetValues("PerCallHeader2", out var perCall2Values)); + Assert.AreEqual(1, perCall2Values.Count()); + Assert.AreEqual("PerCall2", perCall2Values.ElementAt(0)); + + Assert.IsTrue(request.Headers.TryGetValues("PerRetryHeader", out var perRetryValues)); + Assert.AreEqual("PerRetry", perRetryValues.ElementAt(0)); + Assert.AreEqual("PerRetry", perRetryValues.ElementAt(1)); + Assert.AreEqual("PerRetry", perRetryValues.ElementAt(2)); + + Assert.IsTrue(request.Headers.TryGetValues("BeforeTransportHeader", out var beforeTransportValues)); + Assert.AreEqual("BeforeTransport", beforeTransportValues.ElementAt(0)); + Assert.AreEqual("BeforeTransport", beforeTransportValues.ElementAt(1)); + Assert.AreEqual("BeforeTransport", beforeTransportValues.ElementAt(2)); + } + + [Test] + public async Task CanAddPolicies_ThreeWays() + { + var mockTransport = new MockTransport(new MockResponse(200)); + var options = new TestOptions() + { + Transport = mockTransport, + }; + + var perCallPolicies = new HttpPipelinePolicy[] { new AddHeaderPolicy("PerCall", "Builder") }; + var perRetryPolicies = new HttpPipelinePolicy[] { new AddHeaderPolicy("PerRetry", "Builder") }; + + options.AddPolicy(new AddHeaderPolicy("BeforeTransport", "ClientOptions"), HttpPipelinePosition.BeforeTransport); + options.AddPolicy(new AddHeaderPolicy("PerRetry", "ClientOptions"), HttpPipelinePosition.PerRetry); + options.AddPolicy(new AddHeaderPolicy("PerCall", "ClientOptions"), HttpPipelinePosition.PerCall); + + var pipeline = HttpPipelineBuilder.Build(options, perCallPolicies, perRetryPolicies, null); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("PerRetry", "RequestContext"), HttpPipelinePosition.PerRetry); + context.AddPolicy(new AddHeaderPolicy("PerCall", "RequestContext"), HttpPipelinePosition.PerCall); + context.AddPolicy(new AddHeaderPolicy("BeforeTransport", "RequestContext"), HttpPipelinePosition.BeforeTransport); + + var message = pipeline.CreateMessage(context); + await pipeline.SendAsync(message, message.CancellationToken); + + Request request = mockTransport.Requests[0]; + + Assert.IsTrue(request.Headers.TryGetValues("PerCall", out var perCallValues)); + Assert.AreEqual(3, perCallValues.Count()); + Assert.AreEqual("Builder", perCallValues.ElementAt(0)); + Assert.AreEqual("ClientOptions", perCallValues.ElementAt(1)); + Assert.AreEqual("RequestContext", perCallValues.ElementAt(2)); + + Assert.IsTrue(request.Headers.TryGetValues("PerRetry", out var perRetryValues)); + Assert.AreEqual(3, perRetryValues.Count()); + Assert.AreEqual("Builder", perRetryValues.ElementAt(0)); + Assert.AreEqual("ClientOptions", perRetryValues.ElementAt(1)); + Assert.AreEqual("RequestContext", perRetryValues.ElementAt(2)); + + Assert.IsTrue(request.Headers.TryGetValues("BeforeTransport", out var beforeTransportValues)); + Assert.AreEqual(2, beforeTransportValues.Count()); + Assert.AreEqual("ClientOptions", beforeTransportValues.ElementAt(0)); + Assert.AreEqual("RequestContext", beforeTransportValues.ElementAt(1)); + } + + [Test] + public async Task ThrowsIfUsePipelineConstructor() + { + HttpPipeline pipeline = new HttpPipeline(new MockTransport()); + + var context = new RequestContext(); + context.AddPolicy(new AddHeaderPolicy("PerCallHeader", "Value"), HttpPipelinePosition.PerCall); + + var message = pipeline.CreateMessage(context); + + bool throws = false; + try + { + await pipeline.SendAsync(message, context.CancellationToken); + } + catch (InvalidOperationException) + { + throws = true; + } + + Assert.IsTrue(throws); + } + + #region Helpers + public class AddHeaderPolicy : HttpPipelineSynchronousPolicy + { + private string _headerName; + private string _headerVaue; + + public AddHeaderPolicy(string headerName, string headerValue) : base() + { + _headerName = headerName; + _headerVaue = headerValue; + } + + public override void OnSendingRequest(HttpMessage message) + { + message.Request.Headers.Add(_headerName, _headerVaue); + } + } + #endregion + } }