diff --git a/src/Infrastructure/BotSharp.Abstraction/Conversations/Models/IncomingMessageModel.cs b/src/Infrastructure/BotSharp.Abstraction/Conversations/Models/IncomingMessageModel.cs index 294fab04..430c44ad 100644 --- a/src/Infrastructure/BotSharp.Abstraction/Conversations/Models/IncomingMessageModel.cs +++ b/src/Infrastructure/BotSharp.Abstraction/Conversations/Models/IncomingMessageModel.cs @@ -12,6 +12,4 @@ public class IncomingMessageModel : MessageConfig /// Postback message /// public PostbackMessageModel? Postback { get; set; } - - public List Files { get; set; } = new List(); } diff --git a/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs b/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs index 8faed9cc..b6e4e1e0 100644 --- a/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs +++ b/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs @@ -45,7 +45,9 @@ public interface IBotSharpFileService #region Image Task GenerateImage(string? provider, string? model, string text); - Task VarifyImage(string? provider, string? model, BotSharpFile file); + Task VaryImage(string? provider, string? model, BotSharpFile image); + Task EditImage(string? provider, string? model, string text, BotSharpFile image); + Task EditImage(string? provider, string? model, string text, BotSharpFile image, BotSharpFile mask); #endregion #region Pdf diff --git a/src/Infrastructure/BotSharp.Abstraction/Files/Models/InputMessageFiles.cs b/src/Infrastructure/BotSharp.Abstraction/Files/Models/InputMessageFiles.cs new file mode 100644 index 00000000..29749bce --- /dev/null +++ b/src/Infrastructure/BotSharp.Abstraction/Files/Models/InputMessageFiles.cs @@ -0,0 +1,7 @@ +namespace BotSharp.Abstraction.Files.Models; + +public class InputMessageFiles +{ + public List Files { get; set; } = new List(); + public BotSharpFile? Mask { get; set; } +} diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageCompletion.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageCompletion.cs index 13132c63..323fe643 100644 --- a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageCompletion.cs +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageCompletion.cs @@ -18,4 +18,8 @@ public interface IImageCompletion Task GetImageGeneration(Agent agent, RoleDialogModel message); Task GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName); + + Task GetImageEdits(Agent agent, RoleDialogModel message, Stream image, string imageFileName); + + Task GetImageEdits(Agent agent, RoleDialogModel message, Stream image, string imageFileName, Stream mask, string maskFileName); } diff --git a/src/Infrastructure/BotSharp.Abstraction/Models/MessageConfig.cs b/src/Infrastructure/BotSharp.Abstraction/Models/MessageConfig.cs index 357a13aa..24e4152c 100644 --- a/src/Infrastructure/BotSharp.Abstraction/Models/MessageConfig.cs +++ b/src/Infrastructure/BotSharp.Abstraction/Models/MessageConfig.cs @@ -1,6 +1,6 @@ namespace BotSharp.Abstraction.Models; -public class MessageConfig +public class MessageConfig : InputMessageFiles { /// /// Completion Provider diff --git a/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Image.cs b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Image.cs index 01136226..619360d2 100644 --- a/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Image.cs +++ b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Image.cs @@ -14,15 +14,15 @@ public partial class BotSharpFileService return message; } - public async Task VarifyImage(string? provider, string? model, BotSharpFile file) + public async Task VaryImage(string? provider, string? model, BotSharpFile image) { - if (string.IsNullOrWhiteSpace(file?.FileUrl) && string.IsNullOrWhiteSpace(file?.FileData)) + if (string.IsNullOrWhiteSpace(image?.FileUrl) && string.IsNullOrWhiteSpace(image?.FileData)) { - throw new ArgumentException($"Please fill in at least file url or file data!"); + throw new ArgumentException($"Cannot find image url or data!"); } var completion = CompletionProvider.GetImageCompletion(_services, provider: provider ?? "openai", model: model ?? "dall-e-2"); - var bytes = await DownloadFile(file); + var bytes = await DownloadFile(image); using var stream = new MemoryStream(); stream.Write(bytes, 0, bytes.Length); stream.Position = 0; @@ -30,9 +30,61 @@ public partial class BotSharpFileService var message = await completion.GetImageVariation(new Agent() { Id = Guid.Empty.ToString() - }, new RoleDialogModel(AgentRole.User, string.Empty), stream, file.FileName ?? string.Empty); + }, new RoleDialogModel(AgentRole.User, string.Empty), stream, image.FileName ?? string.Empty); + stream.Close(); + return message; + } + public async Task EditImage(string? provider, string? model, string text, BotSharpFile image) + { + if (string.IsNullOrWhiteSpace(image?.FileUrl) && string.IsNullOrWhiteSpace(image?.FileData)) + { + throw new ArgumentException($"Cannot find image url or data!"); + } + + var completion = CompletionProvider.GetImageCompletion(_services, provider: provider ?? "openai", model: model ?? "dall-e-2"); + var bytes = await DownloadFile(image); + using var stream = new MemoryStream(); + stream.Write(bytes, 0, bytes.Length); + stream.Position = 0; + + var message = await completion.GetImageEdits(new Agent() + { + Id = Guid.Empty.ToString() + }, new RoleDialogModel(AgentRole.User, text), stream, image.FileName ?? string.Empty); + + stream.Close(); + return message; + } + + public async Task EditImage(string? provider, string? model, string text, BotSharpFile image, BotSharpFile mask) + { + if ((string.IsNullOrWhiteSpace(image?.FileUrl) && string.IsNullOrWhiteSpace(image?.FileData)) || + (string.IsNullOrWhiteSpace(mask?.FileUrl) && string.IsNullOrWhiteSpace(mask?.FileData))) + { + throw new ArgumentException($"Cannot find image/mask url or data"); + } + + var completion = CompletionProvider.GetImageCompletion(_services, provider: provider ?? "openai", model: model ?? "dall-e-2"); + var imageBytes = await DownloadFile(image); + var maskBytes = await DownloadFile(mask); + + using var imageStream = new MemoryStream(); + imageStream.Write(imageBytes, 0, imageBytes.Length); + imageStream.Position = 0; + + using var maskStream = new MemoryStream(); + maskStream.Write(maskBytes, 0, maskBytes.Length); + maskStream.Position = 0; + + var message = await completion.GetImageEdits(new Agent() + { + Id = Guid.Empty.ToString() + }, new RoleDialogModel(AgentRole.User, text), imageStream, image.FileName ?? string.Empty, maskStream, mask.FileName ?? string.Empty); + + imageStream.Close(); + maskStream.Close(); return message; } diff --git a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs index 8d433123..0b5de0ad 100644 --- a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs +++ b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs @@ -56,6 +56,7 @@ public class InstructModeController : ControllerBase return await textCompletion.GetCompletion(input.Text, Guid.Empty.ToString(), Guid.NewGuid().ToString()); } + #region Chat [HttpPost("/instruct/chat-completion")] public async Task ChatCompletion([FromBody] IncomingMessageModel input) { @@ -75,7 +76,9 @@ public class InstructModeController : ControllerBase }); return message.Content; } + #endregion + #region Read image [HttpPost("/instruct/multi-modal")] public async Task MultiModalCompletion([FromBody] IncomingMessageModel input) { @@ -105,7 +108,9 @@ public class InstructModeController : ControllerBase return error; } } + #endregion + #region Generate image [HttpPost("/instruct/image-generation")] public async Task ImageGeneration([FromBody] IncomingMessageModel input) { @@ -129,7 +134,9 @@ public class InstructModeController : ControllerBase return imageViewModel; } } + #endregion + #region Edit image [HttpPost("/instruct/image-variation")] public async Task ImageVariation([FromBody] IncomingMessageModel input) { @@ -140,12 +147,12 @@ public class InstructModeController : ControllerBase try { - var file = input.Files.FirstOrDefault(x => !string.IsNullOrWhiteSpace(x.FileUrl) || !string.IsNullOrWhiteSpace(x.FileData)); - if (file == null) + var image = input.Files.FirstOrDefault(x => !string.IsNullOrWhiteSpace(x.FileUrl) || !string.IsNullOrWhiteSpace(x.FileData)); + if (image == null) { return new ImageGenerationViewModel { Message = "Error! Cannot find an image!" }; } - var message = await fileService.VarifyImage(input.Provider, input.Model, file); + var message = await fileService.VaryImage(input.Provider, input.Model, image); imageViewModel.Content = message.Content; imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList(); return imageViewModel; @@ -159,6 +166,67 @@ public class InstructModeController : ControllerBase } } + [HttpPost("/instruct/image-edit")] + public async Task ImageEdit([FromBody] IncomingMessageModel input) + { + var fileService = _services.GetRequiredService(); + var state = _services.GetRequiredService(); + input.States.ForEach(x => state.SetState(x.Key, x.Value, activeRounds: x.ActiveRounds, source: StateSource.External)); + var imageViewModel = new ImageGenerationViewModel(); + + try + { + var image = input.Files.FirstOrDefault(x => !string.IsNullOrWhiteSpace(x.FileUrl) || !string.IsNullOrWhiteSpace(x.FileData)); + if (image == null) + { + return new ImageGenerationViewModel { Message = "Error! Cannot find an image!" }; + } + var message = await fileService.EditImage(input.Provider, input.Model, input.Text, image); + imageViewModel.Content = message.Content; + imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList(); + return imageViewModel; + } + catch (Exception ex) + { + var error = $"Error in image edit. {ex.Message}"; + _logger.LogError(error); + imageViewModel.Message = error; + return imageViewModel; + } + } + + [HttpPost("/instruct/image-mask-edit")] + public async Task ImageMaskEdit([FromBody] IncomingMessageModel input) + { + var fileService = _services.GetRequiredService(); + var state = _services.GetRequiredService(); + input.States.ForEach(x => state.SetState(x.Key, x.Value, activeRounds: x.ActiveRounds, source: StateSource.External)); + var imageViewModel = new ImageGenerationViewModel(); + + try + { + var image = input.Files.FirstOrDefault(x => !string.IsNullOrWhiteSpace(x.FileUrl) || !string.IsNullOrWhiteSpace(x.FileData)); + var mask = input.Mask; + if (image == null || mask == null) + { + return new ImageGenerationViewModel { Message = "Error! Cannot find an image or mask!" }; + } + var message = await fileService.EditImage(input.Provider, input.Model, input.Text, image, mask); + imageViewModel.Content = message.Content; + imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList(); + return imageViewModel; + } + catch (Exception ex) + { + var error = $"Error in image mask edit. {ex.Message}"; + _logger.LogError(error); + imageViewModel.Message = error; + return imageViewModel; + } + } + #endregion + + #region Pdf [HttpPost("/instruct/pdf-completion")] public async Task PdfCompletion([FromBody] IncomingMessageModel input) { @@ -181,4 +249,5 @@ public class InstructModeController : ControllerBase return viewModel; } } + #endregion } diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Edit.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Edit.cs new file mode 100644 index 00000000..43272f8e --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Edit.cs @@ -0,0 +1,66 @@ +using OpenAI.Images; + +namespace BotSharp.Plugin.AzureOpenAI.Providers.Image; + +public partial class ImageCompletionProvider +{ + public async Task GetImageEdits(Agent agent, RoleDialogModel message, Stream image, string imageFileName) + { + var client = ProviderHelper.GetClient(Provider, _model, _services); + var (prompt, imageCount, options) = PrepareEditOptions(message); + var imageClient = client.GetImageClient(_model); + + var response = imageClient.GenerateImageEdits(image, imageFileName, prompt, imageCount, options); + var images = response.Value; + + var generatedImages = GetImageGenerations(images, options.ResponseFormat); + var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages + }; + + return await Task.FromResult(responseMessage); + } + + public async Task GetImageEdits(Agent agent, RoleDialogModel message, + Stream image, string imageFileName, Stream mask, string maskFileName) + { + var client = ProviderHelper.GetClient(Provider, _model, _services); + var (prompt, imageCount, options) = PrepareEditOptions(message); + var imageClient = client.GetImageClient(_model); + + var response = imageClient.GenerateImageEdits(image, imageFileName, prompt, mask, maskFileName, imageCount, options); + var images = response.Value; + + var generatedImages = GetImageGenerations(images, options.ResponseFormat); + var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages + }; + + return await Task.FromResult(responseMessage); + } + + private (string, int, ImageEditOptions) PrepareEditOptions(RoleDialogModel message) + { + var prompt = message?.Payload ?? message?.Content ?? string.Empty; + + var state = _services.GetRequiredService(); + var size = GetImageSize(state.GetState("image_size")); + var format = GetImageFormat(state.GetState("image_format")); + var count = GetImageCount(state.GetState("image_count", "1")); + + var options = new ImageEditOptions + { + Size = size, + ResponseFormat = format + }; + return (prompt, count, options); + } +} diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Generation.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Generation.cs index dc2de807..7d608fd1 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Generation.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Generation.cs @@ -7,36 +7,13 @@ public partial class ImageCompletionProvider public async Task GetImageGeneration(Agent agent, RoleDialogModel message) { var client = ProviderHelper.GetClient(Provider, _model, _services); - var (prompt, imageCount, options) = PrepareOptions(message); + var (prompt, imageCount, options) = PrepareGenerationOptions(message); var imageClient = client.GetImageClient(_model); var response = imageClient.GenerateImages(prompt, imageCount, options); - var values = response.Value; - - var generatedImages = new List(); - foreach (var value in values) - { - if (value == null) continue; - - var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; - if (options.ResponseFormat == GeneratedImageFormat.Uri) - { - generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; - } - else if (options.ResponseFormat == GeneratedImageFormat.Bytes) - { - var base64Str = string.Empty; - var bytes = value?.ImageBytes?.ToArray(); - if (!bytes.IsNullOrEmpty()) - { - base64Str = Convert.ToBase64String(bytes); - } - generatedImage.ImageData = base64Str; - } - - generatedImages.Add(generatedImage); - } + var images = response.Value; + var generatedImages = GetImageGenerations(images, options.ResponseFormat); var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) { @@ -48,7 +25,7 @@ public partial class ImageCompletionProvider return await Task.FromResult(responseMessage); } - private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message) + private (string, int, ImageGenerationOptions) PrepareGenerationOptions(RoleDialogModel message) { var prompt = message?.Payload ?? message?.Content ?? string.Empty; diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Variation.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Variation.cs index e7543e8a..88f393e0 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Variation.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.Variation.cs @@ -7,36 +7,13 @@ public partial class ImageCompletionProvider public async Task GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName) { var client = ProviderHelper.GetClient(Provider, _model, _services); - var (imageCount, options) = PrepareOptions(); + var (imageCount, options) = PrepareVariationOptions(); var imageClient = client.GetImageClient(_model); var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options); - var values = response.Value; - - var generatedImages = new List(); - foreach (var value in values) - { - if (value == null) continue; - - var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; - if (options.ResponseFormat == GeneratedImageFormat.Uri) - { - generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; - } - else if (options.ResponseFormat == GeneratedImageFormat.Bytes) - { - var base64Str = string.Empty; - var bytes = value?.ImageBytes?.ToArray(); - if (!bytes.IsNullOrEmpty()) - { - base64Str = Convert.ToBase64String(bytes); - } - generatedImage.ImageData = base64Str; - } - - generatedImages.Add(generatedImage); - } + var images = response.Value; + var generatedImages = GetImageGenerations(images, options.ResponseFormat); var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) { @@ -48,7 +25,7 @@ public partial class ImageCompletionProvider return await Task.FromResult(responseMessage); } - private (int, ImageVariationOptions) PrepareOptions() + private (int, ImageVariationOptions) PrepareVariationOptions() { var state = _services.GetRequiredService(); var size = GetImageSize(state.GetState("image_size")); diff --git a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.cs index 6e19380a..7d064154 100644 --- a/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.AzureOpenAI/Providers/Image/ImageCompletionProvider.cs @@ -31,6 +31,34 @@ public partial class ImageCompletionProvider : IImageCompletion } #region Private methods + private List GetImageGenerations(GeneratedImageCollection images, GeneratedImageFormat? format) + { + var generatedImages = new List(); + foreach (var image in images) + { + if (image == null) continue; + + var generatedImage = new ImageGeneration { Description = image?.RevisedPrompt ?? string.Empty }; + if (format == GeneratedImageFormat.Uri) + { + generatedImage.ImageUrl = image?.ImageUri?.AbsoluteUri ?? string.Empty; + } + else if (format == GeneratedImageFormat.Bytes) + { + var base64Str = string.Empty; + var bytes = image?.ImageBytes?.ToArray(); + if (!bytes.IsNullOrEmpty()) + { + base64Str = Convert.ToBase64String(bytes); + } + generatedImage.ImageData = base64Str; + } + + generatedImages.Add(generatedImage); + } + return generatedImages; + } + private GeneratedImageSize GetImageSize(string size) { var value = !string.IsNullOrEmpty(size) ? size : "1024x1024"; diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Edit.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Edit.cs new file mode 100644 index 00000000..717002fb --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Edit.cs @@ -0,0 +1,67 @@ +using OpenAI.Images; +using static System.Net.Mime.MediaTypeNames; + +namespace BotSharp.Plugin.OpenAI.Providers.Image; + +public partial class ImageCompletionProvider +{ + public async Task GetImageEdits(Agent agent, RoleDialogModel message, Stream image, string imageFileName) + { + var client = ProviderHelper.GetClient(Provider, _model, _services); + var (prompt, imageCount, options) = PrepareEditOptions(message); + var imageClient = client.GetImageClient(_model); + + var response = imageClient.GenerateImageEdits(image, imageFileName, prompt, imageCount, options); + var images = response.Value; + + var generatedImages = GetImageGenerations(images, options.ResponseFormat); + var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages + }; + + return await Task.FromResult(responseMessage); + } + + public async Task GetImageEdits(Agent agent, RoleDialogModel message, + Stream image, string imageFileName, Stream mask, string maskFileName) + { + var client = ProviderHelper.GetClient(Provider, _model, _services); + var (prompt, imageCount, options) = PrepareEditOptions(message); + var imageClient = client.GetImageClient(_model); + + var response = imageClient.GenerateImageEdits(image, imageFileName, prompt, mask, maskFileName, imageCount, options); + var images = response.Value; + + var generatedImages = GetImageGenerations(images, options.ResponseFormat); + var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages + }; + + return await Task.FromResult(responseMessage); + } + + private (string, int, ImageEditOptions) PrepareEditOptions(RoleDialogModel message) + { + var prompt = message?.Payload ?? message?.Content ?? string.Empty; + + var state = _services.GetRequiredService(); + var size = GetImageSize(state.GetState("image_size")); + var format = GetImageFormat(state.GetState("image_format")); + var count = GetImageCount(state.GetState("image_count", "1")); + + var options = new ImageEditOptions + { + Size = size, + ResponseFormat = format + }; + return (prompt, count, options); + } +} diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Generation.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Generation.cs index 18bf7228..85c15686 100644 --- a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Generation.cs +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Generation.cs @@ -7,36 +7,13 @@ public partial class ImageCompletionProvider public async Task GetImageGeneration(Agent agent, RoleDialogModel message) { var client = ProviderHelper.GetClient(Provider, _model, _services); - var (prompt, imageCount, options) = PrepareOptions(message); + var (prompt, imageCount, options) = PrepareGenerationOptions(message); var imageClient = client.GetImageClient(_model); var response = imageClient.GenerateImages(prompt, imageCount, options); - var values = response.Value; - - var generatedImages = new List(); - foreach (var value in values) - { - if (value == null) continue; - - var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; - if (options.ResponseFormat == GeneratedImageFormat.Uri) - { - generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; - } - else if (options.ResponseFormat == GeneratedImageFormat.Bytes) - { - var base64Str = string.Empty; - var bytes = value?.ImageBytes?.ToArray(); - if (!bytes.IsNullOrEmpty()) - { - base64Str = Convert.ToBase64String(bytes); - } - generatedImage.ImageData = base64Str; - } - - generatedImages.Add(generatedImage); - } + var images = response.Value; + var generatedImages = GetImageGenerations(images, options.ResponseFormat); var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) { @@ -48,7 +25,7 @@ public partial class ImageCompletionProvider return await Task.FromResult(responseMessage); } - private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message) + private (string, int, ImageGenerationOptions) PrepareGenerationOptions(RoleDialogModel message) { var prompt = message?.Payload ?? message?.Content ?? string.Empty; diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Variation.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Variation.cs index 0233bad4..8a93df83 100644 --- a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Variation.cs +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.Variation.cs @@ -7,36 +7,13 @@ public partial class ImageCompletionProvider public async Task GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName) { var client = ProviderHelper.GetClient(Provider, _model, _services); - var (imageCount, options) = PrepareOptions(); + var (imageCount, options) = PrepareVariationOptions(); var imageClient = client.GetImageClient(_model); var response = imageClient.GenerateImageVariations(image, imageFileName, imageCount, options); - var values = response.Value; - - var generatedImages = new List(); - foreach (var value in values) - { - if (value == null) continue; - - var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; - if (options.ResponseFormat == GeneratedImageFormat.Uri) - { - generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; - } - else if (options.ResponseFormat == GeneratedImageFormat.Bytes) - { - var base64Str = string.Empty; - var bytes = value?.ImageBytes?.ToArray(); - if (!bytes.IsNullOrEmpty()) - { - base64Str = Convert.ToBase64String(bytes); - } - generatedImage.ImageData = base64Str; - } - - generatedImages.Add(generatedImage); - } + var images = response.Value; + var generatedImages = GetImageGenerations(images, options.ResponseFormat); var content = string.Join("\r\n", generatedImages.Where(x => !string.IsNullOrWhiteSpace(x.Description)).Select(x => x.Description)); var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) { @@ -48,7 +25,7 @@ public partial class ImageCompletionProvider return await Task.FromResult(responseMessage); } - private (int, ImageVariationOptions) PrepareOptions() + private (int, ImageVariationOptions) PrepareVariationOptions() { var state = _services.GetRequiredService(); var size = GetImageSize(state.GetState("image_size")); diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.cs index 13fb197c..a3fddacc 100644 --- a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.cs +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageCompletionProvider.cs @@ -1,3 +1,5 @@ +using Microsoft.Extensions.Options; +using Newtonsoft.Json.Linq; using OpenAI.Images; namespace BotSharp.Plugin.OpenAI.Providers.Image; @@ -31,6 +33,34 @@ public partial class ImageCompletionProvider : IImageCompletion } #region Private methods + private List GetImageGenerations(GeneratedImageCollection images, GeneratedImageFormat? format) + { + var generatedImages = new List(); + foreach (var image in images) + { + if (image == null) continue; + + var generatedImage = new ImageGeneration { Description = image?.RevisedPrompt ?? string.Empty }; + if (format == GeneratedImageFormat.Uri) + { + generatedImage.ImageUrl = image?.ImageUri?.AbsoluteUri ?? string.Empty; + } + else if (format == GeneratedImageFormat.Bytes) + { + var base64Str = string.Empty; + var bytes = image?.ImageBytes?.ToArray(); + if (!bytes.IsNullOrEmpty()) + { + base64Str = Convert.ToBase64String(bytes); + } + generatedImage.ImageData = base64Str; + } + + generatedImages.Add(generatedImage); + } + return generatedImages; + } + private GeneratedImageSize GetImageSize(string size) { var value = !string.IsNullOrEmpty(size) ? size : "1024x1024";