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

102 lines
3.5 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
.Where(x => GetProviderModels(x.Provider).Any())
.Select(x => x.Provider));
var services2 = _services.GetServices<IChatCompletion>();
providers.AddRange(services2
.Where(x => GetProviderModels(x.Provider).Any())
.Select(x => x.Provider));
var services3 = _services.GetServices<ITextEmbedding>();
providers.AddRange(services3
.Where(x => GetProviderModels(x.Provider).Any())
.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-06-24 19:32:52 +00:00
public LlmModelSetting GetProviderModel(string provider, string id, bool? multiModal = null, bool imageGenerate = false)
2024-02-02 22:36:05 +00:00
{
var models = GetProviderModels(provider)
2024-05-16 21:32:58 +00:00
.Where(x => x.Id == id);
if (multiModal.HasValue)
{
models = models.Where(x => x.MultiModal == multiModal);
}
2024-02-02 22:36:05 +00:00
2024-06-24 19:32:52 +00:00
models = models.Where(x => x.ImageGeneration == imageGenerate);
2024-02-02 22:36:05 +00:00
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;
}
}