diff --git a/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs b/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs index 0a50e97..195e34d 100644 --- a/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs +++ b/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs @@ -6,6 +6,7 @@ public sealed class DiscordGatewayClientTests : IDisposable private readonly Mock _mockWebSocketFactory = new(); private readonly Mock> _mockLogger = new(); private readonly DiscordClientOptions _options = new(); +# pragma warning disable CA2213 private readonly DiscordGatewayClient _discordGatewayClient; public DiscordGatewayClientTests() @@ -91,8 +92,82 @@ public sealed class DiscordGatewayClientTests : IDisposable ); } + [Fact] + public async Task ConnectAsync_WhenCalledAndHelloEventReceived_ItShouldStartSendingHeartbeats() + { + _mockDiscordRestClient + .Setup(static x => x.GetGatewayUrlAsync(It.IsAny())) + .ReturnsAsync("wss://gateway.discord.gg"); + + var mockWebSocket = new Mock(); + + mockWebSocket + .Setup(static x => x.State) + .Returns(WebSocketState.Open); + + var messageQueue = new Queue<(WebSocketReceiveResult, byte[])>(); + + var heatbeatInterval = 1000; + + var helloEvent = new + { + op = 10, + d = new + { + heartbeat_interval = heatbeatInterval, + } + }; + + var helloEventJson = JsonSerializer.Serialize(helloEvent); + var helloEventBytes = Encoding.UTF8.GetBytes(helloEventJson); + + var heartbeatAck = new + { + op = 11, + }; + + var heartbeatAckJson = JsonSerializer.Serialize(heartbeatAck); + var heartbeatAckBytes = Encoding.UTF8.GetBytes(heartbeatAckJson); + + messageQueue.Enqueue((new WebSocketReceiveResult(helloEventBytes.Length, WebSocketMessageType.Text, true), helloEventBytes)); + messageQueue.Enqueue((new WebSocketReceiveResult(heartbeatAckBytes.Length, WebSocketMessageType.Text, true), heartbeatAckBytes)); + + mockWebSocket + .Setup(static x => x.SendAsync(It.IsAny>(), It.IsAny(), It.IsAny(), It.IsAny())) + .Returns(Task.CompletedTask); + + mockWebSocket + .Setup(static x => x.ReceiveAsync(It.IsAny>(), It.IsAny())) + .Returns((ArraySegment buffer, CancellationToken token) => + { + if (messageQueue.Count == 0) + { +# pragma warning disable CA2008 + return Task.Delay(-1, token).ContinueWith(_ => new WebSocketReceiveResult(0, WebSocketMessageType.Text, true), token); + } + + var (result, messageBytes) = messageQueue.Dequeue(); + Array.Copy(messageBytes, 0, buffer.Array!, buffer.Offset, Math.Min(messageBytes.Length, buffer.Count)); + return Task.FromResult(result); + }); + + _mockWebSocketFactory + .Setup(static x => x.Create()) + .Returns(mockWebSocket.Object); + + using var cts = new CancellationTokenSource(); + + await _discordGatewayClient.ConnectAsync(CancellationToken.None); + + await Task.Delay(heatbeatInterval + 30000); + + mockWebSocket.Verify( + static x => x.SendAsync(It.IsAny>(), It.IsAny(), It.IsAny(), It.IsAny()), + Times.Once + ); + } + public void Dispose() { - _discordGatewayClient.Dispose(); } } \ No newline at end of file diff --git a/src/StevesBot.Worker/Discord/DiscordGatewayClient.cs b/src/StevesBot.Worker/Discord/DiscordGatewayClient.cs index b696506..3e0e267 100644 --- a/src/StevesBot.Worker/Discord/DiscordGatewayClient.cs +++ b/src/StevesBot.Worker/Discord/DiscordGatewayClient.cs @@ -108,8 +108,8 @@ internal class DiscordGatewayClient : IDiscordGatewayClient if (result.MessageType is WebSocketMessageType.Text) { -# pragma warning disable CA1849 - memoryStream.Write(messageBuffer, 0, result.Count); + await memoryStream.WriteAsync(messageBuffer.AsMemory(0, result.Count), cancellationToken); + await memoryStream.FlushAsync(cancellationToken); } } while (result.EndOfMessage is false);