Skip to content
Closed
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
189 changes: 189 additions & 0 deletions src/Memorizer.UnitTests/EmbeddingApiClientTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,189 @@
using System.Net;
using System.Text;
using System.Text.Json;
using Memorizer.Models;
using Memorizer.Services;
using Memorizer.Settings;
using Microsoft.Extensions.Logging.Abstractions;
using Microsoft.Extensions.Options;

namespace Memorizer.UnitTests;

public class EmbeddingApiClientTests
{
[Fact]
public async Task GenerateAsync_OllamaProvider_PostsToApiEmbeddings_WithModelAndPrompt()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new
{
embedding = new float[] { 0.1f, 0.2f, 0.3f }
}));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = ProviderNames.Ollama,
ApiUrl = new Uri("http://embed.local"),
Model = "all-minilm"
});

var result = await client.GenerateAsync("all-minilm", "hello world");

Assert.Equal([0.1f, 0.2f, 0.3f], result);
Assert.NotNull(handler.LastRequest);
Assert.Equal("/api/embeddings", handler.LastRequest!.RequestUri!.AbsolutePath);

var body = JsonDocument.Parse(handler.LastRequestBody!);
Assert.Equal("all-minilm", body.RootElement.GetProperty("model").GetString());
Assert.Equal("hello world", body.RootElement.GetProperty("prompt").GetString());
Assert.False(body.RootElement.TryGetProperty("input", out _));
Assert.Null(handler.LastRequest.Headers.Authorization);
}

[Fact]
public async Task GenerateAsync_OpenAIProvider_PostsToV1Embeddings_WithModelAndInput_AndBearerAuth()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new
{
data = new[]
{
new { embedding = new float[] { 0.4f, 0.5f, 0.6f }, index = 0 }
}
}));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = ProviderNames.OpenAI,
ApiUrl = new Uri("https://api.openai.com"),
Model = "text-embedding-3-small",
ApiKey = "sk-test-key"
});

var result = await client.GenerateAsync("text-embedding-3-small", "hello world");

Assert.Equal([0.4f, 0.5f, 0.6f], result);
Assert.NotNull(handler.LastRequest);
Assert.Equal("/v1/embeddings", handler.LastRequest!.RequestUri!.AbsolutePath);

var body = JsonDocument.Parse(handler.LastRequestBody!);
Assert.Equal("text-embedding-3-small", body.RootElement.GetProperty("model").GetString());
Assert.Equal("hello world", body.RootElement.GetProperty("input").GetString());
Assert.False(body.RootElement.TryGetProperty("prompt", out _));

Assert.NotNull(handler.LastRequest.Headers.Authorization);
Assert.Equal("Bearer", handler.LastRequest.Headers.Authorization!.Scheme);
Assert.Equal("sk-test-key", handler.LastRequest.Headers.Authorization.Parameter);
}

[Fact]
public async Task GenerateAsync_OpenAIProvider_WithoutApiKey_DoesNotSendAuthHeader()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new
{
data = new[] { new { embedding = new float[] { 1f }, index = 0 } }
}));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = ProviderNames.OpenAI,
ApiUrl = new Uri("http://localhost:8000"),
Model = "local-model",
ApiKey = null
});

await client.GenerateAsync("local-model", "probe");

Assert.Null(handler.LastRequest!.Headers.Authorization);
}

[Fact]
public async Task GenerateAsync_ProviderMatchIsCaseInsensitive()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new
{
data = new[] { new { embedding = new float[] { 1f }, index = 0 } }
}));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = "OpenAI",
ApiUrl = new Uri("https://api.openai.com"),
Model = "text-embedding-3-small"
});

await client.GenerateAsync("text-embedding-3-small", "x");

Assert.Equal("/v1/embeddings", handler.LastRequest!.RequestUri!.AbsolutePath);
}

[Fact]
public async Task GenerateAsync_OllamaProvider_EmptyResponse_Throws()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new { embedding = Array.Empty<float>() }));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = ProviderNames.Ollama,
ApiUrl = new Uri("http://embed.local"),
Model = "m"
});

await Assert.ThrowsAsync<InvalidOperationException>(() => client.GenerateAsync("m", "t"));
}

[Fact]
public async Task GenerateAsync_OpenAIProvider_NoData_Throws()
{
var handler = new RecordingHandler(JsonSerializer.Serialize(new { data = Array.Empty<object>() }));

var client = CreateClient(handler, new EmbeddingSettings
{
Provider = ProviderNames.OpenAI,
ApiUrl = new Uri("https://api.openai.com"),
Model = "m"
});

await Assert.ThrowsAsync<InvalidOperationException>(() => client.GenerateAsync("m", "t"));
}

private static EmbeddingApiClient CreateClient(HttpMessageHandler handler, EmbeddingSettings settings)
{
var httpClient = new HttpClient(handler);
var snapshot = new TestOptionsSnapshot<EmbeddingSettings>(settings);
return new EmbeddingApiClient(httpClient, snapshot, NullLogger<EmbeddingApiClient>.Instance);
}

private sealed class RecordingHandler : HttpMessageHandler
{
private readonly string _responseBody;

public RecordingHandler(string responseBody)
{
_responseBody = responseBody;
}

public HttpRequestMessage? LastRequest { get; private set; }
public string? LastRequestBody { get; private set; }

protected override async Task<HttpResponseMessage> SendAsync(
HttpRequestMessage request,
CancellationToken cancellationToken)
{
LastRequest = request;
LastRequestBody = request.Content is null
? null
: await request.Content.ReadAsStringAsync(cancellationToken);

return new HttpResponseMessage(HttpStatusCode.OK)
{
Content = new StringContent(_responseBody, Encoding.UTF8, "application/json")
};
}
}

private sealed class TestOptionsSnapshot<T> : IOptionsSnapshot<T> where T : class
{
public TestOptionsSnapshot(T value) { Value = value; }
public T Value { get; }
public T Get(string? name) => Value;
}
}
6 changes: 6 additions & 0 deletions src/Memorizer/Controllers/ConfigurationController.cs
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,10 @@ public ActionResult<SystemConfiguration> GetConfiguration()
{
EmbeddingSettings = new EmbeddingConfigDto
{
Provider = EmbeddingSettingsValue.Provider,
ApiUrl = EmbeddingSettingsValue.ApiUrl?.ToString(),
ApiConfigured = EmbeddingSettingsValue.ApiUrl != null,
ApiKeyConfigured = !string.IsNullOrWhiteSpace(EmbeddingSettingsValue.ApiKey),
Model = EmbeddingSettingsValue.Model,
Timeout = EmbeddingSettingsValue.Timeout.ToString()
},
Expand Down Expand Up @@ -80,7 +83,10 @@ public class SystemConfiguration

public class EmbeddingConfigDto
{
public string Provider { get; set; } = string.Empty;
public string? ApiUrl { get; set; }
public bool ApiConfigured { get; set; }
public bool ApiKeyConfigured { get; set; }
public string Model { get; set; } = string.Empty;
public string Timeout { get; set; } = string.Empty;
}
Expand Down
11 changes: 6 additions & 5 deletions src/Memorizer/Extensions/ServiceCollectionExtensions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -38,12 +38,13 @@ public static IServiceCollection AddEmbeddings(
.BindConfiguration("Embeddings")
.ValidateOnStart();

// Register EmbeddingService as Scoped so it gets fresh settings on each request
// Note: HttpClient is configured with base address in the service constructor
services.AddHttpClient<IEmbeddingService, EmbeddingService>();
// Single typed HttpClient that handles provider-specific request/response shapes
// (Ollama and OpenAI-compatible). EmbeddingService and EmbeddingDimensionService
// consume it instead of holding their own HttpClient.
services.AddHttpClient<IEmbeddingApiClient, EmbeddingApiClient>();

// Register EmbeddingDimensionService with its own HttpClient (singleton is fine here)
services.AddHttpClient<IEmbeddingDimensionService, EmbeddingDimensionService>();
services.AddScoped<IEmbeddingService, EmbeddingService>();
services.AddScoped<IEmbeddingDimensionService, EmbeddingDimensionService>();

// Register dimension mismatch state holder for UI warnings
services.AddSingleton<IDimensionMismatchState, DimensionMismatchState>();
Expand Down
27 changes: 27 additions & 0 deletions src/Memorizer/Models/OpenAIEmbeddingRequest.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
using System.Text.Json.Serialization;

namespace Memorizer.Models;

public class OpenAIEmbeddingRequest
{
[JsonPropertyName("model")]
public string Model { get; init; } = string.Empty;

[JsonPropertyName("input")]
public string Input { get; init; } = string.Empty;
}

public class OpenAIEmbeddingResponse
{
[JsonPropertyName("data")]
public OpenAIEmbeddingData[] Data { get; init; } = [];
}

public class OpenAIEmbeddingData
{
[JsonPropertyName("embedding")]
public float[] Embedding { get; init; } = [];

[JsonPropertyName("index")]
public int Index { get; init; }
}
92 changes: 92 additions & 0 deletions src/Memorizer/Services/EmbeddingApiClient.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
using System.Net.Http.Headers;
using Memorizer.Models;
using Memorizer.Settings;
using Microsoft.Extensions.Options;

namespace Memorizer.Services;

/// <summary>
/// Sends embedding requests to the configured provider and normalises the response
/// to a <c>float[]</c>. Supports Ollama's native <c>/api/embeddings</c> shape and the
/// OpenAI-compatible <c>/v1/embeddings</c> shape (also exposed by LiteLLM, vLLM, Azure
/// OpenAI, LocalAI, etc.).
/// </summary>
public interface IEmbeddingApiClient
{
Task<float[]> GenerateAsync(string model, string text, CancellationToken cancellationToken = default);
}

public sealed class EmbeddingApiClient : IEmbeddingApiClient
{
private readonly HttpClient _httpClient;
private readonly IOptionsSnapshot<EmbeddingSettings> _settingsSnapshot;
private readonly ILogger<EmbeddingApiClient> _logger;

private EmbeddingSettings Settings => _settingsSnapshot.Value;

public EmbeddingApiClient(
HttpClient httpClient,
IOptionsSnapshot<EmbeddingSettings> settingsSnapshot,
ILogger<EmbeddingApiClient> logger)
{
_httpClient = httpClient;
_settingsSnapshot = settingsSnapshot;
_logger = logger;

_httpClient.BaseAddress = Settings.ApiUrl;
_httpClient.Timeout = Settings.Timeout;

if (!string.IsNullOrWhiteSpace(Settings.ApiKey))
{
_httpClient.DefaultRequestHeaders.Authorization =
new AuthenticationHeaderValue("Bearer", Settings.ApiKey);
}
}

public async Task<float[]> GenerateAsync(string model, string text, CancellationToken cancellationToken = default)
{
var provider = Settings.Provider;

if (string.Equals(provider, ProviderNames.OpenAI, StringComparison.OrdinalIgnoreCase))
{
return await GenerateOpenAIAsync(model, text, cancellationToken);
}

return await GenerateOllamaAsync(model, text, cancellationToken);
}

private async Task<float[]> GenerateOllamaAsync(string model, string text, CancellationToken cancellationToken)
{
var request = new EmbeddingRequest { Model = model, Prompt = text };

_logger.LogDebug("Sending Ollama embedding request to {ApiUrl}", Settings.ApiUrl);
var response = await _httpClient.PostAsJsonAsync("api/embeddings", request, cancellationToken);
response.EnsureSuccessStatusCode();

var result = await response.Content.ReadFromJsonAsync<EmbeddingResponse>(cancellationToken: cancellationToken);
if (result?.Embedding is null || result.Embedding.Length == 0)
{
throw new InvalidOperationException("Empty embedding response from Ollama API");
}

return result.Embedding;
}

private async Task<float[]> GenerateOpenAIAsync(string model, string text, CancellationToken cancellationToken)
{
var request = new OpenAIEmbeddingRequest { Model = model, Input = text };

_logger.LogDebug("Sending OpenAI-compatible embedding request to {ApiUrl}", Settings.ApiUrl);
var response = await _httpClient.PostAsJsonAsync("v1/embeddings", request, cancellationToken);
response.EnsureSuccessStatusCode();

var result = await response.Content.ReadFromJsonAsync<OpenAIEmbeddingResponse>(cancellationToken: cancellationToken);
var first = result?.Data.FirstOrDefault();
if (first?.Embedding is null || first.Embedding.Length == 0)
{
throw new InvalidOperationException("Empty embedding response from OpenAI-compatible API");
}

return first.Embedding;
}
}
Loading