diff --git a/DOCUMENTATION.md b/DOCUMENTATION.md index 0514f057..94319da1 100644 --- a/DOCUMENTATION.md +++ b/DOCUMENTATION.md @@ -36,6 +36,7 @@ An in depth functional reference to all of Giraffe's default features. - [Content Negotiation](#content-negotiation) - [Streaming](#streaming) - [Redirection](#redirection) + - [Safe Redirection](#safe-redirection) - [Response Caching](#response-caching) - [Response Compression](#response-compression) - [Giraffe View Engine](#giraffe-view-engine) @@ -47,6 +48,7 @@ An in depth functional reference to all of Giraffe's default features. - [Short GUIDs and Short IDs](#short-guids-and-short-ids) - [Common Helper Functions](#common-helper-functions) - [Computation Expressions](#computation-expressions) + - [CSRF Protection Helpers](#csrf-protection-helpers) - [Additional Features](#additional-features) - [Endpoint Routing](#endpoint-routing) - [TokenRouter](#tokenrouter) @@ -2892,6 +2894,14 @@ let webApp = Please note that if the `permanent` flag is set to `true` then the Giraffe web application will send a `301` HTTP status code to browsers which will tell them that the redirection is permanent. This often leads to browsers cache the information and not hit the deprecated URL a second time any more. If this is not desired then please set `permanent` to `false` in order to guarantee that browsers will continue hitting the old URL before redirecting to the (temporary) new one. +#### Safe Redirection + +The `redirectTo` http handler, although giving you more freedom when specifying the redirection logic, does not validate for a common security problem named [open redirect](https://learn.snyk.io/lesson/open-redirect). + +In order to deal with this threat you can either implement your own logic (example from Microsoft docs [Prevent open redirect attacks in ASP.NET Core](https://learn.microsoft.com/en-us/aspnet/core/security/preventing-open-redirects)), or you can leverage the `safeRedirectTo (permanent: bool) (location: string)` http handler, which provides a handler with the necessary validation and a default error handler. + +Furthermore, if you want to use Giraffe's own open redirect validation, although with a different error handler, you can use the `safeRedirectToExt (permanent: bool) (location: string) (invalidRedirectHandler: HttpHandler option)` http handler, which as the signature suggests, accepts a custom `invalidRedirectHandler` that will be executed if the validation fails. + ### Response Caching ASP.NET Core comes with a standard [Response Caching Middleware](https://docs.microsoft.com/en-us/aspnet/core/performance/caching/middleware?view=aspnetcore-2.1) which works out of the box with Giraffe. @@ -3221,6 +3231,8 @@ By default Giraffe uses the `System.Xml.Serialization.XmlSerializer` for (de-)se Customizing Giraffe's XML serialization can either happen via providing a custom object of `XmlWriterSettings` when instantiating the default `SystemXml.Serializer` or swap in an entire different XML library by creating a new class which implements the `Xml.ISerializer` interface. +Notice that Giraffe does secure XML parsing, i.e., when using the `Deserialize<'T>(xml: string)` method, both DTD (Document Type Definition) processing and external entities are disabled to prevent [XXE attacks](https://learn.snyk.io/lesson/xxe). + #### Customizing XmlWriterSettings You can change the default `XmlWriterSettings` of the `SystemXml.Serializer` by registering a new instance of `SystemXml.Serializer` during application startup: @@ -3489,6 +3501,24 @@ let someHttpHandler : HttpHandler = | Error msg -> RequestErrors.BAD_REQUEST msg next ctx ``` +### CSRF Protection Helpers + +CSRF stands for Cross-Site Request Forgery, and according to the OWASP website can be defined as: + +> Cross-Site Request Forgery (CSRF) is an attack that forces an end user to execute unwanted actions on a web application in which they’re currently authenticated. With a little help of social engineering (such as sending a link via email or chat), an attacker may trick the users of a web application into executing actions of the attacker’s choosing. If the victim is a normal user, a successful CSRF attack can force the user to perform state changing requests like transferring funds, changing their email address, and so forth. If the victim is an administrative account, CSRF can compromise the entire web application. +> +> -- Reference [link](https://owasp.org/www-community/attacks/csrf). + +The ASP.NET documentation gives us a tutorial on how to deal with it ([link](https://learn.microsoft.com/en-us/aspnet/core/security/anti-request-forgery)), but you can also leverage the Giraffe's `HttpHandler` helpers from the `Csrf` module: + +- `validateCsrfTokenExt (invalidTokenHandler: HttpHandler option)`: Validates the CSRF token from the request. Checks for token in header (`X-CSRF-TOKEN`) or form field (`__RequestVerificationToken`). +- `requireAntiforgeryTokenExt`: Alias for `validateCsrfTokenExt` - validates anti-forgery tokens from requests with custom error handler. +- `validateCsrfToken`: Validates the CSRF token from the request with default error handling. Checks for token in header (`X-CSRF-TOKEN`) or form field (`__RequestVerificationToken`). Uses default error handling (403 Forbidden) for invalid tokens. +- `requireAntiforgeryToken`: Alias for `validateCsrfToken` - validates anti-forgery tokens from requests. +- `generateCsrfToken`: Generates a CSRF token and adds it to the `HttpContext` items for use in views. The token can be accessed via `ctx.Items["CsrfToken"]` and `ctx.Items["CsrfTokenHeaderName"]`. +- `csrfTokenJson`: Returns the CSRF token as JSON for AJAX requests. Response format: `{ "token": "...", "headerName": "X-CSRF-TOKEN" }`. +- `csrfTokenHtml`: Returns the CSRF token as an HTML hidden input field. Can be included directly in forms. + ## Additional Features There's more features available for Giraffe web applications through additional NuGet packages: diff --git a/src/Giraffe/Core.fs b/src/Giraffe/Core.fs index 89e9a52d..765d11c0 100644 --- a/src/Giraffe/Core.fs +++ b/src/Giraffe/Core.fs @@ -2,6 +2,7 @@ namespace Giraffe [] module Core = + open System open System.Text open System.Threading.Tasks open System.Globalization @@ -242,16 +243,83 @@ module Core = | true -> next ctx | false -> skipPipeline + /// + /// Validates if a redirect URL is safe (prevents open redirect vulnerabilities). + /// Allows only relative URLs or URLs with the same host. + /// + /// The HttpContext to get the request host from. + /// The URL to validate. + /// True if the URL is safe to redirect to, false otherwise. + let isValidRedirectUrl (ctx: HttpContext) (url: string) = + if String.IsNullOrWhiteSpace url then + false + elif url.StartsWith '/' then + true // Relative URL + elif url.StartsWith "~/" then + true // App-relative URL + else + match Uri.TryCreate(url, UriKind.Absolute) with + | true, uri -> + // Only allow redirects to the same host + let requestHost = ctx.Request.Host + uri.Host = requestHost.Host + | false, _ -> false + + /// + /// Redirects to a different location with a `302` or `301` (when permanent) HTTP status code. + /// Validates the redirect URL to prevent open redirect vulnerabilities. + /// + /// If true the redirect is permanent (301), otherwise temporary (302). + /// The URL to redirect the client to. + /// Optional custom handler for invalid redirects. If None, returns 400 Bad Request with logged warning. + /// + /// + /// A Giraffe function which can be composed into a bigger web application. + let safeRedirectToExt + (permanent: bool) + (location: string) + (invalidRedirectHandler: HttpHandler option) + : HttpHandler = + fun (_next: HttpFunc) (ctx: HttpContext) -> + if isValidRedirectUrl ctx location then + ctx.Response.Redirect(location, permanent) + Task.FromResult(Some ctx) + else + let defaultHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + let logger = ctx.GetLogger("Giraffe.Core") + logger.LogWarning("Blocked potential open redirect to: {Location}", location) + ctx.Response.StatusCode <- 400 + Task.FromResult(Some ctx) + + let handler = invalidRedirectHandler |> Option.defaultValue defaultHandler + handler earlyReturn ctx + + /// + /// Redirects to a different location with a `302` or `301` (when permanent) HTTP status code. + /// Validates the redirect URL to prevent **open redirect** vulnerabilities. + /// Uses default error handling (400 Bad Request) for invalid redirects. + /// + /// If true the redirect is permanent (301), otherwise temporary (302). + /// The URL to redirect the client to. + /// + /// + /// A Giraffe function which can be composed into a bigger web application. + let safeRedirectTo (permanent: bool) (location: string) : HttpHandler = + safeRedirectToExt permanent location None + /// /// Redirects to a different location with a `302` or `301` (when permanent) HTTP status code. + /// Does not validate redirection. Consider alternative: safeRedirectTo /// /// If true the redirect is permanent (301), otherwise temporary (302). /// The URL to redirect the client to. /// /// /// A Giraffe function which can be composed into a bigger web application. + [] let redirectTo (permanent: bool) (location: string) : HttpHandler = - fun (next: HttpFunc) (ctx: HttpContext) -> + fun (_next: HttpFunc) (ctx: HttpContext) -> ctx.Response.Redirect(location, permanent) Task.FromResult(Some ctx) diff --git a/src/Giraffe/Csrf.fs b/src/Giraffe/Csrf.fs new file mode 100644 index 00000000..63c5f392 --- /dev/null +++ b/src/Giraffe/Csrf.fs @@ -0,0 +1,155 @@ +namespace Giraffe + +/// +/// CSRF (Cross-Site Request Forgery) protection helpers for Giraffe. +/// Provides anti-forgery token generation and validation. +/// +[] +module Csrf = + open System + open System.Security.Cryptography + open System.Text + open System.Threading.Tasks + open Microsoft.AspNetCore.Http + open Microsoft.Extensions.Logging + open Microsoft.AspNetCore.Antiforgery + + // Defaults are selected to what developers would expect from ASP.NET Core application. + + /// + /// Default CSRF token header name + /// + [] + let DefaultCsrfTokenHeaderName = "X-CSRF-TOKEN" + + /// + /// Default CSRF token form field name + /// + [] + let DefaultCsrfTokenFormFieldName = "__RequestVerificationToken" + + /// + /// Validates the CSRF token from the request. + /// Checks for token in header (X-CSRF-TOKEN) or form field (__RequestVerificationToken). + /// + /// Optional custom handler for invalid tokens. If None, returns 403 Forbidden with logged warning. + /// The next HttpFunc + /// The HttpContext + /// HttpFuncResult + let validateCsrfTokenExt (invalidTokenHandler: HttpHandler option) : HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + task { + let antiforgery = ctx.GetService() + + try + let! isValid = antiforgery.IsRequestValidAsync ctx + + if isValid then + return! next ctx + else + let defaultHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + let logger = ctx.GetLogger("Giraffe.Csrf") + + logger.LogWarning( + "CSRF token validation failed for request to {Path}", + ctx.Request.Path + ) + + ctx.Response.StatusCode <- 403 + Task.FromResult(Some ctx) + + let handler = invalidTokenHandler |> Option.defaultValue defaultHandler + return! handler earlyReturn ctx + with ex -> + let defaultHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + let logger = ctx.GetLogger("Giraffe.Csrf") + logger.LogWarning(ex, "CSRF token validation error for request to {Path}", ctx.Request.Path) + ctx.Response.StatusCode <- 403 + Task.FromResult(Some ctx) + + let handler = invalidTokenHandler |> Option.defaultValue defaultHandler + return! handler earlyReturn ctx + } + + /// + /// Validates the CSRF token from the request with default error handling. + /// Checks for token in header (X-CSRF-TOKEN) or form field (__RequestVerificationToken). + /// Uses default error handling (403 Forbidden) for invalid tokens. + /// + /// The next HttpFunc + /// The HttpContext + /// HttpFuncResult + let validateCsrfToken: HttpHandler = validateCsrfTokenExt None + + /// + /// Alias for validateCsrfToken - validates anti-forgery tokens from requests. + /// + let requireAntiforgeryToken = validateCsrfToken + + /// + /// Alias for validateCsrfTokenExt - validates anti-forgery tokens from requests with custom error handler. + /// + let requireAntiforgeryTokenExt = validateCsrfTokenExt + + /// + /// Generates a CSRF token and adds it to the HttpContext items for use in views. + /// The token can be accessed via ctx.Items["CsrfToken"] and ctx.Items["CsrfTokenHeaderName"]. + /// + /// The next HttpFunc + /// The HttpContext + /// HttpFuncResult + let generateCsrfToken: HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + task { + let antiforgery = ctx.GetService() + let tokens = antiforgery.GetAndStoreTokens ctx + + // Store token for view rendering + ctx.Items.["CsrfToken"] <- tokens.RequestToken + ctx.Items.["CsrfTokenHeaderName"] <- tokens.HeaderName + + return! next ctx + } + + /// + /// Returns the CSRF token as JSON for AJAX requests. + /// Response format: { "token": "...", "headerName": "X-CSRF-TOKEN" } + /// + /// The next HttpFunc + /// The HttpContext + /// HttpFuncResult + let csrfTokenJson: HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + task { + let antiforgery = ctx.GetService() + let tokens = antiforgery.GetAndStoreTokens ctx + + let response = + {| + token = tokens.RequestToken + headerName = tokens.HeaderName + |} + + return! Core.json response next ctx + } + + /// + /// Returns the CSRF token as an HTML hidden input field. + /// Can be included directly in forms. + /// + /// The next HttpFunc + /// The HttpContext + /// HttpFuncResult + let csrfTokenHtml: HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + task { + let antiforgery = ctx.GetService() + let tokens = antiforgery.GetAndStoreTokens(ctx) + + let html = + sprintf "" tokens.HeaderName tokens.RequestToken + + return! Core.htmlString html next ctx + } diff --git a/src/Giraffe/Giraffe.fsproj b/src/Giraffe/Giraffe.fsproj index a417fc33..763ffbfc 100644 --- a/src/Giraffe/Giraffe.fsproj +++ b/src/Giraffe/Giraffe.fsproj @@ -82,6 +82,7 @@ + diff --git a/src/Giraffe/Xml.fs b/src/Giraffe/Xml.fs index 5e37a8f9..d19dc462 100644 --- a/src/Giraffe/Xml.fs +++ b/src/Giraffe/Xml.fs @@ -44,5 +44,14 @@ module SystemXml = member __.Deserialize<'T>(xml: string) = let serializer = XmlSerializer(typeof<'T>) - use reader = new StringReader(xml) - serializer.Deserialize reader :?> 'T + use stringReader = new StringReader(xml) + // Secure XML parsing: disable DTD processing and external entities to prevent XXE attacks + let xmlReaderSettings = + new XmlReaderSettings( + DtdProcessing = DtdProcessing.Prohibit, + XmlResolver = null, + MaxCharactersFromEntities = 1024L * 1024L + ) // 1MB limit + + use xmlReader = XmlReader.Create(stringReader, xmlReaderSettings) + serializer.Deserialize xmlReader :?> 'T diff --git a/tests/Giraffe.Tests/Giraffe.Tests.fsproj b/tests/Giraffe.Tests/Giraffe.Tests.fsproj index 40ce96d0..09dcd321 100644 --- a/tests/Giraffe.Tests/Giraffe.Tests.fsproj +++ b/tests/Giraffe.Tests/Giraffe.Tests.fsproj @@ -22,6 +22,7 @@ + diff --git a/tests/Giraffe.Tests/SecurityTests.fs b/tests/Giraffe.Tests/SecurityTests.fs new file mode 100644 index 00000000..78942fb8 --- /dev/null +++ b/tests/Giraffe.Tests/SecurityTests.fs @@ -0,0 +1,655 @@ +module Giraffe.Tests.SecurityTests + +open System +open System.IO +open System.Text +open System.Collections.Generic +open System.Threading.Tasks +open Microsoft.AspNetCore.Http +open Microsoft.AspNetCore.Antiforgery +open Microsoft.Extensions.Logging +open Xunit +open NSubstitute +open Giraffe + +// --------------------------------- +// URL Redirect Security Tests +// --------------------------------- + +[] +let ``safeRedirectTo allows relative URLs starting with /`` () = + let ctx = Substitute.For() + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false "/safe-path" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected redirect to succeed" + | Some ctx -> ctx.Response.Received().Redirect("/safe-path", false) |> ignore + } + +[] +let ``safeRedirectTo allows app-relative URLs starting with ~/`` () = + let ctx = Substitute.For() + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false "~/app-path" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected redirect to succeed" + | Some ctx -> ctx.Response.Received().Redirect("~/app-path", false) |> ignore + } + +[] +let ``safeRedirectTo allows absolute URLs to same host`` () = + let ctx = Substitute.For() + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false "https://example.com/path" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected redirect to succeed" + | Some ctx -> ctx.Response.Received().Redirect("https://example.com/path", false) |> ignore + } + +[] +let ``safeRedirectTo blocks open redirect to external domain`` () = + let ctx = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false "https://evil.com/phishing" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> + Assert.Equal(400, ctx.Response.StatusCode) + // Verify warning was logged (simplified check - just verify it was called) + logger + .ReceivedWithAnyArgs(1) + .Log( + LogLevel.Warning, + Arg.Any(), + Arg.Any(), + Arg.Any(), + Arg.Any>() + ) + |> ignore + } + +[] +let ``safeRedirectTo blocks javascript protocol XSS attempt`` () = + let ctx = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false "javascript:alert('xss')" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> Assert.Equal(400, ctx.Response.StatusCode) + } + +[] +let ``safeRedirectTo blocks empty or whitespace URLs`` () = + let ctx = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo false " " + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> Assert.Equal(400, ctx.Response.StatusCode) + } + +[] +let ``safeRedirectTo with permanent flag calls Redirect with true`` () = + let ctx = Substitute.For() + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString "example.com" + + let app = safeRedirectTo true "/permanent-redirect" + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected redirect to succeed" + | Some ctx -> ctx.Response.Received().Redirect("/permanent-redirect", true) |> ignore + } + +// --------------------------------- +// CSRF Token Security Tests +// --------------------------------- + +[] +let ``validateCsrfToken succeeds with valid token`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + + antiforgery.IsRequestValidAsync(ctx).Returns(System.Threading.Tasks.Task.FromResult(true)) + |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.validateCsrfToken testNext ctx + + match result with + | None -> assertFail "Expected CSRF validation to succeed" + | Some _ -> + Assert.True(nextCalled, "Expected next handler to be called") + Assert.NotEqual(403, ctx.Response.StatusCode) + } + +[] +let ``validateCsrfToken fails with invalid token`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Path <- PathString("/test") + + antiforgery.IsRequestValidAsync(ctx).Returns(System.Threading.Tasks.Task.FromResult false) + |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.validateCsrfToken testNext ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> + Assert.False(nextCalled, "Expected next handler not to be called") + Assert.Equal(403, ctx.Response.StatusCode) + } + +[] +let ``validateCsrfToken fails on exception`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Path <- PathString "/test" + + antiforgery + .IsRequestValidAsync(ctx) + .Returns(System.Threading.Tasks.Task.FromException(InvalidOperationException("Test error"))) + |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.validateCsrfToken testNext ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> + Assert.False(nextCalled, "Expected next handler not to be called") + Assert.Equal(403, ctx.Response.StatusCode) + } + +[] +let ``generateCsrfToken stores token in context items`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + + let tokens = + AntiforgeryTokenSet("test-request-token", "test-cookie-token", "form-field", "X-CSRF-TOKEN") + + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Items <- Dictionary() :> IDictionary + antiforgery.GetAndStoreTokens(ctx).Returns tokens |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.generateCsrfToken testNext ctx + + match result with + | None -> assertFail "Expected token generation to succeed" + | Some ctx -> + Assert.True(nextCalled, "Expected next handler to be called") + Assert.True(ctx.Items.ContainsKey("CsrfToken")) + Assert.Equal("test-request-token", ctx.Items.["CsrfToken"] :?> string) + Assert.True(ctx.Items.ContainsKey("CsrfTokenHeaderName")) + Assert.Equal("X-CSRF-TOKEN", ctx.Items.["CsrfTokenHeaderName"] :?> string) + } + +[] +let ``csrfTokenJson returns token as JSON`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + + let tokens = + AntiforgeryTokenSet("test-token-value", "cookie-value", "form-field", "X-CSRF-TOKEN") + + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + let jsonSerializer = Json.Serializer(Json.Serializer.DefaultOptions) + + serviceProvider.GetService(typeof).Returns jsonSerializer + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + antiforgery.GetAndStoreTokens(ctx).Returns tokens |> ignore + + task { + let! result = Csrf.csrfTokenJson next ctx + + match result with + | None -> assertFail "Expected JSON response" + | Some ctx -> + let body = getBody ctx + Assert.Contains("test-token-value", body) + Assert.Contains("X-CSRF-TOKEN", body) + Assert.Equal("application/json; charset=utf-8", ctx.Response |> getContentType) + } + +[] +let ``csrfTokenHtml returns token as hidden input`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + + let tokens = + AntiforgeryTokenSet("test-token-value", "cookie-value", "form-field", "X-CSRF-TOKEN") + + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + antiforgery.GetAndStoreTokens(ctx).Returns tokens |> ignore + + task { + let! result = Csrf.csrfTokenHtml next ctx + + match result with + | None -> assertFail "Expected HTML response" + | Some ctx -> + let body = getBody ctx + Assert.Contains("input", body) + Assert.Contains("type=\"hidden\"", body) + Assert.Contains("test-token-value", body) + Assert.Contains("X-CSRF-TOKEN", body) + } + +// --------------------------------- +// XXE Prevention Tests +// --------------------------------- + +[] +type TestXmlData = { Name: string; Value: int } + +[] +let ``XML deserialization works with normal XML`` () = + let serializer = + SystemXml.Serializer(SystemXml.Serializer.DefaultSettings) :> Xml.ISerializer + + let xml = """Test42""" + + let result = serializer.Deserialize xml + + Assert.Equal("Test", result.Name) + Assert.Equal(42, result.Value) + +[] +let ``XML deserialization blocks XXE attack with external entities`` () = + let serializer = + SystemXml.Serializer(SystemXml.Serializer.DefaultSettings) :> Xml.ISerializer + + // XML with DTD and external entity reference (XXE attack) + let maliciousXml = + """ + +]> + + &xxe; + 42 +""" + + // Should throw an exception because DTD processing is prohibited + Assert.Throws(fun () -> serializer.Deserialize maliciousXml |> ignore) + |> ignore + +[] +let ``XML deserialization blocks XXE attack with parameter entities`` () = + let serializer = + SystemXml.Serializer(SystemXml.Serializer.DefaultSettings) :> Xml.ISerializer + + // XML with parameter entity (another XXE attack vector) + let maliciousXml = + """ + + %xxe; +]> + + Attack + 0 +""" + + // Should throw an exception because DTD processing is prohibited + Assert.Throws(fun () -> serializer.Deserialize maliciousXml |> ignore) + |> ignore + +[] +let ``XML deserialization through HTTP context blocks XXE`` () = + task { + let ctx = Substitute.For() + mockXml ctx + + let maliciousXml = + """ + +]> + + &xxe; + 99 +""" + + let stream = new MemoryStream() + let writer = new StreamWriter(stream, Encoding.UTF8) + writer.Write maliciousXml + writer.Flush() + stream.Position <- 0L + + ctx.Request.Body <- stream + + // Attempting to bind the malicious XML should throw + let! result = + task { + try + let! data = ctx.BindXmlAsync() + return Ok data + with ex -> + return Error ex.Message + } + + match result with + | Ok _ -> assertFail "Expected XXE attack to be blocked" + | Error msg -> + // Verify that an error occurred (DTD processing blocked) + Assert.True(true) + } + +[] +let ``XML deserialization through HTTP context allows normal XML`` () = + task { + let ctx = Substitute.For() + mockXml ctx + + let normalXml = """Safe123""" + + let stream = new MemoryStream() + let writer = new StreamWriter(stream, Encoding.UTF8) + writer.Write normalXml + writer.Flush() + stream.Position <- 0L + + ctx.Request.Body <- stream + + let! result = ctx.BindXmlAsync() + + Assert.Equal("Safe", result.Name) + Assert.Equal(123, result.Value) + } + +// --------------------------------- +// Extended API Tests (Custom Error Handlers) +// --------------------------------- + +[] +let ``safeRedirectToExt allows custom error handler for invalid redirects`` () = + let ctx = Substitute.For() + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString("example.com") + + let customHandler: HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + ctx.Response.StatusCode <- 418 // I'm a teapot + System.Threading.Tasks.Task.FromResult(Some ctx) + + let app = safeRedirectToExt false "https://evil.com/phishing" (Some customHandler) + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> Assert.Equal(418, ctx.Response.StatusCode) + } + +[] +let ``safeRedirectToExt with None handler uses default behavior`` () = + let ctx = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Host <- HostString("example.com") + + let app = safeRedirectToExt false "https://evil.com/phishing" None + + task { + let! result = app next ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> Assert.Equal(400, ctx.Response.StatusCode) + } + +[] +let ``isValidRedirectUrl correctly validates safe URLs`` () = + let ctx = Substitute.For() + ctx.Request.Host <- HostString("example.com") + + Assert.True(isValidRedirectUrl ctx "/safe-path") + Assert.True(isValidRedirectUrl ctx "~/app-path") + Assert.True(isValidRedirectUrl ctx "https://example.com/path") + Assert.False(isValidRedirectUrl ctx "https://evil.com/phishing") + Assert.False(isValidRedirectUrl ctx "javascript:alert('xss')") + Assert.False(isValidRedirectUrl ctx " ") + Assert.False(isValidRedirectUrl ctx "") + +[] +let ``validateCsrfTokenExt allows custom error handler for invalid tokens`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Path <- PathString "/test" + + antiforgery.IsRequestValidAsync(ctx).Returns(System.Threading.Tasks.Task.FromResult false) + |> ignore + + let customHandler: HttpHandler = + fun (next: HttpFunc) (ctx: HttpContext) -> + ctx.Response.StatusCode <- 418 // I'm a teapot + System.Threading.Tasks.Task.FromResult(Some ctx) + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.validateCsrfTokenExt (Some customHandler) testNext ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> + Assert.False(nextCalled, "Expected next handler not to be called") + Assert.Equal(418, ctx.Response.StatusCode) + } + +[] +let ``validateCsrfTokenExt with None handler uses default behavior`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let loggerFactory = Substitute.For() + let logger = Substitute.For() + loggerFactory.CreateLogger(Arg.Any()).Returns logger |> ignore + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + + serviceProvider.GetService(typeof).Returns loggerFactory + |> ignore + + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + ctx.Request.Path <- PathString "/test" + + antiforgery.IsRequestValidAsync(ctx).Returns(System.Threading.Tasks.Task.FromResult false) + |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.validateCsrfTokenExt None testNext ctx + + match result with + | None -> assertFail "Expected handler to return context" + | Some ctx -> + Assert.False(nextCalled, "Expected next handler not to be called") + Assert.Equal(403, ctx.Response.StatusCode) + } + +[] +let ``requireAntiforgeryTokenExt is alias for validateCsrfTokenExt`` () = + let ctx = Substitute.For() + let antiforgery = Substitute.For() + let serviceProvider = Substitute.For() + serviceProvider.GetService(typeof).Returns antiforgery |> ignore + ctx.RequestServices.Returns serviceProvider |> ignore + ctx.Response.Body <- new MemoryStream() + + antiforgery.IsRequestValidAsync(ctx).Returns(System.Threading.Tasks.Task.FromResult true) + |> ignore + + let mutable nextCalled = false + + let testNext: HttpFunc = + fun _ -> + nextCalled <- true + System.Threading.Tasks.Task.FromResult(Some ctx) + + task { + let! result = Csrf.requireAntiforgeryTokenExt None testNext ctx + + match result with + | None -> assertFail "Expected CSRF validation to succeed" + | Some _ -> Assert.True(nextCalled, "Expected next handler to be called") + }