feat: add rest client to get gateway url

This commit is contained in:
Stevan Freeborn
2025-05-13 15:55:50 -05:00
parent a25e3194db
commit df9b3cd7fc
7 changed files with 188 additions and 20 deletions
@@ -0,0 +1,73 @@
using System.Net;
using Microsoft.Extensions.Logging;
using Moq;
using RichardSzalay.MockHttp;
namespace StevesBot.Worker.Tests.Unit;
public sealed class DiscordRestClientTests : IDisposable
{
private const string BaseUrl = "https://discord.com/api/v10";
private static string GatewayUrl => $"{BaseUrl}/gateway";
private readonly Mock<ILogger<DiscordRestClient>> _loggerMock = new();
private readonly MockHttpMessageHandler _mockHttpMessageHandler;
private readonly DiscordRestClient _discordRestClient;
public DiscordRestClientTests()
{
_mockHttpMessageHandler = new MockHttpMessageHandler();
var httpClient = _mockHttpMessageHandler.ToHttpClient();
httpClient.BaseAddress = new Uri("https://discord.com/api/v10/");
_discordRestClient = new DiscordRestClient(_loggerMock.Object, httpClient);
}
[Fact]
public async Task GetGatewayUrlAsync_WhenRequestFails_ItShouldThrowException()
{
_mockHttpMessageHandler
.When(GatewayUrl)
.Respond(HttpStatusCode.InternalServerError);
var act = async () => await _discordRestClient.GetGatewayUrlAsync(CancellationToken.None);
await act.Should().ThrowAsync<DiscordRestClientException>();
}
[Fact]
public async Task GetGatewayUrlAsync_WhenResponseIsNull_ItShouldThrowException()
{
_mockHttpMessageHandler
.When(GatewayUrl)
.Respond(HttpStatusCode.OK, "application/json", "null");
var act = async () => await _discordRestClient.GetGatewayUrlAsync(CancellationToken.None);
await act.Should().ThrowAsync<DiscordRestClientException>();
}
[Fact]
public async Task GetGatewayUrlAsync_WhenRequestIsSuccessful_ItShouldReturnGatewayUrl()
{
var expectedUrl = "test";
var jsonResponse = $@"{{
""url"": ""{expectedUrl}""
}}";
_mockHttpMessageHandler
.When(GatewayUrl)
.Respond(HttpStatusCode.OK, "application/json", jsonResponse);
var result = await _discordRestClient.GetGatewayUrlAsync(CancellationToken.None);
result.Should().Be(expectedUrl);
}
public void Dispose()
{
_mockHttpMessageHandler.Dispose();
}
}
@@ -5,44 +5,73 @@ internal class DiscordGatewayClient : IDiscordGatewayClient
private readonly DiscordGatewayClientOptions _options; private readonly DiscordGatewayClientOptions _options;
private readonly IWebSocketFactory _webSocketFactory; private readonly IWebSocketFactory _webSocketFactory;
private readonly ILogger<DiscordGatewayClient> _logger; private readonly ILogger<DiscordGatewayClient> _logger;
private readonly JsonSerializerOptions _jsonSerializerOptions = new() private readonly IDiscordRestClient _discordRestClient;
{ // private readonly JsonSerializerOptions _jsonSerializerOptions = new()
PropertyNamingPolicy = JsonNamingPolicy.CamelCase, // {
ReferenceHandler = ReferenceHandler.IgnoreCycles, // PropertyNamingPolicy = JsonNamingPolicy.CamelCase,
Converters = // ReferenceHandler = ReferenceHandler.IgnoreCycles,
{ // Converters =
new DiscordEventConverter(), // {
}, // new DiscordEventConverter(),
}; // },
// };
private readonly AsyncLock _lock = new(); private readonly AsyncLock _lock = new();
private string _gatewayUrl = string.Empty; private string _gatewayUrl = string.Empty;
private IWebSocket? _webSocket; private IWebSocket? _webSocket;
private Task? _receiveTask; private Task? _receiveTask;
private DateTime _timeLastHeartbeatSent = DateTime.MinValue; // private DateTime _timeLastHeartbeatSent = DateTime.MinValue;
private DateTime _timeLastHeartbeatAcknowledged = DateTime.MinValue; // private DateTime _timeLastHeartbeatAcknowledged = DateTime.MinValue;
private CancellationTokenSource? _heartbeatCts; // private CancellationTokenSource? _heartbeatCts;
private Task? _heartbeatTask; // private Task? _heartbeatTask;
private string _sessionId = string.Empty; // private string _sessionId = string.Empty;
private string _resumeGatewayUrl = string.Empty; // private string _resumeGatewayUrl = string.Empty;
public DiscordGatewayClient( public DiscordGatewayClient(
DiscordGatewayClientOptions options, DiscordGatewayClientOptions options,
IWebSocketFactory webSocketFactory, IWebSocketFactory webSocketFactory,
ILogger<DiscordGatewayClient> logger ILogger<DiscordGatewayClient> logger,
IDiscordRestClient discordRestClient
) )
{ {
_options = options ?? throw new ArgumentNullException(nameof(options)); _options = options ?? throw new ArgumentNullException(nameof(options));
_webSocketFactory = webSocketFactory ?? throw new ArgumentNullException(nameof(webSocketFactory)); _webSocketFactory = webSocketFactory ?? throw new ArgumentNullException(nameof(webSocketFactory));
_logger = logger ?? throw new ArgumentNullException(nameof(logger)); _logger = logger ?? throw new ArgumentNullException(nameof(logger));
_discordRestClient = discordRestClient ?? throw new ArgumentNullException(nameof(discordRestClient));
}
public async Task ConnectAsync(CancellationToken cancellationToken)
{
using (await _lock.LockAsync(cancellationToken))
{
if (string.IsNullOrEmpty(_gatewayUrl))
{
_gatewayUrl = await _discordRestClient.GetGatewayUrlAsync(cancellationToken);
}
_webSocket = _webSocketFactory.Create();
var uri = new Uri(_gatewayUrl);
await _webSocket.ConnectAsync(uri, cancellationToken);
_logger.LogInformation("Connected to Discord Gateway at {GatewayUrl}", _gatewayUrl);
_receiveTask = ReceiveMessagesAsync(cancellationToken);
}
}
private static Task ReceiveMessagesAsync(CancellationToken cancellationToken)
{
return Task.CompletedTask;
} }
public void Dispose() public void Dispose()
{ {
_heartbeatCts?.Cancel(); // _heartbeatCts?.Cancel();
_heartbeatCts?.Dispose(); // _heartbeatCts?.Dispose();
_heartbeatTask?.Dispose(); // _heartbeatTask?.Dispose();
_receiveTask?.Dispose(); _receiveTask?.Dispose();
_webSocket?.Dispose(); _webSocket?.Dispose();
_lock.Dispose();
} }
} }
@@ -0,0 +1,43 @@
namespace StevesBot.Worker.Discord;
internal class DiscordRestClient : IDiscordRestClient
{
private readonly ILogger<DiscordRestClient> _logger;
private readonly HttpClient _httpClient;
public DiscordRestClient(
ILogger<DiscordRestClient> logger,
HttpClient httpClient
)
{
_logger = logger ?? throw new ArgumentNullException(nameof(logger));
_httpClient = httpClient ?? throw new ArgumentNullException(nameof(httpClient));
}
public async Task<string> GetGatewayUrlAsync(CancellationToken cancellationToken)
{
var uri = new Uri("gateway", UriKind.Relative);
var response = await _httpClient.GetAsync(uri, cancellationToken);
if (!response.IsSuccessStatusCode)
{
_logger.LogError("Failed to get gateway URL: {StatusCode}", response.StatusCode);
throw new DiscordRestClientException("Failed to get gateway URL.");
}
var gatewayResponse = await response.Content.ReadFromJsonAsync<GatewayResponse>(cancellationToken);
if (gatewayResponse is null)
{
_logger.LogError("Failed to deserialize gateway response.");
throw new DiscordRestClientException("Failed to deserialize gateway response.");
}
return gatewayResponse.Url;
}
}
internal record GatewayResponse(
[property: JsonPropertyName("url")] string Url
);
@@ -0,0 +1,16 @@
namespace StevesBot.Worker.Discord;
internal class DiscordRestClientException : Exception
{
public DiscordRestClientException()
{
}
public DiscordRestClientException(string message) : base(message)
{
}
public DiscordRestClientException(string message, Exception innerException) : base(message, innerException)
{
}
}
@@ -2,4 +2,5 @@ namespace StevesBot.Worker.Discord;
internal interface IDiscordGatewayClient : IDisposable internal interface IDiscordGatewayClient : IDisposable
{ {
Task ConnectAsync(CancellationToken cancellationToken);
} }
@@ -0,0 +1,6 @@
namespace StevesBot.Worker.Discord;
internal interface IDiscordRestClient
{
Task<string> GetGatewayUrlAsync(CancellationToken cancellationToken);
}
+1 -1
View File
@@ -1,8 +1,8 @@
global using System.Net.Http.Json;
global using System.Net.WebSockets; global using System.Net.WebSockets;
global using System.Text.Json; global using System.Text.Json;
global using System.Text.Json.Serialization; global using System.Text.Json.Serialization;
global using StevesBot.Worker.Discord.Events; global using StevesBot.Worker.Discord.Events;
global using StevesBot.Worker.Threading; global using StevesBot.Worker.Threading;
global using StevesBot.Worker.WebSockets; global using StevesBot.Worker.WebSockets;