refactor: organize things a bit

This commit is contained in:
Stevan Freeborn
2025-05-18 12:16:14 -05:00
parent db221334ca
commit 9af7522c29
37 changed files with 229 additions and 209 deletions
@@ -1,37 +0,0 @@
namespace StevesBot.Worker.Discord.Events;
internal sealed record MessageCreateDiscordEvent : DispatchDiscordEvent
{
[JsonPropertyName("d")]
public new MessageCreateData Data { get; init; } = new MessageCreateData();
public bool IsMessageType(int messageType)
{
return Data.Type == messageType;
}
}
internal sealed record MessageCreateData
{
[JsonPropertyName("id")]
public string Id { get; init; } = string.Empty;
[JsonPropertyName("type")]
public int Type { get; init; }
[JsonPropertyName("channel_id")]
public string ChannelId { get; init; } = string.Empty;
[JsonPropertyName("guild_id")]
public string GuildId { get; init; } = string.Empty;
[JsonPropertyName("author")]
public DiscordUser Author { get; init; } = new DiscordUser();
}
internal sealed record DiscordUser
{
[JsonPropertyName("id")]
public string Id { get; init; } = string.Empty;
}
@@ -0,0 +1,50 @@
namespace StevesBot.Worker.Discord;
internal static class Extensions
{
public static IServiceCollection AddDiscordRestClient(this IServiceCollection services)
{
services
.AddHttpClient<IDiscordRestClient, DiscordRestClient>(static (sp, c) =>
{
var discordOptions = sp.GetRequiredService<DiscordClientOptions>();
c.BaseAddress = new Uri(discordOptions.ApiUrl);
c.DefaultRequestHeaders.Authorization = new("Bot", discordOptions.AppToken);
c.DefaultRequestHeaders.Add("User-Agent", $"DiscordBot (https://github.com/StevanFreeborn/steves-bot, 0.0.0)");
})
.AddStandardResilienceHandler();
return services;
}
public static IServiceCollection AddDiscordGatewayClient(this IServiceCollection services, Action<IDiscordGatewayClient>? configure)
{
services.AddSingleton<IDiscordGatewayClient>(sp =>
{
var discordOptions = sp.GetRequiredService<DiscordClientOptions>();
var discordRestClient = sp.GetRequiredService<IDiscordRestClient>();
var logger = sp.GetRequiredService<ILogger<DiscordGatewayClient>>();
var socketFactory = sp.GetRequiredService<IWebSocketFactory>();
var timeProvider = sp.GetRequiredService<TimeProvider>();
var serviceScopeFactory = sp.GetRequiredService<IServiceScopeFactory>();
var client = new DiscordGatewayClient(
discordOptions,
socketFactory,
logger,
discordRestClient,
timeProvider,
serviceScopeFactory
);
if (configure is not null)
{
configure(client);
}
return client;
});
return services;
}
}
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway;
internal static class DiscordCloseCodes
{
@@ -1,6 +1,4 @@
using System.Text;
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway;
internal sealed class DiscordGatewayClient : IDiscordGatewayClient
{
@@ -129,8 +127,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
private async Task StartReceiveMessagesAsync(CancellationToken cancellationToken)
{
await CancelReceiveMessagesTaskAsync(cancellationToken);
var newReceiveCts = new CancellationTokenSource();
var newReceiveLinkedCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, newReceiveCts.Token);
@@ -164,10 +160,8 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
}
var canResume = IsResumableCloseCode(result.CloseStatus);
var closeStatus = canResume ? WebSocketCloseStatus.MandatoryExtension : WebSocketCloseStatus.NormalClosure;
await SetCanResumeAsync(canResume, _linkedReceiveMessageCts.Token);
await CloseIfOpenAsync(closeStatus, result.CloseStatusDescription, _linkedReceiveMessageCts.Token);
_logger.LogWarning("Reconnecting because of close message: {CloseStatus} - {CloseStatusDescription}", result.CloseStatus, result.CloseStatusDescription);
@@ -213,12 +207,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
_logger.LogWarning("Closing connection and invalidating session.");
await CloseIfOpenAsync(
WebSocketCloseStatus.NormalClosure,
"WebSocket error. Closing connection and invalidating session.",
newReceiveLinkedCts.Token
);
await SetCanResumeAsync(false, newReceiveLinkedCts.Token);
_logger.LogWarning("Reconnecting because of error in receive message task");
@@ -293,12 +281,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
await SetCanResumeAsync(true, cancellationToken);
await CloseIfOpenAsync(
WebSocketCloseStatus.MandatoryExtension,
"Reconnect event received",
cancellationToken
);
_logger.LogWarning("Reconnecting because of reconnect event");
await ReconnectAsync(cancellationToken);
@@ -311,24 +293,12 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
_logger.LogWarning("Session is resumable. Closing connection without invalidating session.");
await SetCanResumeAsync(true, cancellationToken);
await CloseIfOpenAsync(
WebSocketCloseStatus.MandatoryExtension,
"Invalid session event received",
cancellationToken
);
}
else
{
_logger.LogWarning("Session is not resumable. Closing connection and invalidating session.");
await SetCanResumeAsync(false, cancellationToken);
await CloseIfOpenAsync(
WebSocketCloseStatus.NormalClosure,
"Invalid session event received",
cancellationToken
);
}
await ReconnectAsync(cancellationToken);
@@ -341,8 +311,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
private async Task StartHeartbeatAsync(CancellationToken cancellationToken)
{
await CancelHeartbeatTaskAsync(cancellationToken);
var newHeartbeatCts = new CancellationTokenSource();
var newLinkedCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, newHeartbeatCts.Token);
@@ -365,14 +333,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
{
if (IsWebSocketOpen())
{
_logger.LogWarning("Heartbeat not acknowledged. Closing WebSocket.");
await CloseAsync(
WebSocketCloseStatus.ProtocolError,
"Heartbeat not acknowledged",
_linkedHeartbeatCts.Token
);
await SetCanResumeAsync(true, _linkedHeartbeatCts.Token);
}
@@ -397,14 +357,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
{
_logger.LogError(ex, "Error in heartbeat task: {Message}", ex.Message);
_logger.LogWarning("Closing connection and invalidating session.");
await CloseIfOpenAsync(
WebSocketCloseStatus.NormalClosure,
"Heartbeat error. Closing connection and invalidating session.",
newLinkedCts.Token
);
_logger.LogWarning("Reconnecting because of error in heartbeat task");
await ReconnectAsync(cancellationToken);
@@ -415,6 +367,17 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
private async Task ReconnectAsync(CancellationToken cancellationToken)
{
await CancelReceiveMessagesTaskAsync(cancellationToken);
await CancelHeartbeatTaskAsync(cancellationToken);
var closeStatus = _canResume ? WebSocketCloseStatus.MandatoryExtension : WebSocketCloseStatus.NormalClosure;
await CloseIfOpenAsync(
closeStatus,
"Reconnecting",
cancellationToken
);
_webSocket?.Dispose();
if (_canResume)
@@ -638,21 +601,6 @@ internal sealed class DiscordGatewayClient : IDiscordGatewayClient
await _webSocket.CloseAsync(closeStatus, statusDescription, cancellationToken);
}
private async Task CloseAsync(WebSocketCloseStatus closeStatus, string? statusDescription, CancellationToken cancellationToken)
{
if (_webSocket is null)
{
throw new DiscordGatewayClientException("WebSocket is not set. Cannot close.");
}
if (_webSocket.State is not WebSocketState.Open)
{
throw new DiscordGatewayClientException("WebSocket is not open. Cannot close.");
}
await _webSocket.CloseAsync(closeStatus, statusDescription, cancellationToken);
}
private async Task<WebSocketReceiveResult> ReceiveMessageAsync(
ArraySegment<byte> segment,
CancellationToken cancellationToken
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway;
internal sealed class DiscordGatewayClientException : Exception
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway;
internal static class DiscordIntents
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal record DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed class DiscordEventConverter : JsonConverter<DiscordEvent>
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal static class DiscordEventTypes
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal static class DiscordMessageTypes
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal static class DiscordOpCodes
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal record DispatchDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record HeartbeatAckDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record HeartbeatDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record HelloDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record IdentifyDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record InvalidSessionDiscordEvent : DiscordEvent
{
@@ -0,0 +1,12 @@
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record MessageCreateDiscordEvent : DispatchDiscordEvent
{
[JsonPropertyName("d")]
public new DiscordMessage Data { get; init; } = new DiscordMessage();
public bool IsMessageType(int messageType)
{
return Data.Type == messageType;
}
}
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record ReadyDiscordEvent : DispatchDiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record ReconnectDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record ResumeDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord.Events;
namespace StevesBot.Worker.Discord.Gateway.Events;
internal sealed record UpdatePresenceDiscordEvent : DiscordEvent
{
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Gateway;
internal interface IDiscordGatewayClient : IDisposable
{
@@ -1,7 +0,0 @@
namespace StevesBot.Worker.Discord;
internal interface IDiscordRestClient
{
Task<string> GetGatewayUrlAsync(CancellationToken cancellationToken);
Task<MessageCreateData> CreateMessageAsync(string channelId, CreateMessageRequest request, CancellationToken cancellationToken);
}
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Rest;
internal sealed class DiscordRestClient : IDiscordRestClient
{
@@ -14,7 +14,6 @@ internal sealed class DiscordRestClient : IDiscordRestClient
_httpClient = httpClient ?? throw new ArgumentNullException(nameof(httpClient));
}
public async Task<string> GetGatewayUrlAsync(CancellationToken cancellationToken)
{
var gatewayEndpoint = new Uri("gateway", UriKind.Relative);
@@ -37,7 +36,7 @@ internal sealed class DiscordRestClient : IDiscordRestClient
return gatewayResponse.Url;
}
public async Task<MessageCreateData> CreateMessageAsync(string channelId, CreateMessageRequest request, CancellationToken cancellationToken)
public async Task<DiscordMessage> CreateMessageAsync(string channelId, CreateMessageRequest request, CancellationToken cancellationToken)
{
var channelEndpoint = new Uri($"channels/{channelId}/messages", UriKind.Relative);
var response = await _httpClient.PostAsJsonAsync(channelEndpoint, request, cancellationToken);
@@ -48,37 +47,14 @@ internal sealed class DiscordRestClient : IDiscordRestClient
throw new DiscordRestClientException("Failed to create message.");
}
var messageCreateData = await response.Content.ReadFromJsonAsync<MessageCreateData>(cancellationToken);
var discordMessage = await response.Content.ReadFromJsonAsync<DiscordMessage>(cancellationToken);
if (messageCreateData is null)
if (discordMessage is null)
{
_logger.LogError("Failed to deserialize message create response.");
throw new DiscordRestClientException("Failed to deserialize message create response.");
}
return messageCreateData;
return discordMessage;
}
}
internal sealed record GatewayResponse(
[property: JsonPropertyName("url")] string Url
);
internal sealed record CreateMessageRequest(
[property: JsonPropertyName("content")] string Content,
[property: JsonPropertyName("message_reference")] MessageReference? MessageReference
);
internal sealed record MessageReference(
[property: JsonPropertyName("type")] int Type,
[property: JsonPropertyName("message_id")] string MessageId,
[property: JsonPropertyName("channel_id")] string ChannelId,
[property: JsonPropertyName("guild_id")] string GuildId,
[property: JsonPropertyName("fail_if_not_exists")] bool FailIfNotExists
);
internal static class MessageReferenceTypes
{
public const int Default = 0;
public const int Forward = 1;
}
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Rest;
internal sealed class DiscordRestClientException : Exception
{
@@ -0,0 +1,7 @@
namespace StevesBot.Worker.Discord.Rest;
internal interface IDiscordRestClient
{
Task<string> GetGatewayUrlAsync(CancellationToken cancellationToken);
Task<DiscordMessage> CreateMessageAsync(string channelId, CreateMessageRequest request, CancellationToken cancellationToken);
}
@@ -0,0 +1,6 @@
namespace StevesBot.Worker.Discord.Rest.Requests;
internal sealed record CreateMessageRequest(
[property: JsonPropertyName("content")] string Content,
[property: JsonPropertyName("message_reference")] MessageReference? MessageReference
);
@@ -0,0 +1,5 @@
namespace StevesBot.Worker.Discord.Rest.Responses;
internal sealed record GatewayResponse(
[property: JsonPropertyName("url")] string Url
);
@@ -1,4 +1,4 @@
namespace StevesBot.Worker.Discord;
namespace StevesBot.Worker.Discord.Shared;
internal sealed class DiscordClientOptions
{
@@ -0,0 +1,19 @@
namespace StevesBot.Worker.Discord.Shared;
internal sealed record DiscordMessage
{
[JsonPropertyName("id")]
public string Id { get; init; } = string.Empty;
[JsonPropertyName("type")]
public int Type { get; init; }
[JsonPropertyName("channel_id")]
public string ChannelId { get; init; } = string.Empty;
[JsonPropertyName("guild_id")]
public string GuildId { get; init; } = string.Empty;
[JsonPropertyName("author")]
public DiscordUser Author { get; init; } = new DiscordUser();
}
@@ -0,0 +1,15 @@
namespace StevesBot.Worker.Discord.Shared;
internal sealed record MessageReference(
[property: JsonPropertyName("type")] int Type,
[property: JsonPropertyName("message_id")] string MessageId,
[property: JsonPropertyName("channel_id")] string ChannelId,
[property: JsonPropertyName("guild_id")] string GuildId,
[property: JsonPropertyName("fail_if_not_exists")] bool FailIfNotExists
);
internal static class MessageReferenceTypes
{
public const int Default = 0;
public const int Forward = 1;
}
@@ -0,0 +1,7 @@
namespace StevesBot.Worker.Discord.Shared;
internal sealed record DiscordUser
{
[JsonPropertyName("id")]
public string Id { get; init; } = string.Empty;
}
@@ -0,0 +1,53 @@
using System.Globalization;
using System.Text;
namespace StevesBot.Worker.Handlers;
internal static class WelcomeMessageHandler
{
public static async Task HandleAsync(
DiscordEvent discordEvent,
IServiceProvider serviceProvider,
CancellationToken cancellationToken
)
{
var discordRestClient = serviceProvider.GetRequiredService<IDiscordRestClient>();
var logger = serviceProvider.GetRequiredService<ILogger<DiscordGatewayClient>>();
if (discordEvent is not MessageCreateDiscordEvent mcde || mcde.IsMessageType(DiscordMessageTypes.UserJoin) == false)
{
return;
}
logger.LogInformation("Received user join message for user: {UserId}", mcde.Data.Author.Id);
var welcomeMessage = GetWelcomeMessage(mcde.Data.Author.Id);
var request = new CreateMessageRequest(
Content: welcomeMessage,
MessageReference: new(
Type: MessageReferenceTypes.Default,
MessageId: mcde.Data.Id,
ChannelId: mcde.Data.ChannelId,
GuildId: mcde.Data.GuildId,
FailIfNotExists: false
)
);
var message = await discordRestClient.CreateMessageAsync(mcde.Data.ChannelId, request, cancellationToken);
logger.LogInformation("Created welcome message with Id: {MessageId} for user: {UserId}", message.Id, mcde.Data.Author.Id);
}
private static string GetWelcomeMessage(string userId)
{
var builder = new StringBuilder();
builder.AppendLine(CultureInfo.InvariantCulture, $"Hi there <@{userId}>! 👋🏻 Welcome to Stevan's server.");
builder.AppendLine();
builder.AppendLine("We're so stoked to have you join our community here. Feel free to jump right in, tell us a bit about yourself, and explore all the different channels.");
builder.AppendLine();
builder.AppendLine("If you have any questions at all, don't hesitate to ask. We're a friendly bunch and always happy to help out. Glad you're here!");
return builder.ToString();
}
}
+4 -17
View File
@@ -1,7 +1,3 @@
using System.Net.Http.Headers;
using Microsoft.Extensions.Options;
var builder = Host.CreateApplicationBuilder(args);
builder.Services.Configure<HostOptions>(static options => options.ShutdownTimeout = TimeSpan.FromSeconds(30));
@@ -18,20 +14,11 @@ builder.Services.AddSingleton(static sp =>
builder.Services.AddSingleton<IWebSocketFactory, WebSocketFactory>();
builder.Services.AddSingleton(TimeProvider.System);
builder.Services
.AddHttpClient<IDiscordRestClient, DiscordRestClient>(static (sp, c) =>
{
var discordOptions = sp.GetRequiredService<DiscordClientOptions>();
c.BaseAddress = new Uri(discordOptions.ApiUrl);
c.DefaultRequestHeaders.Authorization = new("Bot", discordOptions.AppToken);
c.DefaultRequestHeaders.Add("User-Agent", $"DiscordBot (https://github.com/StevanFreeborn/steves-bot, 0.0.0)");
})
.AddStandardResilienceHandler();
builder.Services.AddDiscordRestClient();
// TODO: Create an extension method that allows adding
// the discord gateway client and allows me to configure
// event handlers
builder.Services.AddSingleton<IDiscordGatewayClient, DiscordGatewayClient>();
builder.Services.AddDiscordGatewayClient(static (client) =>
client.On(DiscordEventTypes.MessageCreate, WelcomeMessageHandler.HandleAsync)
);
builder.Services.AddHostedService<Worker>();
+11 -3
View File
@@ -1,12 +1,20 @@
global using System.Net.Http.Json;
global using System.Net.WebSockets;
global using System.Reflection;
global using System.Text;
global using System.Text.Json;
global using System.Text.Json.Serialization;
global using StevesBot.Worker.Discord.Events;
global using StevesBot.Worker.Threading;
global using StevesBot.Worker.WebSockets;
global using Microsoft.Extensions.Options;
global using StevesBot.Worker;
global using StevesBot.Worker.Discord;
global using StevesBot.Worker.Discord.Gateway;
global using StevesBot.Worker.Discord.Gateway.Events;
global using StevesBot.Worker.Discord.Rest;
global using StevesBot.Worker.Discord.Rest.Requests;
global using StevesBot.Worker.Discord.Rest.Responses;
global using StevesBot.Worker.Discord.Shared;
global using StevesBot.Worker.Handlers;
global using StevesBot.Worker.Threading;
global using StevesBot.Worker.WebSockets;
-29
View File
@@ -15,35 +15,6 @@ internal class Worker : IHostedService
public Task StartAsync(CancellationToken cancellationToken)
{
_logger.LogInformation("Connecting Discord Gateway Client");
_discordGatewayClient.On(DiscordEventTypes.MessageCreate, static async (discordEvent, sp, cancellationToken) =>
{
var discordRestClient = sp.GetRequiredService<IDiscordRestClient>();
var logger = sp.GetRequiredService<ILogger<DiscordGatewayClient>>();
if (discordEvent is not MessageCreateDiscordEvent mcde || mcde.IsMessageType(DiscordMessageTypes.UserJoin) == false)
{
return;
}
logger.LogInformation("Received user join message for user: {UserId}", mcde.Data.Author.Id);
var request = new CreateMessageRequest(
Content: $"Welcome to the server <@{mcde.Data.Author.Id}>! We're glad to have you here.",
MessageReference: new(
Type: MessageReferenceTypes.Default,
MessageId: mcde.Data.Id,
ChannelId: mcde.Data.ChannelId,
GuildId: mcde.Data.GuildId,
FailIfNotExists: false
)
);
var message = await discordRestClient.CreateMessageAsync(mcde.Data.ChannelId, request, cancellationToken);
logger.LogInformation("Created welcome message with Id: {MessageId} for user: {UserId}", message.Id, mcde.Data.Author.Id);
});
return _discordGatewayClient.ConnectAsync(cancellationToken);
}