merge with master

This commit is contained in:
Jicheng Lu 2023-09-09 11:28:52 -05:00
commit 56eb31dd1d
22 changed files with 117 additions and 311 deletions

View file

@ -1,5 +1,3 @@
using BotSharp.Abstraction.MLTasks;
namespace BotSharp.Abstraction.Conversations;
public interface IConversationService
@ -11,8 +9,6 @@ public interface IConversationService
Task<List<Conversation>> GetConversations();
Task DeleteConversation(string id);
IChatCompletion GetChatCompletion();
/// <summary>
/// Send message to LLM
/// </summary>

View file

@ -11,6 +11,8 @@ public class RoleDialogModel
public DateTime CreatedAt { get; set; } = DateTime.UtcNow;
public string Content { get; set; }
public string CurrentAgentId { get; set; }
public string ModelName { get; set; } = "gpt-3.5-turbo";
public float Temperature { get; set; } = 0.5f;
/// <summary>
/// Function name if LLM response function call

View file

@ -1,9 +1,8 @@
using BotSharp.Abstraction.Conversations.Models;
namespace BotSharp.Abstraction.MLTasks;
public interface IChatCompletion
{
string ModelName { get; }
Task<bool> GetChatCompletionsAsync(Agent agent,
List<RoleDialogModel> conversations,
Func<RoleDialogModel, Task> onMessageReceived,

View file

@ -7,13 +7,13 @@ public class RoutingItem
public string Id { get; set; }
[JsonPropertyName("agent_id")]
public string AgentId { get; set; }
public string AgentId { get; set; } = string.Empty;
[JsonPropertyName("name")]
public string Name { get; set; }
public string Name { get; set; } = string.Empty;
[JsonPropertyName("description")]
public string Description { get; set; }
public string Description { get; set; } = string.Empty;
[JsonPropertyName("required")]
public List<string> RequiredFields { get; set; } = new List<string>();

View file

@ -86,4 +86,9 @@
<ProjectReference Include="..\BotSharp.Abstraction\BotSharp.Abstraction.csproj" />
</ItemGroup>
<ItemGroup>
<Folder Include="Functions\" />
<Folder Include="Hooks\" />
</ItemGroup>
</Project>

View file

@ -9,13 +9,14 @@ public partial class ConversationService
{
int currentRecursiveDepth = 0;
private async Task<bool> GetChatCompletionsAsyncRecursively(IChatCompletion chatCompletion,
Agent agent,
private async Task<bool> GetChatCompletionsAsyncRecursively(Agent agent,
List<RoleDialogModel> wholeDialogs,
Func<RoleDialogModel, Task> onMessageReceived,
Func<RoleDialogModel, Task> onFunctionExecuting,
Func<RoleDialogModel, Task> onFunctionExecuted)
{
var chatCompletion = CompletionProvider.GetChatCompletion(_services, wholeDialogs.Last().ModelName);
currentRecursiveDepth++;
if (currentRecursiveDepth > _settings.MaxRecursiveDepth)
{
@ -28,11 +29,16 @@ public partial class ConversationService
text = latestResponse.Content.Split("=>").Last();
}
await HandleAssistantMessage(new RoleDialogModel(AgentRole.Assistant, text)
var msg = new RoleDialogModel(AgentRole.Assistant, text)
{
CurrentAgentId = agent.Id,
Channel = wholeDialogs.Last().Channel
}, onMessageReceived);
};
await HandleAssistantMessage(msg, onMessageReceived);
// Add to dialog history
_storage.Append(_conversationId, agent.Id, msg);
return false;
}
@ -85,8 +91,7 @@ public partial class ConversationService
wholeDialogs.Add(fn);
await GetChatCompletionsAsyncRecursively(chatCompletion,
agent,
await GetChatCompletionsAsyncRecursively(agent,
wholeDialogs,
onMessageReceived,
onFunctionExecuting,
@ -115,8 +120,7 @@ public partial class ConversationService
// After function is executed, pass the result to LLM to get a natural response
wholeDialogs.Add(fn);
await GetChatCompletionsAsyncRecursively(chatCompletion,
agent,
await GetChatCompletionsAsyncRecursively(agent,
wholeDialogs,
onMessageReceived,
onFunctionExecuting,

View file

@ -1,6 +1,5 @@
using BotSharp.Abstraction.Agents.Enums;
using BotSharp.Abstraction.Agents.Models;
using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.Routing.Settings;
using BotSharp.Core.Routing;
@ -92,9 +91,7 @@ public partial class ConversationService
});
}
var chatCompletion = GetChatCompletion();
var result = await GetChatCompletionsAsyncRecursively(chatCompletion,
agent,
var result = await GetChatCompletionsAsyncRecursively(agent,
wholeDialogs,
onMessageReceived,
onFunctionExecuting,
@ -136,16 +133,4 @@ public partial class ConversationService
}
}
}
public IChatCompletion GetChatCompletion()
{
var completions = _services.GetServices<IChatCompletion>();
return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.ChatCompletion));
}
public IChatCompletion GetGpt4ChatCompletion()
{
var completions = _services.GetServices<IChatCompletion>();
return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith("GPT4CompletionProvider"));
}
}

View file

@ -0,0 +1,14 @@
using BotSharp.Abstraction.MLTasks;
namespace BotSharp.Core.Infrastructures;
public class CompletionProvider
{
public static IChatCompletion GetChatCompletion(IServiceProvider services, string modelName = "gpt-3.5-turbo")
{
var completions = services.GetServices<IChatCompletion>();
var settings = services.GetRequiredService<ConversationSetting>();
// completions.FirstOrDefault(x => x.GetType().FullName.EndsWith(settings.ChatCompletion));
return completions.FirstOrDefault(x => x.ModelName == modelName);
}
}

View file

@ -29,7 +29,7 @@ public partial class InstructService : IInstructService
var wholeDialogs = new List<RoleDialogModel>
{
new RoleDialogModel("user", message.Content)
message
};
// Trigger before completion hooks
@ -71,7 +71,7 @@ public partial class InstructService : IInstructService
Func<RoleDialogModel, Task> onFunctionExecuting,
Func<RoleDialogModel, Task> onFunctionExecuted)
{
var chatCompletion = GetChatCompletion();
var chatCompletion = CompletionProvider.GetChatCompletion(_services, wholeDialogs.Last().ModelName);
var result = await chatCompletion.GetChatCompletionsAsync(agent, wholeDialogs, async msg =>
{
@ -124,11 +124,4 @@ public partial class InstructService : IInstructService
await CallFunctions(msg);
await onFunctionExecuted(msg);
}
public IChatCompletion GetChatCompletion()
{
var completions = _services.GetServices<IChatCompletion>();
var settings = _services.GetRequiredService<ConversationSetting>();
return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith(settings.ChatCompletion));
}
}

View file

@ -1,4 +1,4 @@
namespace BotSharp.Core.Hooks;
namespace BotSharp.Core.Routing;
public class ReasoningHook : AgentHookBase
{

View file

@ -1,9 +1,8 @@
using BotSharp.Abstraction.Conversations.Models;
using BotSharp.Abstraction.Functions;
using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.Routing.Models;
using System.IO;
namespace BotSharp.Core.Functions;
namespace BotSharp.Core.Routing;
/// <summary>
/// Router calls this function to set the Active Agent according to the context
@ -63,26 +62,36 @@ public class RouteToAgentFn : IFunctionCallback
agentId = routingRule.AgentId;
// Check required fields
var jo = JsonSerializer.Deserialize<object>(message.FunctionArgs);
var root = JsonSerializer.Deserialize<JsonElement>(message.FunctionArgs);
bool hasMissingField = false;
string missingFieldName = "";
foreach (var field in routingRule.RequiredFields)
{
if (jo is JsonElement root)
if (!root.EnumerateObject().Any(x => x.Name == field))
{
if (!root.EnumerateObject().Any(x => x.Name == field))
{
message.ExecutionResult = $"missing {field}.";
hasMissingField = true;
break;
}
else if (root.EnumerateObject().Any(x => x.Name == field) &&
string.IsNullOrEmpty(root.EnumerateObject().FirstOrDefault(x => x.Name == field).Value.ToString()))
{
message.ExecutionResult = $"missing {field}.";
hasMissingField = true;
break;
}
message.ExecutionResult = $"missing {field}.";
hasMissingField = true;
missingFieldName = field;
break;
}
else if (root.EnumerateObject().Any(x => x.Name == field) &&
string.IsNullOrEmpty(root.EnumerateObject().FirstOrDefault(x => x.Name == field).Value.ToString()))
{
message.ExecutionResult = $"missing {field}.";
hasMissingField = true;
missingFieldName = field;
break;
}
}
// Check if states contains the field according conversation context.
var states = _services.GetRequiredService<IConversationStateService>();
if (!string.IsNullOrEmpty(states.GetState(missingFieldName)))
{
var value = states.GetState(missingFieldName);
message.FunctionArgs = message.FunctionArgs.Substring(0, message.FunctionArgs.Length - 1) + $", \"{missingFieldName}\": \"{value}\"" + "}";
hasMissingField = false;
missingFieldName = "";
}
if (hasMissingField && !string.IsNullOrEmpty(routingRule.RedirectTo))

View file

@ -1,4 +1,4 @@
namespace BotSharp.Core.Hooks;
namespace BotSharp.Core.Routing;
public class RoutingHook : AgentHookBase
{

View file

@ -64,7 +64,7 @@ public class Simulator
new RoleDialogModel(AgentRole.User, @"What's the next step, your response must be in JSON format with ""function"" and ""parameters"". ")
};
var chatCompletion = GetGpt4ChatCompletion();
var chatCompletion = CompletionProvider.GetChatCompletion(_services, "gpt-4");
RoleDialogModel response = null;
await chatCompletion.GetChatCompletionsAsync(reasoner, wholeDialogs, async msg
@ -111,7 +111,7 @@ public class Simulator
var agentService = _services.GetRequiredService<IAgentService>();
var agent = await agentService.LoadAgent(agentId);
var chatCompletion = GetChatCompletion();
var chatCompletion = CompletionProvider.GetChatCompletion(_services, wholeDialogs.Last().ModelName);
RoleDialogModel response = null;
await chatCompletion.GetChatCompletionsAsync(agent, wholeDialogs, async msg
@ -132,19 +132,6 @@ public class Simulator
return response;
}
public IChatCompletion GetChatCompletion()
{
var completions = _services.GetServices<IChatCompletion>();
var settings = _services.GetRequiredService<ConversationSetting>();
return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith(settings.ChatCompletion));
}
public IChatCompletion GetGpt4ChatCompletion()
{
var completions = _services.GetServices<IChatCompletion>();
return completions.FirstOrDefault(x => x.GetType().FullName.EndsWith("GPT4CompletionProvider"));
}
private void SaveStateByArgs(JsonDocument args)
{
var stateService = _services.GetRequiredService<IConversationStateService>();

View file

@ -40,11 +40,10 @@ public class ConversationController : ControllerBase, IApiAdapter
[HttpPost("/conversation/{agentId}/{conversationId}")]
public async Task<MessageResponseModel> SendMessage([FromRoute] string agentId,
[FromRoute] string conversationId,
[FromBody] NewMessageModel input,
[FromQuery] string? channel = "openapi")
[FromBody] NewMessageModel input)
{
var conv = _services.GetRequiredService<IConversationService>();
conv.SetConversationId(conversationId, channel);
conv.SetConversationId(conversationId, input.Channel);
input.States.ForEach(x => conv.States.SetState(x.Split('=')[0], x.Split('=')[1]));
var response = new MessageResponseModel();
@ -53,7 +52,8 @@ public class ConversationController : ControllerBase, IApiAdapter
await conv.SendMessage(agentId,
new RoleDialogModel("user", input.Text)
{
Channel = channel
Channel = input.Channel,
ModelName = input.ModelName
},
async msg =>
{

View file

@ -36,7 +36,10 @@ public class InstructModeController : ControllerBase, IApiAdapter
}
return await instructor.ExecuteInstruction(agent,
new RoleDialogModel(AgentRole.User, input.Text),
new RoleDialogModel(AgentRole.User, input.Text)
{
ModelName = input.ModelName
},
fn => Task.CompletedTask,
fn => Task.CompletedTask,
fn => Task.CompletedTask);

View file

@ -3,6 +3,8 @@ namespace BotSharp.OpenAPI.ViewModels.Conversations;
public class NewMessageModel
{
public string Text { get; set; }
public string ModelName { get; set; } = "gpt-3.5-turbo";
public string Channel { get; set; } = "openapi";
/// <summary>
/// Conversation states from input

View file

@ -23,6 +23,8 @@ public class ChatCompletionProvider : IChatCompletion
private readonly IServiceProvider _services;
private readonly ILogger _logger;
public virtual string ModelName => "gpt-3.5-turbo";
public ChatCompletionProvider(AzureOpenAiSettings settings,
ILogger<ChatCompletionProvider> logger,
IServiceProvider services)
@ -32,10 +34,10 @@ public class ChatCompletionProvider : IChatCompletion
_services = services;
}
private OpenAIClient GetClient()
protected virtual (OpenAIClient, string) GetClient()
{
var client = new OpenAIClient(new Uri(_settings.Endpoint), new AzureKeyCredential(_settings.ApiKey));
return client;
return (client, _settings.DeploymentModel.ChatCompletionModel);
}
public List<RoleDialogModel> GetChatSamples(string sampleText)
@ -85,10 +87,10 @@ public class ChatCompletionProvider : IChatCompletion
Func<RoleDialogModel, Task> onMessageReceived,
Func<RoleDialogModel, Task> onFunctionExecuting)
{
var client = GetClient();
var (client, deploymentModel) = GetClient();
var chatCompletionsOptions = PrepareOptions(agent, conversations);
var response = await client.GetChatCompletionsAsync(_settings.DeploymentModel.ChatCompletionModel, chatCompletionsOptions);
var response = await client.GetChatCompletionsAsync(deploymentModel, chatCompletionsOptions);
var choice = response.Value.Choices[0];
var message = choice.Message;
@ -106,6 +108,12 @@ public class ChatCompletionProvider : IChatCompletion
Channel = conversations.Last().Channel
};
// Somethings LLM will generate a function name with agent name.
if (!string.IsNullOrEmpty(funcContextIn.FunctionName))
{
funcContextIn.FunctionName = funcContextIn.FunctionName.Split('.').Last();
}
// Execute functions
await onFunctionExecuting(funcContextIn);
}
@ -171,7 +179,7 @@ public class ChatCompletionProvider : IChatCompletion
}
private ChatCompletionsOptions PrepareOptions(Agent agent, List<RoleDialogModel> conversations)
protected ChatCompletionsOptions PrepareOptions(Agent agent, List<RoleDialogModel> conversations)
{
var chatCompletionsOptions = new ChatCompletionsOptions();

View file

@ -1,236 +1,31 @@
using Azure;
using Azure.AI.OpenAI;
using BotSharp.Abstraction.Agents.Enums;
using BotSharp.Abstraction.Agents.Models;
using BotSharp.Abstraction.Conversations.Models;
using BotSharp.Abstraction.Conversations.Settings;
using BotSharp.Abstraction.Functions.Models;
using BotSharp.Abstraction.MLTasks;
using BotSharp.Plugin.AzureOpenAI.Settings;
using Microsoft.Extensions.DependencyInjection;
using Microsoft.Extensions.Logging;
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading.Tasks;
namespace BotSharp.Plugin.AzureOpenAI.Providers;
public class GPT4CompletionProvider : IChatCompletion
public class GPT4CompletionProvider : ChatCompletionProvider
{
private readonly AzureOpenAiSettings _settings;
private readonly IServiceProvider _services;
private readonly ILogger _logger;
public override string ModelName => "gpt-4";
public GPT4CompletionProvider(AzureOpenAiSettings settings,
ILogger<GPT4CompletionProvider> logger,
IServiceProvider services)
IServiceProvider services) : base(settings, logger, services)
{
_settings = settings;
_logger = logger;
_services = services;
}
private OpenAIClient GetClient()
protected override (OpenAIClient, string) GetClient()
{
var client = new OpenAIClient(new Uri(_settings.GPT4.Endpoint), new AzureKeyCredential(_settings.GPT4.ApiKey));
return client;
}
public List<RoleDialogModel> GetChatSamples(string sampleText)
{
var samples = new List<RoleDialogModel>();
if (string.IsNullOrEmpty(sampleText))
{
return samples;
}
var lines = sampleText.Split('\n');
for (int i = 0; i < lines.Length; i++)
{
var line = lines[i];
if (string.IsNullOrEmpty(line.Trim()))
{
continue;
}
var role = line.Substring(0, line.IndexOf(' ') - 1).Trim();
var content = line.Substring(line.IndexOf(' ') + 1).Trim();
// comments
if (role == "##")
{
continue;
}
samples.Add(new RoleDialogModel(role, content));
}
return samples;
}
public List<FunctionDef> GetFunctions(List<string> functionsJson)
{
var functions = functionsJson?.Select(x => JsonSerializer.Deserialize<FunctionDef>(x, new JsonSerializerOptions
{
PropertyNameCaseInsensitive = true,
AllowTrailingCommas = true
}))?.ToList() ?? new List<FunctionDef>();
return functions;
}
public async Task<bool> GetChatCompletionsAsync(Agent agent,
List<RoleDialogModel> conversations,
Func<RoleDialogModel, Task> onMessageReceived,
Func<RoleDialogModel, Task> onFunctionExecuting)
{
var client = GetClient();
var chatCompletionsOptions = PrepareOptions(agent, conversations);
var response = await client.GetChatCompletionsAsync(_settings.GPT4.DeploymentModel, chatCompletionsOptions);
var choice = response.Value.Choices[0];
var message = choice.Message;
if (choice.FinishReason == CompletionsFinishReason.FunctionCall)
{
_logger.LogInformation($"[{agent.Name}]: {message.FunctionCall.Name} => {message.FunctionCall.Arguments}");
var funcContextIn = new RoleDialogModel(AgentRole.Function, message.Content)
{
CurrentAgentId = agent.Id,
FunctionName = message.FunctionCall.Name,
FunctionArgs = message.FunctionCall.Arguments,
Channel = conversations.Last().Channel
};
// Execute functions
await onFunctionExecuting(funcContextIn);
}
else
{
_logger.LogInformation($"[{agent.Name}] {message.Role}: {message.Content}");
var msg = new RoleDialogModel(AgentRole.Assistant, message.Content)
{
CurrentAgentId= agent.Id,
Channel = conversations.Last().Channel
};
// Text response received
await onMessageReceived(msg);
}
return true;
}
public async Task<bool> GetChatCompletionsStreamingAsync(Agent agent, List<RoleDialogModel> conversations, Func<RoleDialogModel, Task> onMessageReceived)
{
var client = new OpenAIClient(new Uri(_settings.Endpoint), new AzureKeyCredential(_settings.ApiKey));
var chatCompletionsOptions = PrepareOptions(agent, conversations);
var response = await client.GetChatCompletionsStreamingAsync(_settings.DeploymentModel.ChatCompletionModel, chatCompletionsOptions);
using StreamingChatCompletions streaming = response.Value;
string output = "";
await foreach (var choice in streaming.GetChoicesStreaming())
{
if (choice.FinishReason == CompletionsFinishReason.FunctionCall)
{
var args = "";
await foreach (var message in choice.GetMessageStreaming())
{
if (message.FunctionCall == null || message.FunctionCall.Arguments == null)
continue;
Console.Write(message.FunctionCall.Arguments);
args += message.FunctionCall.Arguments;
}
await onMessageReceived(new RoleDialogModel(ChatRole.Assistant.ToString(), args));
continue;
}
await foreach (var message in choice.GetMessageStreaming())
{
if (message.Content == null)
continue;
Console.Write(message.Content);
output += message.Content;
_logger.LogInformation(message.Content);
await onMessageReceived(new RoleDialogModel(message.Role.ToString(), message.Content));
}
output = "";
}
return true;
}
private ChatCompletionsOptions PrepareOptions(Agent agent, List<RoleDialogModel> conversations)
{
var chatCompletionsOptions = new ChatCompletionsOptions();
if (!string.IsNullOrEmpty(agent.Instruction))
{
chatCompletionsOptions.Messages.Add(new ChatMessage(ChatRole.System, agent.Instruction));
}
if (!string.IsNullOrEmpty(agent.Knowledges))
{
chatCompletionsOptions.Messages.Add(new ChatMessage(ChatRole.System, agent.Knowledges));
}
var samples = GetChatSamples(agent.Samples);
foreach (var message in samples)
{
chatCompletionsOptions.Messages.Add(new ChatMessage(message.Role, message.Content));
}
var functions = GetFunctions(agent.Functions);
foreach (var function in functions)
{
chatCompletionsOptions.Functions.Add(new FunctionDefinition
{
Name = function.Name,
Description = function.Description,
Parameters = BinaryData.FromObjectAsJson(function.Parameters)
});
}
foreach (var message in conversations)
{
if (message.Role == ChatRole.Function)
{
chatCompletionsOptions.Messages.Add(new ChatMessage(message.Role, message.Content)
{
Name = message.FunctionName
});
}
else
{
chatCompletionsOptions.Messages.Add(new ChatMessage(message.Role, message.Content));
}
}
// https://community.openai.com/t/cheat-sheet-mastering-temperature-and-top-p-in-chatgpt-api-a-few-tips-and-tricks-on-controlling-the-creativity-deterministic-output-of-prompt-responses/172683
chatCompletionsOptions.Temperature = 0.5f;
chatCompletionsOptions.NucleusSamplingFactor = 0.5f;
var convSetting = _services.GetRequiredService<ConversationSetting>();
if (convSetting.ShowVerboseLog)
{
var verbose = string.Join("\n", chatCompletionsOptions.Messages.Select(x =>
{
return x.Role == ChatRole.Function ?
$"{x.Role}: {x.Name} {x.Content}" :
$"{x.Role}: {x.Content}";
}));
_logger.LogInformation(verbose);
}
return chatCompletionsOptions;
return (client, _settings.GPT4.DeploymentModel);
}
}

View file

@ -2,7 +2,7 @@ namespace BotSharp.Plugin.AzureOpenAI.Settings;
public class DeploymentModelSetting
{
public string? ChatCompletionModel { get; set; }
public string ChatCompletionModel { get; set; } = string.Empty;
public string? TextCompletionModel { get; set; }
public override string ToString()

View file

@ -44,9 +44,17 @@ public class ChatbotUiController : ControllerBase, IApiAdapter
{
Id = "gpt-3.5-turbo",
Model = "gpt-3.5-turbo",
Name = "Default (GPT-3.5)",
MaxLength = 4000,
TokenLimit = 4000
Name = "GPT-3.5 Turbo",
MaxLength = 4 * 1024,
TokenLimit = 4 * 1024
},
new AiModel
{
Id = "gpt-4",
Model = "gpt-4",
Name = "GPT-4",
MaxLength = 8 * 1024,
TokenLimit = 8 * 1024
}
}
};
@ -62,7 +70,7 @@ public class ChatbotUiController : ControllerBase, IApiAdapter
var outputStream = Response.Body;
var channel = "webchat";
var conversation = input.Messages
var message = input.Messages
.Where(x => x.Role == AgentRole.User)
.Select(x => new RoleDialogModel(x.Role, x.Content)
{
@ -74,7 +82,7 @@ public class ChatbotUiController : ControllerBase, IApiAdapter
input.States.ForEach(x => conv.States.SetState(x.Split('=')[0], x.Split('=')[1]));
var result = await conv.SendMessage(input.AgentId,
conversation,
message,
async msg =>
await OnChunkReceived(outputStream, msg),
async fn

View file

@ -27,10 +27,8 @@ public class ChatCompletionProvider : IChatCompletion
_logger = logger;
}
public string GetChatCompletions(Agent agent, List<RoleDialogModel> conversations, Func<RoleDialogModel, Task> onMessageReceived)
{
throw new NotImplementedException();
}
public string ModelName => "llama-2";
public async Task<bool> GetChatCompletionsAsync(Agent agent,
List<RoleDialogModel> conversations,

View file

@ -26,9 +26,7 @@
"Conversation": {
"DataDir": "conversations",
"ShowVerboseLog": false,
"ChatCompletion": "AzureOpenAI.Providers.ChatCompletionProvider"
// "ChatCompletion": "LLamaSharp.ChatCompletionProvider"
"ShowVerboseLog": false
},
"LlamaSharp": {