diff --git a/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs b/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs index 3f0f0c5c..e81bf6b8 100644 --- a/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs +++ b/src/Infrastructure/BotSharp.Abstraction/Files/IBotSharpFileService.cs @@ -2,8 +2,7 @@ namespace BotSharp.Abstraction.Files; public interface IBotSharpFileService { - string GetDirectory(string conversationId); - + #region Conversation /// /// Get the files that have been uploaded in the chat. /// If includeScreenShot is true, it will take the screenshots of non-image files, such as pdf, and return the screenshots instead of the original file. @@ -19,14 +18,19 @@ public interface IBotSharpFileService IEnumerable conversations, IEnumerable contentTypes, bool includeScreenShot = false, int? offset = null); + /// + /// Get the files that have been uploaded in the chat. No screenshot images are included. + /// + /// + /// + /// + /// + /// IEnumerable GetMessageFiles(string conversationId, IEnumerable messageIds, string source, bool imageOnly = false); string GetMessageFile(string conversationId, string messageId, string source, string index, string fileName); IEnumerable GetMessagesWithFile(string conversationId, IEnumerable messageIds); bool SaveMessageFiles(string conversationId, string messageId, string source, List files); - string GetUserAvatar(); - bool SaveUserAvatar(BotSharpFile file); - /// /// Delete files under messages /// @@ -37,7 +41,13 @@ public interface IBotSharpFileService /// bool DeleteMessageFiles(string conversationId, IEnumerable messageIds, string targetMessageId, string? newMessageId = null); bool DeleteConversationFiles(IEnumerable conversationIds); + #endregion + #region Image + + #endregion + + #region Pdf /// /// Take screenshots of pdf pages and get response from llm /// @@ -45,13 +55,21 @@ public interface IBotSharpFileService /// Pdf files /// Task InstructPdf(string? provider, string? model, string? modelId, string prompt, List files); + #endregion + #region User + string GetUserAvatar(); + bool SaveUserAvatar(BotSharpFile file); + #endregion + + #region Common /// /// Get file bytes and content type from data, e.g., "data:image/png;base64,aaaaaaaaa" /// /// /// (string, byte[]) GetFileInfoFromData(string data); - + string GetDirectory(string conversationId); string GetFileContentType(string filePath); + #endregion } diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageEdit.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageEdit.cs new file mode 100644 index 00000000..78f44d23 --- /dev/null +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageEdit.cs @@ -0,0 +1,5 @@ +namespace BotSharp.Abstraction.MLTasks; + +public interface IImageEdit +{ +} diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageGeneration.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageGeneration.cs index b257e114..2b345862 100644 --- a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageGeneration.cs +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageGeneration.cs @@ -13,5 +13,5 @@ public interface IImageGeneration /// deployment name void SetModelName(string model); - Task GetImageGeneration(Agent agent, List conversations); + Task GetImageGeneration(Agent agent, RoleDialogModel message); } diff --git a/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageVariation.cs b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageVariation.cs new file mode 100644 index 00000000..c679b327 --- /dev/null +++ b/src/Infrastructure/BotSharp.Abstraction/MLTasks/IImageVariation.cs @@ -0,0 +1,19 @@ +using System.IO; + +namespace BotSharp.Abstraction.MLTasks; + +public interface IImageVariation +{ + /// + /// The LLM provider like Microsoft Azure, OpenAI, ClaudAI + /// + string Provider { get; } + + /// + /// Set model name, one provider can consume different model or version(s) + /// + /// deployment name + void SetModelName(string model); + + RoleDialogModel GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName); +} diff --git a/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Common.cs b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Common.cs new file mode 100644 index 00000000..f29bf67f --- /dev/null +++ b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.Common.cs @@ -0,0 +1,5 @@ +namespace BotSharp.Core.Files.Services; + +public partial class BotSharpFileService +{ +} diff --git a/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageGeneration.cs b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageGeneration.cs new file mode 100644 index 00000000..f29bf67f --- /dev/null +++ b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageGeneration.cs @@ -0,0 +1,5 @@ +namespace BotSharp.Core.Files.Services; + +public partial class BotSharpFileService +{ +} diff --git a/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageVariation.cs b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageVariation.cs new file mode 100644 index 00000000..f29bf67f --- /dev/null +++ b/src/Infrastructure/BotSharp.Core/Files/Services/BotSharpFileService.ImageVariation.cs @@ -0,0 +1,5 @@ +namespace BotSharp.Core.Files.Services; + +public partial class BotSharpFileService +{ +} diff --git a/src/Infrastructure/BotSharp.Core/Infrastructures/CompletionProvider.cs b/src/Infrastructure/BotSharp.Core/Infrastructures/CompletionProvider.cs index f78edf8e..22675d3a 100644 --- a/src/Infrastructure/BotSharp.Core/Infrastructures/CompletionProvider.cs +++ b/src/Infrastructure/BotSharp.Core/Infrastructures/CompletionProvider.cs @@ -5,7 +5,7 @@ namespace BotSharp.Core.Infrastructures; public class CompletionProvider { - public static object GetCompletion(IServiceProvider services, + public static object? GetCompletion(IServiceProvider services, string? provider = null, string? model = null, AgentLlmConfig? agentConfig = null) @@ -23,13 +23,21 @@ public class CompletionProvider model: model, agentConfig: agentConfig); } - else + else if (settings.Type == LlmModelType.Embedding) { - return GetChatCompletion(services, - provider: provider, - model: model, + return GetTextEmbedding(services, + provider: provider, + model: model); + } + else if (settings.Type == LlmModelType.Chat) + { + return GetChatCompletion(services, + provider: provider, + model: model, agentConfig: agentConfig); } + + return null; } public static IChatCompletion GetChatCompletion(IServiceProvider services, @@ -51,7 +59,6 @@ public class CompletionProvider } completer?.SetModelName(model); - return completer; } @@ -72,7 +79,6 @@ public class CompletionProvider } completer.SetModelName(model); - return completer; } @@ -95,7 +101,26 @@ public class CompletionProvider } completer?.SetModelName(model); + return completer; + } + public static IImageVariation GetImageVariation(IServiceProvider services, + string? provider = null, + string? model = null, + string? modelId = null, + bool imageGenerate = false) + { + var completions = services.GetServices(); + (provider, model) = GetProviderAndModel(services, provider: provider, model: model, modelId: modelId, imageGenerate: imageGenerate); + + var completer = completions.FirstOrDefault(x => x.Provider == provider); + if (completer == null) + { + var logger = services.GetRequiredService>(); + logger.LogError($"Can't resolve completion provider by {provider}"); + } + + completer?.SetModelName(model); return completer; } @@ -152,7 +177,6 @@ public class CompletionProvider state.SetState("provider", provider); state.SetState("model", model); - return (provider, model); } } diff --git a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs index a19eb6df..c1fd1f8e 100644 --- a/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs +++ b/src/Infrastructure/BotSharp.OpenAPI/Controllers/InstructModeController.cs @@ -119,10 +119,7 @@ public class InstructModeController : ControllerBase var message = await completion.GetImageGeneration(new Agent() { Id = Guid.Empty.ToString(), - }, new List - { - new RoleDialogModel(AgentRole.User, input.Text) - }); + }, new RoleDialogModel(AgentRole.User, input.Text)); imageViewModel.Content = message.Content; imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList(); @@ -137,6 +134,34 @@ public class InstructModeController : ControllerBase } } + //[HttpPost("/instruct/image-variation")] + //public ImageGenerationViewModel ImageVariation([FromBody] IncomingMessageModel input) + //{ + // 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 completion = CompletionProvider.GetImageVariation(_services, provider: input.Provider ?? "openai", model: input.Model ?? "dall-e-2"); + // var message = completion.GetImageVariation(new Agent() + // { + // Id = Guid.Empty.ToString(), + // }, new RoleDialogModel(AgentRole.User, input.Text)); + + // imageViewModel.Content = message.Content; + // imageViewModel.Images = message.GeneratedImages.Select(x => ImageViewModel.ToViewModel(x)).ToList(); + // return imageViewModel; + // } + // catch (Exception ex) + // { + // var error = $"Error in image generation. {ex.Message}"; + // _logger.LogError(error); + // imageViewModel.Message = error; + // return imageViewModel; + // } + //} + [HttpPost("/instruct/pdf-completion")] public async Task PdfCompletion([FromBody] IncomingMessageModel input) { diff --git a/src/Plugins/BotSharp.Plugin.FileHandler/Functions/GenerateImageFn.cs b/src/Plugins/BotSharp.Plugin.FileHandler/Functions/GenerateImageFn.cs index 50482e66..e787da84 100644 --- a/src/Plugins/BotSharp.Plugin.FileHandler/Functions/GenerateImageFn.cs +++ b/src/Plugins/BotSharp.Plugin.FileHandler/Functions/GenerateImageFn.cs @@ -62,7 +62,7 @@ public class GenerateImageFn : IFunctionCallback var completion = CompletionProvider.GetImageGeneration(_services, provider: "openai", model: "dall-e-3"); var text = !string.IsNullOrWhiteSpace(description) ? description : message.Content; var dialog = RoleDialogModel.From(message, AgentRole.User, text); - var result = await completion.GetImageGeneration(agent, new List { dialog }); + var result = await completion.GetImageGeneration(agent, dialog); SaveGeneratedImages(result?.GeneratedImages); return result?.Content ?? string.Empty; } diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/OpenAiPlugin.cs b/src/Plugins/BotSharp.Plugin.OpenAI/OpenAiPlugin.cs index ce77e589..b2f66387 100644 --- a/src/Plugins/BotSharp.Plugin.OpenAI/OpenAiPlugin.cs +++ b/src/Plugins/BotSharp.Plugin.OpenAI/OpenAiPlugin.cs @@ -28,7 +28,8 @@ public class OpenAiPlugin : IBotSharpPlugin services.AddScoped(); services.AddScoped(); - services.AddScoped(); services.AddScoped(); + services.AddScoped(); + services.AddScoped(); } } \ No newline at end of file diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageGenerationProvider.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageGenerationProvider.cs index b8f512ae..f8159e34 100644 --- a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageGenerationProvider.cs +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageGenerationProvider.cs @@ -6,7 +6,7 @@ public class ImageGenerationProvider : IImageGeneration { protected readonly OpenAiSettings _settings; protected readonly IServiceProvider _services; - protected readonly ILogger _logger; + protected readonly ILogger _logger; private const int DEFAULT_IMAGE_COUNT = 1; private const int IMAGE_COUNT_LIMIT = 5; @@ -26,32 +26,24 @@ public class ImageGenerationProvider : IImageGeneration } - public async Task GetImageGeneration(Agent agent, List conversations) + public async Task GetImageGeneration(Agent agent, RoleDialogModel message) { - var contentHooks = _services.GetServices().ToList(); - - // Before - foreach (var hook in contentHooks) - { - await hook.BeforeGenerating(agent, conversations); - } - var client = ProviderHelper.GetClient(Provider, _model, _services); - var (prompt, imageCount, options) = PrepareOptions(conversations); + var (prompt, imageCount, options) = PrepareOptions(message); var imageClient = client.GetImageClient(_model); var response = imageClient.GenerateImages(prompt, imageCount, options); var values = response.Value; - var images = new List(); + var generatedImages = new List(); foreach (var value in values) { if (value == null) continue; - var image = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; + var generatedImage = new ImageGeneration { Description = value?.RevisedPrompt ?? string.Empty }; if (options.ResponseFormat == GeneratedImageFormat.Uri) { - image.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; + generatedImage.ImageUrl = value?.ImageUri?.AbsoluteUri ?? string.Empty; } else if (options.ResponseFormat == GeneratedImageFormat.Bytes) { @@ -61,21 +53,22 @@ public class ImageGenerationProvider : IImageGeneration { base64Str = Convert.ToBase64String(bytes); } - image.ImageData = base64Str; + generatedImage.ImageData = base64Str; } - images.Add(image); + generatedImages.Add(generatedImage); } - var content = string.Join("\r\n", images.Select(x => x.Description)); + var content = string.Join("\r\n", generatedImages.Select(x => x.Description)); var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) { CurrentAgentId = agent.Id, - MessageId = conversations.LastOrDefault()?.MessageId ?? string.Empty, - GeneratedImages = images + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages }; // After + var contentHooks = _services.GetServices().ToList(); foreach (var hook in contentHooks) { await hook.AfterGenerated(responseMessage, new TokenStatsModel @@ -91,9 +84,14 @@ public class ImageGenerationProvider : IImageGeneration return responseMessage; } - private (string, int, ImageGenerationOptions) PrepareOptions(List conversations) + public void SetModelName(string model) { - var prompt = conversations.LastOrDefault()?.Payload ?? conversations.LastOrDefault()?.Content ?? string.Empty; + _model = model; + } + + private (string, int, ImageGenerationOptions) PrepareOptions(RoleDialogModel message) + { + var prompt = message?.Payload ?? message?.Content ?? string.Empty; var state = _services.GetRequiredService(); var size = state.GetState("image_size"); @@ -112,11 +110,6 @@ public class ImageGenerationProvider : IImageGeneration return (prompt, count, options); } - public void SetModelName(string model) - { - _model = model; - } - private GeneratedImageSize GetImageSize(string size) { var value = !string.IsNullOrEmpty(size) ? size : "1024x1024"; diff --git a/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageVariationProvider.cs b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageVariationProvider.cs new file mode 100644 index 00000000..78e6a33f --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.OpenAI/Providers/Image/ImageVariationProvider.cs @@ -0,0 +1,154 @@ +using OpenAI.Images; + +namespace BotSharp.Plugin.OpenAI.Providers.Image; + +public class ImageVariationProvider : IImageVariation +{ + protected readonly OpenAiSettings _settings; + protected readonly IServiceProvider _services; + protected readonly ILogger _logger; + + private const int DEFAULT_IMAGE_COUNT = 1; + private const int IMAGE_COUNT_LIMIT = 5; + + protected string _model; + + public virtual string Provider => "openai"; + + public ImageVariationProvider( + OpenAiSettings settings, + ILogger logger, + IServiceProvider services) + { + _settings = settings; + _services = services; + _logger = logger; + } + + public RoleDialogModel GetImageVariation(Agent agent, RoleDialogModel message, Stream image, string imageFileName) + { + var client = ProviderHelper.GetClient(Provider, _model, _services); + var (imageCount, options) = PrepareOptions(); + 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 content = string.Join("\r\n", generatedImages.Select(x => x.Description)); + var responseMessage = new RoleDialogModel(AgentRole.Assistant, content) + { + CurrentAgentId = agent.Id, + MessageId = message?.MessageId ?? string.Empty, + GeneratedImages = generatedImages + }; + + return responseMessage; + } + + public void SetModelName(string model) + { + _model = model; + } + + private (int, ImageVariationOptions) PrepareOptions() + { + var state = _services.GetRequiredService(); + var size = state.GetState("image_size"); + var quality = state.GetState("image_quality"); + var style = state.GetState("image_style"); + var format = state.GetState("image_format"); + var count = GetImageCount(state.GetState("image_count", "1")); + + var options = new ImageVariationOptions + { + Size = GetImageSize(size), + ResponseFormat = GetImageFormat(format) + }; + return (count, options); + } + + 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 GeneratedImageFormat GetImageFormat(string format) + { + var value = !string.IsNullOrEmpty(format) ? format : "uri"; + + GeneratedImageFormat retFormat; + switch (value) + { + case "uri": + retFormat = GeneratedImageFormat.Uri; + break; + case "bytes": + retFormat = GeneratedImageFormat.Bytes; + break; + default: + retFormat = GeneratedImageFormat.Uri; + break; + } + + return retFormat; + } + + private int GetImageCount(string count) + { + if (!int.TryParse(count, out var retCount)) + { + return DEFAULT_IMAGE_COUNT; + } + + return retCount > 0 && retCount <= IMAGE_COUNT_LIMIT ? retCount : DEFAULT_IMAGE_COUNT; + } +}