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,
|
|
|
|
|
bool multiModal = false,
|
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);
|
2023-09-19 16:29:25 +00:00
|
|
|
|
|
|
|
|
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);
|
2023-09-19 16:29:25 +00:00
|
|
|
|
|
|
|
|
return completer;
|
2023-09-09 15:37:38 +00:00
|
|
|
}
|
2023-10-09 22:28:17 +00:00
|
|
|
|
2024-03-25 16:22:43 +00:00
|
|
|
private static (string, string) GetProviderAndModel(IServiceProvider services,
|
|
|
|
|
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,
|
|
|
|
|
bool multiModal = 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"))
|
|
|
|
|
{
|
|
|
|
|
model = state.GetState("model", model ?? "gpt-35-turbo-4k");
|
|
|
|
|
}
|
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-05-14 18:34:10 +00:00
|
|
|
model = llmProviderService.GetProviderModel(provider, modelIdentity, multiModal)?.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);
|
|
|
|
|
}
|
|
|
|
|
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);
|
|
|
|
|
|
2023-10-09 22:28:17 +00:00
|
|
|
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;
|
|
|
|
|
}
|
2023-09-09 15:37:38 +00:00
|
|
|
}
|