BotSharp/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/TextCompletionProvider.cs

120 lines
3.6 KiB
C#
Raw Normal View History

2023-06-17 02:42:35 +00:00
using Azure.AI.OpenAI;
using BotSharp.Abstraction.MLTasks;
using System;
using System.Threading.Tasks;
2023-06-19 18:32:49 +00:00
using BotSharp.Plugin.AzureOpenAI.Settings;
using Microsoft.Extensions.Logging;
2023-10-09 22:28:17 +00:00
using BotSharp.Abstraction.Conversations;
using Microsoft.Extensions.DependencyInjection;
using BotSharp.Abstraction.Conversations.Models;
2023-10-13 20:08:59 +00:00
using BotSharp.Abstraction.Agents.Enums;
2023-10-16 20:07:07 +00:00
using System.Linq;
using System.Collections.Generic;
using BotSharp.Abstraction.Agents.Models;
2023-10-25 12:48:11 +00:00
using BotSharp.Abstraction.Conversations.Settings;
2023-06-17 02:42:35 +00:00
2023-06-17 13:32:39 +00:00
namespace BotSharp.Plugin.AzureOpenAI.Providers;
2023-06-17 02:42:35 +00:00
public class TextCompletionProvider : ITextCompletion
{
2023-10-09 22:28:17 +00:00
private readonly IServiceProvider _services;
2023-06-17 02:42:35 +00:00
private readonly AzureOpenAiSettings _settings;
private readonly ILogger _logger;
2023-10-08 20:46:42 +00:00
private string _model;
public string Provider => "azure-openai";
2023-06-17 02:42:35 +00:00
2023-10-09 22:28:17 +00:00
public TextCompletionProvider(IServiceProvider services,
AzureOpenAiSettings settings,
2023-10-16 20:07:07 +00:00
ILogger<TextCompletionProvider> logger)
2023-06-17 02:42:35 +00:00
{
2023-10-09 22:28:17 +00:00
_services = services;
2023-06-17 02:42:35 +00:00
_settings = settings;
_logger = logger;
2023-06-17 02:42:35 +00:00
}
2023-10-30 16:48:18 +00:00
public async Task<string> GetCompletion(string text, string agentId, string messageId)
2023-06-17 02:42:35 +00:00
{
2023-10-16 20:07:07 +00:00
var hooks = _services.GetServices<IContentGeneratingHook>().ToList();
// Before chat completion hook
2023-10-30 16:48:18 +00:00
var agent = new Agent()
{
Id = agentId,
};
var message = new RoleDialogModel(AgentRole.User, text)
{
CurrentAgentId = agentId,
MessageId = messageId
};
2023-10-16 20:07:07 +00:00
Task.WaitAll(hooks.Select(hook =>
2023-10-30 16:48:18 +00:00
hook.BeforeGenerating(agent,
new List<RoleDialogModel>
{
message
2023-10-23 23:20:18 +00:00
})).ToArray());
2023-10-16 20:07:07 +00:00
2023-10-23 23:20:18 +00:00
var client = ProviderHelper.GetClient(_model, _settings);
2023-10-09 22:28:17 +00:00
2023-06-17 02:42:35 +00:00
var completionsOptions = new CompletionsOptions()
{
Prompts =
{
text
},
2023-10-13 20:08:59 +00:00
MaxTokens = 256,
2023-06-17 02:42:35 +00:00
};
2023-10-13 20:08:59 +00:00
completionsOptions.StopSequences.Add($"{AgentRole.Assistant}:");
2023-06-17 02:42:35 +00:00
2023-10-25 12:48:11 +00:00
var setting = _services.GetRequiredService<ConversationSetting>();
if (setting.ShowVerboseLog)
{
_logger.LogInformation(text);
}
2023-10-09 22:28:17 +00:00
var state = _services.GetRequiredService<IConversationStateService>();
var temperature = float.Parse(state.GetState("temperature", "0.5"));
var samplingFactor = float.Parse(state.GetState("sampling_factor", "0.5"));
completionsOptions.Temperature = temperature;
completionsOptions.NucleusSamplingFactor = samplingFactor;
2023-06-17 02:42:35 +00:00
var response = await client.GetCompletionsAsync(
2023-10-23 23:20:18 +00:00
deploymentOrModelName: _model,
2023-06-17 02:42:35 +00:00
completionsOptions);
// OpenAI
var completion = "";
foreach (var t in response.Value.Choices)
{
completion += t.Text;
};
2023-10-25 12:48:11 +00:00
if (setting.ShowVerboseLog)
{
_logger.LogInformation(completion);
}
2023-10-16 20:07:07 +00:00
// After chat completion hook
2023-10-30 16:48:18 +00:00
var responseMessage = new RoleDialogModel(AgentRole.Assistant, completion)
{
CurrentAgentId = agentId,
MessageId = messageId
};
2023-10-16 20:07:07 +00:00
Task.WaitAll(hooks.Select(hook =>
2023-10-30 16:48:18 +00:00
hook.AfterGenerated(responseMessage, new TokenStatsModel
2023-10-16 20:07:07 +00:00
{
2023-10-30 16:48:18 +00:00
Prompt = text,
2023-10-16 20:07:07 +00:00
Model = _model,
PromptCount = response.Value.Usage.PromptTokens,
CompletionCount = response.Value.Usage.CompletionTokens
})).ToArray());
2023-06-27 19:17:53 +00:00
return completion.Trim();
2023-06-17 02:42:35 +00:00
}
2023-10-08 20:46:42 +00:00
public void SetModelName(string model)
{
_model = model;
}
2023-06-17 02:42:35 +00:00
}