Merge pull request #807 from iceljc/master

refine ai client
This commit is contained in:
iceljc 2024-12-26 09:57:22 -06:00 committed by GitHub
commit a8a9b77aa6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 16 additions and 14 deletions

View file

@ -13,7 +13,7 @@ public class GeminiChatCompletionProvider : IChatCompletion
private string _model;
public string Provider => "google-gemini";
public string Provider => "google-ai";
public GeminiChatCompletionProvider(
IServiceProvider services,
@ -33,7 +33,7 @@ public class GeminiChatCompletionProvider : IChatCompletion
await hook.BeforeGenerating(agent, conversations);
}
var client = ProviderHelper.GetGeminiClient(_services);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var aiModel = client.GenerativeModel(_model);
var (prompt, request) = PrepareOptions(aiModel, agent, conversations);

View file

@ -16,7 +16,7 @@ public class PalmChatCompletionProvider : IChatCompletion
private string _model;
public string Provider => "google-ai";
public string Provider => "google-palm";
public PalmChatCompletionProvider(
IServiceProvider services,
@ -36,7 +36,7 @@ public class PalmChatCompletionProvider : IChatCompletion
await hook.BeforeGenerating(agent, conversations);
}
var client = ProviderHelper.GetPalmClient(_services);
var client = ProviderHelper.GetPalmClient(Provider, _model, _services);
var (prompt, messages, hasFunctions) = PrepareOptions(agent, conversations);
RoleDialogModel msg;

View file

@ -5,17 +5,19 @@ namespace BotSharp.Plugin.GoogleAi.Providers;
public static class ProviderHelper
{
public static GoogleAI GetGeminiClient(IServiceProvider services)
public static GoogleAI GetGeminiClient(string provider, string model, IServiceProvider services)
{
var settings = services.GetRequiredService<GoogleAiSettings>();
var client = new GoogleAI(settings.Gemini.ApiKey);
var settingsService = services.GetRequiredService<ILlmProviderService>();
var settings = settingsService.GetSetting(provider, model);
var client = new GoogleAI(settings.ApiKey);
return client;
}
public static GooglePalmClient GetPalmClient(IServiceProvider services)
public static GooglePalmClient GetPalmClient(string provider, string model, IServiceProvider services)
{
var settings = services.GetRequiredService<GoogleAiSettings>();
var client = new GooglePalmClient(settings.PaLM.ApiKey);
var settingsService = services.GetRequiredService<ILlmProviderService>();
var settings = settingsService.GetSetting(provider, model);
var client = new GooglePalmClient(settings.ApiKey);
return client;
}
}

View file

@ -13,7 +13,7 @@ public class GeminiTextCompletionProvider : ITextCompletion
private readonly ITokenStatistics _tokenStatistics;
private string _model;
public string Provider => "google-gemini";
public string Provider => "google-ai";
public GeminiTextCompletionProvider(
IServiceProvider services,
@ -45,7 +45,7 @@ public class GeminiTextCompletionProvider : ITextCompletion
await hook.BeforeGenerating(agent, new List<RoleDialogModel> { userMessage });
}
var client = ProviderHelper.GetGeminiClient(_services);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var aiModel = client.GenerativeModel(_model);
PrepareOptions(aiModel);

View file

@ -13,7 +13,7 @@ public class PalmTextCompletionProvider : ITextCompletion
private string _model;
public string Provider => "google-ai";
public string Provider => "google-palm";
public PalmTextCompletionProvider(
IServiceProvider services,
@ -38,7 +38,7 @@ public class PalmTextCompletionProvider : ITextCompletion
await hook.BeforeGenerating(agent, new List<RoleDialogModel> { userMessage });
}
var client = ProviderHelper.GetPalmClient(_services);
var client = ProviderHelper.GetPalmClient(Provider, _model, _services);
_tokenStatistics.StartTimer();
var response = await client.GenerateTextAsync(text, null);
_tokenStatistics.StopTimer();