diff --git a/Morpheus.Tests/ChannelServiceConcurrencyTests.cs b/Morpheus.Tests/ChannelServiceConcurrencyTests.cs new file mode 100644 index 0000000..40145cc --- /dev/null +++ b/Morpheus.Tests/ChannelServiceConcurrencyTests.cs @@ -0,0 +1,56 @@ +using Microsoft.Data.Sqlite; +using Microsoft.EntityFrameworkCore; +using Morpheus.Database; +using Morpheus.Database.Models; +using Morpheus.Services; + +namespace Morpheus.Tests; + +public class ChannelServiceConcurrencyTests +{ + [Fact] + public async Task TryGetCreateChannel_WhenAnotherHandlerCreatesChannel_ReturnsPersistedChannel() + { + await using SqliteConnection connection = new("Data Source=:memory:"); + await connection.OpenAsync(); + + DbContextOptions options = new DbContextOptionsBuilder() + .UseSqlite(connection) + .Options; + await using (DB setup = new(options)) + await setup.Database.EnsureCreatedAsync(); + + await using RacingDb db = new(options); + db.InsertCompetingChannelOnNextSave = true; + ChannelService service = new(db, new LogsService(new LogQueue())); + + Channel result = await service.TryGetCreateChannel(123, "current-name"); + + Assert.Equal((ulong)123, result.DiscordId); + Assert.Equal("current-name", result.Name); + Assert.Equal(1, await db.Channels.CountAsync()); + } + + private sealed class RacingDb(DbContextOptions options) : DB(options) + { + public bool InsertCompetingChannelOnNextSave { get; set; } + + public override async Task SaveChangesAsync(CancellationToken cancellationToken = default) + { + if (InsertCompetingChannelOnNextSave) + { + InsertCompetingChannelOnNextSave = false; + ChangeTracker.Clear(); + Channels.Add(new Channel + { + DiscordId = 123, + Name = "stale-name" + }); + await base.SaveChangesAsync(cancellationToken); + throw new DbUpdateException("Simulated concurrent unique-key conflict."); + } + + return await base.SaveChangesAsync(cancellationToken); + } + } +} diff --git a/Services/ChannelService.cs b/Services/ChannelService.cs index 56f5d7d..68c76ef 100644 --- a/Services/ChannelService.cs +++ b/Services/ChannelService.cs @@ -28,7 +28,27 @@ public async Task TryGetCreateChannel(ulong discordId, string name) }; await dbContext.Channels.AddAsync(channel); - await dbContext.SaveChangesAsync(); + try + { + await dbContext.SaveChangesAsync(); + } + catch (DbUpdateException) + { + // Another handler may have created the same Discord channel after our initial lookup. + // Clear the failed insert and use the row protected by the unique DiscordId index. + dbContext.ChangeTracker.Clear(); + Channel? concurrentChannel = await dbContext.Channels.FirstOrDefaultAsync(c => c.DiscordId == discordId); + if (concurrentChannel == null) + throw; + + if (concurrentChannel.Name != name) + { + concurrentChannel.Name = name; + await dbContext.SaveChangesAsync(); + } + + return concurrentChannel; + } logsService.Log($"New channel created {name}", Discord.LogSeverity.Verbose);