diff --git a/src/SIL.Harmony.Tests/Adapter/CustomObjectAdapterTests.cs b/src/SIL.Harmony.Tests/Adapter/CustomObjectAdapterTests.cs index f94b37d..8014d54 100644 --- a/src/SIL.Harmony.Tests/Adapter/CustomObjectAdapterTests.cs +++ b/src/SIL.Harmony.Tests/Adapter/CustomObjectAdapterTests.cs @@ -194,6 +194,6 @@ await dataModel.AddChange(Guid.NewGuid(), myClass2.MyNumber.Should().Be(123.45m); myClass2.DeletedTime.Should().BeNull(); - dataModel.QueryLatest().Should().NotBeEmpty(); + dataModel.QueryLatest().ToBlockingEnumerable().Should().NotBeEmpty(); } } \ No newline at end of file diff --git a/src/SIL.Harmony.Tests/DataModelReferenceTests.cs b/src/SIL.Harmony.Tests/DataModelReferenceTests.cs index 9dc8232..5072362 100644 --- a/src/SIL.Harmony.Tests/DataModelReferenceTests.cs +++ b/src/SIL.Harmony.Tests/DataModelReferenceTests.cs @@ -2,9 +2,8 @@ using SIL.Harmony.Changes; using SIL.Harmony.Sample.Changes; using SIL.Harmony.Sample.Models; -using SIL.Harmony.Tests; -namespace Tests; +namespace SIL.Harmony.Tests; public class DataModelReferenceTests : DataModelTestBase { @@ -158,7 +157,7 @@ public async Task CanCreate2TagsWithTheSameNameOutOfOrder() var commitA = await WriteNextChange(SetTag(Guid.NewGuid(), tagText)); //represents someone syncing in a tag with the same name await WriteChangeBefore(commitA, SetTag(Guid.NewGuid(), tagText)); - DataModel.QueryLatest().Where(t => t.Text == tagText).Should().ContainSingle(); + DataModel.QueryLatest().ToBlockingEnumerable().Where(t => t.Text == tagText).Should().ContainSingle(); } [Fact] @@ -170,6 +169,6 @@ public async Task CanUpdateTagWithTheSameNameOutOfOrder() var commitA = await WriteNextChange(SetTag(Guid.NewGuid(), tagText)); //represents someone syncing in a tag with the same name await WriteNextChange(SetTag(renameTagId, tagText)); - DataModel.QueryLatest().Where(t => t.Text == tagText).Should().ContainSingle(); + DataModel.QueryLatest().ToBlockingEnumerable().Where(t => t.Text == tagText).Should().ContainSingle(); } } \ No newline at end of file diff --git a/src/SIL.Harmony.Tests/DataModelSimpleChanges.cs b/src/SIL.Harmony.Tests/DataModelSimpleChanges.cs index b545d4e..284d2cc 100644 --- a/src/SIL.Harmony.Tests/DataModelSimpleChanges.cs +++ b/src/SIL.Harmony.Tests/DataModelSimpleChanges.cs @@ -1,12 +1,10 @@ -using SIL.Harmony; using SIL.Harmony.Changes; using SIL.Harmony.Db; using SIL.Harmony.Sample.Changes; using SIL.Harmony.Sample.Models; -using SIL.Harmony.Tests; using Microsoft.EntityFrameworkCore; -namespace Tests; +namespace SIL.Harmony.Tests; public class DataModelSimpleChanges : DataModelTestBase { @@ -79,7 +77,7 @@ public async Task WriteMultipleCommits() await WriteNextChange(SetWord(Guid.NewGuid(), "change 3")); DbContext.Snapshots.Should().HaveCount(3); - DataModel.QueryLatest().Should().HaveCount(3); + DataModel.QueryLatest().ToBlockingEnumerable().Should().HaveCount(3); } [Fact] diff --git a/src/SIL.Harmony.Tests/DataModelTestBase.cs b/src/SIL.Harmony.Tests/DataModelTestBase.cs index ea4b305..09e5b2c 100644 --- a/src/SIL.Harmony.Tests/DataModelTestBase.cs +++ b/src/SIL.Harmony.Tests/DataModelTestBase.cs @@ -18,7 +18,6 @@ public class DataModelTestBase : IAsyncLifetime private readonly bool _performanceTest; public readonly DataModel DataModel; public readonly SampleDbContext DbContext; - internal readonly CrdtRepository CrdtRepository; protected readonly MockTimeProvider MockTimeProvider = new(); public DataModelTestBase(bool saveToDisk = false, bool alwaysValidate = true, @@ -47,7 +46,6 @@ public DataModelTestBase(SqliteConnection connection, bool alwaysValidate = true DbContext.Database.OpenConnection(); DbContext.Database.EnsureCreated(); DataModel = _services.GetRequiredService(); - CrdtRepository = _services.GetRequiredService(); } public DataModelTestBase ForkDatabase(bool alwaysValidate = true) @@ -170,6 +168,7 @@ public virtual Task InitializeAsync() public async Task DisposeAsync() { + await _services.DisposeAsync(); } diff --git a/src/SIL.Harmony.Tests/PersistExtraDataTests.cs b/src/SIL.Harmony.Tests/PersistExtraDataTests.cs index efc5eb4..35e8815 100644 --- a/src/SIL.Harmony.Tests/PersistExtraDataTests.cs +++ b/src/SIL.Harmony.Tests/PersistExtraDataTests.cs @@ -76,7 +76,7 @@ public async Task CanPersistExtraData() { var entityId = Guid.NewGuid(); var commit = await _dataModelTestBase.WriteNextChange(new CreateExtraDataModelChange(entityId)); - var extraDataModel = _dataModelTestBase.DataModel.QueryLatest().Should().ContainSingle().Subject; + var extraDataModel = _dataModelTestBase.DataModel.QueryLatest().ToBlockingEnumerable().Should().ContainSingle().Subject; extraDataModel.Id.Should().Be(entityId); extraDataModel.CommitId.Should().Be(commit.Id); extraDataModel.DateTime.Should().Be(commit.HybridDateTime.DateTime); diff --git a/src/SIL.Harmony.Tests/RepositoryTests.cs b/src/SIL.Harmony.Tests/RepositoryTests.cs index 1cf787d..b3e23ad 100644 --- a/src/SIL.Harmony.Tests/RepositoryTests.cs +++ b/src/SIL.Harmony.Tests/RepositoryTests.cs @@ -21,7 +21,7 @@ public RepositoryTests() .AddCrdtDataSample(":memory:") .BuildServiceProvider(); - _repository = _services.GetRequiredService(); + _repository = _services.GetRequiredService().CreateRepositorySync(); _crdtDbContext = _services.GetRequiredService(); } @@ -34,6 +34,7 @@ public async Task InitializeAsync() public async Task DisposeAsync() { + await _repository.DisposeAsync(); await _services.DisposeAsync(); } diff --git a/src/SIL.Harmony.Tests/SnapshotTests.cs b/src/SIL.Harmony.Tests/SnapshotTests.cs index e203b2f..248d574 100644 --- a/src/SIL.Harmony.Tests/SnapshotTests.cs +++ b/src/SIL.Harmony.Tests/SnapshotTests.cs @@ -116,8 +116,8 @@ await WriteNextChange( TagWord(wordId, tagId), ]); - var word = await DataModel.QueryLatest().Include(w => w.Tags) - .Where(w => w.Id == wordId).FirstOrDefaultAsync(); + var word = await DataModel.QueryLatest(q => q.Include(w => w.Tags) + .Where(w => w.Id == wordId)).FirstOrDefaultAsync(); word.Should().NotBeNull(); word.Tags.Should().BeEquivalentTo([new Tag { Id = tagId, Text = "tag-1" }]); } @@ -135,8 +135,8 @@ await WriteNextChange( var tagCreation = await WriteNextChange(TagWord(wordId, tagId)); await WriteChangeBefore(tagCreation, TagWord(wordId, tagId)); - var word = await DataModel.QueryLatest().Include(w => w.Tags) - .Where(w => w.Id == wordId).FirstOrDefaultAsync(); + var word = await DataModel.QueryLatest(q=> q.Include(w => w.Tags) + .Where(w => w.Id == wordId)).FirstOrDefaultAsync(); word.Should().NotBeNull(); word.Tags.Should().BeEquivalentTo([new Tag { Id = tagId, Text = "tag-1" }]); } diff --git a/src/SIL.Harmony.Tests/SyncTests.cs b/src/SIL.Harmony.Tests/SyncTests.cs index 385cba1..bc9985a 100644 --- a/src/SIL.Harmony.Tests/SyncTests.cs +++ b/src/SIL.Harmony.Tests/SyncTests.cs @@ -116,7 +116,7 @@ public async Task CanSync_AddDependentWithMultipleChanges() await _client2.DataModel.SyncWith(_client1.DataModel); - _client2.DataModel.QueryLatest().Should() - .BeEquivalentTo(_client1.DataModel.QueryLatest()); + _client2.DataModel.QueryLatest().ToBlockingEnumerable().Should() + .BeEquivalentTo(_client1.DataModel.QueryLatest().ToBlockingEnumerable()); } } \ No newline at end of file diff --git a/src/SIL.Harmony/CrdtKernel.cs b/src/SIL.Harmony/CrdtKernel.cs index 0345057..2d1c95d 100644 --- a/src/SIL.Harmony/CrdtKernel.cs +++ b/src/SIL.Harmony/CrdtKernel.cs @@ -1,4 +1,5 @@ using System.Text.Json; +using Microsoft.EntityFrameworkCore; using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Options; @@ -8,21 +9,34 @@ namespace SIL.Harmony; public static class CrdtKernel { + public static IServiceCollection AddCrdtDataDbFactory(this IServiceCollection services, + Action configureCrdt) where TContext : DbContext, ICrdtDbContext + { + + services.AddCrdtDataCore(configureCrdt); + services.AddScoped>(); + return services; + } + public static IServiceCollection AddCrdtData(this IServiceCollection services, - Action configureCrdt) where TContext: ICrdtDbContext + Action configureCrdt) where TContext : DbContext, ICrdtDbContext + { + services.AddCrdtDataCore(configureCrdt); + services.AddScoped>(); + return services; + } + public static IServiceCollection AddCrdtDataCore(this IServiceCollection services, Action configureCrdt) { services.AddLogging(); - services.AddOptions().Configure(configureCrdt).PostConfigure(crdtConfig => crdtConfig.ObjectTypeListBuilder.Freeze()); + services.AddOptions().Configure(configureCrdt) + .PostConfigure(crdtConfig => crdtConfig.ObjectTypeListBuilder.Freeze()); services.AddSingleton(sp => sp.GetRequiredService>().Value.JsonSerializerOptions); services.AddSingleton(TimeProvider.System); services.AddScoped(NewTimeProvider); - //must use factory, otherwise one context will be created for this registration, and one for the application. - //we want to have one context per application - services.AddScoped(p => p.GetRequiredService()); - services.AddScoped(); + services.AddScoped(); //must use factory method because DataModel constructor is internal services.AddScoped(provider => new DataModel( - provider.GetRequiredService(), + provider.GetRequiredService(), provider.GetRequiredService(), provider.GetRequiredService(), provider.GetRequiredService>(), @@ -30,7 +44,7 @@ public static IServiceCollection AddCrdtData(this IServiceCollection s )); //must use factory method because ResourceService constructor is internal services.AddScoped(provider => new ResourceService( - provider.GetRequiredService(), + provider.GetRequiredService(), provider.GetRequiredService>(), provider.GetRequiredService(), provider.GetRequiredService>() @@ -43,8 +57,11 @@ public static HybridDateTimeProvider NewTimeProvider(IServiceProvider servicePro //todo, if this causes issues getting the order correct, we can update the last date time after the db is created //as long as it's before we get a date time from the provider //todo use IMemoryCache to store the last date time, possibly based on the current project - var hybridDateTime = serviceProvider.GetRequiredService().GetLatestDateTime(); + using var repo = serviceProvider.GetRequiredService().CreateRepositorySync(); + var hybridDateTime = repo.GetLatestDateTime(); hybridDateTime ??= HybridDateTimeProvider.DefaultLastDateTime; return ActivatorUtilities.CreateInstance(serviceProvider, hybridDateTime); } } + + diff --git a/src/SIL.Harmony/DataModel.cs b/src/SIL.Harmony/DataModel.cs index c0c0915..9235389 100644 --- a/src/SIL.Harmony/DataModel.cs +++ b/src/SIL.Harmony/DataModel.cs @@ -18,28 +18,24 @@ public class DataModel : ISyncable, IAsyncDisposable /// private bool AlwaysValidate => _crdtConfig.Value.AlwaysValidateCommits; - private static readonly ConcurrentDictionary Locks = new(); - private AsyncLock _lock; - - private readonly CrdtRepository _crdtRepository; + private readonly CrdtRepositoryFactory _crdtRepositoryFactory; private readonly JsonSerializerOptions _serializerOptions; private readonly IHybridDateTimeProvider _timeProvider; private readonly IOptions _crdtConfig; private readonly ILogger _logger; //constructor must be internal because CrdtRepository is internal - internal DataModel(CrdtRepository crdtRepository, + internal DataModel(CrdtRepositoryFactory crdtRepositoryFactory, JsonSerializerOptions serializerOptions, IHybridDateTimeProvider timeProvider, IOptions crdtConfig, ILogger logger) { - _crdtRepository = crdtRepository; + _crdtRepositoryFactory = crdtRepositoryFactory; _serializerOptions = serializerOptions; _timeProvider = timeProvider; _crdtConfig = crdtConfig; _logger = logger; - _lock = Locks.GetOrAdd(crdtRepository.DatabaseIdentifier, new AsyncLock()); } @@ -70,67 +66,73 @@ public async Task AddChange( return await AddChanges(clientId, [change], commitId, commitMetadata); } + public async Task AddManyChanges(Guid clientId, + IEnumerable changes, + Func commitMetadata, + int changesPerCommitMax = 100) + { + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var commits = changes + .Chunk(changesPerCommitMax) + .Select(chunk => NewCommit(Guid.NewGuid(), clientId, commitMetadata(), chunk)) + .ToArray(); + if (commits is []) return; + using var locked = await repo.Lock(); + repo.ClearChangeTracker(); + + await using var transaction = await repo.BeginTransactionAsync(); + await repo.AddCommits(commits); + await UpdateSnapshots(repo, commits.First(), commits); + await ValidateCommits(repo); + await transaction.CommitAsync(); + } + /// public async Task AddChanges( Guid clientId, IEnumerable changes, Guid commitId = default, - CommitMetadata? commitMetadata = null, - bool deferCommit = false) + CommitMetadata? commitMetadata = null) + { + var commit = NewCommit(commitId, clientId, commitMetadata, changes); + await Add(commit); + return commit; + } + + private Commit NewCommit(Guid commitId, Guid clientId, CommitMetadata? commitMetadata, IEnumerable changes) { commitId = commitId == default ? Guid.NewGuid() : commitId; - var commit = new Commit(commitId) + return new Commit(commitId) { ClientId = clientId, HybridDateTime = _timeProvider.GetDateTime(), ChangeEntities = [..changes.Select(ToChangeEntity)], Metadata = commitMetadata ?? new() }; - await Add(commit, deferCommit); - return commit; } - private List _deferredCommits = []; - - private async Task Add(Commit commit, bool deferSnapshotUpdates) + private async Task Add(Commit commit) { - if (await _crdtRepository.HasCommit(commit.Id)) return; - using var locked = await _lock.LockAsync(); - _crdtRepository.ClearChangeTracker(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + if (await repo.HasCommit(commit.Id)) return; + using var locked = await repo.Lock(); + repo.ClearChangeTracker(); - await using var transaction = _crdtRepository.IsInTransaction ? null : await _crdtRepository.BeginTransactionAsync(); - await _crdtRepository.AddCommit(commit); - if (!deferSnapshotUpdates) - { - //if there are deferred commits, update snapshots with them first - if (_deferredCommits is not []) await FlushDeferredCommits(); - await UpdateSnapshots(commit, [commit]); + await using var transaction = repo.IsInTransaction ? null : await repo.BeginTransactionAsync(); + await repo.AddCommit(commit); + await UpdateSnapshots(repo, commit, [commit]); - if (AlwaysValidate) await ValidateCommits(); - } - else - { - _deferredCommits.Add(commit); - } - if (transaction is not null) await transaction.CommitAsync(); - } + if (AlwaysValidate) await ValidateCommits(repo); - public async ValueTask DisposeAsync() - { - if (_deferredCommits is []) return; - await FlushDeferredCommits(); + + if (transaction is not null) await transaction.CommitAsync(); } - public async Task FlushDeferredCommits() + public ValueTask DisposeAsync() { - var commits = Interlocked.Exchange(ref _deferredCommits, []); - var oldestChange = commits.MinBy(c => c.CompareKey); - if (oldestChange is null) return; - await UpdateSnapshots(oldestChange, commits.ToArray()); - if (AlwaysValidate) await ValidateCommits(); + return ValueTask.CompletedTask; } - private static ChangeEntity ToChangeEntity(IChange change, int index) { return new ChangeEntity() @@ -144,20 +146,19 @@ async Task ISyncable.AddRangeFromSync(IEnumerable commits) commits = commits.ToArray(); try { - using var locked = await _lock.LockAsync(); - _crdtRepository.ClearChangeTracker(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + using var locked = await repo.Lock(); + repo.ClearChangeTracker(); _timeProvider.TakeLatestTime(commits.Select(c => c.HybridDateTime)); - var (oldestChange, newCommits) = await _crdtRepository.FilterExistingCommits(commits.ToArray()); + var (oldestChange, newCommits) = await repo.FilterExistingCommits(commits.ToArray()); //no changes added if (oldestChange is null || newCommits is []) return; - await using var transaction = await _crdtRepository.BeginTransactionAsync(); - //if there are deferred commits, update snapshots with them first - if (_deferredCommits is not []) await FlushDeferredCommits(); + await using var transaction = await repo.BeginTransactionAsync(); //don't save since UpdateSnapshots will also modify newCommits with hashes, so changes will be saved once that's done - await _crdtRepository.AddCommits(newCommits, false); - await UpdateSnapshots(oldestChange, newCommits); - await ValidateCommits(); + await repo.AddCommits(newCommits, false); + await UpdateSnapshots(repo, oldestChange, newCommits); + await ValidateCommits(repo); await transaction.CommitAsync(); } catch (DbUpdateException e) @@ -193,14 +194,14 @@ ValueTask ISyncable.ShouldSync() return ValueTask.FromResult(true); } - private async Task UpdateSnapshots(Commit oldestAddedCommit, Commit[] newCommits) + private async Task UpdateSnapshots(CrdtRepository repo, Commit oldestAddedCommit, Commit[] newCommits) { - await _crdtRepository.DeleteStaleSnapshots(oldestAddedCommit); + await repo.DeleteStaleSnapshots(oldestAddedCommit); Dictionary snapshotLookup; if (newCommits.Length > 10) { var entityIds = newCommits.SelectMany(c => c.ChangeEntities.Select(ce => ce.EntityId)); - snapshotLookup = await _crdtRepository.CurrentSnapshots() + snapshotLookup = await repo.CurrentSnapshots() .Where(s => entityIds.Contains(s.EntityId)) .Select(s => new KeyValuePair(s.EntityId, s.Id)) .ToDictionaryAsync(s => s.Key, s => s.Value); @@ -210,14 +211,14 @@ private async Task UpdateSnapshots(Commit oldestAddedCommit, Commit[] newCommits snapshotLookup = []; } - var snapshotWorker = new SnapshotWorker(snapshotLookup, _crdtRepository, _crdtConfig.Value); + var snapshotWorker = new SnapshotWorker(snapshotLookup, repo, _crdtConfig.Value); await snapshotWorker.UpdateSnapshots(oldestAddedCommit, newCommits); } - private async Task ValidateCommits() + private async Task ValidateCommits(CrdtRepository repo) { Commit? parentCommit = null; - await foreach (var commit in _crdtRepository.CurrentCommits().AsNoTracking().AsAsyncEnumerable()) + await foreach (var commit in repo.CurrentCommits().AsNoTracking().AsAsyncEnumerable()) { var parentHash = parentCommit?.Hash ?? CommitBase.NullParentHash; var expectedHash = commit.GenerateHash(parentHash); @@ -227,8 +228,8 @@ private async Task ValidateCommits() continue; } - var actualParentCommit = await _crdtRepository.FindCommitByHash(commit.ParentHash); - var commitWithSnapshots = await _crdtRepository.CurrentCommits().Include(c => c.Snapshots).SingleAsync(c => c.Id == commit.Id); + var actualParentCommit = await repo.FindCommitByHash(commit.ParentHash); + var commitWithSnapshots = await repo.CurrentCommits().Include(c => c.Snapshots).SingleAsync(c => c.Id == commit.Id); throw new CommitValidationException( $"Commit {commit} does not match expected hash, parent hash [{commit.ParentHash}] !== [{parentHash}], expected parent {parentCommit?.ToString() ?? "null"} and actual parent {actualParentCommit?.ToString() ?? "null"}, with snapshots: {string.Join(", ", commitWithSnapshots.Snapshots.Select(s => s.Entity.DbObject))}"); } @@ -236,48 +237,63 @@ private async Task ValidateCommits() public async Task RegenerateSnapshots() { - await _crdtRepository.DeleteSnapshotsAndProjectedTables(); - _crdtRepository.ClearChangeTracker(); - var allCommits = await _crdtRepository.CurrentCommits().AsNoTracking().ToArrayAsync(); - await UpdateSnapshots(allCommits.First(), allCommits); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + await repo.DeleteSnapshotsAndProjectedTables(); + repo.ClearChangeTracker(); + var allCommits = await repo.CurrentCommits().AsNoTracking().ToArrayAsync(); + await UpdateSnapshots(repo, allCommits.First(), allCommits); } public async Task GetLatestSnapshotByObjectId(Guid entityId) { - return await _crdtRepository.GetCurrentSnapshotByObjectId(entityId) ?? + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await repo.GetCurrentSnapshotByObjectId(entityId) ?? throw new ArgumentException($"unable to find snapshot for entity {entityId}"); } public async Task GetLatest(Guid objectId) where T : class { - return await _crdtRepository.GetCurrent(objectId); + return await _crdtRepositoryFactory.Execute(repo => repo.GetCurrent(objectId)); } - public async Task GetProjectSnapshot(bool includeDeleted = false) + + public IAsyncEnumerable QueryLatest(Func, IQueryable>? apply = null) + where T : class { - return new ModelSnapshot(await _crdtRepository.CurrenSimpleSnapshots(includeDeleted).ToArrayAsync()); + return QueryLatest(apply ?? (static q => q)); } - public IQueryable QueryLatest() where T : class + public async IAsyncEnumerable QueryLatest(Func, IQueryable> apply) where T : class { - var q = _crdtRepository.GetCurrentObjects(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var q = repo.GetCurrentObjects(); if (q is IQueryable) { q = q.OrderBy(o => EF.Property(o, nameof(IOrderableCrdt.Order))) .ThenBy(o => EF.Property(o, nameof(IOrderableCrdt.Id))); } - return q; + await foreach (var result in apply(q).AsAsyncEnumerable()) + { + yield return result; + } + } + + public async Task GetProjectSnapshot(bool includeDeleted = false) + { + var snapshots = await _crdtRepositoryFactory.Execute(repo => repo.CurrenSimpleSnapshots(includeDeleted).ToArrayAsync()); + return new ModelSnapshot(snapshots); } public async Task GetBySnapshotId(Guid snapshotId) { - return await _crdtRepository.GetObjectBySnapshotId(snapshotId); + return await _crdtRepositoryFactory.Execute(repo => repo.GetObjectBySnapshotId(snapshotId)); } public async Task> GetSnapshotsAtCommit(Commit commit) { - var repository = _crdtRepository.GetScopedRepository(commit); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var repository = repo.GetScopedRepository(commit); var (snapshots, pendingCommits) = await repository.GetCurrentSnapshotsAndPendingCommits(); if (pendingCommits.Length != 0) @@ -293,20 +309,23 @@ public async Task> GetSnapshotsAtCommit(Commit public async Task GetAtTime(DateTimeOffset time, Guid entityId) { - var commitBefore = await _crdtRepository.CurrentCommits().LastOrDefaultAsync(c => c.HybridDateTime.DateTime <= time); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var commitBefore = await repo.CurrentCommits().LastOrDefaultAsync(c => c.HybridDateTime.DateTime <= time); if (commitBefore is null) throw new ArgumentException("unable to find any commits"); return await GetAtCommit(commitBefore, entityId); } public async Task GetAtCommit(Guid commitId, Guid entityId) { - return await GetAtCommit(await _crdtRepository.CurrentCommits().SingleAsync(c => c.Id == commitId), + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await GetAtCommit(await repo.CurrentCommits().SingleAsync(c => c.Id == commitId), entityId); } public async Task GetAtCommit(Commit commit, Guid entityId) { - var repository = _crdtRepository.GetScopedRepository(commit); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var repository = repo.GetScopedRepository(commit); var snapshot = await repository.GetCurrentSnapshotByObjectId(entityId, false); ArgumentNullException.ThrowIfNull(snapshot); var newCommits = await repository.CurrentCommits() @@ -330,12 +349,14 @@ public async Task GetAtCommit(Commit commit, Guid entityId) public async Task GetSyncState() { - return await _crdtRepository.GetCurrentSyncState(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await repo.GetCurrentSyncState(); } public async Task> GetChanges(SyncState remoteState) { - return await _crdtRepository.GetChanges(remoteState); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await repo.GetChanges(remoteState); } public async Task SyncWith(ISyncable remoteModel) diff --git a/src/SIL.Harmony/Db/CrdtDbContextFactory.cs b/src/SIL.Harmony/Db/CrdtDbContextFactory.cs new file mode 100644 index 0000000..e808eca --- /dev/null +++ b/src/SIL.Harmony/Db/CrdtDbContextFactory.cs @@ -0,0 +1,97 @@ +using Microsoft.EntityFrameworkCore; +using Microsoft.EntityFrameworkCore.ChangeTracking; +using Microsoft.EntityFrameworkCore.Infrastructure; + +namespace SIL.Harmony.Db; + +public class CrdtDbContextFactory(IDbContextFactory dbContextFactory) : ICrdtDbContextFactory + where TContext : DbContext, ICrdtDbContext +{ + public async Task CreateDbContextAsync(CancellationToken cancellationToken = new CancellationToken()) + { + return await dbContextFactory.CreateDbContextAsync(cancellationToken); + } + + public ICrdtDbContext CreateDbContext() + { + return dbContextFactory.CreateDbContext(); + } +} + +public interface ICrdtDbContextFactory +{ + Task CreateDbContextAsync(CancellationToken cancellationToken = new CancellationToken()); + ICrdtDbContext CreateDbContext(); +} + +public class CrdtDbContextNoDisposeFactory(TContext dbContext) : ICrdtDbContextFactory + where TContext : ICrdtDbContext +{ + public Task CreateDbContextAsync(CancellationToken cancellationToken = new CancellationToken()) + { + return Task.FromResult(new NoDisposeWrapper(dbContext)); + } + + public ICrdtDbContext CreateDbContext() + { + return new NoDisposeWrapper(dbContext); + } + + private class NoDisposeWrapper(ICrdtDbContext context): ICrdtDbContext + { + public void Dispose() + { + //noop, don't dispose because the context is owned by the ioc container + } + + public ValueTask DisposeAsync() + { + //noop, don't dispose because the context is owned by the ioc container + return ValueTask.CompletedTask; + } + + public Task SaveChangesAsync(CancellationToken cancellationToken = default) + { + return context.SaveChangesAsync(cancellationToken); + } + + public ValueTask FindAsync(Type entityType, params object?[]? keyValues) + { + return context.FindAsync(entityType, keyValues); + } + + public DbSet Set() where TEntity : class + { + return context.Set(); + } + + public DatabaseFacade Database => context.Database; + + public ChangeTracker ChangeTracker => context.ChangeTracker; + + public EntityEntry Entry(TEntity entity) where TEntity : class + { + return context.Entry(entity); + } + + public EntityEntry Entry(object entity) + { + return context.Entry(entity); + } + + public EntityEntry Add(object entity) + { + return context.Add(entity); + } + + public void AddRange(IEnumerable entities) + { + context.AddRange(entities); + } + + public EntityEntry Remove(object entity) + { + return context.Remove(entity); + } + } +} diff --git a/src/SIL.Harmony/Db/CrdtRepository.cs b/src/SIL.Harmony/Db/CrdtRepository.cs index e9a7644..768613b 100644 --- a/src/SIL.Harmony/Db/CrdtRepository.cs +++ b/src/SIL.Harmony/Db/CrdtRepository.cs @@ -1,17 +1,48 @@ +using System.Collections.Concurrent; using System.Reflection; using Microsoft.EntityFrameworkCore; using Microsoft.EntityFrameworkCore.ChangeTracking; using Microsoft.EntityFrameworkCore.Infrastructure; using Microsoft.EntityFrameworkCore.Storage; +using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Options; +using Nito.AsyncEx; using SIL.Harmony.Changes; using SIL.Harmony.Resource; namespace SIL.Harmony.Db; -internal class CrdtRepository +internal class CrdtRepositoryFactory(IServiceProvider serviceProvider, ICrdtDbContextFactory dbContextFactory) { + public async Task CreateRepository() + { + return ActivatorUtilities.CreateInstance(serviceProvider, await dbContextFactory.CreateDbContextAsync()); + } + + public CrdtRepository CreateRepositorySync() + { + return ActivatorUtilities.CreateInstance(serviceProvider, dbContextFactory.CreateDbContext()); + } + + public async Task Execute(Func> func) + { + await using var repo = await CreateRepository(); + return await func(repo); + } + + public async ValueTask Execute(Func> func) + { + await using var repo = await CreateRepository(); + return await func(repo); + } +} + +internal class CrdtRepository : IDisposable, IAsyncDisposable +{ + private static readonly ConcurrentDictionary Locks = new(); + + private readonly AsyncLock _lock; private readonly ICrdtDbContext _dbContext; private readonly IOptions _crdtConfig; private readonly ILogger _logger; @@ -26,6 +57,12 @@ public CrdtRepository(ICrdtDbContext dbContext, IOptions crdtConfig, //we can't use the scoped db context is it prevents access to the DbSet for the Snapshots, //but since we're using a custom query, we can use it directly and apply the scoped filters manually _currentSnapshotsQueryable = MakeCurrentSnapshotsQuery(dbContext, ignoreChangesAfter); + _lock = Locks.GetOrAdd(DatabaseIdentifier, _ => new AsyncLock()); + } + + public AwaitableDisposable Lock() + { + return _lock.LockAsync(); } /// @@ -33,7 +70,7 @@ public CrdtRepository(ICrdtDbContext dbContext, IOptions crdtConfig, /// may be the connection string so it could contain sensitive information /// if it's in memory we'll just use a random guid /// - internal string DatabaseIdentifier + private string DatabaseIdentifier { get { @@ -43,6 +80,8 @@ internal string DatabaseIdentifier } } + //doesn't really do anything when using a dbcontext factory since it will likely just have been created + //but when not using the factory it is still useful internal void ClearChangeTracker() { _dbContext.ChangeTracker.Clear(); @@ -372,6 +411,16 @@ public IQueryable LocalResourceIds() { return await _dbContext.Set().FindAsync(resourceId); } + + public void Dispose() + { + _dbContext.Dispose(); + } + + public async ValueTask DisposeAsync() + { + await _dbContext.DisposeAsync(); + } } internal class ScopedDbContext(ICrdtDbContext inner, Commit ignoreChangesAfter) : ICrdtDbContext @@ -422,4 +471,14 @@ public EntityEntry Remove(object entity) { return inner.Remove(entity); } + + public void Dispose() + { + inner.Dispose(); + } + + public ValueTask DisposeAsync() + { + return inner.DisposeAsync(); + } } diff --git a/src/SIL.Harmony/Db/ICrdtDbContext.cs b/src/SIL.Harmony/Db/ICrdtDbContext.cs index 68f3281..438a225 100644 --- a/src/SIL.Harmony/Db/ICrdtDbContext.cs +++ b/src/SIL.Harmony/Db/ICrdtDbContext.cs @@ -4,7 +4,7 @@ namespace SIL.Harmony.Db; -public interface ICrdtDbContext +public interface ICrdtDbContext : IDisposable, IAsyncDisposable { IQueryable Commits => Set(); IQueryable Snapshots => Set(); @@ -18,4 +18,4 @@ public interface ICrdtDbContext EntityEntry Add(object entity); void AddRange(IEnumerable entities); EntityEntry Remove(object entity); -} \ No newline at end of file +} diff --git a/src/SIL.Harmony/ResourceService.cs b/src/SIL.Harmony/ResourceService.cs index b0f1774..c4fae85 100644 --- a/src/SIL.Harmony/ResourceService.cs +++ b/src/SIL.Harmony/ResourceService.cs @@ -10,14 +10,14 @@ namespace SIL.Harmony; public class ResourceService { - private readonly CrdtRepository _crdtRepository; + private readonly CrdtRepositoryFactory _crdtRepositoryFactory; private readonly IOptions _crdtConfig; private readonly DataModel _dataModel; private readonly ILogger _logger; - internal ResourceService(CrdtRepository crdtRepository, IOptions crdtConfig, DataModel dataModel, ILogger logger) + internal ResourceService(CrdtRepositoryFactory crdtRepositoryFactory, IOptions crdtConfig, DataModel dataModel, ILogger logger) { - _crdtRepository = crdtRepository; + _crdtRepositoryFactory = crdtRepositoryFactory; _crdtConfig = crdtConfig; _dataModel = dataModel; _logger = logger; @@ -34,14 +34,15 @@ public async Task AddLocalResource(string resourcePath, IRemoteResourceService? resourceService = null) { ValidateResourcesSetup(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); var localResource = new LocalResource { Id = id == default ? Guid.NewGuid() : id, LocalPath = Path.GetFullPath(resourcePath) }; if (!localResource.FileExists()) throw new FileNotFoundException(localResource.LocalPath); - await using var transaction = await _crdtRepository.BeginTransactionAsync(); - await _crdtRepository.AddLocalResource(localResource); + await using var transaction = await repo.BeginTransactionAsync(); + await repo.AddLocalResource(localResource); UploadResult? uploadResult = null; if (resourceService is not null) { @@ -77,8 +78,9 @@ public async Task AddLocalResource(string resourcePath, public async Task ListResourcesPendingUpload() { ValidateResourcesSetup(); - var remoteResources = await _dataModel.QueryLatest().Where(r => r.RemoteId == null).ToArrayAsync(); - var localResource = _crdtRepository.LocalResourcesByIds(remoteResources.Select(r => r.Id)); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var remoteResources = await repo.GetCurrentObjects().Where(r => r.RemoteId == null).ToArrayAsync(); + var localResource = repo.LocalResourcesByIds(remoteResources.Select(r => r.Id)); return await localResource.ToArrayAsync(); } @@ -104,7 +106,8 @@ public async Task UploadPendingResources(Guid clientId, IRemoteResourceService r public async Task UploadPendingResource(Guid resourceId, Guid clientId, IRemoteResourceService remoteResourceService) { - var localResource = await _crdtRepository.GetLocalResource(resourceId) ?? + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var localResource = await repo.GetLocalResource(resourceId) ?? throw new ArgumentException($"unable to find local resource with id {resourceId}"); ValidateResourcesSetup(); await UploadPendingResource(localResource, clientId, remoteResourceService); @@ -120,8 +123,9 @@ public async Task UploadPendingResource(LocalResource localResource, Guid client public async Task ListResourcesPendingDownload() { ValidateResourcesSetup(); - var localResourceIds = _crdtRepository.LocalResourceIds(); - var remoteResources = await _dataModel.QueryLatest() + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var localResourceIds = repo.LocalResourceIds(); + var remoteResources = await repo.GetCurrentObjects() .Where(r => r.RemoteId != null && !localResourceIds.Contains(r.Id)) .ToArrayAsync(); return remoteResources; @@ -130,14 +134,24 @@ public async Task ListResourcesPendingDownload() public async Task DownloadResource(Guid resourceId, IRemoteResourceService remoteResourceService) { ValidateResourcesSetup(); - return await DownloadResource( - await _dataModel.GetLatest(resourceId) ?? + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await DownloadResourceInternal(repo, + await repo.GetCurrent(resourceId) ?? throw new EntityNotFoundException("Unable to find remote resource"), remoteResourceService ); } - public async Task DownloadResource(RemoteResource remoteResource, IRemoteResourceService remoteResourceService) + public async Task DownloadResource(RemoteResource remoteResource, + IRemoteResourceService remoteResourceService) + { + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + return await DownloadResourceInternal(repo, remoteResource, remoteResourceService); + } + + private async Task DownloadResourceInternal(CrdtRepository repo, + RemoteResource remoteResource, + IRemoteResourceService remoteResourceService) { ValidateResourcesSetup(); ArgumentNullException.ThrowIfNull(remoteResource.RemoteId); @@ -147,13 +161,13 @@ public async Task DownloadResource(RemoteResource remoteResource, Id = remoteResource.Id, LocalPath = downloadResult.LocalPath }; - await _crdtRepository.AddLocalResource(localResource); + await repo.AddLocalResource(localResource); return localResource; } public async Task GetLocalResource(Guid resourceId) { - return await _crdtRepository.GetLocalResource(resourceId); + return await _crdtRepositoryFactory.Execute(repo => repo.GetLocalResource(resourceId)); } public async Task AllResources() @@ -163,8 +177,9 @@ public async Task AllResources() private async Task> AllResourcesInternal() { - var remoteResources = await _dataModel.QueryLatest().ToArrayAsync(); - var localResources = await _crdtRepository.LocalResources().ToArrayAsync(); + await using var repo = await _crdtRepositoryFactory.CreateRepository(); + var remoteResources = await repo.GetCurrentObjects().ToArrayAsync(); + var localResources = await repo.LocalResources().ToArrayAsync(); return remoteResources.FullOuterJoin(localResources, r => r.Id, l => l.Id, @@ -181,4 +196,4 @@ private async Task> AllResourcesInternal() var resources = await AllResourcesInternal(); return resources.FirstOrDefault(r => r.Id == resourceId); } -} \ No newline at end of file +}