BotSharp/src/Plugins/BotSharp.Plugin.SqlDriver/Functions/ExecuteQueryFn.cs

81 lines
2.9 KiB
C#
Raw Normal View History

2024-10-01 21:43:20 +00:00
using BotSharp.Abstraction.Agents.Enums;
using BotSharp.Core.Infrastructures;
2024-09-17 11:32:11 +00:00
using BotSharp.Plugin.SqlDriver.Models;
using Dapper;
using Microsoft.Data.SqlClient;
using MySqlConnector;
2024-02-19 22:55:41 +00:00
namespace BotSharp.Plugin.SqlDriver.Functions;
2024-01-13 21:17:40 +00:00
2024-02-19 22:55:41 +00:00
public class ExecuteQueryFn : IFunctionCallback
2024-01-13 21:17:40 +00:00
{
public string Name => "execute_sql";
2024-09-18 02:27:47 +00:00
public string Indication => "Performing data retrieval operation.";
2024-01-18 11:30:28 +00:00
private readonly SqlDriverSetting _setting;
2024-09-17 11:32:11 +00:00
private readonly IServiceProvider _services;
2024-01-13 21:17:40 +00:00
2024-09-17 11:32:11 +00:00
public ExecuteQueryFn(IServiceProvider services, SqlDriverSetting setting)
2024-01-13 21:17:40 +00:00
{
2024-09-17 11:32:11 +00:00
_services = services;
2024-01-13 21:17:40 +00:00
_setting = setting;
}
public async Task<bool> Execute(RoleDialogModel message)
{
2024-09-17 11:32:11 +00:00
var args = JsonSerializer.Deserialize<ExecuteQueryArgs>(message.FunctionArgs);
var settings = _services.GetRequiredService<SqlDriverSetting>();
var results = settings.DatabaseType switch
2024-01-13 21:17:40 +00:00
{
2024-09-17 11:32:11 +00:00
"MySql" => RunQueryInMySql(args.SqlStatements),
"SqlServer" => RunQueryInSqlServer(args.SqlStatements),
_ => throw new NotImplementedException($"Database type {settings.DatabaseType} is not supported.")
};
2024-10-02 02:34:23 +00:00
2024-09-24 01:33:42 +00:00
if (results.Count() == 0)
{
message.Content = "No record found";
2024-10-02 02:34:23 +00:00
return true;
2024-09-24 01:33:42 +00:00
}
2024-09-17 11:32:11 +00:00
message.Content = JsonSerializer.Serialize(results);
2024-10-01 21:43:20 +00:00
if (args.FormattingResult)
2024-09-24 01:33:42 +00:00
{
2024-10-01 21:43:20 +00:00
var conv = _services.GetRequiredService<IConversationService>();
var sqlAgent = await _services.GetRequiredService<IAgentService>().LoadAgent(BuiltInAgentId.SqlDriver);
var prompt = sqlAgent.Templates.FirstOrDefault(x => x.Name == "query_result_formatting");
var completion = CompletionProvider.GetChatCompletion(_services,
provider: sqlAgent.LlmConfig.Provider,
model: sqlAgent.LlmConfig.Model);
var result = await completion.GetChatCompletions(new Agent
{
Id = sqlAgent.Id,
Instruction = prompt.Content,
}, new List<RoleDialogModel>
{
new RoleDialogModel(AgentRole.User, message.Content)
});
message.Content = result.Content;
2024-09-24 01:33:42 +00:00
}
2024-10-01 21:43:20 +00:00
2024-01-13 21:17:40 +00:00
return true;
}
2024-09-17 11:32:11 +00:00
private IEnumerable<dynamic> RunQueryInMySql(string[] sqlTexts)
{
var settings = _services.GetRequiredService<SqlDriverSetting>();
2024-09-18 19:09:04 +00:00
using var connection = new MySqlConnection(settings.MySqlExecutionConnectionString ?? settings.MySqlConnectionString);
2024-09-17 11:32:11 +00:00
return connection.Query(string.Join(";\r\n", sqlTexts));
}
private IEnumerable<dynamic> RunQueryInSqlServer(string[] sqlTexts)
{
var settings = _services.GetRequiredService<SqlDriverSetting>();
using var connection = new SqlConnection(settings.SqlServerExecutionConnectionString ?? settings.SqlServerConnectionString);
return connection.Query(string.Join("\r\n", sqlTexts));
}
2024-01-13 21:17:40 +00:00
}