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

177 lines
5.3 KiB
C#
Raw Normal View History

2024-06-26 03:38:01 +00:00
using OpenAI.Images;
2024-06-24 19:32:52 +00:00
namespace BotSharp.Plugin.AzureOpenAI.Providers;
public class ImageGenerationProvider : IImageGeneration
{
protected readonly AzureOpenAiSettings _settings;
protected readonly IServiceProvider _services;
protected readonly ILogger _logger;
protected string _model;
public virtual string Provider => "azure-openai";
public ImageGenerationProvider(
AzureOpenAiSettings settings,
ILogger<ImageGenerationProvider> logger,
IServiceProvider services)
{
_settings = settings;
_services = services;
_logger = logger;
}
public async Task<RoleDialogModel> GetImageGeneration(Agent agent, List<RoleDialogModel> conversations)
{
var contentHooks = _services.GetServices<IContentGeneratingHook>().ToList();
// Before
foreach (var hook in contentHooks)
{
await hook.BeforeGenerating(agent, conversations);
}
var client = ProviderHelper.GetClient(Provider, _model, _services);
2024-06-26 03:38:01 +00:00
var (prompt, options) = PrepareOptions(conversations);
var imageClient = client.GetImageClient(_model);
ImageGenerationOptions myoptions = new()
{
Quality = GeneratedImageQuality.High,
Size = GeneratedImageSize.W1792xH1024,
Style = GeneratedImageStyle.Vivid,
ResponseFormat = GeneratedImageFormat.Bytes
};
var response = imageClient.GenerateImage(prompt, myoptions);
var imageUri = response.Value.ImageUri;
var revisedPrompt = response.Value.RevisedPrompt;
2024-06-24 19:32:52 +00:00
var content = string.Empty;
2024-06-26 03:38:01 +00:00
if (!string.IsNullOrEmpty(revisedPrompt))
2024-06-24 19:32:52 +00:00
{
2024-06-26 03:38:01 +00:00
content = revisedPrompt;
2024-06-24 19:32:52 +00:00
}
var responseMessage = new RoleDialogModel(AgentRole.Assistant, content)
{
CurrentAgentId = agent.Id,
MessageId = conversations.LastOrDefault()?.MessageId ?? string.Empty,
2024-06-26 03:38:01 +00:00
Data = imageUri.AbsoluteUri
2024-06-24 19:32:52 +00:00
};
// After
foreach (var hook in contentHooks)
{
await hook.AfterGenerated(responseMessage, new TokenStatsModel
{
2024-06-26 03:38:01 +00:00
Prompt = prompt,
2024-06-24 19:32:52 +00:00
Provider = Provider,
Model = _model,
2024-06-26 03:38:01 +00:00
PromptCount = prompt.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count(),
2024-06-24 19:32:52 +00:00
CompletionCount = content.Split(' ', StringSplitOptions.RemoveEmptyEntries).Count()
});
}
return responseMessage;
}
2024-06-26 03:38:01 +00:00
private (string, ImageGenerationOptions) PrepareOptions(List<RoleDialogModel> conversations)
2024-06-24 19:32:52 +00:00
{
2024-06-26 03:38:01 +00:00
var prompt = conversations.LastOrDefault()?.Payload ?? conversations.LastOrDefault()?.Content ?? string.Empty;
2024-06-24 19:32:52 +00:00
2024-06-26 03:38:01 +00:00
var state = _services.GetRequiredService<IConversationStateService>();
var size = state.GetState("image_size");
var quality = state.GetState("image_quality");
var style = state.GetState("image_style");
2024-06-24 19:32:52 +00:00
var options = new ImageGenerationOptions
{
2024-06-26 03:38:01 +00:00
//Size = GetImageSize(size),
//Quality = GetImageQuality(quality),
//Style = GetImageStyle(style),
//ResponseFormat = GeneratedImageFormat.Uri
2024-06-24 19:32:52 +00:00
};
2024-06-26 03:38:01 +00:00
return (prompt, options);
2024-06-24 19:32:52 +00:00
}
public void SetModelName(string model)
{
_model = model;
}
2024-06-26 03:38:01 +00:00
private GeneratedImageSize GetImageSize(string size)
{
var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";
GeneratedImageSize retSize;
switch (value)
{
case "256x256":
retSize = GeneratedImageSize.W256xH256;
break;
case "512x512":
retSize = GeneratedImageSize.W512xH512;
break;
case "1024x1024":
retSize = GeneratedImageSize.W1024xH1024;
break;
case "1024x1792":
retSize = GeneratedImageSize.W1024xH1792;
break;
case "1792x1024":
retSize = GeneratedImageSize.W1792xH1024;
break;
default:
retSize = GeneratedImageSize.W1024xH1024;
break;
}
return retSize;
}
private GeneratedImageQuality GetImageQuality(string quality)
{
var value = !string.IsNullOrEmpty(quality) ? quality : "standard";
GeneratedImageQuality retQuality;
switch (value)
{
case "standard":
retQuality = GeneratedImageQuality.Standard;
break;
case "hd":
retQuality = GeneratedImageQuality.High;
break;
default:
retQuality = GeneratedImageQuality.Standard;
break;
}
return retQuality;
}
private GeneratedImageStyle GetImageStyle(string style)
{
var value = !string.IsNullOrEmpty(style) ? style : "natural";
GeneratedImageStyle retStyle;
switch (value)
{
case "standard":
retStyle = GeneratedImageStyle.Natural;
break;
case "vivid":
retStyle = GeneratedImageStyle.Vivid;
break;
default:
retStyle = GeneratedImageStyle.Natural;
break;
}
return retStyle;
}
2024-06-24 19:32:52 +00:00
}