BotSharp/src/Infrastructure/BotSharp.Core/Infrastructures/CompletionProvider.cs

161 lines
5.7 KiB
C#
Raw Normal View History

2023-09-09 15:37:38 +00:00
using BotSharp.Abstraction.MLTasks;
2023-12-15 16:31:11 +00:00
using BotSharp.Abstraction.MLTasks.Settings;
2023-09-09 15:37:38 +00:00
namespace BotSharp.Core.Infrastructures;
public class CompletionProvider
{
2023-12-15 17:54:44 +00:00
public static object GetCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
AgentLlmConfig? agentConfig = null)
2023-12-15 16:31:11 +00:00
{
2024-03-25 16:22:43 +00:00
var settingsService = services.GetRequiredService<ILlmProviderService>();
2023-12-15 16:31:11 +00:00
2024-03-25 16:22:43 +00:00
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, agentConfig: agentConfig);
2023-12-15 16:31:11 +00:00
var settings = settingsService.GetSetting(provider, model);
2023-12-15 17:54:44 +00:00
if (settings.Type == LlmModelType.Text)
2023-12-15 16:31:11 +00:00
{
2024-03-25 16:22:43 +00:00
return GetTextCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
2023-12-15 16:31:11 +00:00
}
else
{
2024-03-25 16:22:43 +00:00
return GetChatCompletion(services,
provider: provider,
model: model,
agentConfig: agentConfig);
2023-12-15 16:31:11 +00:00
}
}
2023-12-15 17:54:44 +00:00
public static IChatCompletion GetChatCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
2024-05-14 18:34:10 +00:00
string? modelId = null,
2024-05-16 21:32:58 +00:00
bool? multiModal = null,
2023-12-15 17:54:44 +00:00
AgentLlmConfig? agentConfig = null)
2023-09-09 15:37:38 +00:00
{
var completions = services.GetServices<IChatCompletion>();
2024-05-14 18:34:10 +00:00
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, modelId: modelId,
multiModal: multiModal, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{
var logger = services.GetRequiredService<ILogger<CompletionProvider>>();
logger.LogError($"Can't resolve completion provider by {provider}");
}
2024-05-14 16:51:39 +00:00
completer?.SetModelName(model);
return completer;
2023-09-09 15:37:38 +00:00
}
2023-10-09 22:28:17 +00:00
2024-06-24 19:32:52 +00:00
public static ITextCompletion GetTextCompletion(IServiceProvider services,
string? provider = null,
string? model = null,
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<ITextCompletion>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{
var logger = services.GetRequiredService<ILogger<CompletionProvider>>();
logger.LogError($"Can't resolve completion provider by {provider}");
}
completer.SetModelName(model);
return completer;
}
public static IImageGeneration GetImageGeneration(IServiceProvider services,
string? provider = null,
string? model = null,
string? modelId = null,
bool imageGenerate = false,
AgentLlmConfig? agentConfig = null)
{
var completions = services.GetServices<IImageGeneration>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model, modelId: modelId,
imageGenerate: imageGenerate, agentConfig: agentConfig);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{
var logger = services.GetRequiredService<ILogger<CompletionProvider>>();
logger.LogError($"Can't resolve completion provider by {provider}");
}
completer?.SetModelName(model);
return completer;
}
2024-07-01 02:14:02 +00:00
public static ITextEmbedding GetTextEmbedding(IServiceProvider services,
string? provider = null,
string? model = null)
{
var completions = services.GetServices<ITextEmbedding>();
(provider, model) = GetProviderAndModel(services, provider: provider, model: model);
var completer = completions.FirstOrDefault(x => x.Provider == provider);
if (completer == null)
{
var logger = services.GetRequiredService<ILogger<CompletionProvider>>();
logger.LogError($"Can't resolve completion provider by {provider}");
}
completer.SetModelName(model);
return completer;
}
2024-06-24 19:32:52 +00:00
private static (string, string) GetProviderAndModel(IServiceProvider services,
2024-03-25 16:22:43 +00:00
string? provider = null,
2023-12-15 17:54:44 +00:00
string? model = null,
2024-05-14 18:34:10 +00:00
string? modelId = null,
2024-05-16 21:32:58 +00:00
bool? multiModal = null,
2024-06-24 19:32:52 +00:00
bool imageGenerate = false,
2023-12-15 17:54:44 +00:00
AgentLlmConfig? agentConfig = null)
2023-10-09 22:28:17 +00:00
{
2023-12-15 17:54:44 +00:00
var agentSetting = services.GetRequiredService<AgentSettings>();
2023-10-09 22:28:17 +00:00
var state = services.GetRequiredService<IConversationStateService>();
if (string.IsNullOrEmpty(provider))
{
2023-12-15 17:54:44 +00:00
provider = agentConfig?.Provider ?? agentSetting.LlmConfig?.Provider;
provider = state.GetState("provider", provider ?? "azure-openai");
2023-10-09 22:28:17 +00:00
}
if (string.IsNullOrEmpty(model))
{
2023-12-15 17:54:44 +00:00
model = agentConfig?.Model ?? agentSetting.LlmConfig?.Model;
2024-03-25 16:22:43 +00:00
if (state.ContainsState("model"))
{
2024-06-24 19:32:52 +00:00
model = state.GetState("model", model ?? "dall-e-3");
2024-03-25 16:22:43 +00:00
}
2024-05-14 18:34:10 +00:00
else if (state.ContainsState("model_id") || !string.IsNullOrEmpty(modelId))
2024-03-25 16:22:43 +00:00
{
2024-05-14 18:34:10 +00:00
var modelIdentity = state.ContainsState("model_id") ? state.GetState("model_id") : modelId;
2024-03-25 16:22:43 +00:00
var llmProviderService = services.GetRequiredService<ILlmProviderService>();
2024-06-24 19:32:52 +00:00
model = llmProviderService.GetProviderModel(provider, modelIdentity,
multiModal: multiModal, imageGenerate: imageGenerate)?.Name;
2024-03-25 16:22:43 +00:00
}
2023-10-09 22:28:17 +00:00
}
2024-03-25 16:22:43 +00:00
state.SetState("provider", provider);
state.SetState("model", model);
return (provider, model);
}
2023-09-09 15:37:38 +00:00
}