// Copyright (c) .NET Foundation. All rights reserved.
// Licensed under the Apache License, Version 2.0. See License.txt in the project root for license information.
using System.Diagnostics.CodeAnalysis;
using System.Globalization;
using System.Net;
using System.Net.Http;
using System.Threading;
using System.Threading.Tasks;
using System.Web.Cors;
using System.Web.Http.Cors.Properties;
namespace System.Web.Http.Cors
{
///
/// Custom for handling CORS requests.
///
public class CorsMessageHandler : DelegatingHandler
{
private HttpConfiguration _httpConfiguration;
private bool _rethrowExceptions;
///
/// Initializes a new instance of the class.
///
/// The .
/// httpConfiguration
public CorsMessageHandler(HttpConfiguration httpConfiguration) : this(httpConfiguration, false)
{
}
///
/// Initializes a new instance of the class.
///
/// The .
/// Indicates whether upstream exceptions should be rethrown
/// httpConfiguration
public CorsMessageHandler(HttpConfiguration httpConfiguration, bool rethrowExceptions)
{
if (httpConfiguration == null)
{
throw new ArgumentNullException("httpConfiguration");
}
_httpConfiguration = httpConfiguration;
_rethrowExceptions = rethrowExceptions;
}
///
/// Sends an HTTP request to the inner handler to send to the server as an asynchronous operation.
///
/// The HTTP request message to send to the server.
/// The token to monitor for cancellation requests.
///
/// Returns . The task object representing the asynchronous operation.
///
protected async override Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken)
{
CorsRequestContext corsRequestContext = request.GetCorsRequestContext();
if (corsRequestContext != null)
{
try
{
if (corsRequestContext.IsPreflight)
{
return await HandleCorsPreflightRequestAsync(request, corsRequestContext, cancellationToken);
}
else
{
return await HandleCorsRequestAsync(request, corsRequestContext, cancellationToken);
}
}
catch (Exception exception)
{
if (_rethrowExceptions)
{
throw;
}
return HandleException(request, exception);
}
}
else
{
return await base.SendAsync(request, cancellationToken);
}
}
///
/// Handles the actual CORS request.
///
/// The .
/// The .
/// The .
/// The .
///
/// request
/// or
/// corsRequestContext
///
public virtual async Task HandleCorsRequestAsync(HttpRequestMessage request, CorsRequestContext corsRequestContext, CancellationToken cancellationToken)
{
if (request == null)
{
throw new ArgumentNullException("request");
}
if (corsRequestContext == null)
{
throw new ArgumentNullException("corsRequestContext");
}
HttpResponseMessage response = await base.SendAsync(request, cancellationToken);
CorsPolicy corsPolicy = await GetCorsPolicyAsync(request, cancellationToken);
if (corsPolicy != null)
{
CorsResult result;
if (TryEvaluateCorsPolicy(corsRequestContext, corsPolicy, out result))
{
if (response != null)
{
response.WriteCorsHeaders(result);
}
}
}
return response;
}
///
/// Handles the preflight request specified by CORS.
///
/// The request.
/// The cors request context.
/// The token to monitor for cancellation requests.
/// The
///
/// request
/// or
/// corsRequestContext
///
public virtual async Task HandleCorsPreflightRequestAsync(HttpRequestMessage request, CorsRequestContext corsRequestContext, CancellationToken cancellationToken)
{
if (request == null)
{
throw new ArgumentNullException("request");
}
if (corsRequestContext == null)
{
throw new ArgumentNullException("corsRequestContext");
}
try
{
// Make sure Access-Control-Request-Method is valid.
new HttpMethod(corsRequestContext.AccessControlRequestMethod);
}
catch (ArgumentException)
{
return request.CreateErrorResponse(HttpStatusCode.BadRequest,
SRResources.AccessControlRequestMethodCannotBeNullOrEmpty);
}
catch (FormatException)
{
return request.CreateErrorResponse(HttpStatusCode.BadRequest,
String.Format(CultureInfo.CurrentCulture,
SRResources.InvalidAccessControlRequestMethod,
corsRequestContext.AccessControlRequestMethod));
}
CorsPolicy corsPolicy = await GetCorsPolicyAsync(request, cancellationToken);
if (corsPolicy != null)
{
HttpResponseMessage response = null;
CorsResult result;
if (TryEvaluateCorsPolicy(corsRequestContext, corsPolicy, out result))
{
response = request.CreateResponse(HttpStatusCode.OK);
response.WriteCorsHeaders(result);
}
else
{
response = result != null ?
request.CreateErrorResponse(HttpStatusCode.BadRequest, String.Join(" | ", result.ErrorMessages)) :
request.CreateResponse(HttpStatusCode.BadRequest);
}
return response;
}
else
{
return await base.SendAsync(request, cancellationToken);
}
}
[SuppressMessage("Microsoft.Reliability", "CA2000:Dispose objects before losing scope", Justification = "Caller owns HttpRequestMessage instance.")]
private static HttpResponseMessage HandleException(HttpRequestMessage request, Exception exception)
{
HttpResponseException httpResponseException = exception as HttpResponseException;
if (httpResponseException != null)
{
return httpResponseException.Response;
}
return request.CreateErrorResponse(HttpStatusCode.InternalServerError, exception);
}
private async Task GetCorsPolicyAsync(HttpRequestMessage request, CancellationToken cancellationToken)
{
CorsPolicy corsPolicy = null;
ICorsPolicyProviderFactory corsPolicyProviderFactory = _httpConfiguration.GetCorsPolicyProviderFactory();
ICorsPolicyProvider corsPolicyProvider = corsPolicyProviderFactory.GetCorsPolicyProvider(request);
if (corsPolicyProvider != null)
{
corsPolicy = await corsPolicyProvider.GetCorsPolicyAsync(request, cancellationToken);
}
return corsPolicy;
}
private bool TryEvaluateCorsPolicy(CorsRequestContext requestContext, CorsPolicy corsPolicy, out CorsResult corsResult)
{
ICorsEngine engine = _httpConfiguration.GetCorsEngine();
corsResult = engine.EvaluatePolicy(requestContext, corsPolicy);
return corsResult != null && corsResult.IsValid;
}
}
}