fix: correct locking issue that was preventing progress

This commit is contained in:
Stevan Freeborn
2025-05-14 16:45:44 -05:00
parent 3458890960
commit 16aa4b5775
9 changed files with 147 additions and 52 deletions
+21
View File
@@ -0,0 +1,21 @@
{
"version": "0.2.0",
"configurations": [
{
"name": "Worker",
"type": "coreclr",
"request": "launch",
"preLaunchTask": "build",
"program": "${workspaceFolder}/src/StevesBot.Worker/bin/Debug/net9.0/StevesBot.Worker.dll",
"args": [],
"cwd": "${workspaceFolder}/src/StevesBot.Worker",
"console": "internalConsole",
"stopAtEntry": false
},
{
"name": ".NET Core Attach",
"type": "coreclr",
"request": "attach"
}
]
}
+7
View File
@@ -0,0 +1,7 @@
{
"editor.formatOnSave": true,
"dotnet.defaultSolution": "src/StevesBot.sln",
"cSpell.words": [
"Steves"
],
}
+41
View File
@@ -0,0 +1,41 @@
{
"version": "2.0.0",
"tasks": [
{
"label": "build",
"command": "dotnet",
"type": "process",
"args": [
"build",
"${workspaceFolder}/src/StevesBot.sln",
"/property:GenerateFullPaths=true",
"/consoleloggerparameters:NoSummary;ForceNoAlign"
],
"problemMatcher": "$msCompile"
},
{
"label": "publish",
"command": "dotnet",
"type": "process",
"args": [
"publish",
"${workspaceFolder}/src/StevesBot.sln",
"/property:GenerateFullPaths=true",
"/consoleloggerparameters:NoSummary;ForceNoAlign"
],
"problemMatcher": "$msCompile"
},
{
"label": "watch",
"command": "dotnet",
"type": "process",
"args": [
"watch",
"run",
"--project",
"${workspaceFolder}/src/StevesBot.sln"
],
"problemMatcher": "$msCompile"
}
]
}
@@ -107,67 +107,80 @@ public sealed class DiscordGatewayClientTests : IDisposable
var messageQueue = new Queue<(WebSocketReceiveResult, byte[])>(); var messageQueue = new Queue<(WebSocketReceiveResult, byte[])>();
var heatbeatInterval = 1000; var heartbeatInterval = 1000;
var helloEvent = new var helloEventPayload = CreateEventPayload(new
{ {
op = 10, op = 10,
d = new d = new
{ {
heartbeat_interval = heatbeatInterval, heartbeat_interval = heartbeatInterval,
} }
}; });
var helloEventJson = JsonSerializer.Serialize(helloEvent); messageQueue.Enqueue((
var helloEventBytes = Encoding.UTF8.GetBytes(helloEventJson); new(helloEventPayload.Bytes.Length, WebSocketMessageType.Text, true),
helloEventPayload.Bytes
));
var heartbeatAck = new SetupReceiveMessageSequence(mockWebSocket, messageQueue);
_mockWebSocketFactory
.Setup(static x => x.Create())
.Returns(mockWebSocket.Object);
var cts = new CancellationTokenSource();
await _discordGatewayClient.ConnectAsync(cts.Token);
await Task.Delay((int)(heartbeatInterval * 1.5));
await cts.CancelAsync();
var expectedHeartbeatPayload = CreateEventPayload(new HeartbeatDiscordEvent(null));
mockWebSocket.Verify(
x => x.SendAsync(
It.Is<ArraySegment<byte>>(b => expectedHeartbeatPayload.Bytes.SequenceEqual(b)),
It.IsAny<WebSocketMessageType>(),
It.IsAny<bool>(),
It.IsAny<CancellationToken>()
),
Times.Once
);
}
private static void SetupReceiveMessageSequence(
Mock<IWebSocket> mockWebSocket,
Queue<(WebSocketReceiveResult, byte[])> messageQueue
)
{ {
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<ArraySegment<byte>>(), It.IsAny<WebSocketMessageType>(), It.IsAny<bool>(), It.IsAny<CancellationToken>()))
.Returns(Task.CompletedTask);
mockWebSocket mockWebSocket
.Setup(static x => x.ReceiveAsync(It.IsAny<ArraySegment<byte>>(), It.IsAny<CancellationToken>())) .Setup(static x => x.ReceiveAsync(It.IsAny<ArraySegment<byte>>(), It.IsAny<CancellationToken>()))
.Returns((ArraySegment<byte> buffer, CancellationToken token) => .Returns((ArraySegment<byte> buffer, CancellationToken token) =>
{ {
if (messageQueue.Count == 0) if (messageQueue.Count == 0)
{ {
# pragma warning disable CA2008 return Task.Delay(-1, token)
return Task.Delay(-1, token).ContinueWith(_ => new WebSocketReceiveResult(0, WebSocketMessageType.Text, true), token); .ContinueWith(
_ => new WebSocketReceiveResult(0, WebSocketMessageType.Text, true),
TaskScheduler.Default
);
} }
var (result, messageBytes) = messageQueue.Dequeue(); var (result, messageBytes) = messageQueue.Dequeue();
Array.Copy(messageBytes, 0, buffer.Array!, buffer.Offset, Math.Min(messageBytes.Length, buffer.Count)); Array.Copy(messageBytes, 0, buffer.Array!, buffer.Offset, Math.Min(messageBytes.Length, buffer.Count));
return Task.FromResult(result); return Task.FromResult(result);
}); });
}
_mockWebSocketFactory private static (byte[] Bytes, string Json) CreateEventPayload(object e)
.Setup(static x => x.Create()) {
.Returns(mockWebSocket.Object); var json = JsonSerializer.Serialize(e);
var bytes = Encoding.UTF8.GetBytes(json);
using var cts = new CancellationTokenSource(); return (bytes, json);
await _discordGatewayClient.ConnectAsync(CancellationToken.None);
await Task.Delay(heatbeatInterval + 30000);
mockWebSocket.Verify(
static x => x.SendAsync(It.IsAny<ArraySegment<byte>>(), It.IsAny<WebSocketMessageType>(), It.IsAny<bool>(), It.IsAny<CancellationToken>()),
Times.Once
);
} }
public void Dispose() public void Dispose()
{ {
_discordGatewayClient.Dispose();
} }
} }
@@ -2,7 +2,7 @@ using System.Text;
namespace StevesBot.Worker.Discord; namespace StevesBot.Worker.Discord;
internal class DiscordGatewayClient : IDiscordGatewayClient internal sealed class DiscordGatewayClient : IDiscordGatewayClient
{ {
private readonly DiscordClientOptions _options; private readonly DiscordClientOptions _options;
private readonly IWebSocketFactory _webSocketFactory; private readonly IWebSocketFactory _webSocketFactory;
@@ -80,6 +80,12 @@ internal class DiscordGatewayClient : IDiscordGatewayClient
} }
} }
if (_webSocket?.State is not WebSocketState.Open)
{
_logger.LogWarning("WebSocket is not open. Cannot receive messages.");
return;
}
// websocket message might be larger than the // websocket message might be larger than the
// size of the buffer so we need to loop until // size of the buffer so we need to loop until
// we receive the end of the message and // we receive the end of the message and
@@ -244,15 +250,12 @@ internal class DiscordGatewayClient : IDiscordGatewayClient
} }
private async Task SendJsonAsync(object data, CancellationToken cancellationToken) private async Task SendJsonAsync(object data, CancellationToken cancellationToken)
{
using (await _lock.LockAsync(cancellationToken))
{ {
if (_webSocket?.State is not WebSocketState.Open) if (_webSocket?.State is not WebSocketState.Open)
{ {
_logger.LogWarning("WebSocket is not open. Cannot send message."); _logger.LogWarning("WebSocket is not open. Cannot send message.");
return; return;
} }
}
try try
{ {
@@ -275,10 +278,20 @@ internal class DiscordGatewayClient : IDiscordGatewayClient
public void Dispose() public void Dispose()
{ {
_heartbeatCts?.Cancel(); _heartbeatCts?.Cancel();
_linkedCts?.Cancel();
if (_heartbeatTask is not null && _heartbeatTask.IsCompleted)
{
_heartbeatTask.Dispose();
}
if (_receiveTask is not null && _receiveTask.IsCompleted)
{
_receiveTask.Dispose();
}
_heartbeatCts?.Dispose(); _heartbeatCts?.Dispose();
_linkedCts?.Dispose(); _linkedCts?.Dispose();
_heartbeatTask?.Dispose();
_receiveTask?.Dispose();
_webSocket?.Dispose(); _webSocket?.Dispose();
_lock.Dispose(); _lock.Dispose();
} }
@@ -1,6 +1,6 @@
namespace StevesBot.Worker.Discord; namespace StevesBot.Worker.Discord;
internal class DiscordRestClient : IDiscordRestClient internal sealed class DiscordRestClient : IDiscordRestClient
{ {
private readonly ILogger<DiscordRestClient> _logger; private readonly ILogger<DiscordRestClient> _logger;
private readonly HttpClient _httpClient; private readonly HttpClient _httpClient;
@@ -1,6 +1,6 @@
namespace StevesBot.Worker.Discord; namespace StevesBot.Worker.Discord;
internal class DiscordRestClientException : Exception internal sealed class DiscordRestClientException : Exception
{ {
public DiscordRestClientException() public DiscordRestClientException()
{ {
@@ -1,5 +1,5 @@
namespace StevesBot.Worker.Discord.Events; namespace StevesBot.Worker.Discord.Events;
internal record HeartbeatAckDiscordEvent : DiscordEvent internal sealed record HeartbeatAckDiscordEvent : DiscordEvent
{ {
} }
@@ -1,6 +1,6 @@
namespace StevesBot.Worker.Discord.Events; namespace StevesBot.Worker.Discord.Events;
internal record HeartbeatDiscordEvent : DiscordEvent internal sealed record HeartbeatDiscordEvent : DiscordEvent
{ {
public HeartbeatDiscordEvent(int? sequence) public HeartbeatDiscordEvent(int? sequence)
{ {