diff --git a/src/modules/Elsa.Http/Activities/WriteFileHttpResponse.cs b/src/modules/Elsa.Http/Activities/WriteFileHttpResponse.cs index 1497cb527..ce13a9e50 100644 --- a/src/modules/Elsa.Http/Activities/WriteFileHttpResponse.cs +++ b/src/modules/Elsa.Http/Activities/WriteFileHttpResponse.cs @@ -1,6 +1,7 @@ using System.Security.Cryptography; using Elsa.Extensions; using Elsa.Http.Contracts; +using Elsa.Http.Exceptions; using Elsa.Http.Models; using Elsa.Http.Options; using Elsa.Http.Services; @@ -256,8 +257,8 @@ public class WriteFileHttpResponse : Activity var manager = context.GetRequiredService(); var headers = httpContext.Request.Headers; - var eTag = headers.TryGetValue(HeaderNames.IfMatch, out var header) ? new EntityTagHeaderValue(header.ToString()) : default; - var range = headers.TryGetValue(HeaderNames.Range, out header) ? RangeHeaderValue.Parse(header.ToString()) : default; + var eTag = GetIfMatchHeaderValue(headers); + var range = GetRangeHeaderHeaderValue(headers); var options = new DownloadableOptions { ETag = eTag, Range = range }; return manager.GetDownloadablesAsync(content, options, context.CancellationToken); } @@ -267,6 +268,32 @@ public class WriteFileHttpResponse : Activity var provider = context.GetRequiredService(); return provider.TryGetContentType(filename, out var contentType) ? contentType : System.Net.Mime.MediaTypeNames.Application.Octet; } + + private static RangeHeaderValue? GetRangeHeaderHeaderValue(IHeaderDictionary headers) + { + try + { + return headers.TryGetValue(HeaderNames.Range, out var header) ? RangeHeaderValue.Parse(header.ToString()) : default; + + } + catch (Exception e) + { + throw new HttpBadRequestException("Failed to parse Range header value", e); + } + } + + private static EntityTagHeaderValue? GetIfMatchHeaderValue(IHeaderDictionary headers) + { + try + { + return headers.TryGetValue(HeaderNames.IfMatch, out var header) ? new EntityTagHeaderValue(header.ToString()) : default; + + } + catch (Exception e) + { + throw new HttpBadRequestException("Failed to parse If-Match header value", e); + } + } private async ValueTask OnResumeAsync(ActivityExecutionContext context) { diff --git a/src/modules/Elsa.Http/Exceptions/HttpBadRequestException.cs b/src/modules/Elsa.Http/Exceptions/HttpBadRequestException.cs new file mode 100644 index 000000000..1a2e152f2 --- /dev/null +++ b/src/modules/Elsa.Http/Exceptions/HttpBadRequestException.cs @@ -0,0 +1,12 @@ +namespace Elsa.Http.Exceptions; + +/// +/// Exception thrown when a bad request is received. +/// +public class HttpBadRequestException : Exception +{ + /// + public HttpBadRequestException(string message, Exception exception) : base(message, exception) + { + } +} \ No newline at end of file diff --git a/src/modules/Elsa.Http/Handlers/DefaultHttpEndpointFaultHandler.cs b/src/modules/Elsa.Http/Handlers/DefaultHttpEndpointFaultHandler.cs index 2725d55b1..54d43f781 100644 --- a/src/modules/Elsa.Http/Handlers/DefaultHttpEndpointFaultHandler.cs +++ b/src/modules/Elsa.Http/Handlers/DefaultHttpEndpointFaultHandler.cs @@ -2,6 +2,7 @@ using Elsa.Http.Models; using Microsoft.AspNetCore.Http; using System.Net.Mime; using Elsa.Http.Contracts; +using Elsa.Http.Exceptions; using Elsa.Workflows.Core.Contracts; namespace Elsa.Http.Handlers; @@ -16,16 +17,30 @@ public sealed class DefaultHttpEndpointFaultHandler : IHttpEndpointFaultHandler { var httpContext = context.HttpContext; var isTimeoutIncident = GetIsTimeoutFault(context); - var statusCode = isTimeoutIncident ? StatusCodes.Status408RequestTimeout : StatusCodes.Status500InternalServerError; - + var isBadRequest = GetIsBadRequestFault(context); + var statusCode = isTimeoutIncident + ? StatusCodes.Status408RequestTimeout + : isBadRequest + ? StatusCodes.Status400BadRequest + : StatusCodes.Status500InternalServerError; + httpContext.Response.StatusCode = statusCode; return ValueTask.CompletedTask; } private bool GetIsTimeoutFault(HttpEndpointFaultContext context) + { + return ContainsException(context, typeof(OperationCanceledException), typeof(TaskCanceledException), typeof(TimeoutException)); + } + + private bool GetIsBadRequestFault(HttpEndpointFaultContext context) + { + return ContainsException(context, typeof(HttpBadRequestException)); + } + + private bool ContainsException(HttpEndpointFaultContext context, params Type[] exceptionTypes) { var workflowState = context.WorkflowState; - var exceptionTypes = new[] { typeof(OperationCanceledException), typeof(TaskCanceledException), typeof(TimeoutException) }; var timeoutIncident = workflowState.Incidents.FirstOrDefault(x => exceptionTypes.Contains(x.Exception?.Type)); return timeoutIncident != null;