Merge pull request #783 from iceljc/master

add common agent hook
This commit is contained in:
iceljc 2024-12-06 00:58:57 -06:00 committed by GitHub
commit d758e4b5f3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 85 additions and 73 deletions

View file

@ -1,10 +1,5 @@
using BotSharp.Abstraction.Agents.Settings;
using BotSharp.Abstraction.Conversations;
using BotSharp.Abstraction.Functions.Models;
using BotSharp.Abstraction.Repositories;
using BotSharp.Abstraction.Routing;
using Microsoft.Extensions.DependencyInjection;
using System.Data;
namespace BotSharp.Abstraction.Agents;
@ -60,73 +55,5 @@ public abstract class AgentHookBase : IAgentHook
public virtual void OnAgentUtilityLoaded(Agent agent)
{
if (agent.Type == AgentType.Routing) return;
var conv = _services.GetRequiredService<IConversationService>();
var isConvMode = conv.IsConversationMode();
if (!isConvMode) return;
agent.Functions ??= [];
agent.Utilities ??= [];
var (functions, templates) = GetUtilityContent(agent);
foreach (var fn in functions)
{
if (!agent.Functions.Any(x => x.Name.Equals(fn.Name, StringComparison.OrdinalIgnoreCase)))
{
agent.Functions.Add(fn);
}
}
foreach (var prompt in templates)
{
agent.Instruction += $"\r\n\r\n{prompt}\r\n\r\n";
}
}
private (IEnumerable<FunctionDef>, IEnumerable<string>) GetUtilityContent(Agent agent)
{
var db = _services.GetRequiredService<IBotSharpRepository>();
var (functionNames, templateNames) = GetUniqueContent(agent.Utilities);
if (agent.MergeUtility)
{
var routing = _services.GetRequiredService<IRoutingContext>();
var entryAgentId = routing.EntryAgentId;
if (!string.IsNullOrEmpty(entryAgentId))
{
var entryAgent = db.GetAgent(entryAgentId);
var (fns, tps) = GetUniqueContent(entryAgent?.Utilities);
functionNames = functionNames.Concat(fns).Distinct().ToList();
templateNames = templateNames.Concat(tps).Distinct().ToList();
}
}
var ua = db.GetAgent(BuiltInAgentId.UtilityAssistant);
var functions = ua?.Functions?.Where(x => functionNames.Contains(x.Name, StringComparer.OrdinalIgnoreCase))?.ToList() ?? [];
var templates = ua?.Templates?.Where(x => templateNames.Contains(x.Name, StringComparer.OrdinalIgnoreCase))?.Select(x => x.Content)?.ToList() ?? [];
return (functions, templates);
}
private (IEnumerable<string>, IEnumerable<string>) GetUniqueContent(IEnumerable<AgentUtility>? utilities)
{
if (utilities.IsNullOrEmpty())
{
return ([], []);
}
var prefix = "util-";
utilities = utilities?.Where(x => !string.IsNullOrEmpty(x.Name) && !x.Disabled)?.ToList() ?? [];
var functionNames = utilities.SelectMany(x => x.Functions)
.Where(x => !string.IsNullOrEmpty(x.Name) && x.Name.StartsWith(prefix))
.Select(x => x.Name)
.Distinct().ToList();
var templateNames = utilities.SelectMany(x => x.Templates)
.Where(x => !string.IsNullOrEmpty(x.Name) && x.Name.StartsWith(prefix))
.Select(x => x.Name)
.Distinct().ToList();
return (functionNames, templateNames);
}
}

View file

@ -3,6 +3,7 @@ using BotSharp.Abstraction.Plugins.Models;
using BotSharp.Abstraction.Settings;
using BotSharp.Abstraction.Templating;
using BotSharp.Abstraction.Users.Enums;
using BotSharp.Core.Agents.Hooks;
using Microsoft.Extensions.Configuration;
namespace BotSharp.Core.Agents;
@ -30,6 +31,7 @@ public class AgentPlugin : IBotSharpPlugin
{
services.AddScoped<ILlmProviderService, LlmProviderService>();
services.AddScoped<IAgentService, AgentService>();
services.AddScoped<IAgentHook, CommonAgentHook>();
services.AddScoped(provider =>
{

View file

@ -0,0 +1,83 @@
namespace BotSharp.Core.Agents.Hooks;
public class CommonAgentHook : AgentHookBase
{
public override string SelfId => string.Empty;
public CommonAgentHook(IServiceProvider services, AgentSettings settings)
: base(services, settings)
{
}
public override void OnAgentUtilityLoaded(Agent agent)
{
if (agent.Type == AgentType.Routing) return;
var conv = _services.GetRequiredService<IConversationService>();
var isConvMode = conv.IsConversationMode();
if (!isConvMode) return;
agent.Functions ??= [];
agent.Utilities ??= [];
var (functions, templates) = GetUtilityContent(agent);
foreach (var fn in functions)
{
if (!agent.Functions.Any(x => x.Name.Equals(fn.Name, StringComparison.OrdinalIgnoreCase)))
{
agent.Functions.Add(fn);
}
}
foreach (var prompt in templates)
{
agent.Instruction += $"\r\n\r\n{prompt}\r\n\r\n";
}
}
private (IEnumerable<FunctionDef>, IEnumerable<string>) GetUtilityContent(Agent agent)
{
var db = _services.GetRequiredService<IBotSharpRepository>();
var (functionNames, templateNames) = GetUniqueContent(agent.Utilities);
if (agent.MergeUtility)
{
var routing = _services.GetRequiredService<IRoutingContext>();
var entryAgentId = routing.EntryAgentId;
if (!string.IsNullOrEmpty(entryAgentId))
{
var entryAgent = db.GetAgent(entryAgentId);
var (fns, tps) = GetUniqueContent(entryAgent?.Utilities);
functionNames = functionNames.Concat(fns).Distinct().ToList();
templateNames = templateNames.Concat(tps).Distinct().ToList();
}
}
var ua = db.GetAgent(BuiltInAgentId.UtilityAssistant);
var functions = ua?.Functions?.Where(x => functionNames.Contains(x.Name, StringComparer.OrdinalIgnoreCase))?.ToList() ?? [];
var templates = ua?.Templates?.Where(x => templateNames.Contains(x.Name, StringComparer.OrdinalIgnoreCase))?.Select(x => x.Content)?.ToList() ?? [];
return (functions, templates);
}
private (IEnumerable<string>, IEnumerable<string>) GetUniqueContent(IEnumerable<AgentUtility>? utilities)
{
if (utilities.IsNullOrEmpty())
{
return ([], []);
}
var prefix = "util-";
utilities = utilities?.Where(x => !string.IsNullOrEmpty(x.Name) && !x.Disabled)?.ToList() ?? [];
var functionNames = utilities.SelectMany(x => x.Functions)
.Where(x => !string.IsNullOrEmpty(x.Name) && x.Name.StartsWith(prefix))
.Select(x => x.Name)
.Distinct().ToList();
var templateNames = utilities.SelectMany(x => x.Templates)
.Where(x => !string.IsNullOrEmpty(x.Name) && x.Name.StartsWith(prefix))
.Select(x => x.Name)
.Distinct().ToList();
return (functionNames, templateNames);
}
}