tests: man i'm so proud of figuring out how to effectively unit test the gateway client =D

This commit is contained in:
Stevan Freeborn
2025-05-23 12:09:41 -05:00
parent 86b317ab4e
commit 6258d02998
6 changed files with 214 additions and 17 deletions
@@ -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<CancellationToken>()))
.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<ArraySegment<byte>>(b => heartbeatPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<ArraySegment<byte>>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<WebSocketMessageType>(m => m == WebSocketMessageType.Text),
true,
It.IsAny<CancellationToken>()
@@ -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<ArraySegment<byte>>(b => identifyPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<ArraySegment<byte>>(b => expectedIdentifyPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<WebSocketMessageType>(m => m == WebSocketMessageType.Text),
true,
It.IsAny<CancellationToken>()
@@ -240,6 +237,180 @@ public sealed class DiscordGatewayClientTests : IDisposable
);
}
[Fact]
public async Task ConnectAsync_OnceConnected_ItShouldStopSendingHeartbeatsIfTheyAreNotAcknowledgedAndAttemptToResume()
{
_mockDiscordRestClient
.Setup(static x => x.GetGatewayUrlAsync(It.IsAny<CancellationToken>()))
.ReturnsAsync("wss://gateway.discord.gg");
var initialSocketState = WebSocketState.Closed;
var initialWebSocket = new Mock<IWebSocket>();
initialWebSocket
.Setup(static x => x.State)
.Returns(() => initialSocketState);
initialWebSocket
.Setup(static x => x.ConnectAsync(It.IsAny<Uri>(), It.IsAny<CancellationToken>()))
.Callback(() => initialSocketState = WebSocketState.Open)
.Returns(Task.CompletedTask);
initialWebSocket
.Setup(static x => x.CloseAsync(It.IsAny<WebSocketCloseStatus>(), It.IsAny<string>(), It.IsAny<CancellationToken>()))
.Callback(() => initialSocketState = WebSocketState.Closed)
.Returns(Task.CompletedTask);
var resumingSocketState = WebSocketState.Closed;
var resumingWebSocket = new Mock<IWebSocket>();
resumingWebSocket
.Setup(static x => x.State)
.Returns(() => resumingSocketState);
resumingWebSocket
.Setup(static x => x.ConnectAsync(It.IsAny<Uri>(), It.IsAny<CancellationToken>()))
.Callback(() => resumingSocketState = WebSocketState.Open)
.Returns(Task.CompletedTask);
resumingWebSocket
.Setup(static x => x.CloseAsync(It.IsAny<WebSocketCloseStatus>(), It.IsAny<string>(), It.IsAny<CancellationToken>()))
.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<ArraySegment<byte>>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<WebSocketMessageType>(m => m == WebSocketMessageType.Text),
true,
It.IsAny<CancellationToken>()
),
Times.AtMostOnce
);
var expectedUri = new Uri($"{resumeGatewayUrl}/?v=10&encoding=json");
resumingWebSocket
.Verify(
x => x.ConnectAsync(
It.Is<Uri>(uri => uri.Equals(expectedUri)),
It.IsAny<CancellationToken>()
),
Times.Once
);
}
[Fact]
public async Task ConnectAsync_OnceConnected_ItShouldContinueSendingHeartbeatsIfTheyAreAcknowledged()
{
_mockDiscordRestClient
.Setup(static x => x.GetGatewayUrlAsync(It.IsAny<CancellationToken>()))
.ReturnsAsync("wss://gateway.discord.gg");
var mockWebSocket = new Mock<IWebSocket>();
var socketState = WebSocketState.Closed;
mockWebSocket
.Setup(static x => x.State)
.Returns(() => socketState);
mockWebSocket
.Setup(static x => x.ConnectAsync(It.IsAny<Uri>(), It.IsAny<CancellationToken>()))
.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<ArraySegment<byte>>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b.Array!)),
It.Is<WebSocketMessageType>(m => m == WebSocketMessageType.Text),
true,
It.IsAny<CancellationToken>()
),
Times.AtLeast(2)
);
}
private static void SetupReceiveMessageSequence(
Mock<IWebSocket> mockWebSocket,
Queue<(WebSocketReceiveResult, byte[])> messageQueue
@@ -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();
}
}
@@ -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);
@@ -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<DiscordEvent>(
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;
@@ -2,4 +2,8 @@ namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record HeartbeatAckDiscordEvent : DiscordEvent
{
public HeartbeatAckDiscordEvent()
{
OpCode = DiscordOpCodes.HeartbeatAck;
}
}
@@ -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();
}