using BotSharp.Abstraction.Agents.Enums; using BotSharp.Abstraction.Routing; using BotSharp.Core.Infrastructures; using BotSharp.Plugin.SqlDriver.Interfaces; using BotSharp.Plugin.SqlDriver.Models; using Dapper; using Microsoft.Data.SqlClient; using Microsoft.Extensions.Logging; using MySqlConnector; using Npgsql; using System.Data.Common; namespace BotSharp.Plugin.SqlDriver.Functions; public class ExecuteQueryFn : IFunctionCallback { public string Name => "execute_sql"; public string Indication => "Performing data retrieval operation."; private readonly SqlDriverSetting _setting; private readonly IServiceProvider _services; private readonly ILogger _logger; public ExecuteQueryFn(IServiceProvider services, SqlDriverSetting setting, ILogger logger) { _services = services; _setting = setting; _logger = logger; } public async Task Execute(RoleDialogModel message) { var args = JsonSerializer.Deserialize(message.FunctionArgs); var refinedArgs = await RefineSqlStatement(message, args); var dbHook = _services.GetRequiredService(); var dbType = dbHook.GetDatabaseType(message); try { var results = dbType.ToLower() switch { "mysql" => RunQueryInMySql(refinedArgs.SqlStatements), "sqlserver" => RunQueryInSqlServer(refinedArgs.SqlStatements), "redshift" => RunQueryInRedshift(refinedArgs.SqlStatements), _ => throw new NotImplementedException($"Database type {dbType} is not supported.") }; if (refinedArgs.SqlStatements.Length == 1 && refinedArgs.SqlStatements[0].StartsWith("DROP TABLE")) { message.Content = "Drop table successfully"; return true; } if (results.Count() == 0) { message.Content = "No record found"; return true; } message.Content = JsonSerializer.Serialize(results); } catch (DbException ex) { _logger.LogError(ex, "Error occurred while executing SQL query."); message.Content = $"Error occurred while executing SQL query: {ex.Message}"; message.Data = ex; message.StopCompletion = true; return false; } catch (Exception ex) { _logger.LogError(ex, "Error occurred while executing SQL query."); message.Content = $"Error occurred while executing SQL query: {ex.Message}"; message.StopCompletion = true; return false; } if (args.FormattingResult) { var conv = _services.GetRequiredService(); var sqlAgent = await _services.GetRequiredService().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 { new RoleDialogModel(AgentRole.User, message.Content) }); message.Content = result.Content; message.StopCompletion = true; } return true; } private IEnumerable RunQueryInMySql(string[] sqlTexts) { var settings = _services.GetRequiredService(); using var connection = new MySqlConnection(settings.MySqlExecutionConnectionString ?? settings.MySqlConnectionString); return connection.Query(string.Join(";\r\n", sqlTexts)); } private IEnumerable RunQueryInSqlServer(string[] sqlTexts) { var settings = _services.GetRequiredService(); using var connection = new SqlConnection(settings.SqlServerExecutionConnectionString ?? settings.SqlServerConnectionString); return connection.Query(string.Join("\r\n", sqlTexts)); } private IEnumerable RunQueryInRedshift(string[] sqlTexts) { var settings = _services.GetRequiredService(); using var connection = new NpgsqlConnection(settings.RedshiftConnectionString); return connection.Query(string.Join("\r\n", sqlTexts)); } private async Task RefineSqlStatement(RoleDialogModel message, ExecuteQueryArgs args) { if (args.Tables == null || args.Tables.Length == 0) { return args; } // get table DDL var fn = _services.GetRequiredService(); var msg = RoleDialogModel.From(message); await fn.InvokeFunction("sql_table_definition", msg); // refine SQL var agentService = _services.GetRequiredService(); 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(), LlmConfig = currentAgent.LlmConfig }; var completion = CompletionProvider.GetChatCompletion(_services, provider: agent.LlmConfig.Provider, model: agent.LlmConfig.Model); var refinedMessage = await completion.GetChatCompletions(agent, new List { new RoleDialogModel(AgentRole.User, "Check and output the correct SQL statements") }); return refinedMessage.Content.JsonContent(); } private async Task GetDictionarySQLPrompt(string originalSql, string tableStructure) { var agentService = _services.GetRequiredService(); var render = _services.GetRequiredService(); var knowledgeHooks = _services.GetServices(); 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 { { "original_sql", originalSql }, { "table_structure", tableStructure }, { "response_format", responseFormat } }); } }