|
| 1 | +using Avalonia.Controls.Notifications; |
| 2 | +using NSubstitute; |
| 3 | +using StabilityMatrix.Avalonia.Services; |
| 4 | +using StabilityMatrix.Core.Models; |
| 5 | +using StabilityMatrix.Core.Models.FileInterfaces; |
| 6 | +using StabilityMatrix.Core.Models.Progress; |
| 7 | +using StabilityMatrix.Core.Services; |
| 8 | + |
| 9 | +namespace StabilityMatrix.Tests.Avalonia; |
| 10 | + |
| 11 | +[TestClass] |
| 12 | +public class ModelImportServiceTests |
| 13 | +{ |
| 14 | + [TestMethod] |
| 15 | + public async Task DoCustomImport_StartsTrackedDownloadBeforePreviewDownloadCompletes() |
| 16 | + { |
| 17 | + var downloadService = Substitute.For<IDownloadService>(); |
| 18 | + var notificationService = Substitute.For<INotificationService>(); |
| 19 | + var trackedDownloadService = Substitute.For<ITrackedDownloadService>(); |
| 20 | + var service = new ModelImportService(downloadService, notificationService, trackedDownloadService); |
| 21 | + |
| 22 | + var tempDir = Directory.CreateTempSubdirectory(); |
| 23 | + var modelUri = new Uri("https://example.org/model.safetensors"); |
| 24 | + var previewUri = new Uri("https://example.org/preview.webp"); |
| 25 | + var previewDownload = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); |
| 26 | + var trackedDownload = new TrackedDownload |
| 27 | + { |
| 28 | + Id = Guid.NewGuid(), |
| 29 | + SourceUrl = modelUri, |
| 30 | + DownloadDirectory = new DirectoryPath(tempDir.FullName), |
| 31 | + FileName = "model.safetensors", |
| 32 | + TempFileName = "model.safetensors.tmp", |
| 33 | + }; |
| 34 | + |
| 35 | + try |
| 36 | + { |
| 37 | + downloadService |
| 38 | + .DownloadToFileAsync( |
| 39 | + previewUri.ToString(), |
| 40 | + Arg.Any<string>(), |
| 41 | + Arg.Any<IProgress<ProgressReport>?>(), |
| 42 | + Arg.Any<string?>(), |
| 43 | + Arg.Any<CancellationToken>() |
| 44 | + ) |
| 45 | + .Returns(previewDownload.Task); |
| 46 | + |
| 47 | + notificationService |
| 48 | + .TryAsync(Arg.Any<Task>(), Arg.Any<string>(), Arg.Any<string?>(), Arg.Any<NotificationType>()) |
| 49 | + .Returns(call => AwaitTask(call.Arg<Task>())); |
| 50 | + |
| 51 | + trackedDownloadService.NewDownload(modelUri, Arg.Any<FilePath>()).Returns(trackedDownload); |
| 52 | + trackedDownloadService.TryStartDownload(trackedDownload).Returns(Task.CompletedTask); |
| 53 | + |
| 54 | + var importTask = service.DoCustomImport( |
| 55 | + [modelUri], |
| 56 | + "model.safetensors", |
| 57 | + new DirectoryPath(tempDir.FullName), |
| 58 | + previewUri, |
| 59 | + ".webp" |
| 60 | + ); |
| 61 | + |
| 62 | + var completedTask = await Task.WhenAny(importTask, Task.Delay(TimeSpan.FromSeconds(1))); |
| 63 | + |
| 64 | + Assert.AreSame( |
| 65 | + importTask, |
| 66 | + completedTask, |
| 67 | + "The model import should not wait for preview image download completion." |
| 68 | + ); |
| 69 | + |
| 70 | + await trackedDownloadService.Received(1).TryStartDownload(trackedDownload); |
| 71 | + Assert.IsFalse(previewDownload.Task.IsCompleted); |
| 72 | + } |
| 73 | + finally |
| 74 | + { |
| 75 | + previewDownload.TrySetResult(); |
| 76 | + tempDir.Delete(recursive: true); |
| 77 | + } |
| 78 | + } |
| 79 | + |
| 80 | + private static async Task<TaskResult<bool>> AwaitTask(Task task) |
| 81 | + { |
| 82 | + await task; |
| 83 | + return new TaskResult<bool>(true); |
| 84 | + } |
| 85 | +} |
0 commit comments