Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
156 changes: 156 additions & 0 deletions ImageDescriber.Test/ScanTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,12 @@

namespace ktsu.ImageDescriber.Tests;

using System.Collections.Concurrent;
using System.Net;
using System.Net.Sockets;
using System.Text;
using System.Text.Json;

using ktsu.ImageDescriber.Verbs;
using ktsu.Semantics.Paths;
using ktsu.Semantics.Strings;
Expand Down Expand Up @@ -80,4 +86,154 @@ public void SanitizeFileNameHandlesSimpleName()

Assert.AreEqual("sunset-over-ocean.webp", result.WeakString);
}

[TestMethod]
public void HashFilesSkipsFilesThatCannotBeRead()
{
string tempDir = Path.Combine(Path.GetTempPath(), Guid.NewGuid().ToString());
Directory.CreateDirectory(tempDir);

try
{
AbsoluteFilePath real = Path.Combine(tempDir, "real.jpg").As<AbsoluteFilePath>();
AbsoluteFilePath missing = Path.Combine(tempDir, "broken.jpg").As<AbsoluteFilePath>();
File.WriteAllBytes(real.WeakString, [0xFF, 0xD8]);

Dictionary<AbsoluteFilePath, string> hashes = ImageHasher.HashFiles([real, missing]);

Assert.HasCount(1, hashes);
Assert.IsTrue(hashes.ContainsKey(real));
}
finally
{
Directory.Delete(tempDir, true);
}
}

[TestMethod]
public void DescribeImagesContinuesPastPerImageFailures()
{
string tempDir = Path.Combine(Path.GetTempPath(), Guid.NewGuid().ToString());
Directory.CreateDirectory(tempDir);

using HttpListener listener = StartFakeOllama(out string endpoint);
try
{
AbsoluteFilePath good = Path.Combine(tempDir, "good.jpg").As<AbsoluteFilePath>();
AbsoluteFilePath bad = Path.Combine(tempDir, "bad.jpg").As<AbsoluteFilePath>();
AbsoluteFilePath gone = Path.Combine(tempDir, "gone.jpg").As<AbsoluteFilePath>();
File.WriteAllBytes(good.WeakString, [0xFF, 0xD8]);
File.WriteAllBytes(bad.WeakString, [0xFF, 0xD8]);

Dictionary<string, List<AbsoluteFilePath>> newHashPaths = new()
{
["good"] = [good],
["bad"] = [bad],
["gone"] = [gone],
};
ConcurrentBag<ImageDescription> stored = [];

IReadOnlyList<string> failures = Scan.DescribeImages(
newHashPaths,
endpoint.As<OllamaEndpoint>(),
"test-model".As<OllamaModelName>(),
"Describe.",
"Name it.",
maxConcurrency: 1,
stored.Add);

Assert.HasCount(1, stored);
Assert.AreEqual("good", stored.Single().Hash);
Assert.AreEqual("a dog on a beach", stored.Single().Description);
Assert.HasCount(2, failures);
Assert.Contains(f => f.Contains("bad.jpg", StringComparison.Ordinal) && f.Contains(nameof(JsonException), StringComparison.Ordinal), failures);
Assert.Contains(f => f.Contains("gone.jpg", StringComparison.Ordinal), failures);
}
finally
{
listener.Stop();
Directory.Delete(tempDir, true);
}
}

[TestMethod]
public void ScanReportsFailedImagesAndCompletes()
{
string tempDir = Path.Combine(Path.GetTempPath(), Guid.NewGuid().ToString());
Directory.CreateDirectory(tempDir);
PersistentState originalSettings = Program.Settings;
TextWriter originalOut = Console.Out;

using HttpListener listener = StartFakeOllama(out string endpoint);
using StringWriter output = new();
try
{
// Every image fails, so nothing is described and nothing is saved to the real app data.
File.WriteAllBytes(Path.Combine(tempDir, "bad.jpg"), [0xFF, 0xD8]);
Program.Settings = new PersistentState();
Console.SetOut(output);

Scan scan = new() { PathString = tempDir, EndpointString = endpoint, ModelString = "test-model" };
scan.Run(scan);
}
finally
{
Console.SetOut(originalOut);
Program.Settings = originalSettings;
listener.Stop();
Directory.Delete(tempDir, true);
}

string text = output.ToString();
Assert.Contains("Failed to describe 1 image(s):", text);
Assert.Contains("bad.jpg: JsonException", text);
Assert.Contains("Scan complete.", text);
}

/// <summary>
/// Serves /api/generate like Ollama, except that a request mentioning bad.jpg gets an HTML
/// page, as a proxy or the wrong service on the endpoint would return.
/// </summary>
private static HttpListener StartFakeOllama(out string endpoint)
{
using TcpListener portFinder = new(IPAddress.Loopback, 0);
portFinder.Start();
int port = ((IPEndPoint)portFinder.LocalEndpoint).Port;
portFinder.Stop();

endpoint = $"http://localhost:{port}";
HttpListener listener = new();
listener.Prefixes.Add($"{endpoint}/");
listener.Start();

_ = Task.Run(async () =>
{
while (listener.IsListening)
{
HttpListenerContext context;
try
{
context = await listener.GetContextAsync().ConfigureAwait(false);
}
catch (HttpListenerException)
{
return;
}
catch (ObjectDisposedException)
{
return;
}

using StreamReader reader = new(context.Request.InputStream);
string body = await reader.ReadToEndAsync().ConfigureAwait(false);
bool isBad = body.Contains("bad.jpg", StringComparison.Ordinal);
context.Response.ContentType = isBad ? "text/html" : "application/json";
byte[] response = Encoding.UTF8.GetBytes(isBad ? "<html>Not Ollama</html>" : "{\"response\":\"a dog on a beach\"}");
await context.Response.OutputStream.WriteAsync(response).ConfigureAwait(false);
context.Response.Close();
}
});

return listener;
}
}
17 changes: 16 additions & 1 deletion ImageDescriber/ImageHasher.cs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,22 @@ internal static Dictionary<AbsoluteFilePath, string> HashFiles(IReadOnlyList<Abs

Parallel.ForEach(filePaths, filePath =>
{
string hash = ComputeHash(filePath);
string hash;
try
{
hash = ComputeHash(filePath);
}
catch (Exception ex) when (ex is IOException or UnauthorizedAccessException)
{
// One unreadable, locked or dangling file must not abort hashing the rest.
lock (ConsoleLock)
{
Console.WriteLine($" Warning: could not hash {filePath}: {ex.Message}");
}

return;
}

results[filePath] = hash;

lock (ConsoleLock)
Expand Down
90 changes: 72 additions & 18 deletions ImageDescriber/Verbs/Scan.cs
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,11 @@

namespace ktsu.ImageDescriber.Verbs;

using System.Collections.Concurrent;
using System.Collections.Generic;
using System.IO;
using System.Net.Http;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;

Expand Down Expand Up @@ -118,17 +121,72 @@ internal override void Run(Scan options)
Console.WriteLine();

// Step 5: Describe new images with configurable concurrency
string descriptionPrompt = Program.Settings.DescriptionPrompt;
string fileNamePrompt = Program.Settings.SuggestedFileNamePrompt;
int maxConcurrency = Math.Max(1, Program.Settings.MaxConcurrentRequests);
int current = 0;
int total = newHashPaths.Count;
Lock consoleLock = new();
Lock saveLock = new();

Console.WriteLine($"Processing with {maxConcurrency} concurrent request(s)...");
Console.WriteLine();

IReadOnlyList<string> failures = DescribeImages(
newHashPaths,
options.Endpoint,
options.Model,
Program.Settings.DescriptionPrompt,
Program.Settings.SuggestedFileNamePrompt,
maxConcurrency,
entry =>
{
lock (saveLock)
{
Program.Settings.Descriptions[entry.Hash] = entry;
Program.Settings.Save();
}
});

PrintFailureSummary(failures);

Console.WriteLine();
Console.WriteLine("Scan complete.");
Console.WriteLine($"Total descriptions in database: {Program.Settings.Descriptions.Count}");

PathString = ".";
}

private static void PrintFailureSummary(IReadOnlyList<string> failures)
{
if (failures.Count == 0)
{
return;
}

Console.WriteLine();
Console.WriteLine($"Failed to describe {failures.Count} image(s):");
foreach (string failure in failures)
{
Console.WriteLine($" {failure}");
}
}

/// <summary>
/// Describes each new image and hands every finished entry to <paramref name="store"/>.
/// A failure on one image is logged and returned rather than thrown, so one bad file or one
/// bad or timed-out model response does not abort a run that may have been going for hours.
/// </summary>
/// <returns>One line per image that could not be described.</returns>
internal static IReadOnlyList<string> DescribeImages(
Dictionary<string, List<AbsoluteFilePath>> newHashPaths,
OllamaEndpoint endpoint,
OllamaModelName model,
string descriptionPrompt,
string fileNamePrompt,
int maxConcurrency,
Action<ImageDescription> store)
{
int current = 0;
int total = newHashPaths.Count;
Lock consoleLock = new();
ConcurrentQueue<string> failures = new();

ParallelOptions parallelOptions = new() { MaxDegreeOfParallelism = maxConcurrency };

Parallel.ForEach(newHashPaths, parallelOptions, kvp =>
Expand All @@ -146,10 +204,10 @@ internal override void Run(Scan options)
{
string pathContext = string.Join("\n", paths.Select(p => p.WeakString));
string fullPrompt = $"Known file paths for this image:\n{pathContext}\n\n{descriptionPrompt}";
string description = OllamaClient.DescribeImageAsync(options.Endpoint, options.Model, fullPrompt, filePath).GetAwaiter().GetResult();
string description = OllamaClient.DescribeImageAsync(endpoint, model, fullPrompt, filePath).GetAwaiter().GetResult();

string combinedFileNamePrompt = $"Image description: {description}\n\n{fileNamePrompt}";
string rawSuggestion = OllamaClient.GenerateAsync(options.Endpoint, options.Model, combinedFileNamePrompt).GetAwaiter().GetResult();
string rawSuggestion = OllamaClient.GenerateAsync(endpoint, model, combinedFileNamePrompt).GetAwaiter().GetResult();
FileName suggestedFileName = SanitizeFileName(rawSuggestion, filePath.FileExtension);

ImageDescription entry = new()
Expand All @@ -158,36 +216,32 @@ internal override void Run(Scan options)
KnownPaths = [.. paths],
Description = description,
SuggestedFileName = suggestedFileName,
Model = options.Model,
Model = model,
DescribedAt = DateTime.UtcNow,
FileSizeBytes = new FileInfo(filePath.WeakString).Length,
};

lock (saveLock)
{
Program.Settings.Descriptions[hash] = entry;
Program.Settings.Save();
}
store(entry);

lock (consoleLock)
{
Console.WriteLine($" [{index}/{total}] Suggested: {suggestedFileName}");
Console.WriteLine($" [{index}/{total}] Done: {description[..Math.Min(80, description.Length)]}...");
}
}
catch (HttpRequestException ex)
catch (Exception ex) when (ex is HttpRequestException or JsonException or OperationCanceledException or IOException or UnauthorizedAccessException)
{
// OperationCanceledException covers the TaskCanceledException HttpClient throws on timeout.
string failure = $"{filePath}: {ex.GetType().Name}: {ex.Message}";
failures.Enqueue(failure);
lock (consoleLock)
{
Console.WriteLine($" [{index}/{total}] Error describing {filePath.FileName}: {ex.Message}");
}
}
});

Console.WriteLine("Scan complete.");
Console.WriteLine($"Total descriptions in database: {Program.Settings.Descriptions.Count}");

PathString = ".";
return [.. failures];
}

internal static FileName SanitizeFileName(string rawSuggestion, FileExtension extension)
Expand Down
Loading