diff --git a/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs b/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs index 3206121..4d4ef41 100644 --- a/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs +++ b/src/StevesBot.Worker.Tests/Unit/DiscordGatewayClientTests.cs @@ -1,3 +1,5 @@ +using System.ComponentModel; + namespace StevesBot.Worker.Tests.Unit; public sealed class DiscordGatewayClientTests : IDisposable @@ -150,11 +152,6 @@ public sealed class DiscordGatewayClientTests : IDisposable [Fact] public async Task ConnectAsync_WhenConnectedAndHelloEventIsReceived_ItShouldStartSendingHeartbeatsAndIdentify() { - // TODO: This does not work! You need - // to figure it out! When this runs - // concurrently with test above that - // test starts to fail - _mockDiscordRestClient .Setup(static x => x.GetGatewayUrlAsync(It.IsAny())) .ReturnsAsync("wss://gateway.discord.gg"); @@ -181,9 +178,9 @@ public sealed class DiscordGatewayClientTests : IDisposable HeartbeatInterval = heartbeatInterval, } }; - var payload = CreateEventPayload(helloEvent); - var result = new WebSocketReceiveResult(payload.Bytes.Length, WebSocketMessageType.Text, true); - messagesToReceive.Enqueue((result, payload.Bytes)); + var helloPayload = CreateEventPayload(helloEvent); + var helloResult = new WebSocketReceiveResult(helloPayload.Bytes.Length, WebSocketMessageType.Text, true); + messagesToReceive.Enqueue((helloResult, helloPayload.Bytes)); SetupReceiveMessageSequence(mockWebSocket, messagesToReceive); @@ -197,13 +194,13 @@ public sealed class DiscordGatewayClientTests : IDisposable await Task.Delay((int)(heartbeatInterval * 1.5)); await cts.CancelAsync(); - var heartbeatEvent = new HeartbeatDiscordEvent(helloEvent.Sequence); - var heartbeatPayload = CreateEventPayload(heartbeatEvent); + var expectedHeartbeatEvent = new HeartbeatDiscordEvent(helloEvent.Sequence); + var expectedHeartbeatPayload = CreateEventPayload(expectedHeartbeatEvent); mockWebSocket .Verify( x => x.SendAsync( - It.Is>(b => heartbeatPayload.Bytes.SequenceEqual(b.Array!)), + It.Is>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)), It.Is(m => m == WebSocketMessageType.Text), true, It.IsAny() @@ -211,7 +208,7 @@ public sealed class DiscordGatewayClientTests : IDisposable Times.Once ); - var identifyEvent = new IdentifyDiscordEvent( + var expectedIdentifyEvent = new IdentifyDiscordEvent( _options.AppToken, _options.Intents, new UpdatePresenceData @@ -226,12 +223,12 @@ public sealed class DiscordGatewayClientTests : IDisposable ], } ); - var identifyPayload = CreateEventPayload(identifyEvent); + var expectedIdentifyPayload = CreateEventPayload(expectedIdentifyEvent); mockWebSocket .Verify( x => x.SendAsync( - It.Is>(b => identifyPayload.Bytes.SequenceEqual(b.Array!)), + It.Is>(b => expectedIdentifyPayload.Bytes.SequenceEqual(b.Array!)), It.Is(m => m == WebSocketMessageType.Text), true, It.IsAny() @@ -240,6 +237,180 @@ public sealed class DiscordGatewayClientTests : IDisposable ); } + [Fact] + public async Task ConnectAsync_OnceConnected_ItShouldStopSendingHeartbeatsIfTheyAreNotAcknowledgedAndAttemptToResume() + { + _mockDiscordRestClient + .Setup(static x => x.GetGatewayUrlAsync(It.IsAny())) + .ReturnsAsync("wss://gateway.discord.gg"); + + var initialSocketState = WebSocketState.Closed; + + var initialWebSocket = new Mock(); + + initialWebSocket + .Setup(static x => x.State) + .Returns(() => initialSocketState); + + initialWebSocket + .Setup(static x => x.ConnectAsync(It.IsAny(), It.IsAny())) + .Callback(() => initialSocketState = WebSocketState.Open) + .Returns(Task.CompletedTask); + + initialWebSocket + .Setup(static x => x.CloseAsync(It.IsAny(), It.IsAny(), It.IsAny())) + .Callback(() => initialSocketState = WebSocketState.Closed) + .Returns(Task.CompletedTask); + + var resumingSocketState = WebSocketState.Closed; + + var resumingWebSocket = new Mock(); + + resumingWebSocket + .Setup(static x => x.State) + .Returns(() => resumingSocketState); + + resumingWebSocket + .Setup(static x => x.ConnectAsync(It.IsAny(), It.IsAny())) + .Callback(() => resumingSocketState = WebSocketState.Open) + .Returns(Task.CompletedTask); + + resumingWebSocket + .Setup(static x => x.CloseAsync(It.IsAny(), It.IsAny(), It.IsAny())) + .Callback(() => resumingSocketState = WebSocketState.Closed) + .Returns(Task.CompletedTask); + + var messagesToReceive = new Queue<(WebSocketReceiveResult, byte[])>(); + var heartbeatInterval = 100; + var helloEvent = new HelloDiscordEvent() + { + Data = new() + { + HeartbeatInterval = heartbeatInterval, + } + }; + var helloPayload = CreateEventPayload(helloEvent); + var helloResult = new WebSocketReceiveResult(helloPayload.Bytes.Length, WebSocketMessageType.Text, true); + messagesToReceive.Enqueue((helloResult, helloPayload.Bytes)); + + var sessionId = "session_id"; + var resumeGatewayUrl = "wss://resume.discord.gg"; + var readyEvent = new ReadyDiscordEvent() + { + Sequence = 1, + Data = new ReadyData() + { + SessionId = sessionId, + ResumeGatewayUrl = resumeGatewayUrl, + } + }; + var readyPayload = CreateEventPayload(readyEvent); + var readyResult = new WebSocketReceiveResult(readyPayload.Bytes.Length, WebSocketMessageType.Text, true); + messagesToReceive.Enqueue((readyResult, readyPayload.Bytes)); + + SetupReceiveMessageSequence(initialWebSocket, messagesToReceive); + + _mockWebSocketFactory + .SetupSequence(static x => x.Create()) + .Returns(initialWebSocket.Object) + .Returns(resumingWebSocket.Object); + + using var cts = new CancellationTokenSource(); + + await _discordGatewayClient.ConnectAsync(CancellationToken.None); + await Task.Delay((int)(heartbeatInterval * 2.5)); + + var expectedHeartbeatEvent = new HeartbeatDiscordEvent(helloEvent.Sequence); + var expectedHeartbeatPayload = CreateEventPayload(expectedHeartbeatEvent); + + initialWebSocket + .Verify( + x => x.SendAsync( + It.Is>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)), + It.Is(m => m == WebSocketMessageType.Text), + true, + It.IsAny() + ), + Times.AtMostOnce + ); + + var expectedUri = new Uri($"{resumeGatewayUrl}/?v=10&encoding=json"); + + resumingWebSocket + .Verify( + x => x.ConnectAsync( + It.Is(uri => uri.Equals(expectedUri)), + It.IsAny() + ), + Times.Once + ); + } + + [Fact] + public async Task ConnectAsync_OnceConnected_ItShouldContinueSendingHeartbeatsIfTheyAreAcknowledged() + { + _mockDiscordRestClient + .Setup(static x => x.GetGatewayUrlAsync(It.IsAny())) + .ReturnsAsync("wss://gateway.discord.gg"); + + var mockWebSocket = new Mock(); + + var socketState = WebSocketState.Closed; + + mockWebSocket + .Setup(static x => x.State) + .Returns(() => socketState); + + mockWebSocket + .Setup(static x => x.ConnectAsync(It.IsAny(), It.IsAny())) + .Callback(() => socketState = WebSocketState.Open) + .Returns(Task.CompletedTask); + + var messagesToReceive = new Queue<(WebSocketReceiveResult, byte[])>(); + var heartbeatInterval = 100; + var helloEvent = new HelloDiscordEvent() + { + Data = new() + { + HeartbeatInterval = heartbeatInterval, + } + }; + var helloPayload = CreateEventPayload(helloEvent); + var helloResult = new WebSocketReceiveResult(helloPayload.Bytes.Length, WebSocketMessageType.Text, true); + messagesToReceive.Enqueue((helloResult, helloPayload.Bytes)); + + var heartbeatAckEvent = new HeartbeatAckDiscordEvent(); + var heartbeatAckPayload = CreateEventPayload(heartbeatAckEvent); + var heartbeatAckResult = new WebSocketReceiveResult(heartbeatAckPayload.Bytes.Length, WebSocketMessageType.Text, true); + messagesToReceive.Enqueue((heartbeatAckResult, heartbeatAckPayload.Bytes)); + + SetupReceiveMessageSequence(mockWebSocket, messagesToReceive); + + _mockWebSocketFactory + .Setup(static x => x.Create()) + .Returns(mockWebSocket.Object); + + using var cts = new CancellationTokenSource(); + + await _discordGatewayClient.ConnectAsync(cts.Token); + await Task.Delay((int)(heartbeatInterval * 2.5)); + await cts.CancelAsync(); + + var expectedHeartbeatEvent = new HeartbeatDiscordEvent(helloEvent.Sequence); + var expectedHeartbeatPayload = CreateEventPayload(expectedHeartbeatEvent); + + mockWebSocket + .Verify( + x => x.SendAsync( + It.Is>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)), + It.Is(m => m == WebSocketMessageType.Text), + true, + It.IsAny() + ), + Times.AtLeast(2) + ); + } + private static void SetupReceiveMessageSequence( Mock mockWebSocket, Queue<(WebSocketReceiveResult, byte[])> messageQueue diff --git a/src/StevesBot.Worker.Tests/Unit/HeartbeatAckDiscordEventTests.cs b/src/StevesBot.Worker.Tests/Unit/HeartbeatAckDiscordEventTests.cs index e69de29..75f8ef0 100644 --- a/src/StevesBot.Worker.Tests/Unit/HeartbeatAckDiscordEventTests.cs +++ b/src/StevesBot.Worker.Tests/Unit/HeartbeatAckDiscordEventTests.cs @@ -0,0 +1,15 @@ +namespace StevesBot.Worker.Tests.Unit; + +public class HeartbeatAckDiscordEventTests +{ + [Fact] + public void Constructor_WhenCalled_ItShouldReturnAnInstance() + { + var result = new HeartbeatAckDiscordEvent(); + + result.OpCode.Should().Be(DiscordOpCodes.HeartbeatAck); + result.Sequence.Should().BeNull(); + result.Type.Should().BeNull(); + result.Data.Should().BeNull(); + } +} \ No newline at end of file diff --git a/src/StevesBot.Worker.Tests/Unit/ReadyDiscordEventTests.cs b/src/StevesBot.Worker.Tests/Unit/ReadyDiscordEventTests.cs index b7b3cde..6847123 100644 --- a/src/StevesBot.Worker.Tests/Unit/ReadyDiscordEventTests.cs +++ b/src/StevesBot.Worker.Tests/Unit/ReadyDiscordEventTests.cs @@ -10,7 +10,7 @@ public class ReadyDiscordEventTests readyEvent.OpCode.Should().Be(0); readyEvent.Sequence.Should().BeNull(); - readyEvent.Type.Should().BeNull(); + readyEvent.Type.Should().Be(DiscordEventTypes.Ready); readyEvent.Data.Should().BeEquivalentTo(readyData); readyData.Version.Should().Be(0); readyData.SessionId.Should().Be(string.Empty); diff --git a/src/StevesBot.Worker/Discord/Gateway/DiscordGatewayClient.cs b/src/StevesBot.Worker/Discord/Gateway/DiscordGatewayClient.cs index 15d4a1e..471e665 100644 --- a/src/StevesBot.Worker/Discord/Gateway/DiscordGatewayClient.cs +++ b/src/StevesBot.Worker/Discord/Gateway/DiscordGatewayClient.cs @@ -180,6 +180,10 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient memoryStream.Seek(0, SeekOrigin.Begin); + var msg = Encoding.UTF8.GetString(messageBuffer, 0, result.Count); + + _logger.LogDebug("Received message: {Message}", msg); + var e = await JsonSerializer.DeserializeAsync( memoryStream, _jsonSerializerOptions, @@ -192,7 +196,7 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient continue; } - await HandleEventAsync(e, _linkedReceiveMessageCts.Token); + await HandleEventAsync(e, cancellationToken); } } catch (OperationCanceledException ex) @@ -366,8 +370,8 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient private async Task ReconnectAsync(CancellationToken cancellationToken) { - await CancelReceiveMessagesTaskAsync(cancellationToken); await CancelHeartbeatTaskAsync(cancellationToken); + await CancelReceiveMessagesTaskAsync(cancellationToken); var closeStatus = _canResume ? WebSocketCloseStatus.MandatoryExtension : WebSocketCloseStatus.NormalClosure; diff --git a/src/StevesBot.Worker/Discord/Gateway/Events/HeartbeatAckDiscordEvent.cs b/src/StevesBot.Worker/Discord/Gateway/Events/HeartbeatAckDiscordEvent.cs index e005a88..b5a8483 100644 --- a/src/StevesBot.Worker/Discord/Gateway/Events/HeartbeatAckDiscordEvent.cs +++ b/src/StevesBot.Worker/Discord/Gateway/Events/HeartbeatAckDiscordEvent.cs @@ -2,4 +2,8 @@ namespace StevesBot.Worker.Discord.Gateway.Events; internal sealed record HeartbeatAckDiscordEvent : DiscordEvent { + public HeartbeatAckDiscordEvent() + { + OpCode = DiscordOpCodes.HeartbeatAck; + } } \ No newline at end of file diff --git a/src/StevesBot.Worker/Discord/Gateway/Events/ReadyDiscordEvent.cs b/src/StevesBot.Worker/Discord/Gateway/Events/ReadyDiscordEvent.cs index 69deb59..0bcaa2a 100644 --- a/src/StevesBot.Worker/Discord/Gateway/Events/ReadyDiscordEvent.cs +++ b/src/StevesBot.Worker/Discord/Gateway/Events/ReadyDiscordEvent.cs @@ -2,6 +2,9 @@ namespace StevesBot.Worker.Discord.Gateway.Events; internal sealed record ReadyDiscordEvent : DispatchDiscordEvent { + [JsonPropertyName("t")] + public new string Type { get; init; } = DiscordEventTypes.Ready; + [JsonPropertyName("d")] public new ReadyData Data { get; init; } = new ReadyData(); }