2024-10-01 21:43:20 +00:00
|
|
|
using BotSharp.Abstraction.Agents.Enums;
|
2024-10-31 16:35:01 +00:00
|
|
|
using BotSharp.Abstraction.Repositories;
|
2024-10-12 01:18:45 +00:00
|
|
|
using BotSharp.Abstraction.Routing;
|
2024-10-01 21:43:20 +00:00
|
|
|
using BotSharp.Core.Infrastructures;
|
2024-09-17 11:32:11 +00:00
|
|
|
using BotSharp.Plugin.SqlDriver.Models;
|
|
|
|
|
using Dapper;
|
|
|
|
|
using Microsoft.Data.SqlClient;
|
2024-10-12 01:18:45 +00:00
|
|
|
using Microsoft.Extensions.Logging;
|
2024-09-17 11:32:11 +00:00
|
|
|
using MySqlConnector;
|
2024-10-31 16:35:01 +00:00
|
|
|
using Npgsql;
|
2024-09-17 11:32:11 +00:00
|
|
|
|
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-10-12 01:18:45 +00:00
|
|
|
private readonly ILogger _logger;
|
2024-01-13 21:17:40 +00:00
|
|
|
|
2024-10-12 01:18:45 +00:00
|
|
|
public ExecuteQueryFn(IServiceProvider services, SqlDriverSetting setting, ILogger<ExecuteQueryFn> logger)
|
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;
|
2024-10-12 01:18:45 +00:00
|
|
|
_logger = logger;
|
2024-01-13 21:17:40 +00:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
public async Task<bool> Execute(RoleDialogModel message)
|
|
|
|
|
{
|
2024-09-17 11:32:11 +00:00
|
|
|
var args = JsonSerializer.Deserialize<ExecuteQueryArgs>(message.FunctionArgs);
|
2024-10-12 01:18:45 +00:00
|
|
|
var refinedArgs = await RefineSqlStatement(message, args);
|
2024-10-31 16:35:01 +00:00
|
|
|
var dbHook = _services.GetRequiredService<IDatabaseHook>();
|
|
|
|
|
var dbType = dbHook.GetDatabaseType(message);
|
2024-10-12 01:18:45 +00:00
|
|
|
|
|
|
|
|
try
|
2024-01-13 21:17:40 +00:00
|
|
|
{
|
2024-10-31 16:35:01 +00:00
|
|
|
var results = dbType.ToLower() switch
|
2024-10-12 01:18:45 +00:00
|
|
|
{
|
2024-10-31 16:35:01 +00:00
|
|
|
"mysql" => RunQueryInMySql(refinedArgs.SqlStatements),
|
|
|
|
|
"sqlserver" => RunQueryInSqlServer(refinedArgs.SqlStatements),
|
|
|
|
|
"redshift" => RunQueryInRedshift(refinedArgs.SqlStatements),
|
|
|
|
|
_ => throw new NotImplementedException($"Database type {dbType} is not supported.")
|
2024-10-12 01:18:45 +00:00
|
|
|
};
|
|
|
|
|
|
2024-10-23 21:56:52 +00:00
|
|
|
if (refinedArgs.SqlStatements.Length == 1 && refinedArgs.SqlStatements[0].StartsWith("DROP TABLE"))
|
|
|
|
|
{
|
|
|
|
|
message.Content = "Drop table successfully";
|
|
|
|
|
return true;
|
|
|
|
|
}
|
|
|
|
|
|
2024-10-12 01:18:45 +00:00
|
|
|
if (results.Count() == 0)
|
|
|
|
|
{
|
|
|
|
|
message.Content = "No record found";
|
|
|
|
|
return true;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
message.Content = JsonSerializer.Serialize(results);
|
|
|
|
|
}
|
|
|
|
|
catch (Exception ex)
|
2024-09-24 01:33:42 +00:00
|
|
|
{
|
2024-10-12 01:18:45 +00:00
|
|
|
_logger.LogError(ex, "Error occurred while executing SQL query.");
|
|
|
|
|
message.Content = "Error occurred while retrieving information.";
|
|
|
|
|
message.StopCompletion = true;
|
|
|
|
|
return false;
|
2024-09-24 01:33:42 +00:00
|
|
|
}
|
2024-09-17 11:32:11 +00:00
|
|
|
|
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-10-12 01:18:45 +00:00
|
|
|
message.StopCompletion = true;
|
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-10-31 16:35:01 +00:00
|
|
|
private IEnumerable<dynamic> RunQueryInRedshift(string[] sqlTexts)
|
|
|
|
|
{
|
|
|
|
|
var settings = _services.GetRequiredService<SqlDriverSetting>();
|
|
|
|
|
using var connection = new NpgsqlConnection(settings.RedshiftConnectionString);
|
|
|
|
|
return connection.Query(string.Join("\r\n", sqlTexts));
|
|
|
|
|
}
|
2024-10-12 01:18:45 +00:00
|
|
|
|
|
|
|
|
private async Task<ExecuteQueryArgs> RefineSqlStatement(RoleDialogModel message, ExecuteQueryArgs args)
|
|
|
|
|
{
|
2024-10-16 19:44:18 +00:00
|
|
|
if (args.Tables == null || args.Tables.Length == 0)
|
|
|
|
|
{
|
|
|
|
|
return args;
|
|
|
|
|
}
|
|
|
|
|
|
2024-10-12 01:18:45 +00:00
|
|
|
// get table DDL
|
|
|
|
|
var fn = _services.GetRequiredService<IRoutingService>();
|
|
|
|
|
var msg = RoleDialogModel.From(message);
|
|
|
|
|
await fn.InvokeFunction("sql_table_definition", msg);
|
|
|
|
|
|
|
|
|
|
// refine SQL
|
|
|
|
|
var agentService = _services.GetRequiredService<IAgentService>();
|
|
|
|
|
var currentAgent = await agentService.LoadAgent(message.CurrentAgentId);
|
|
|
|
|
var dictionarySqlPrompt = await GetDictionarySQLPrompt(string.Join("\r\n\r\n", args.SqlStatements), msg.Content);
|
|
|
|
|
var agent = new Agent
|
|
|
|
|
{
|
|
|
|
|
Id = message.CurrentAgentId ?? string.Empty,
|
|
|
|
|
Name = "sqlDriver_ExecuteQuery",
|
|
|
|
|
Instruction = dictionarySqlPrompt,
|
|
|
|
|
TemplateDict = new Dictionary<string, object>(),
|
|
|
|
|
LlmConfig = currentAgent.LlmConfig
|
|
|
|
|
};
|
|
|
|
|
|
|
|
|
|
var completion = CompletionProvider.GetChatCompletion(_services,
|
|
|
|
|
provider: agent.LlmConfig.Provider,
|
|
|
|
|
model: agent.LlmConfig.Model);
|
|
|
|
|
|
|
|
|
|
var refinedMessage = await completion.GetChatCompletions(agent, new List<RoleDialogModel>
|
|
|
|
|
{
|
|
|
|
|
new RoleDialogModel(AgentRole.User, "Check and output the correct SQL statements")
|
|
|
|
|
});
|
|
|
|
|
|
|
|
|
|
return refinedMessage.Content.JsonContent<ExecuteQueryArgs>();
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
private async Task<string> GetDictionarySQLPrompt(string originalSql, string tableStructure)
|
|
|
|
|
{
|
|
|
|
|
var agentService = _services.GetRequiredService<IAgentService>();
|
|
|
|
|
var render = _services.GetRequiredService<ITemplateRender>();
|
|
|
|
|
var knowledgeHooks = _services.GetServices<IKnowledgeHook>();
|
|
|
|
|
|
|
|
|
|
var agent = await agentService.GetAgent(BuiltInAgentId.SqlDriver);
|
|
|
|
|
var template = agent.Templates.FirstOrDefault(x => x.Name == "sql_statement_correctness")?.Content ?? string.Empty;
|
|
|
|
|
var responseFormat = JsonSerializer.Serialize(new ExecuteQueryArgs { });
|
|
|
|
|
|
|
|
|
|
return render.Render(template, new Dictionary<string, object>
|
|
|
|
|
{
|
|
|
|
|
{ "original_sql", originalSql },
|
|
|
|
|
{ "table_structure", tableStructure },
|
|
|
|
|
{ "response_format", responseFormat }
|
|
|
|
|
});
|
|
|
|
|
}
|
2024-01-13 21:17:40 +00:00
|
|
|
}
|