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

View file

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

View file

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

View file

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

View file

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