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

90 lines
3 KiB
C#
Raw Normal View History

2023-12-13 18:12:25 +00:00
using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.MLTasks.Settings;
using BotSharp.Abstraction.Settings;
2023-12-13 18:12:25 +00:00
namespace BotSharp.Core.Infrastructures;
public class LlmProviderService : ILlmProviderService
2023-12-13 18:12:25 +00:00
{
private readonly IServiceProvider _services;
private readonly ILogger _logger;
public LlmProviderService(IServiceProvider services, ILogger<LlmProviderService> logger)
2023-12-13 18:12:25 +00:00
{
_services = services;
_logger = logger;
}
public List<string> GetProviders()
{
var providers = new List<string>();
var services1 = _services.GetServices<ITextCompletion>();
providers.AddRange(services1.Select(x => x.Provider));
var services2 = _services.GetServices<IChatCompletion>();
providers.AddRange(services2.Select(x => x.Provider));
var services3 = _services.GetServices<ITextEmbedding>();
providers.AddRange(services3.Select(x => x.Provider));
return providers.Distinct().ToList();
}
public List<LlmModelSetting> GetProviderModels(string provider)
{
var settingService = _services.GetRequiredService<ISettingService>();
return settingService.Bind<List<LlmProviderSetting>>($"LlmProviders")
.FirstOrDefault(x => x.Provider.Equals(provider))
?.Models ?? new List<LlmModelSetting>();
}
2024-02-02 22:36:05 +00:00
public LlmModelSetting GetProviderModel(string provider, string id)
{
var models = GetProviderModels(provider)
.Where(x => x.Id == id)
.ToList();
var random = new Random();
var index = random.Next(0, models.Count());
var modelSetting = models.ElementAt(index);
return modelSetting;
}
2023-12-13 18:12:25 +00:00
public LlmModelSetting? GetSetting(string provider, string model)
{
var settings = _services.GetRequiredService<List<LlmProviderSetting>>();
2024-01-22 02:39:13 +00:00
var providerSetting = settings.FirstOrDefault(p =>
p.Provider.Equals(provider, StringComparison.CurrentCultureIgnoreCase));
2023-12-13 18:12:25 +00:00
if (providerSetting == null)
{
_logger.LogError($"Can't find provider settings for {provider}");
return null;
2024-01-22 02:39:13 +00:00
}
2023-12-13 18:12:25 +00:00
2024-01-22 02:39:13 +00:00
var modelSetting = providerSetting.Models.FirstOrDefault(m =>
m.Name.Equals(model, StringComparison.CurrentCultureIgnoreCase));
2023-12-13 18:12:25 +00:00
if (modelSetting == null)
{
_logger.LogError($"Can't find model settings for {provider}.{model}");
return null;
}
2024-01-22 02:39:13 +00:00
// load balancing
if (!string.IsNullOrEmpty(modelSetting.Group))
{
// find the models in the same group
var models = providerSetting.Models
.Where(m => !string.IsNullOrEmpty(m.Group) &&
m.Group.Equals(modelSetting.Group, StringComparison.CurrentCultureIgnoreCase))
.ToList();
// pick one model randomly
var random = new Random();
var index = random.Next(0, models.Count());
modelSetting = models.ElementAt(index);
}
2023-12-13 18:12:25 +00:00
return modelSetting;
}
}