Migrated the get secret calls from Tauri JS to the runtime API

This commit is contained in:
tsommer committed 2024-08-28 21:00:19 +02:00
1 parent 273376ad97
commit 97854e7eaa
53 files changed
+299 -243

No files matched your search

@@ -25,7 +25,7 @@ public sealed class ProviderAnthropic(ILogger logger) : BaseProvider("https://ap
public async IAsyncEnumerable<string> StreamChatCompletion(IJSRuntime jsRuntime, SettingsManager settings, Model chatModel, ChatThread chatThread, [EnumeratorCancellation] CancellationToken token = default)
{
// Get the API key:
var requestedSecret = await settings.GetAPIKey(jsRuntime, this);
var requestedSecret = await RUST_SERVICE.GetAPIKey(this);
if(!requestedSecret.Success)
yield break;
@@ -64,7 +64,7 @@ public sealed class ProviderAnthropic(ILogger logger) : BaseProvider("https://ap
var request = new HttpRequestMessage(HttpMethod.Post, "messages");
// Set the authorization header:
request.Headers.Add("x-api-key", requestedSecret.Secret);
request.Headers.Add("x-api-key", await requestedSecret.Secret.Decrypt(ENCRYPTION));
// Set the Anthropic version:
request.Headers.Add("anthropic-version", "2023-06-01");
@@ -1,3 +1,5 @@
using RustService = AIStudio.Tools.RustService;
namespace AIStudio.Provider;
/// <summary>
@@ -15,6 +17,16 @@ public abstract class BaseProvider
/// </summary>
protected readonly ILogger logger;
static BaseProvider()
{
RUST_SERVICE = Program.RUST_SERVICE;
ENCRYPTION = Program.ENCRYPTION;
}
protected static readonly RustService RUST_SERVICE;
protected static readonly Encryption ENCRYPTION;
/// <summary>
/// Constructor for the base provider.
/// </summary>
@@ -27,7 +27,7 @@ public class ProviderFireworks(ILogger logger) : BaseProvider("https://api.firew
public async IAsyncEnumerable<string> StreamChatCompletion(IJSRuntime jsRuntime, SettingsManager settings, Model chatModel, ChatThread chatThread, [EnumeratorCancellation] CancellationToken token = default)
{
// Get the API key:
var requestedSecret = await settings.GetAPIKey(jsRuntime, this);
var requestedSecret = await RUST_SERVICE.GetAPIKey(this);
if(!requestedSecret.Success)
yield break;
@@ -73,7 +73,7 @@ public class ProviderFireworks(ILogger logger) : BaseProvider("https://api.firew
var request = new HttpRequestMessage(HttpMethod.Post, "chat/completions");
// Set the authorization header:
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", requestedSecret.Secret);
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", await requestedSecret.Secret.Decrypt(ENCRYPTION));
// Set the content:
request.Content = new StringContent(fireworksChatRequest, Encoding.UTF8, "application/json");
@@ -26,7 +26,7 @@ public sealed class ProviderMistral(ILogger logger) : BaseProvider("https://api.
public async IAsyncEnumerable<string> StreamChatCompletion(IJSRuntime jsRuntime, SettingsManager settings, Provider.Model chatModel, ChatThread chatThread, [EnumeratorCancellation] CancellationToken token = default)
{
// Get the API key:
var requestedSecret = await settings.GetAPIKey(jsRuntime, this);
var requestedSecret = await RUST_SERVICE.GetAPIKey(this);
if(!requestedSecret.Success)
yield break;
@@ -75,7 +75,7 @@ public sealed class ProviderMistral(ILogger logger) : BaseProvider("https://api.
var request = new HttpRequestMessage(HttpMethod.Post, "chat/completions");
// Set the authorization header:
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", requestedSecret.Secret);
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", await requestedSecret.Secret.Decrypt(ENCRYPTION));
// Set the content:
request.Content = new StringContent(mistralChatRequest, Encoding.UTF8, "application/json");
@@ -153,9 +153,9 @@ public sealed class ProviderMistral(ILogger logger) : BaseProvider("https://api.
var secretKey = apiKeyProvisional switch
{
not null => apiKeyProvisional,
_ => await settings.GetAPIKey(jsRuntime, this) switch
_ => await RUST_SERVICE.GetAPIKey(this) switch
{
{ Success: true } result => result.Secret,
{ Success: true } result => await result.Secret.Decrypt(ENCRYPTION),
_ => null,
}
};
@@ -30,7 +30,7 @@ public sealed class ProviderOpenAI(ILogger logger) : BaseProvider("https://api.o
public async IAsyncEnumerable<string> StreamChatCompletion(IJSRuntime jsRuntime, SettingsManager settings, Model chatModel, ChatThread chatThread, [EnumeratorCancellation] CancellationToken token = default)
{
// Get the API key:
var requestedSecret = await settings.GetAPIKey(jsRuntime, this);
var requestedSecret = await RUST_SERVICE.GetAPIKey(this);
if(!requestedSecret.Success)
yield break;
@@ -79,7 +79,7 @@ public sealed class ProviderOpenAI(ILogger logger) : BaseProvider("https://api.o
var request = new HttpRequestMessage(HttpMethod.Post, "chat/completions");
// Set the authorization header:
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", requestedSecret.Secret);
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", await requestedSecret.Secret.Decrypt(ENCRYPTION));
// Set the content:
request.Content = new StringContent(openAIChatRequest, Encoding.UTF8, "application/json");
@@ -154,25 +154,25 @@ public sealed class ProviderOpenAI(ILogger logger) : BaseProvider("https://api.o
/// <inheritdoc />
public Task<IEnumerable<Model>> GetTextModels(IJSRuntime jsRuntime, SettingsManager settings, string? apiKeyProvisional = null, CancellationToken token = default)
{
return this.LoadModels(jsRuntime, settings, "gpt-", token, apiKeyProvisional);
return this.LoadModels("gpt-", token, apiKeyProvisional);
}
/// <inheritdoc />
public Task<IEnumerable<Model>> GetImageModels(IJSRuntime jsRuntime, SettingsManager settings, string? apiKeyProvisional = null, CancellationToken token = default)
{
return this.LoadModels(jsRuntime, settings, "dall-e-", token, apiKeyProvisional);
return this.LoadModels("dall-e-", token, apiKeyProvisional);
}
#endregion
private async Task<IEnumerable<Model>> LoadModels(IJSRuntime jsRuntime, SettingsManager settings, string prefix, CancellationToken token, string? apiKeyProvisional = null)
private async Task<IEnumerable<Model>> LoadModels(string prefix, CancellationToken token, string? apiKeyProvisional = null)
{
var secretKey = apiKeyProvisional switch
{
not null => apiKeyProvisional,
_ => await settings.GetAPIKey(jsRuntime, this) switch
_ => await RUST_SERVICE.GetAPIKey(this) switch
{
{ Success: true } result => result.Secret,
{ Success: true } result => await result.Secret.Decrypt(ENCRYPTION),
_ => null,
}
};
@@ -4,6 +4,8 @@ using AIStudio.Provider.Mistral;
using AIStudio.Provider.OpenAI;
using AIStudio.Provider.SelfHosted;
using RustService = AIStudio.Tools.RustService;
namespace AIStudio.Provider;
public static class ProvidersExtensions
@@ -33,8 +35,9 @@ public static class ProvidersExtensions
/// </summary>
/// <param name="providerSettings">The provider settings.</param>
/// <param name="logger">The logger to use.</param>
/// <param name="rustService">The Rust instance to use.</param>
/// <returns>The provider instance.</returns>
public static IProvider CreateProvider(this Settings.Provider providerSettings, ILogger logger)
public static IProvider CreateProvider(this Settings.Provider providerSettings, ILogger logger, RustService rustService)
{
try
{