Skip to content

Commit 6454d50

Browse files
committed
Fix CivArchive download startup delay
1 parent cab41cc commit 6454d50

4 files changed

Lines changed: 174 additions & 10 deletions

File tree

StabilityMatrix.Avalonia/Services/ModelImportService.cs

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -299,7 +299,9 @@ public async Task DoCustomImport(
299299

300300
var downloadPath = downloadFolder.JoinFile(modelBaseFileName + modelFileExtension);
301301

302-
// Save model info and preview image first if available
302+
// Save model info first if available. Preview image downloads can be slow
303+
// or hosted on flaky third-party mirrors, so start the model download before
304+
// fetching the preview in the background.
303305
var cleanupFilePaths = new List<string>();
304306
if (connectedModelInfo is not null)
305307
{
@@ -308,6 +310,8 @@ public async Task DoCustomImport(
308310
downloadFolder.JoinFile(modelBaseFileName + ConnectedModelInfo.FileExtension)
309311
);
310312
}
313+
314+
FilePath? previewImageDownloadPath = null;
311315
if (previewImageUri is not null)
312316
{
313317
if (previewImageFileExtension is null)
@@ -321,15 +325,10 @@ public async Task DoCustomImport(
321325
}
322326
}
323327

324-
var previewImageDownloadPath = downloadFolder.JoinFile(
328+
previewImageDownloadPath = downloadFolder.JoinFile(
325329
modelBaseFileName + ".preview" + previewImageFileExtension
326330
);
327331

328-
await notificationService.TryAsync(
329-
downloadService.DownloadToFileAsync(previewImageUri.ToString(), previewImageDownloadPath),
330-
"Could not download preview image"
331-
);
332-
333332
cleanupFilePaths.Add(previewImageDownloadPath);
334333
}
335334

@@ -351,6 +350,19 @@ await notificationService.TryAsync(
351350
// download.ContextAction = CivitPostDownloadContextAction.FromCivitFile(modelFile);
352351

353352
await trackedDownloadService.TryStartDownload(download);
353+
354+
if (previewImageUri is not null && previewImageDownloadPath is not null)
355+
{
356+
DownloadPreviewImageAsync(previewImageUri, previewImageDownloadPath).SafeFireAndForget();
357+
}
358+
}
359+
360+
private async Task DownloadPreviewImageAsync(Uri previewImageUri, FilePath previewImageDownloadPath)
361+
{
362+
await notificationService.TryAsync(
363+
downloadService.DownloadToFileAsync(previewImageUri.ToString(), previewImageDownloadPath),
364+
"Could not download preview image"
365+
);
354366
}
355367

356368
private string GenerateUniqueFileName(string folder, string fileName)

StabilityMatrix.Avalonia/ViewModels/CheckpointBrowser/CivArchiveDetailsPageViewModel.cs

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -310,7 +310,7 @@ private void PopulateVersionData(CivArchiveModelVersion? version)
310310
HasDownloadUrl = GetDownloadUris(version).Count > 0;
311311

312312
Images.Clear();
313-
foreach (var image in version?.Images.Where(i => !string.IsNullOrWhiteSpace(i.Url)) ?? [])
313+
foreach (var image in version?.Images.Where(IsUsableImage) ?? [])
314314
{
315315
Images.Add(image);
316316
}
@@ -689,7 +689,7 @@ private async Task ExecuteDownloadAsync(
689689

690690
Uri? previewImageUri = null;
691691
string? previewImageExtension = null;
692-
var firstImage = version.Images.FirstOrDefault(i => !string.IsNullOrWhiteSpace(i.Url));
692+
var firstImage = version.Images.FirstOrDefault(IsUsableImage);
693693
if (firstImage?.Url is not null)
694694
{
695695
previewImageUri = new Uri(firstImage.Url);
@@ -879,7 +879,7 @@ string sourceUrl
879879
ImportedAt = DateTimeOffset.UtcNow,
880880
Hashes = new CivitFileHashes { SHA256 = primaryFile?.Sha256 },
881881
TrainedWords = version.Trigger.ToArray(),
882-
ThumbnailImageUrl = version.Images.FirstOrDefault(i => !string.IsNullOrWhiteSpace(i.Url))?.Url,
882+
ThumbnailImageUrl = version.Images.FirstOrDefault(IsUsableImage)?.Url,
883883
Source = ConnectedModelSource.CivArchive,
884884
SourceUrl = sourceUrl,
885885
Stats = new CivitModelStats
@@ -893,6 +893,15 @@ string sourceUrl
893893
};
894894
}
895895

896+
private static bool IsUsableImage(CivArchiveModelImage image)
897+
{
898+
return !string.IsNullOrWhiteSpace(image.Url)
899+
&& (
900+
string.IsNullOrWhiteSpace(image.Type)
901+
|| string.Equals(image.Type, "image", StringComparison.OrdinalIgnoreCase)
902+
);
903+
}
904+
896905
[RelayCommand]
897906
private void OpenVersionMirror(CivArchiveVersionMirror? mirror)
898907
{

StabilityMatrix.Tests/Avalonia/CivArchiveBrowserViewModelTests.cs

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -274,6 +274,64 @@ public async Task DownloadModel_UsesFileMirrorUrlWhenDirectUrlIsMissing()
274274
Assert.AreEqual("https://example.org/mirror/mirror-only.safetensors", capturedUris?[0].ToString());
275275
}
276276

277+
[TestMethod]
278+
public async Task DownloadModel_UsesImagePreviewAndSkipsVideoMedia()
279+
{
280+
var apiClient = Substitute.For<ICivArchiveApiClient>();
281+
var modelImportService = Substitute.For<IModelImportService>();
282+
var settingsManager = Substitute.For<ISettingsManager>();
283+
var model = CreateDetailsModel(
284+
new CivArchiveModelFile
285+
{
286+
Name = "model.safetensors",
287+
DownloadUrl = "https://example.org/download/model.safetensors",
288+
IsPrimary = true,
289+
}
290+
);
291+
model.Version!.Images =
292+
[
293+
new CivArchiveModelImage { Url = "https://c.genur.art/video-id", Type = "video" },
294+
new CivArchiveModelImage
295+
{
296+
Url = "https://img.genur.art/sig/width:450/quality:85/image-id",
297+
Type = "image",
298+
},
299+
];
300+
301+
Uri? capturedPreviewUri = null;
302+
303+
apiClient
304+
.GetModelDetailsAsync(Arg.Any<string>(), Arg.Any<CancellationToken>())
305+
.Returns(new CivArchiveModelDetailsResponse { Model = model });
306+
apiClient.GetAbsoluteUri(Arg.Any<string>()).Returns(call => new Uri(call.Arg<string>()));
307+
settingsManager.IsLibraryDirSet.Returns(true);
308+
settingsManager.ModelsDirectory.Returns(Path.GetTempPath());
309+
modelImportService
310+
.DoCustomImport(
311+
Arg.Any<IEnumerable<Uri>>(),
312+
Arg.Any<string>(),
313+
Arg.Any<DirectoryPath>(),
314+
Arg.Do<Uri?>(uri => capturedPreviewUri = uri),
315+
Arg.Any<string?>(),
316+
Arg.Any<ConnectedModelInfo?>(),
317+
Arg.Any<Action<TrackedDownload>?>()
318+
)
319+
.Returns(Task.CompletedTask);
320+
321+
var vm = CreateDetailsViewModel(apiClient, modelImportService, settingsManager);
322+
vm.RelativeUrl = "/models/1?modelVersionId=2";
323+
324+
await vm.OnLoadedAsync();
325+
await vm.DownloadModelCommand.ExecuteAsync(null);
326+
327+
Assert.AreEqual(1, vm.Images.Count);
328+
Assert.AreEqual("https://img.genur.art/sig/width:450/quality:85/image-id", vm.Images[0].Url);
329+
Assert.AreEqual(
330+
"https://img.genur.art/sig/width:450/quality:85/image-id",
331+
capturedPreviewUri?.ToString()
332+
);
333+
}
334+
277335
[TestMethod]
278336
public void ParseSearchQuery_PlainQuery_ReturnsQueryOnly()
279337
{
Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
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

Comments
 (0)