From 4faa3383dba0ac01b1e7cf66f84b2d6a475d9c56 Mon Sep 17 00:00:00 2001 From: Andrii Chebukin Date: Fri, 18 Sep 2026 19:32:57 +0200 Subject: [PATCH 1/6] Harden WebSocket streaming lifecycle Add the WebSocket-specific transport, subscription lifecycle, and payload/error serialization changes on top of the generic streaming and data-contract fixes: the SubscriptionExecutionResult.Data contract change to obj voption Skippable and related lifecycle handling in Shared/WebSockets.fs; GraphQLSubscriptionsManagement.fs subscription bookkeeping; the ObservableErrorHandling sanitization/deduplication module and remaining lifecycle/serialization changes in GraphQLWebsocketMiddleware.fs; the RELEASE_NOTES.md entry documenting the SubscriptionExecutionResult.Data change; and the new/updated WebSocket wire-format and error-sanitization tests in SerializationTests.fs. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- RELEASE_NOTES.md | 2 +- .../GraphQLSubscriptionsManagement.fs | 37 +- .../GraphQLWebsocketMiddleware.fs | 340 +++++++++++------- src/FSharp.Data.GraphQL.Shared/WebSockets.fs | 24 +- .../AspNetCore/SerializationTests.fs | 100 ++++-- 5 files changed, 329 insertions(+), 174 deletions(-) diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index 8508dc49f..64e3b9627 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -288,7 +288,7 @@ * **Breaking Change** Migrated to .NET 10 * **Breaking Change** Made Relay `Edge` a read-only struct -* **Breaking Change** `SubscriptionExecutionResult.Data` is now `obj Skippable`, and the record has new `Path` and `HasNext` fields for incremental delivery +* **Breaking Change** `SubscriptionExecutionResult.Data` is now `obj voption Skippable`, and the record has new `Path` and `HasNext` fields for incremental delivery * **Breaking Change** `BufferedStreamOptions.Interval` and `BufferedStreamOptions.PreferredBatchSize` are now `int voption` * **Breaking Change** `ServerMessage.Error` and `ServerRawPayload.ErrorMessages` now carry `GQLProblemDetails list` instead of `NameValueLookup list`, so an `error` message's `payload` is a standard GraphQL error array as the `graphql-transport-ws` protocol requires * **Breaking Change** A query or mutation whose non-null root field fails during execution now produces a `Direct` (execution) result with `null` data instead of a `RequestError`, which is now only ever produced for a request rejected before execution (validation, planning, variable or inline argument coercion, a middleware, or the executor itself failing); HTTP and `graphql-transport-ws` responses for such a failure now carry `data: null` as the spec requires, instead of omitting `data` entirely. This also changes the public `GQLResponse.Data`, `GQLResponseContent.Direct.Data`, `DeferredErrors.Data`, and `SubscriptionErrors.Data` signatures to use `voption` diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs index 7cd4ba431..95c078b1d 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs @@ -6,9 +6,9 @@ let addSubscription (id : SubscriptionId, unsubscriber : SubscriptionUnsubscriber, onUnsubscribe : OnUnsubscribeAction) (subscriptions : SubscriptionsDict) = - subscriptions.Add (id, (unsubscriber, onUnsubscribe)) + lock subscriptions (fun () -> subscriptions.Add (id, (unsubscriber, onUnsubscribe))) -let isIdTaken (id : SubscriptionId) (subscriptions : SubscriptionsDict) = subscriptions.ContainsKey (id) +let isIdTaken (id : SubscriptionId) (subscriptions : SubscriptionsDict) = lock subscriptions (fun () -> subscriptions.ContainsKey (id)) let executeOnUnsubscribeAndDispose (id : SubscriptionId) (subscription : SubscriptionUnsubscriber * OnUnsubscribeAction) = match subscription with @@ -19,15 +19,28 @@ let executeOnUnsubscribeAndDispose (id : SubscriptionId) (subscription : Subscri unsubscriber.Dispose () let removeSubscription (id : SubscriptionId) (subscriptions : SubscriptionsDict) = - match subscriptions.TryGetValue id with - | true, sub -> - sub |> executeOnUnsubscribeAndDispose id - subscriptions.Remove (id) |> ignore - | false, _ -> () + let subscription = + lock subscriptions (fun () -> + match subscriptions.TryGetValue id with + | true, sub -> + subscriptions.Remove (id) |> ignore + ValueSome sub + | false, _ -> ValueNone) + + match subscription with + | ValueSome sub -> sub |> executeOnUnsubscribeAndDispose id + | ValueNone -> () let removeAllSubscriptions (subscriptions : SubscriptionsDict) = - subscriptions - |> Seq.iter (fun subscription -> - subscription.Value - |> executeOnUnsubscribeAndDispose subscription.Key) - subscriptions.Clear () + let subscriptionsToDispose = + lock subscriptions (fun () -> + let snapshot = + subscriptions + |> Seq.map (fun subscription -> struct (subscription.Key, subscription.Value)) + |> Seq.toArray + + subscriptions.Clear () + snapshot) + + subscriptionsToDispose + |> Array.iter (fun struct (id, subscription) -> subscription |> executeOnUnsubscribeAndDispose id) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs index d399a879b..53d37fea7 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs @@ -40,14 +40,10 @@ module internal IncrementalPayloadSplitting = // Written as `obj list`, not the (internal, and here inaccessible) `FieldPath` abbreviation it stands for: // a type abbreviation is erased, so this is the exact same type and unifies fine with FieldPath-typed values. - let trySkipPathPrefix (prefix : obj list) (path : obj list) = - let rec loop prefix path = - match prefix, path with - | [], remainingPath -> ValueSome remainingPath - | _ :: _, [] -> ValueNone - | prefixHead :: prefixTail, pathHead :: pathTail when prefixHead = pathHead -> loop prefixTail pathTail - | _ -> ValueNone - loop prefix path + let pathStartsWith (prefix : obj list) (path : obj list) = + let prefixLength = List.length prefix + List.length path >= prefixLength + && List.truncate prefixLength path = prefix /// Matches a path ending in a list of indices, such as the path of a batched deferred payload, returning the /// path of the batch's own field and the indices of its items. @@ -62,23 +58,49 @@ module internal IncrementalPayloadSplitting = /// item) into one (data, errors, path) triple per item, addressed at that item's own path. let splitBatch (fieldPath : obj list) (indices : obj list) (data : obj) (errors : GQLProblemDetails list) = let items = data :?> obj[] - let errorsByIndex = - errors - |> Seq.vchoose (fun error -> - error.Path - |> Skippable.toValueOption - |> ValueOption.bind (trySkipPathPrefix fieldPath) - |> ValueOption.bind (function - | itemIndex :: _ -> ValueSome struct (itemIndex, error) - | [] -> ValueNone)) - |> _.ToLookup((fun struct (itemIndex, _) -> itemIndex), (fun struct (_, error) -> error)) (indices, List.ofArray items) ||> List.map2 (fun index item -> let itemPath = [ yield! fieldPath; yield index ] - let itemErrors = errorsByIndex[index] |> List.ofSeq + let itemErrors = + errors + |> List.filter (fun error -> + error.Path + |> Skippable.toValueOption + |> ValueOption.map (pathStartsWith itemPath) + |> ValueOption.defaultValue false) box [| item |], itemErrors, itemPath) +module internal ObservableErrorHandling = + + [] + let UnexpectedObservableErrorMessage = "Unexpected error during subscription" + + let private deduplicationKey (problem : GQLProblemDetails) = + let extensions = + problem.Extensions + |> Skippable.toValueOption + |> ValueOption.map ( + Seq.sortBy _.Key + >> Seq.map (fun kvp -> kvp.Key, kvp.Value) + >> Seq.toList + ) + + problem.Message, problem.Path, problem.Locations, extensions + + let rec problemDetailsOfObservableError (ex : exn) = + match ex with + | :? AggregateException as aggregate -> + aggregate.Flatten().InnerExceptions + |> Seq.collect problemDetailsOfObservableError + |> Seq.distinctBy deduplicationKey + |> Seq.toList + | _ -> + match box ex with + | :? IGQLError as error -> [ GQLProblemDetails.OfError error ] + | _ -> [ GQLProblemDetails.Create UnexpectedObservableErrorMessage ] + open IncrementalPayloadSplitting +open ObservableErrorHandling type GraphQLWebSocketMiddleware<'Root> ( @@ -100,30 +122,41 @@ type GraphQLWebSocketMiddleware<'Root> | ConnectionAck -> { Id = ValueNone; Type = "connection_ack"; Payload = ValueNone } | ServerPing -> { Id = ValueNone; Type = "ping"; Payload = ValueNone } | ServerPong p -> { Id = ValueNone; Type = "pong"; Payload = p |> ValueOption.map CustomResponse } - | Next (id, payload) -> { Id = ValueSome id; Type = "next"; Payload = ValueSome <| ExecutionResult payload } + | Next (id, payload) -> { + Id = ValueSome id + Type = "next" + Payload = ValueSome <| ExecutionResult payload + } | Complete id -> { Id = ValueSome id; Type = "complete"; Payload = ValueNone } - | Error (id, errMessages) -> { Id = ValueSome id; Type = "error"; Payload = ValueSome <| ErrorMessages errMessages } + | Error (id, errMessages) -> { + Id = ValueSome id + Type = "error" + Payload = ValueSome <| ErrorMessages errMessages + } return JsonSerializer.Serialize (raw, jsonSerializerOptions) } static let invalidJsonInClientMessageError = - Result.Error <| InvalidMessage (4400, "Invalid json in client message") + Result.Error + <| InvalidMessage (4400, "Invalid json in client message") let deserializeClientMessage (serializerOptions : JsonSerializerOptions) (msg : IReadOnlyPooledList) = taskResult { try - return JsonSerializer.Deserialize (msg.Span, serializerOptions) + return JsonSerializer.Deserialize(msg.Span, serializerOptions) with | :? InvalidWebsocketMessageException as ex -> - logger.LogError(ex, "Invalid websocket message:\n{payload}", msg) - return! Result.Error <| InvalidMessage (4400, ex.Message.ToString ()) - | :? JsonException as ex when logger.IsEnabled(LogLevel.Trace) -> - logger.LogError(ex, "Cannot deserialize WebSocket message:\n{payload}", msg) + logger.LogError (ex, "Invalid websocket message:\n{payload}", msg) + return! + Result.Error + <| InvalidMessage (4400, ex.Message.ToString ()) + | :? JsonException as ex when logger.IsEnabled (LogLevel.Trace) -> + logger.LogError (ex, "Cannot deserialize WebSocket message:\n{payload}", msg) return! invalidJsonInClientMessageError | :? JsonException as ex -> - logger.LogError(ex, "Cannot deserialize WebSocket message") + logger.LogError (ex, "Cannot deserialize WebSocket message") return! invalidJsonInClientMessageError | ex -> - logger.LogError(ex, $"Unexpected exception '{ex.GetType().Name}' in GraphQLWebsocketMiddleware") + logger.LogError (ex, $"Unexpected exception '{ex.GetType().Name}' in GraphQLWebsocketMiddleware") return! invalidJsonInClientMessageError } @@ -156,10 +189,10 @@ type GraphQLWebSocketMiddleware<'Root> let message = completeMessage |> Seq.filter (fun x -> x > 0uy) - |> Array.ofSeq + |> Seq.toArray |> System.Text.Encoding.UTF8.GetString logger.LogInformation ("-> Request: {request}", message) - if completeMessage.All(fun b -> b = 0uy) then + if completeMessage.All (fun b -> b = 0uy) then return ValueNone else let! result = deserializeClientMessage serializerOptions completeMessage @@ -170,13 +203,19 @@ type GraphQLWebSocketMiddleware<'Root> let sendMessageViaSocket (jsonSerializerOptions) (socket : WebSocket) (message : ServerMessage) : Task = task { if not (socket.State = WebSocketState.Open) then - logger.LogTrace ($"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", socket.State) + logger.LogTrace ( + $"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", + socket.State + ) else // TODO: Allocate string only if a debugger is attached let! serializedMessage = message |> serializeServerMessage jsonSerializerOptions - let segment = ArraySegment(System.Text.Encoding.UTF8.GetBytes (serializedMessage)) + let segment = ArraySegment(System.Text.Encoding.UTF8.GetBytes serializedMessage) if not (socket.State = WebSocketState.Open) then - logger.LogTrace ($"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", socket.State) + logger.LogTrace ( + $"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", + socket.State + ) else do! socket.SendAsync (segment, WebSocketMessageType.Text, endOfMessage = true, cancellationToken = CancellationToken.None) @@ -186,20 +225,40 @@ type GraphQLWebSocketMiddleware<'Root> let addClientSubscription (id : SubscriptionId) (howToSendDataOnNext : SubscriptionId -> 'ResponseContent -> Task) - (subscriptions : SubscriptionsDict, - socket : WebSocket, - streamSource : IObservable<'ResponseContent>, - jsonSerializerOptions : JsonSerializerOptions) - = + ( + subscriptions : SubscriptionsDict, + socket : WebSocket, + streamSource : IObservable<'ResponseContent>, + jsonSerializerOptions : JsonSerializerOptions + ) = + let sendTerminalError (ex : exn) = + sendMessageViaSocket jsonSerializerOptions socket (Error (id, problemDetailsOfObservableError ex)) + let observer = new Reactive.AnonymousObserver<'ResponseContent> ( - onNext = (fun theOutput -> (howToSendDataOnNext id theOutput).Wait ()), - onError = (fun ex -> logger.LogError (ex, "Error on subscription with Id = '{id}'", id)), + onNext = + (fun theOutput -> + try + (howToSendDataOnNext id theOutput).Wait() + with _ -> + subscriptions + |> GraphQLSubscriptionsManagement.removeSubscription id + reraise ()), + onError = + (fun ex -> + logger.LogError (ex, "Error on subscription with Id = '{id}'", id) + try + (sendTerminalError ex).Wait() + finally + subscriptions + |> GraphQLSubscriptionsManagement.removeSubscription (id)), onCompleted = (fun () -> - (sendMessageViaSocket jsonSerializerOptions socket (Complete id)).Wait () - subscriptions - |> GraphQLSubscriptionsManagement.removeSubscription (id)) + try + (sendMessageViaSocket jsonSerializerOptions socket (Complete id)).Wait() + finally + subscriptions + |> GraphQLSubscriptionsManagement.removeSubscription id) ) // Registered before subscribing, so a stream that completes synchronously (from inside Subscribe) still @@ -215,7 +274,8 @@ type GraphQLWebSocketMiddleware<'Root> with _ -> // Nothing will ever complete this subscription now, so the id is freed here instead; a no-op if the // synchronous completion above already removed it. Rethrown for the caller to report the failure. - subscriptions |> GraphQLSubscriptionsManagement.removeSubscription id + subscriptions + |> GraphQLSubscriptionsManagement.removeSubscription id reraise () let tryToGracefullyCloseSocket (code, message) theSocket = @@ -237,18 +297,24 @@ type GraphQLWebSocketMiddleware<'Root> let sendMsg = sendMessageViaSocket serializerOptions socket let rcv () = socket |> rcvMsgViaSocket serializerOptions - let sendOutput id (output : SubscriptionExecutionResult) = - sendMsg (Next (id, output)) + let sendOutput id (output : SubscriptionExecutionResult) = sendMsg (Next (id, output)) let sendSubscriptionResponseOutput id subscriptionResult = match subscriptionResult with - | SubscriptionResult output -> SubscriptionExecutionResult.Create (output, []) |> sendOutput id + | SubscriptionResult output -> + SubscriptionExecutionResult.Create (output, []) + |> sendOutput id | SubscriptionErrors (output, errors) -> + // TODO: Use StringBuilder logger.LogWarning ("Subscription errors: {subscriptionErrors}", (String.Join ('\n', errors |> Seq.map (fun x -> $"- %s{x.Message}")))) // The executor may still have resolved partial data alongside the field errors; forward it as-is match output with - | ValueNone -> SubscriptionExecutionResult.CreateErrors errors |> sendOutput id - | ValueSome output -> SubscriptionExecutionResult.Create (output, errors) |> sendOutput id + | ValueNone -> + SubscriptionExecutionResult.CreateErrors errors + |> sendOutput id + | ValueSome output -> + SubscriptionExecutionResult.Create (output, errors) + |> sendOutput id // Incremental payloads are sent as soon as they are produced, with their path inside the initial result, // so a client can merge them. The completion marker becomes a final payload with hasNext set to false. @@ -258,23 +324,36 @@ type GraphQLWebSocketMiddleware<'Root> match deferredResult with | ValueSome (DeferredResult (data, BatchPath (fieldPath, indices))) -> for itemData, _, itemPath in splitBatch fieldPath indices data [] do - do! SubscriptionExecutionResult.CreateIncremental (itemData, [], itemPath) |> sendOutput id + do! + SubscriptionExecutionResult.CreateIncremental (itemData, [], itemPath) + |> sendOutput id | ValueSome (DeferredResult (data, path)) -> - do! SubscriptionExecutionResult.CreateIncremental (data, [], path) |> sendOutput id + do! + SubscriptionExecutionResult.CreateIncremental (data, [], path) + |> sendOutput id | ValueSome (DeferredErrors (ValueSome data, errors, BatchPath (fieldPath, indices))) -> logger.LogWarning ( "Deferred response errors: {deferredErrors}", + // TODO: Use StringBuilder (String.Join ('\n', errors |> Seq.map (fun x -> $"- %s{x.Message}"))) ) for itemData, itemErrors, itemPath in splitBatch fieldPath indices data errors do - do! SubscriptionExecutionResult.CreateIncremental (itemData, itemErrors, itemPath) |> sendOutput id + do! + SubscriptionExecutionResult.CreateIncremental (itemData, itemErrors, itemPath) + |> sendOutput id | ValueSome (DeferredErrors (data, errors, path)) -> logger.LogWarning ( "Deferred response errors: {deferredErrors}", + // TODO: Use StringBuilder (String.Join ('\n', errors |> Seq.map (fun x -> $"- %s{x.Message}"))) ) - do! SubscriptionExecutionResult.CreateIncremental (data |> ValueOption.toObj, errors, path) |> sendOutput id - | ValueNone -> do! SubscriptionExecutionResult.CreateCompleted () |> sendOutput id + do! + SubscriptionExecutionResult.CreateIncremental (data |> ValueOption.toObj, errors, path) + |> sendOutput id + | ValueNone -> + do! + SubscriptionExecutionResult.CreateCompleted () + |> sendOutput id } let applyPlanExecutionResult (id : SubscriptionId) (socket) (executionResult : GQLExecutionResult) : Task = task { @@ -283,7 +362,9 @@ type GraphQLWebSocketMiddleware<'Root> (subscriptions, socket, observableOutput, serializerOptions) |> addClientSubscription id sendSubscriptionResponseOutput | Deferred (data, errors, observableOutput) -> - do! SubscriptionExecutionResult.CreateInitial (data, errors) |> sendOutput id + do! + SubscriptionExecutionResult.CreateInitial (data, errors) + |> sendOutput id (subscriptions, socket, observableOutput |> Observable.withCompletionMarker, serializerOptions) |> addClientSubscription id sendDeferredResponseOutput | Direct (data, errors) -> @@ -292,11 +373,13 @@ type GraphQLWebSocketMiddleware<'Root> // message below if not errors.IsEmpty then logger.LogWarning ("Execution errors:\n{errors}", errors) - do! SubscriptionExecutionResult.Create (data |> ValueOption.toObj, errors) |> sendOutput id + do! + SubscriptionExecutionResult.Create (data |> ValueOption.toObj, errors) + |> sendOutput id // The graphql-transport-ws protocol requires Complete after the single Next of a query or mutation do! sendMsg (Complete id) | RequestError problemDetails -> - logger.LogWarning("Request errors:\n{errors}", problemDetails) + logger.LogWarning ("Request errors:\n{errors}", problemDetails) // The request was rejected before execution, so it is not a result: the protocol requires it to be // sent as the terminal Error message instead of a Next followed by Complete, or a client would // read it as a successful result with null data @@ -319,68 +402,72 @@ type GraphQLWebSocketMiddleware<'Root> // -------> task { try - while not cancellationToken.IsCancellationRequested - && socket |> isSocketOpen do - let! receivedMessage = rcv () - match receivedMessage with - | Result.Error failureMessages -> - nameof InvalidMessage - |> logMsgReceivedWithOptionalPayload ValueNone - match failureMessages with - | InvalidMessage (code, explanation) -> do! socket.CloseAsync (enum code, explanation, CancellationToken.None) - | Ok ValueNone -> logger.LogTrace ("WebSocket received empty message! State = '{socketState}'", socket.State) - | Ok (ValueSome msg) -> - match msg with - | ConnectionInit p -> - nameof ConnectionInit |> logMsgReceivedWithOptionalPayload p - do! - socket.CloseAsync ( - enum CustomWebSocketStatus.TooManyInitializationRequests, - "Too many initialization requests", - CancellationToken.None - ) - | ClientPing p -> - nameof ClientPing |> logMsgReceivedWithOptionalPayload p - match pingHandler with - | ValueSome func -> - let! customP = p |> func serviceProvider - do! ServerPong customP |> sendMsg - | ValueNone -> do! ServerPong p |> sendMsg - | ClientPong p -> nameof ClientPong |> logMsgReceivedWithOptionalPayload p - | Subscribe (id, query) -> - try - nameof Subscribe |> logMsgWithIdReceived id - if subscriptions |> GraphQLSubscriptionsManagement.isIdTaken id then - do! - let warningMsg : FormattableString = $"Subscriber for Id = '{id}' already exists" - logger.LogWarning (String.Format (warningMsg.Format, "id"), id) - socket.CloseAsync ( - enum CustomWebSocketStatus.SubscriberAlreadyExists, - warningMsg.ToString (), - CancellationToken.None - ) - else - let variables = query.Variables |> Skippable.toValueOption - let getInputContext() = httpContext.RequestServices.GetRequiredService() - let! planExecutionResult = - let root = options.RootFactory httpContext - options.SchemaExecutor.AsyncExecute (query.Query, getInputContext, root, ?variables = variables) - do! planExecutionResult |> applyPlanExecutionResult id socket - with ex -> - logger.LogError (ex, "Unexpected error during subscription with id '{id}'", id) - do! sendMsg (Error (id, [ GQLProblemDetails.Create "Unexpected error during subscription" ])) - | ClientComplete id -> - "ClientComplete" |> logMsgWithIdReceived id - subscriptions - |> GraphQLSubscriptionsManagement.removeSubscription (id) - logger.LogTrace "Leaving the 'graphql-ws' connection loop..." - do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior - with ex -> - logger.LogError (ex, "Cannot handle a message; dropping a websocket connection") - // At this point, only something really weird must have happened. - // In order to avoid faulty state scenarios and unimagined damages, - // just close the socket without further ado. - do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + try + while not cancellationToken.IsCancellationRequested + && socket |> isSocketOpen do + let! receivedMessage = rcv () + match receivedMessage with + | Result.Error failureMessages -> + nameof InvalidMessage + |> logMsgReceivedWithOptionalPayload ValueNone + match failureMessages with + | InvalidMessage (code, explanation) -> do! socket.CloseAsync (enum code, explanation, CancellationToken.None) + | Ok ValueNone -> logger.LogTrace ("WebSocket received empty message! State = '{socketState}'", socket.State) + | Ok (ValueSome msg) -> + match msg with + | ConnectionInit p -> + nameof ConnectionInit |> logMsgReceivedWithOptionalPayload p + do! + socket.CloseAsync ( + enum CustomWebSocketStatus.TooManyInitializationRequests, + "Too many initialization requests", + CancellationToken.None + ) + | ClientPing p -> + nameof ClientPing |> logMsgReceivedWithOptionalPayload p + match pingHandler with + | ValueSome func -> + let! customP = p |> func serviceProvider + do! ServerPong customP |> sendMsg + | ValueNone -> do! ServerPong p |> sendMsg + | ClientPong p -> nameof ClientPong |> logMsgReceivedWithOptionalPayload p + | Subscribe (id, query) -> + try + nameof Subscribe |> logMsgWithIdReceived id + if subscriptions |> GraphQLSubscriptionsManagement.isIdTaken id then + do! + let warningMsg : FormattableString = $"Subscriber for Id = '{id}' already exists" + logger.LogWarning (String.Format (warningMsg.Format, "id"), id) + socket.CloseAsync ( + enum CustomWebSocketStatus.SubscriberAlreadyExists, + warningMsg.ToString (), + CancellationToken.None + ) + else + let variables = query.Variables |> Skippable.toValueOption + let getInputContext () = httpContext.RequestServices.GetRequiredService() + let! planExecutionResult = + let root = options.RootFactory httpContext + options.SchemaExecutor.AsyncExecute (query.Query, getInputContext, root, ?variables = variables) + do! planExecutionResult |> applyPlanExecutionResult id socket + with ex -> + logger.LogError (ex, "Unexpected error during subscription with id '{id}'", id) + do! sendMsg (Error (id, [ GQLProblemDetails.Create UnexpectedObservableErrorMessage ])) + | ClientComplete id -> + "ClientComplete" |> logMsgWithIdReceived id + subscriptions + |> GraphQLSubscriptionsManagement.removeSubscription (id) + logger.LogTrace "Leaving the 'graphql-ws' connection loop..." + do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + with ex -> + logger.LogError (ex, "Cannot handle a message; dropping a websocket connection") + // At this point, only something really weird must have happened. + // In order to avoid faulty state scenarios and unimagined damages, + // just close the socket without further ado. + do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + finally + subscriptions + |> GraphQLSubscriptionsManagement.removeAllSubscriptions } // <-------- @@ -394,10 +481,10 @@ type GraphQLWebSocketMiddleware<'Root> timerTokenSource.Token.Register (fun _ -> (socket |> tryToGracefullyCloseSocket (enum CustomWebSocketStatus.ConnectionTimeout, "Connection initialization timeout")) - .Wait ()) + .Wait()) let! connectionInitSucceeded = - TaskResult.Run ( + TaskResult.Run( (fun _ -> task { logger.LogDebug ($"Waiting for {nameof ConnectionInit}...") let! receivedMessage = receiveMessageViaSocket CancellationToken.None serializerOptions socket @@ -443,10 +530,8 @@ type GraphQLWebSocketMiddleware<'Root> | Result.Error errMsg -> logger.LogWarning errMsg | Ok _ -> let longRunningCancellationToken = - (CancellationTokenSource - .CreateLinkedTokenSource(ctx.RequestAborted, applicationLifetime.ApplicationStopping) - .Token) - longRunningCancellationToken.Register (fun _ -> (socket |> tryToGracefullyCloseSocketWithDefaultBehavior).Wait ()) + (CancellationTokenSource.CreateLinkedTokenSource(ctx.RequestAborted, applicationLifetime.ApplicationStopping).Token) + longRunningCancellationToken.Register (fun _ -> (socket |> tryToGracefullyCloseSocketWithDefaultBehavior).Wait()) |> ignore try do! socket |> handleMessages longRunningCancellationToken ctx @@ -458,5 +543,6 @@ type GraphQLWebSocketMiddleware<'Root> title = "WebSocket connection expected.", detail = $"'{options.WebsocketOptions.EndpointUrl}' endpoint only accepts WebSocket connections.", statusCode = StatusCodes.Status400BadRequest - ) :> IResult + ) + :> IResult |> _.ExecuteAsync(ctx) diff --git a/src/FSharp.Data.GraphQL.Shared/WebSockets.fs b/src/FSharp.Data.GraphQL.Shared/WebSockets.fs index 923edb7b7..bd27bfdd7 100644 --- a/src/FSharp.Data.GraphQL.Shared/WebSockets.fs +++ b/src/FSharp.Data.GraphQL.Shared/WebSockets.fs @@ -27,7 +27,7 @@ type RawMessage = { Id : string voption; Type : string; Payload : JsonDocument v type SubscriptionExecutionResult = { /// Result data: an object for complete and initial payloads, or a deferred or streamed value for incremental payloads. /// It is omitted from the final payload of an incremental delivery. - Data : obj Skippable + Data : obj voption Skippable /// Errors raised while producing the payload. Errors : GQLProblemDetails list /// Path of a deferred or streamed value inside the initial result. @@ -37,19 +37,29 @@ type SubscriptionExecutionResult = { } with /// Creates a payload of a complete execution result. - static member Create (data : Output, errors : GQLProblemDetails list) = { - Data = Include (box data) + static member Create (data : Output | null, errors : GQLProblemDetails list) = { + Data = + Include ( + Option.ofObj data + |> ValueOption.ofOption + |> ValueOption.map box + ) Errors = errors Path = Skip HasNext = Skip } /// Creates a payload that carries only errors. - static member CreateErrors (errors : GQLProblemDetails list) = { Data = Include null; Errors = errors; Path = Skip; HasNext = Skip } + static member CreateErrors (errors : GQLProblemDetails list) = { Data = Include ValueNone; Errors = errors; Path = Skip; HasNext = Skip } /// Creates the initial payload of an incremental delivery, which is always followed by incremental payloads. - static member CreateInitial (data : Output, errors : GQLProblemDetails list) = { - Data = Include (box data) + static member CreateInitial (data : Output | null, errors : GQLProblemDetails list) = { + Data = + Include ( + Option.ofObj data + |> ValueOption.ofOption + |> ValueOption.map box + ) Errors = errors Path = Skip HasNext = Include true @@ -60,7 +70,7 @@ type SubscriptionExecutionResult = { /// More payloads may follow, so is . /// static member CreateIncremental (data : objnull, errors : GQLProblemDetails list, path : FieldPath) = { - Data = Include data + Data = Include (data |> ValueOption.ofObj) Errors = errors Path = Include path HasNext = Include true diff --git a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs index 2c1316c99..9438a545e 100644 --- a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs +++ b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs @@ -1,10 +1,13 @@ module FSharp.Data.GraphQL.Tests.AspNetCore.SerializationTests +open System +open System.Collections.Generic open Xunit open System.Text.Json open FSharp.Data.GraphQL.Ast open FSharp.Data.GraphQL.Shared open FSharp.Data.GraphQL.Shared.WebSockets +open FSharp.Data.GraphQL.Server.AspNetCore.ObservableErrorHandling open System.Text.Json.Serialization [] @@ -12,7 +15,7 @@ let ``Deserializes ConnectionInit correctly`` () = let input = "{\"type\":\"connection_init\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ConnectionInit ValueNone -> () // <-- expected @@ -23,7 +26,7 @@ let ``Deserializes ConnectionInit with payload correctly`` () = let input = "{\"type\":\"connection_init\", \"payload\":\"hello\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ConnectionInit _ -> () // <-- expected @@ -34,7 +37,7 @@ let ``Deserializes ClientPing correctly`` () = let input = "{\"type\":\"ping\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ClientPing ValueNone -> () // <-- expected @@ -45,7 +48,7 @@ let ``Deserializes ClientPing with payload correctly`` () = let input = "{\"type\":\"ping\", \"payload\":\"ping!\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ClientPing _ -> () // <-- expected @@ -56,7 +59,7 @@ let ``Deserializes ClientPong correctly`` () = let input = "{\"type\":\"pong\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ClientPong ValueNone -> () // <-- expected @@ -67,7 +70,7 @@ let ``Deserializes ClientPong with payload correctly`` () = let input = "{\"type\":\"pong\", \"payload\": \"pong!\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ClientPong _ -> () // <-- expected @@ -78,7 +81,7 @@ let ``Deserializes ClientComplete correctly`` () = let input = "{\"id\": \"65fca2b5-f149-4a70-a055-5123dea4628f\", \"type\":\"complete\"}" - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | ClientComplete id -> Assert.Equal ("65fca2b5-f149-4a70-a055-5123dea4628f", id) @@ -97,7 +100,7 @@ let ``Deserializes client subscription correctly`` () = } """ - let result = JsonSerializer.Deserialize (input, serializerOptions) + let result = JsonSerializer.Deserialize(input, serializerOptions) match result with | Subscribe (id, payload) -> @@ -110,7 +113,11 @@ let ``Deserializes client subscription correctly`` () = open FSharp.Data.GraphQL let private serializePayload (payload : SubscriptionExecutionResult) = - let message : RawServerMessage = { Id = ValueSome "1"; Type = "next"; Payload = ValueSome (ExecutionResult payload) } + let message : RawServerMessage = { + Id = ValueSome "1" + Type = "next" + Payload = ValueSome (ExecutionResult payload) + } JsonSerializer.Serialize (message, serializerOptions) let private hasProperty (name : string) (element : JsonElement) = @@ -119,65 +126,104 @@ let private hasProperty (name : string) (element : JsonElement) = [] let ``Serializes incremental payload with path and hasNext`` () = - let json = serializePayload (SubscriptionExecutionResult.CreateIncremental (box [| box 1 |], [], [ box "numbers"; box 0 ])) + let json = + serializePayload (SubscriptionExecutionResult.CreateIncremental (box [| box 1 |], [], [ box "numbers"; box 0 ])) use document = JsonDocument.Parse json let payload = document.RootElement.GetProperty "payload" let data = payload.GetProperty "data" Assert.Equal (JsonValueKind.Array, data.ValueKind) - Assert.Equal (1, data[0].GetInt32 ()) + Assert.Equal (1, data[0].GetInt32()) let path = payload.GetProperty "path" - Assert.Equal ("numbers", path[0].GetString ()) - Assert.Equal (0, path[1].GetInt32 ()) - Assert.True (payload.GetProperty("hasNext").GetBoolean (), $"Expected hasNext to be true in {json}") + Assert.Equal ("numbers", path[0].GetString()) + Assert.Equal (0, path[1].GetInt32()) + Assert.True (payload.GetProperty("hasNext").GetBoolean(), $"Expected hasNext to be true in {json}") [] let ``Serializes final incremental payload with hasNext only`` () = let json = serializePayload (SubscriptionExecutionResult.CreateCompleted ()) use document = JsonDocument.Parse json let payload = document.RootElement.GetProperty "payload" - Assert.False (payload.GetProperty("hasNext").GetBoolean (), $"Expected hasNext to be false in {json}") + Assert.False (payload.GetProperty("hasNext").GetBoolean(), $"Expected hasNext to be false in {json}") Assert.False (hasProperty "data" payload, $"Expected no data in {json}") Assert.False (hasProperty "path" payload, $"Expected no path in {json}") [] let ``Serializes complete payload without path and hasNext`` () = - let json = serializePayload (SubscriptionExecutionResult.Create (NameValueLookup.ofList [ "name", upcast "R2-D2" ], [])) + let json = + serializePayload (SubscriptionExecutionResult.Create (NameValueLookup.ofList [ "name", upcast "R2-D2" ], [])) use document = JsonDocument.Parse json let payload = document.RootElement.GetProperty "payload" - Assert.Equal ("R2-D2", payload.GetProperty("data").GetProperty("name").GetString ()) + Assert.Equal ("R2-D2", payload.GetProperty("data").GetProperty("name").GetString()) Assert.False (hasProperty "path" payload, $"Expected no path in {json}") Assert.False (hasProperty "hasNext" payload, $"Expected no hasNext in {json}") [] let ``Serializes errors payload with null data as before`` () = - let json = serializePayload (SubscriptionExecutionResult.CreateErrors [ GQLProblemDetails.CreateWithKind ("Boom", Execution, [ box "numbers" ]) ]) + let json = + serializePayload (SubscriptionExecutionResult.CreateErrors [ GQLProblemDetails.CreateWithKind ("Boom", Execution, [ box "numbers" ]) ]) use document = JsonDocument.Parse json let payload = document.RootElement.GetProperty "payload" Assert.Equal (JsonValueKind.Null, payload.GetProperty("data").ValueKind) - Assert.Equal ("Boom", (payload.GetProperty "errors").Item(0).GetProperty("message").GetString ()) + Assert.Equal ("Boom", (payload.GetProperty "errors").Item(0).GetProperty("message").GetString()) [] let ``Serializes an error message with its problem details as the payload`` () = // Regression test: RawServerMessageConverter used to write the ErrorMessages payload without a preceding // WritePropertyName ("payload"), which Utf8JsonWriter rejects, so every "error" message failed to serialize - let message : RawServerMessage = - { Id = ValueSome "1"; Type = "error"; Payload = ValueSome (ErrorMessages [ GQLProblemDetails.Create "Boom" ]) } + let message : RawServerMessage = { + Id = ValueSome "1" + Type = "error" + Payload = ValueSome (ErrorMessages [ GQLProblemDetails.Create "Boom" ]) + } let json = JsonSerializer.Serialize (message, serializerOptions) use document = JsonDocument.Parse json let root = document.RootElement - Assert.Equal ("error", root.GetProperty("type").GetString ()) - Assert.Equal ("1", root.GetProperty("id").GetString ()) + Assert.Equal ("error", root.GetProperty("type").GetString()) + Assert.Equal ("1", root.GetProperty("id").GetString()) let payload = root.GetProperty "payload" Assert.Equal (JsonValueKind.Array, payload.ValueKind) - Assert.Equal ("Boom", payload[0].GetProperty("message").GetString ()) + Assert.Equal ("Boom", payload[0].GetProperty("message").GetString()) [] let ``Serializes a pong message with its payload`` () = // Regression test: the same missing WritePropertyName ("payload") affected a pong carrying a custom response use responseDocument = JsonDocument.Parse "\"pong!\"" - let message : RawServerMessage = { Id = ValueNone; Type = "pong"; Payload = ValueSome (CustomResponse responseDocument) } + let message : RawServerMessage = { + Id = ValueNone + Type = "pong" + Payload = ValueSome (CustomResponse responseDocument) + } let json = JsonSerializer.Serialize (message, serializerOptions) use document = JsonDocument.Parse json let root = document.RootElement - Assert.Equal ("pong", root.GetProperty("type").GetString ()) - Assert.Equal ("pong!", root.GetProperty("payload").GetString ()) + Assert.Equal ("pong", root.GetProperty("type").GetString()) + Assert.Equal ("pong!", root.GetProperty("payload").GetString()) + +[] +let ``Observable error details sanitize non-GraphQL exception messages`` () = + let actual = problemDetailsOfObservableError (Exception "sensitive backend failure") + let error = Assert.Single actual + Assert.Equal (UnexpectedObservableErrorMessage, error.Message) + +[] +let ``Observable error details preserve GraphQL-facing messages inside aggregates`` () = + let actual = + AggregateException [| Exception "sensitive backend failure"; GQLMessageException "Visible to client" |] + |> problemDetailsOfObservableError + |> List.map _.Message + + Assert.Contains (UnexpectedObservableErrorMessage, actual) + Assert.Contains ("Visible to client", actual) + Assert.DoesNotContain ("sensitive backend failure", actual) + +[] +let ``Observable error details do not duplicate repeated aggregate errors`` () = + let actual = + AggregateException [| + GQLMessageException ("Visible to client", Dictionary(dict [ "a", box 1; "b", box 2 ])) :> exn + GQLMessageException ("Visible to client", Dictionary(dict [ "b", box 2; "a", box 1 ])) :> exn + |] + |> problemDetailsOfObservableError + + let error = Assert.Single actual + Assert.Equal ("Visible to client", error.Message) From 92db2a89fafa83ebf18086bf7d80537f9dcc4ba9 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sat, 19 Sep 2026 00:08:49 +0000 Subject: [PATCH 2/6] Fix WebSocket review feedback Co-authored-by: xperiandri <2365592+xperiandri@users.noreply.github.com> --- .../GraphQLWebsocketMiddleware.fs | 120 +++++++++++------- .../AspNetCore/SerializationTests.fs | 18 +++ 2 files changed, 94 insertions(+), 44 deletions(-) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs index 53d37fea7..eda7c0e5a 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs @@ -45,9 +45,17 @@ module internal IncrementalPayloadSplitting = List.length path >= prefixLength && List.truncate prefixLength path = prefix + let tryGetPathItemIndex (fieldPath : obj list) (path : obj list) = + let fieldPathLength = List.length fieldPath + + if pathStartsWith fieldPath path then + path |> List.tryItem fieldPathLength |> ValueOption.ofOption + else + ValueNone + /// Matches a path ending in a list of indices, such as the path of a batched deferred payload, returning the /// path of the batch's own field and the indices of its items. - [] + [] let (|BatchPath|_|) (path : obj list) = match List.rev path with | (:? (obj list) as indices) :: fieldPathRev -> ValueSome (List.rev fieldPathRev, indices) @@ -58,16 +66,20 @@ module internal IncrementalPayloadSplitting = /// item) into one (data, errors, path) triple per item, addressed at that item's own path. let splitBatch (fieldPath : obj list) (indices : obj list) (data : obj) (errors : GQLProblemDetails list) = let items = data :?> obj[] + let errorsByItemIndex = + errors + |> Seq.choose (fun error -> + error.Path + |> Skippable.toValueOption + |> ValueOption.bind (tryGetPathItemIndex fieldPath) + |> ValueOption.map (fun index -> struct (index, error)) + |> ValueOption.toOption) + |> _.ToLookup((fun struct (index, _) -> index), (fun struct (_, error) -> error)) + (indices, List.ofArray items) ||> List.map2 (fun index item -> let itemPath = [ yield! fieldPath; yield index ] - let itemErrors = - errors - |> List.filter (fun error -> - error.Path - |> Skippable.toValueOption - |> ValueOption.map (pathStartsWith itemPath) - |> ValueOption.defaultValue false) + let itemErrors = errorsByItemIndex[index] |> Seq.toList box [| item |], itemErrors, itemPath) module internal ObservableErrorHandling = @@ -90,15 +102,30 @@ module internal ObservableErrorHandling = let rec problemDetailsOfObservableError (ex : exn) = match ex with | :? AggregateException as aggregate -> - aggregate.Flatten().InnerExceptions - |> Seq.collect problemDetailsOfObservableError - |> Seq.distinctBy deduplicationKey - |> Seq.toList + let problemDetails = + aggregate.Flatten().InnerExceptions + |> Seq.collect problemDetailsOfObservableError + |> Seq.distinctBy deduplicationKey + |> Seq.toList + + match problemDetails with + | [] -> [ GQLProblemDetails.Create UnexpectedObservableErrorMessage ] + | _ -> problemDetails | _ -> match box ex with | :? IGQLError as error -> [ GQLProblemDetails.OfError error ] | _ -> [ GQLProblemDetails.Create UnexpectedObservableErrorMessage ] + let sanitizeRequestError (problemDetails : GQLProblemDetails) = + match + problemDetails.Exception + |> ValueOption.map box + |> ValueOption.toObj + with + | :? IGQLError -> problemDetails + | :? exn -> GQLProblemDetails.Create UnexpectedObservableErrorMessage + | _ -> problemDetails + open IncrementalPayloadSplitting open ObservableErrorHandling @@ -201,38 +228,39 @@ type GraphQLWebSocketMiddleware<'Root> ArrayPool.Shared.Return buffer } - let sendMessageViaSocket (jsonSerializerOptions) (socket : WebSocket) (message : ServerMessage) : Task = task { - if not (socket.State = WebSocketState.Open) then - logger.LogTrace ( - $"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", - socket.State - ) - else - // TODO: Allocate string only if a debugger is attached - let! serializedMessage = message |> serializeServerMessage jsonSerializerOptions - let segment = ArraySegment(System.Text.Encoding.UTF8.GetBytes serializedMessage) + let sendMessageViaSocket (sendGate : SemaphoreSlim) (jsonSerializerOptions) (socket : WebSocket) (message : ServerMessage) : Task = task { + do! sendGate.WaitAsync () + + try + logger.LogTrace ("<- Response: {response}", message) + if not (socket.State = WebSocketState.Open) then logger.LogTrace ( $"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", socket.State ) else - do! socket.SendAsync (segment, WebSocketMessageType.Text, endOfMessage = true, cancellationToken = CancellationToken.None) - - logger.LogTrace ("<- Response: {response}", message) + // TODO: Allocate string only if a debugger is attached + let! serializedMessage = message |> serializeServerMessage jsonSerializerOptions + let segment = ArraySegment(System.Text.Encoding.UTF8.GetBytes serializedMessage) + + if not (socket.State = WebSocketState.Open) then + logger.LogTrace ( + $"Ignoring message to be sent via socket, since its state is not '{nameof WebSocketState.Open}', but '{{state}}'", + socket.State + ) + else + do! socket.SendAsync (segment, WebSocketMessageType.Text, endOfMessage = true, cancellationToken = CancellationToken.None) + finally + sendGate.Release () |> ignore } let addClientSubscription (id : SubscriptionId) (howToSendDataOnNext : SubscriptionId -> 'ResponseContent -> Task) - ( - subscriptions : SubscriptionsDict, - socket : WebSocket, - streamSource : IObservable<'ResponseContent>, - jsonSerializerOptions : JsonSerializerOptions - ) = - let sendTerminalError (ex : exn) = - sendMessageViaSocket jsonSerializerOptions socket (Error (id, problemDetailsOfObservableError ex)) + (subscriptions : SubscriptionsDict, streamSource : IObservable<'ResponseContent>, sendMsg : ServerMessage -> Task) + = + let sendTerminalError (ex : exn) = sendMsg (Error (id, problemDetailsOfObservableError ex)) let observer = new Reactive.AnonymousObserver<'ResponseContent> ( @@ -255,7 +283,7 @@ type GraphQLWebSocketMiddleware<'Root> onCompleted = (fun () -> try - (sendMessageViaSocket jsonSerializerOptions socket (Complete id)).Wait() + (sendMsg (Complete id)).Wait() finally subscriptions |> GraphQLSubscriptionsManagement.removeSubscription id) @@ -287,14 +315,14 @@ type GraphQLWebSocketMiddleware<'Root> let tryToGracefullyCloseSocketWithDefaultBehavior = tryToGracefullyCloseSocket (WebSocketCloseStatus.NormalClosure, "Normal Closure") - let handleMessages (cancellationToken : CancellationToken) (httpContext : HttpContext) (socket : WebSocket) : Task = + let handleMessages (sendGate : SemaphoreSlim) (cancellationToken : CancellationToken) (httpContext : HttpContext) (socket : WebSocket) : Task = let subscriptions = Dictionary() // ----------> // Helpers --> // ----------> let rcvMsgViaSocket = receiveMessageViaSocket (CancellationToken.None) - let sendMsg = sendMessageViaSocket serializerOptions socket + let sendMsg = sendMessageViaSocket sendGate serializerOptions socket let rcv () = socket |> rcvMsgViaSocket serializerOptions let sendOutput id (output : SubscriptionExecutionResult) = sendMsg (Next (id, output)) @@ -359,13 +387,13 @@ type GraphQLWebSocketMiddleware<'Root> let applyPlanExecutionResult (id : SubscriptionId) (socket) (executionResult : GQLExecutionResult) : Task = task { match executionResult with | Stream observableOutput -> - (subscriptions, socket, observableOutput, serializerOptions) + (subscriptions, observableOutput, sendMsg) |> addClientSubscription id sendSubscriptionResponseOutput | Deferred (data, errors, observableOutput) -> do! SubscriptionExecutionResult.CreateInitial (data, errors) |> sendOutput id - (subscriptions, socket, observableOutput |> Observable.withCompletionMarker, serializerOptions) + (subscriptions, observableOutput |> Observable.withCompletionMarker, sendMsg) |> addClientSubscription id sendDeferredResponseOutput | Direct (data, errors) -> // An execution result, whose data is null when a non-null root field failed during execution; @@ -379,11 +407,12 @@ type GraphQLWebSocketMiddleware<'Root> // The graphql-transport-ws protocol requires Complete after the single Next of a query or mutation do! sendMsg (Complete id) | RequestError problemDetails -> - logger.LogWarning ("Request errors:\n{errors}", problemDetails) + let sanitizedProblemDetails = problemDetails |> List.map sanitizeRequestError + logger.LogWarning ("Request errors:\n{errors}", sanitizedProblemDetails) // The request was rejected before execution, so it is not a result: the protocol requires it to be // sent as the terminal Error message instead of a Next followed by Complete, or a client would // read it as a successful result with null data - do! sendMsg (Error (id, problemDetails)) + do! sendMsg (Error (id, sanitizedProblemDetails)) } let logMsgReceivedWithOptionalPayload optionalPayload (msgAsStr : string) = @@ -474,7 +503,7 @@ type GraphQLWebSocketMiddleware<'Root> // <-- Main // <-------- - let waitForConnectionInitAndRespondToClient (socket : WebSocket) : TaskResult = task { + let waitForConnectionInitAndRespondToClient (sendGate : SemaphoreSlim) (socket : WebSocket) : TaskResult = task { let timerTokenSource = new CancellationTokenSource () timerTokenSource.CancelAfter connectionInitTimeout let detonationRegistration = @@ -494,7 +523,7 @@ type GraphQLWebSocketMiddleware<'Root> detonationRegistration.Unregister () |> ignore do! ConnectionAck - |> sendMessageViaSocket serializerOptions socket + |> sendMessageViaSocket sendGate serializerOptions socket return true | Ok (ValueSome (Subscribe _)) -> do! @@ -525,7 +554,8 @@ type GraphQLWebSocketMiddleware<'Root> if ctx.WebSockets.IsWebSocketRequest then task { use! socket = ctx.WebSockets.AcceptWebSocketAsync ("graphql-transport-ws") - let! connectionInitResult = socket |> waitForConnectionInitAndRespondToClient + let sendGate = new SemaphoreSlim (1, 1) + let! connectionInitResult = socket |> waitForConnectionInitAndRespondToClient sendGate match connectionInitResult with | Result.Error errMsg -> logger.LogWarning errMsg | Ok _ -> @@ -534,7 +564,9 @@ type GraphQLWebSocketMiddleware<'Root> longRunningCancellationToken.Register (fun _ -> (socket |> tryToGracefullyCloseSocketWithDefaultBehavior).Wait()) |> ignore try - do! socket |> handleMessages longRunningCancellationToken ctx + do! + socket + |> handleMessages sendGate longRunningCancellationToken ctx with ex -> logger.LogError (ex, "Cannot handle WebSocket message.") } diff --git a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs index 9438a545e..39321121d 100644 --- a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs +++ b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs @@ -216,6 +216,12 @@ let ``Observable error details preserve GraphQL-facing messages inside aggregate Assert.Contains ("Visible to client", actual) Assert.DoesNotContain ("sensitive backend failure", actual) +[] +let ``Observable error details fall back to the generic message for empty aggregates`` () = + let actual = problemDetailsOfObservableError (AggregateException ()) + let error = Assert.Single actual + Assert.Equal (UnexpectedObservableErrorMessage, error.Message) + [] let ``Observable error details do not duplicate repeated aggregate errors`` () = let actual = @@ -227,3 +233,15 @@ let ``Observable error details do not duplicate repeated aggregate errors`` () = let error = Assert.Single actual Assert.Equal ("Visible to client", error.Message) + +[] +let ``Request error sanitization replaces backend exception messages`` () = + let actual = + sanitizeRequestError (GQLProblemDetails.Create ("sensitive backend failure", Exception "sensitive backend failure")) + Assert.Equal (UnexpectedObservableErrorMessage, actual.Message) + +[] +let ``Request error sanitization preserves GraphQL-facing errors`` () = + let expected = GQLProblemDetails.OfError (GQLMessageException "Visible to client") + let actual = sanitizeRequestError expected + Assert.Equal (expected, actual) From ea8628a7b7ea997ef2be1c8a63269f37c016ad86 Mon Sep 17 00:00:00 2001 From: Andrii Chebukin Date: Sat, 19 Sep 2026 02:29:28 +0200 Subject: [PATCH 3/6] Added `vtryItem` for collections; updated `ValueOption` usage * Introduced `vtryItem` for Seq, List, and Array to provide ValueOption-based safe indexing. * Updated tryGetPathItemIndex to use `List.vtryItem` for improved safety and consistency. * Replaced `Seq.choose` with `Seq.vchoose` in `splitBatch`. * Added `InternalsVisibleTo` for `FSharp.Data.GraphQL.Server.AspNetCore` in the shared project. --- .../GraphQLWebsocketMiddleware.fs | 7 ++-- .../FSharp.Data.GraphQL.Shared.fsproj | 1 + .../Helpers/ObjAndStructConversions.fs | 34 +++++++++++++++++++ 3 files changed, 38 insertions(+), 4 deletions(-) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs index eda7c0e5a..c890cf3dd 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs @@ -49,7 +49,7 @@ module internal IncrementalPayloadSplitting = let fieldPathLength = List.length fieldPath if pathStartsWith fieldPath path then - path |> List.tryItem fieldPathLength |> ValueOption.ofOption + path |> List.vtryItem fieldPathLength else ValueNone @@ -68,12 +68,11 @@ module internal IncrementalPayloadSplitting = let items = data :?> obj[] let errorsByItemIndex = errors - |> Seq.choose (fun error -> + |> Seq.vchoose (fun error -> error.Path |> Skippable.toValueOption |> ValueOption.bind (tryGetPathItemIndex fieldPath) - |> ValueOption.map (fun index -> struct (index, error)) - |> ValueOption.toOption) + |> ValueOption.map (fun index -> struct (index, error))) |> _.ToLookup((fun struct (index, _) -> index), (fun struct (_, error) -> error)) (indices, List.ofArray items) diff --git a/src/FSharp.Data.GraphQL.Shared/FSharp.Data.GraphQL.Shared.fsproj b/src/FSharp.Data.GraphQL.Shared/FSharp.Data.GraphQL.Shared.fsproj index 270d1ea22..3bd336b2c 100644 --- a/src/FSharp.Data.GraphQL.Shared/FSharp.Data.GraphQL.Shared.fsproj +++ b/src/FSharp.Data.GraphQL.Shared/FSharp.Data.GraphQL.Shared.fsproj @@ -15,6 +15,7 @@ + diff --git a/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs b/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs index b68d5a9d4..eb7923080 100644 --- a/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs +++ b/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs @@ -39,6 +39,24 @@ module Seq = |> Seq.map ValueSome |> _.FirstOrDefault() + let vtryItem index (source : 'T seq) = + if index < 0 then + ValueNone + else + use enumerator = source.GetEnumerator () + let mutable currentIndex = 0 + let mutable result = ValueNone + let mutable found = false + + while not found && enumerator.MoveNext () do + if currentIndex = index then + result <- ValueSome enumerator.Current + found <- true + else + currentIndex <- currentIndex + 1 + + result + let vtryHead (source : 'T seq) = use enumerator = source.GetEnumerator () if not (enumerator.MoveNext ()) then @@ -66,10 +84,26 @@ module internal List = let vtryFind predicate list = list |> Seq.ofList |> Seq.vtryFind predicate + let vtryItem index list = + let rec loop currentIndex list = + match currentIndex, list with + | _, [] -> ValueNone + | 0, head :: _ -> ValueSome head + | currentIndex, _ :: tail when currentIndex > 0 -> loop (currentIndex - 1) tail + | _ -> ValueNone + + loop index list + module internal Array = let vchoose mapping array = array |> Seq.vchoose mapping |> Array.ofSeq + let vtryItem index (array : 'T array) = + if index < 0 || index >= array.Length then + ValueNone + else + ValueSome array.[index] + module internal Map = let vtryFind key (map : Map<_, _>) = From 075abb6887a669ea91fe9700ac81749615e6272d Mon Sep 17 00:00:00 2001 From: Andrii Chebukin Date: Sat, 19 Sep 2026 02:38:40 +0200 Subject: [PATCH 4/6] Stop treating null elements as missing in vtryHead and vtryLast --- .../Helpers/ObjAndStructConversions.fs | 53 ++++++++----------- 1 file changed, 22 insertions(+), 31 deletions(-) diff --git a/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs b/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs index eb7923080..b2eeaec9e 100644 --- a/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs +++ b/src/FSharp.Data.GraphQL.Shared/Helpers/ObjAndStructConversions.fs @@ -1,6 +1,5 @@ namespace rec FSharp.Data.GraphQL -open System.Linq open System.Collections.Generic open FsToolkit.ErrorHandling @@ -27,17 +26,30 @@ module internal ValueTuple = [] module Seq = + let vtryHead (source : 'T seq) = + use enumerator = source.GetEnumerator () + if not (enumerator.MoveNext ()) then + ValueNone + else + ValueSome enumerator.Current + + let vtryLast (source : 'T seq) = + use enumerator = source.GetEnumerator () + if not (enumerator.MoveNext ()) then + ValueNone + else + let mutable last = enumerator.Current + while enumerator.MoveNext () do + last <- enumerator.Current + ValueSome last + let vchoose mapping seq = seq |> Seq.map mapping |> Seq.where ValueOption.isSome |> Seq.map ValueOption.get - let vtryFind predicate seq = - seq - |> Seq.where predicate - |> Seq.map ValueSome - |> _.FirstOrDefault() + let vtryFind predicate (source : 'T seq) = source |> Seq.where predicate |> Seq.vtryHead let vtryItem index (source : 'T seq) = if index < 0 then @@ -57,32 +69,11 @@ module Seq = result - let vtryHead (source : 'T seq) = - use enumerator = source.GetEnumerator () - if not (enumerator.MoveNext ()) then - ValueNone - else - match enumerator.Current with - | null -> ValueNone - | head -> ValueSome head - - let vtryLast (source : 'T seq) = - use enumerator = source.GetEnumerator () - if not (enumerator.MoveNext ()) then - ValueNone - else - let mutable last = enumerator.Current - while enumerator.MoveNext () do - last <- enumerator.Current - match last with - | null -> ValueNone - | last -> ValueSome last - module internal List = - let vchoose mapping list = list |> Seq.ofList |> Seq.vchoose mapping |> Seq.toList + let vchoose mapping list = list |> Seq.vchoose mapping |> Seq.toList - let vtryFind predicate list = list |> Seq.ofList |> Seq.vtryFind predicate + let vtryFind predicate list = list |> Seq.where predicate |> Seq.vtryHead let vtryItem index list = let rec loop currentIndex list = @@ -96,13 +87,13 @@ module internal List = module internal Array = - let vchoose mapping array = array |> Seq.vchoose mapping |> Array.ofSeq + let vchoose mapping array = array |> Seq.vchoose mapping |> Seq.toArray let vtryItem index (array : 'T array) = if index < 0 || index >= array.Length then ValueNone else - ValueSome array.[index] + ValueSome array[index] module internal Map = From 9f66d7a8991d22b0c4bf3bb91a59281d9ebfffd4 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sat, 19 Sep 2026 00:50:30 +0000 Subject: [PATCH 5/6] Serialize websocket closes and cleanup all subscriptions Co-authored-by: xperiandri <2365592+xperiandri@users.noreply.github.com> --- .../GraphQLSubscriptionsManagement.fs | 13 +++- .../GraphQLWebsocketMiddleware.fs | 69 ++++++++++++------- .../AspNetCore/SerializationTests.fs | 50 +++++++++++++- 3 files changed, 103 insertions(+), 29 deletions(-) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs index 95c078b1d..8c80e63b4 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLSubscriptionsManagement.fs @@ -1,5 +1,7 @@ module internal FSharp.Data.GraphQL.Server.AspNetCore.GraphQLSubscriptionsManagement +open System + open FSharp.Data.GraphQL.Shared.WebSockets let addSubscription @@ -42,5 +44,14 @@ let removeAllSubscriptions (subscriptions : SubscriptionsDict) = subscriptions.Clear () snapshot) + let exceptions = ResizeArray () + subscriptionsToDispose - |> Array.iter (fun struct (id, subscription) -> subscription |> executeOnUnsubscribeAndDispose id) + |> Array.iter (fun struct (id, subscription) -> + try + subscription |> executeOnUnsubscribeAndDispose id + with ex -> + exceptions.Add ex) + + if exceptions.Count > 0 then + raise (AggregateException ("One or more subscriptions failed to unsubscribe.", exceptions)) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs index c890cf3dd..3e4dba193 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs @@ -305,14 +305,23 @@ type GraphQLWebSocketMiddleware<'Root> |> GraphQLSubscriptionsManagement.removeSubscription id reraise () - let tryToGracefullyCloseSocket (code, message) theSocket = - if theSocket |> canCloseSocket then - theSocket.CloseAsync (code, message, CancellationToken.None) - else - Task.CompletedTask + let tryToGracefullyCloseSocket (sendGate : SemaphoreSlim) (code, message) (theSocket : WebSocket) : Task = task { + do! sendGate.WaitAsync () - let tryToGracefullyCloseSocketWithDefaultBehavior = - tryToGracefullyCloseSocket (WebSocketCloseStatus.NormalClosure, "Normal Closure") + try + if theSocket |> canCloseSocket then + do! theSocket.CloseAsync (code, message, CancellationToken.None) + else + logger.LogTrace ( + $"Ignoring socket close request, since its state is neither writable nor closeable, but '{{state}}'", + theSocket.State + ) + finally + sendGate.Release () |> ignore + } + + let tryToGracefullyCloseSocketWithDefaultBehavior sendGate = + tryToGracefullyCloseSocket sendGate (WebSocketCloseStatus.NormalClosure, "Normal Closure") let handleMessages (sendGate : SemaphoreSlim) (cancellationToken : CancellationToken) (httpContext : HttpContext) (socket : WebSocket) : Task = let subscriptions = Dictionary() @@ -439,18 +448,20 @@ type GraphQLWebSocketMiddleware<'Root> nameof InvalidMessage |> logMsgReceivedWithOptionalPayload ValueNone match failureMessages with - | InvalidMessage (code, explanation) -> do! socket.CloseAsync (enum code, explanation, CancellationToken.None) + | InvalidMessage (code, explanation) -> + do! + socket + |> tryToGracefullyCloseSocket sendGate (enum code, explanation) | Ok ValueNone -> logger.LogTrace ("WebSocket received empty message! State = '{socketState}'", socket.State) | Ok (ValueSome msg) -> match msg with | ConnectionInit p -> nameof ConnectionInit |> logMsgReceivedWithOptionalPayload p do! - socket.CloseAsync ( - enum CustomWebSocketStatus.TooManyInitializationRequests, - "Too many initialization requests", - CancellationToken.None - ) + socket + |> tryToGracefullyCloseSocket + sendGate + (enum CustomWebSocketStatus.TooManyInitializationRequests, "Too many initialization requests") | ClientPing p -> nameof ClientPing |> logMsgReceivedWithOptionalPayload p match pingHandler with @@ -466,11 +477,10 @@ type GraphQLWebSocketMiddleware<'Root> do! let warningMsg : FormattableString = $"Subscriber for Id = '{id}' already exists" logger.LogWarning (String.Format (warningMsg.Format, "id"), id) - socket.CloseAsync ( - enum CustomWebSocketStatus.SubscriberAlreadyExists, - warningMsg.ToString (), - CancellationToken.None - ) + socket + |> tryToGracefullyCloseSocket + sendGate + (enum CustomWebSocketStatus.SubscriberAlreadyExists, warningMsg.ToString ()) else let variables = query.Variables |> Skippable.toValueOption let getInputContext () = httpContext.RequestServices.GetRequiredService() @@ -486,13 +496,17 @@ type GraphQLWebSocketMiddleware<'Root> subscriptions |> GraphQLSubscriptionsManagement.removeSubscription (id) logger.LogTrace "Leaving the 'graphql-ws' connection loop..." - do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + do! + socket + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate with ex -> logger.LogError (ex, "Cannot handle a message; dropping a websocket connection") // At this point, only something really weird must have happened. // In order to avoid faulty state scenarios and unimagined damages, // just close the socket without further ado. - do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + do! + socket + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate finally subscriptions |> GraphQLSubscriptionsManagement.removeAllSubscriptions @@ -508,7 +522,7 @@ type GraphQLWebSocketMiddleware<'Root> let detonationRegistration = timerTokenSource.Token.Register (fun _ -> (socket - |> tryToGracefullyCloseSocket (enum CustomWebSocketStatus.ConnectionTimeout, "Connection initialization timeout")) + |> tryToGracefullyCloseSocket sendGate (enum CustomWebSocketStatus.ConnectionTimeout, "Connection initialization timeout")) .Wait()) let! connectionInitSucceeded = @@ -527,15 +541,17 @@ type GraphQLWebSocketMiddleware<'Root> | Ok (ValueSome (Subscribe _)) -> do! socket - |> tryToGracefullyCloseSocket (enum CustomWebSocketStatus.Unauthorized, "Unauthorized") + |> tryToGracefullyCloseSocket sendGate (enum CustomWebSocketStatus.Unauthorized, "Unauthorized") return false | Result.Error (InvalidMessage (code, explanation)) -> do! socket - |> tryToGracefullyCloseSocket (enum code, explanation) + |> tryToGracefullyCloseSocket sendGate (enum code, explanation) return false | _ -> - do! socket |> tryToGracefullyCloseSocketWithDefaultBehavior + do! + socket + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate return false }), timerTokenSource.Token @@ -560,7 +576,10 @@ type GraphQLWebSocketMiddleware<'Root> | Ok _ -> let longRunningCancellationToken = (CancellationTokenSource.CreateLinkedTokenSource(ctx.RequestAborted, applicationLifetime.ApplicationStopping).Token) - longRunningCancellationToken.Register (fun _ -> (socket |> tryToGracefullyCloseSocketWithDefaultBehavior).Wait()) + longRunningCancellationToken.Register (fun _ -> + (socket + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate) + .Wait()) |> ignore try do! diff --git a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs index 39321121d..490e481a6 100644 --- a/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs +++ b/tests/FSharp.Data.GraphQL.Tests/AspNetCore/SerializationTests.fs @@ -1,14 +1,18 @@ module FSharp.Data.GraphQL.Tests.AspNetCore.SerializationTests open System +open System.Collections.Concurrent open System.Collections.Generic -open Xunit open System.Text.Json +open System.Text.Json.Serialization + +open Xunit + open FSharp.Data.GraphQL.Ast open FSharp.Data.GraphQL.Shared -open FSharp.Data.GraphQL.Shared.WebSockets +open FSharp.Data.GraphQL.Server.AspNetCore.GraphQLSubscriptionsManagement open FSharp.Data.GraphQL.Server.AspNetCore.ObservableErrorHandling -open System.Text.Json.Serialization +open FSharp.Data.GraphQL.Shared.WebSockets [] let ``Deserializes ConnectionInit correctly`` () = @@ -245,3 +249,43 @@ let ``Request error sanitization preserves GraphQL-facing errors`` () = let expected = GQLProblemDetails.OfError (GQLMessageException "Visible to client") let actual = sanitizeRequestError expected Assert.Equal (expected, actual) + +type private TrackingSubscription (onDispose : unit -> unit) = + interface IDisposable with + member _.Dispose () = onDispose () + +[] +let ``Removing all subscriptions attempts every disposal before raising aggregate failure`` () = + let disposedIds = ConcurrentQueue () + let unsubscribedIds = ConcurrentQueue () + let subscriptions = + Dictionary() :> SubscriptionsDict + + let createSubscription id shouldThrow = + let subscription = + new TrackingSubscription (fun () -> + disposedIds.Enqueue id + + if shouldThrow then + raise (InvalidOperationException $"Dispose failed for {id}")) + + let onUnsubscribe removedId = + unsubscribedIds.Enqueue removedId + + if shouldThrow then + raise (InvalidOperationException $"Unsubscribe failed for {removedId}") + + id, (subscription :> SubscriptionUnsubscriber), onUnsubscribe + + subscriptions + |> addSubscription (createSubscription "first" true) + subscriptions + |> addSubscription (createSubscription "second" false) + + let error = Assert.Throws(fun () -> subscriptions |> removeAllSubscriptions) + + Assert.False (subscriptions.ContainsKey "first") + Assert.False (subscriptions.ContainsKey "second") + Assert.Equal(set [ "first"; "second" ], set disposedIds) + Assert.Equal(set [ "first"; "second" ], set unsubscribedIds) + Assert.Single error.InnerExceptions From 6486c2d73f1b1525b69b42c5ce4ebca0aa454727 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sat, 19 Sep 2026 01:05:27 +0000 Subject: [PATCH 6/6] Bound websocket closes and preserve request logs Co-authored-by: xperiandri <2365592+xperiandri@users.noreply.github.com> --- .../GraphQLWebsocketMiddleware.fs | 84 ++++++++++++------- 1 file changed, 54 insertions(+), 30 deletions(-) diff --git a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs index 3e4dba193..b4fcd6cc3 100644 --- a/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs +++ b/src/FSharp.Data.GraphQL.Server.AspNetCore/GraphQLWebsocketMiddleware.fs @@ -141,6 +141,7 @@ type GraphQLWebSocketMiddleware<'Root> let serializerOptions = options.SerializerOptions let pingHandler = options.WebsocketOptions.CustomPingHandler let connectionInitTimeout = options.WebsocketOptions.ConnectionInitTimeout + let gracefulCloseTimeout : TimeSpan = TimeSpan.FromSeconds 5.0 let serializeServerMessage (jsonSerializerOptions : JsonSerializerOptions) (serverMessage : ServerMessage) = task { let raw = @@ -305,23 +306,34 @@ type GraphQLWebSocketMiddleware<'Root> |> GraphQLSubscriptionsManagement.removeSubscription id reraise () - let tryToGracefullyCloseSocket (sendGate : SemaphoreSlim) (code, message) (theSocket : WebSocket) : Task = task { - do! sendGate.WaitAsync () + let tryToGracefullyCloseSocket (sendGate : SemaphoreSlim) (cancellationToken : CancellationToken) (code, message) (theSocket : WebSocket) : Task = + task { + do! sendGate.WaitAsync () - try - if theSocket |> canCloseSocket then - do! theSocket.CloseAsync (code, message, CancellationToken.None) - else - logger.LogTrace ( - $"Ignoring socket close request, since its state is neither writable nor closeable, but '{{state}}'", - theSocket.State - ) - finally - sendGate.Release () |> ignore - } + try + if theSocket |> canCloseSocket then + use closeCancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource cancellationToken + closeCancellationTokenSource.CancelAfter gracefulCloseTimeout + + try + do! theSocket.CloseAsync (code, message, closeCancellationTokenSource.Token) + with :? OperationCanceledException -> + logger.LogWarning ( + "Aborting WebSocket after graceful close did not complete before cancellation. State = '{state}'", + theSocket.State + ) + theSocket.Abort () + else + logger.LogTrace ( + $"Ignoring socket close request, since its state is neither writable nor closeable, but '{{state}}'", + theSocket.State + ) + finally + sendGate.Release () |> ignore + } - let tryToGracefullyCloseSocketWithDefaultBehavior sendGate = - tryToGracefullyCloseSocket sendGate (WebSocketCloseStatus.NormalClosure, "Normal Closure") + let tryToGracefullyCloseSocketWithDefaultBehavior sendGate cancellationToken = + tryToGracefullyCloseSocket sendGate cancellationToken (WebSocketCloseStatus.NormalClosure, "Normal Closure") let handleMessages (sendGate : SemaphoreSlim) (cancellationToken : CancellationToken) (httpContext : HttpContext) (socket : WebSocket) : Task = let subscriptions = Dictionary() @@ -416,7 +428,7 @@ type GraphQLWebSocketMiddleware<'Root> do! sendMsg (Complete id) | RequestError problemDetails -> let sanitizedProblemDetails = problemDetails |> List.map sanitizeRequestError - logger.LogWarning ("Request errors:\n{errors}", sanitizedProblemDetails) + logger.LogWarning ("Request errors:\n{errors}", problemDetails) // The request was rejected before execution, so it is not a result: the protocol requires it to be // sent as the terminal Error message instead of a Next followed by Complete, or a client would // read it as a successful result with null data @@ -451,7 +463,7 @@ type GraphQLWebSocketMiddleware<'Root> | InvalidMessage (code, explanation) -> do! socket - |> tryToGracefullyCloseSocket sendGate (enum code, explanation) + |> tryToGracefullyCloseSocket sendGate cancellationToken (enum code, explanation) | Ok ValueNone -> logger.LogTrace ("WebSocket received empty message! State = '{socketState}'", socket.State) | Ok (ValueSome msg) -> match msg with @@ -461,6 +473,7 @@ type GraphQLWebSocketMiddleware<'Root> socket |> tryToGracefullyCloseSocket sendGate + cancellationToken (enum CustomWebSocketStatus.TooManyInitializationRequests, "Too many initialization requests") | ClientPing p -> nameof ClientPing |> logMsgReceivedWithOptionalPayload p @@ -480,6 +493,7 @@ type GraphQLWebSocketMiddleware<'Root> socket |> tryToGracefullyCloseSocket sendGate + cancellationToken (enum CustomWebSocketStatus.SubscriberAlreadyExists, warningMsg.ToString ()) else let variables = query.Variables |> Skippable.toValueOption @@ -498,7 +512,7 @@ type GraphQLWebSocketMiddleware<'Root> logger.LogTrace "Leaving the 'graphql-ws' connection loop..." do! socket - |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate cancellationToken with ex -> logger.LogError (ex, "Cannot handle a message; dropping a websocket connection") // At this point, only something really weird must have happened. @@ -506,7 +520,7 @@ type GraphQLWebSocketMiddleware<'Root> // just close the socket without further ado. do! socket - |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate cancellationToken finally subscriptions |> GraphQLSubscriptionsManagement.removeAllSubscriptions @@ -516,13 +530,20 @@ type GraphQLWebSocketMiddleware<'Root> // <-- Main // <-------- - let waitForConnectionInitAndRespondToClient (sendGate : SemaphoreSlim) (socket : WebSocket) : TaskResult = task { + let waitForConnectionInitAndRespondToClient + (sendGate : SemaphoreSlim) + (cancellationToken : CancellationToken) + (socket : WebSocket) + : TaskResult = task { let timerTokenSource = new CancellationTokenSource () timerTokenSource.CancelAfter connectionInitTimeout let detonationRegistration = timerTokenSource.Token.Register (fun _ -> (socket - |> tryToGracefullyCloseSocket sendGate (enum CustomWebSocketStatus.ConnectionTimeout, "Connection initialization timeout")) + |> tryToGracefullyCloseSocket + sendGate + cancellationToken + (enum CustomWebSocketStatus.ConnectionTimeout, "Connection initialization timeout")) .Wait()) let! connectionInitSucceeded = @@ -541,17 +562,17 @@ type GraphQLWebSocketMiddleware<'Root> | Ok (ValueSome (Subscribe _)) -> do! socket - |> tryToGracefullyCloseSocket sendGate (enum CustomWebSocketStatus.Unauthorized, "Unauthorized") + |> tryToGracefullyCloseSocket sendGate cancellationToken (enum CustomWebSocketStatus.Unauthorized, "Unauthorized") return false | Result.Error (InvalidMessage (code, explanation)) -> do! socket - |> tryToGracefullyCloseSocket sendGate (enum code, explanation) + |> tryToGracefullyCloseSocket sendGate cancellationToken (enum code, explanation) return false | _ -> do! socket - |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate cancellationToken return false }), timerTokenSource.Token @@ -570,21 +591,24 @@ type GraphQLWebSocketMiddleware<'Root> task { use! socket = ctx.WebSockets.AcceptWebSocketAsync ("graphql-transport-ws") let sendGate = new SemaphoreSlim (1, 1) - let! connectionInitResult = socket |> waitForConnectionInitAndRespondToClient sendGate + use connectionLifetimeCancellationTokenSource = + CancellationTokenSource.CreateLinkedTokenSource (ctx.RequestAborted, applicationLifetime.ApplicationStopping) + let connectionLifetimeCancellationToken = connectionLifetimeCancellationTokenSource.Token + let! connectionInitResult = + socket + |> waitForConnectionInitAndRespondToClient sendGate connectionLifetimeCancellationToken match connectionInitResult with | Result.Error errMsg -> logger.LogWarning errMsg | Ok _ -> - let longRunningCancellationToken = - (CancellationTokenSource.CreateLinkedTokenSource(ctx.RequestAborted, applicationLifetime.ApplicationStopping).Token) - longRunningCancellationToken.Register (fun _ -> + connectionLifetimeCancellationToken.Register (fun _ -> (socket - |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate) + |> tryToGracefullyCloseSocketWithDefaultBehavior sendGate connectionLifetimeCancellationToken) .Wait()) |> ignore try do! socket - |> handleMessages sendGate longRunningCancellationToken ctx + |> handleMessages sendGate connectionLifetimeCancellationToken ctx with ex -> logger.LogError (ex, "Cannot handle WebSocket message.") }