Merge pull request #776 from Joannall/master

add OnSourceCodeGenerated Hook
This commit is contained in:
Haiping 2024-12-03 21:32:23 +00:00 committed by GitHub
commit be687bd916
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 37 additions and 12 deletions

View file

@ -5,6 +5,9 @@ public interface IPlanningHook
Task<string> GetSummaryAdditionalRequirements(string planner, RoleDialogModel message)
=> Task.FromResult(string.Empty);
Task OnSourceCodeGenerated(string planner, RoleDialogModel msg, string language)
=> Task.CompletedTask;
Task OnPlanningCompleted(string planner, RoleDialogModel msg)
=> Task.CompletedTask;
}

View file

@ -34,7 +34,7 @@ public class OneStepForwardReasoner : IRoutingReasoner
private readonly IServiceProvider _services;
private readonly ILogger _logger;
public OneStepForwardReasoner(IServiceProvider services, ILogger<NaiveReasoner> logger)
public OneStepForwardReasoner(IServiceProvider services, ILogger<OneStepForwardReasoner> logger)
{
_services = services;
_logger = logger;
@ -116,7 +116,7 @@ public class OneStepForwardReasoner : IRoutingReasoner
}
else
{
context.Empty(reason: $"Agent queue is cleared by {nameof(NaiveReasoner)}");
context.Empty(reason: $"Agent queue is cleared by {nameof(OneStepForwardReasoner)}");
// context.Push(inst.OriginalAgent, "Push user goal agent");
}
return true;

View file

@ -33,7 +33,7 @@ public class SequentialReasoner : IRoutingReasoner
public int MaxLoopCount => 100;
private FunctionCallFromLlm _lastInst;
public SequentialReasoner(IServiceProvider services, ILogger<NaiveReasoner> logger)
public SequentialReasoner(IServiceProvider services, ILogger<SequentialReasoner> logger)
{
_services = services;
_logger = logger;

View file

@ -71,13 +71,15 @@ public class SummaryPlanFn : IFunctionCallback
var summary = await GetAiResponse(plannerAgent);
message.Content = summary.Content;
// Validate the sql result
// Emit event if the sql statement is generated by planner
var args = JsonSerializer.Deserialize<SummaryPlan>(message.FunctionArgs);
if (args.IsSqlTemplate == false)
if (args != null && !args.IsSqlTemplate && args.ContainsSqlStatements)
{
await fn.InvokeFunction("validate_sql", message);
await HookEmitter.Emit<IPlanningHook>(_services, async hook =>
await hook.OnSourceCodeGenerated(nameof(TwoStageTaskPlanner), message, "sql")
);
}
await HookEmitter.Emit<IPlanningHook>(_services, async hook =>
await hook.OnPlanningCompleted(nameof(TwoStageTaskPlanner), message)
);

View file

@ -4,4 +4,7 @@ public class SummaryPlan
{
[JsonPropertyName("is_sql_template")]
public bool IsSqlTemplate { get; set; } = false;
[JsonPropertyName("contains_sql_statements")]
public bool ContainsSqlStatements { get; set; } = false;
}

View file

@ -1,7 +1,7 @@
{
"id": "282a7128-69a1-44b0-878c-a9159b88f3b9",
"name": "Planner",
"description": "Plan feasible implementation steps for user task request",
"description": "Plan feasible implementation steps for complex user task request",
"type": "task",
"createdDateTime": "2023-08-27T10:39:00Z",
"updatedDateTime": "2023-08-27T14:39:00Z",

View file

@ -8,6 +8,10 @@
"type": "boolean",
"description": "If user request is to generate sql template instead of actual sql statement."
},
"contains_sql_statements": {
"type": "boolean",
"description": "Set to true if the response contains sql statements."
},
"related_tables": {
"type": "array",
"description": "table name in planning steps",
@ -17,6 +21,6 @@
}
}
},
"required": [ "related_tables", "is_sql_template" ]
"required": [ "related_tables", "is_sql_template", "contains_sql_statements" ]
}
}

View file

@ -20,14 +20,23 @@ public class SqlDriverPlanningHook : IPlanningHook
_services = services;
}
public async Task OnPlanningCompleted(string planner, RoleDialogModel msg)
public async Task OnSourceCodeGenerated(string planner, RoleDialogModel msg, string language)
{
// envoke validate
if (language != "sql")
{
return;
}
var routing = _services.GetRequiredService<IRoutingService>();
await routing.InvokeFunction("validate_sql", msg);
await HookEmitter.Emit<ISqlDriverHook>(_services, async (hook) =>
{
await hook.SqlGenerated(msg);
});
var settings = _services.GetRequiredService<SqlDriverSetting>();
var settings = _services.GetRequiredService<SqlDriverSetting>();
if (!settings.ExecuteSqlSelectAutonomous)
{
var conversationStateService = _services.GetRequiredService<IConversationStateService>();
@ -51,7 +60,6 @@ public class SqlDriverPlanningHook : IPlanningHook
var response = await completion.GetChatCompletions(agent, wholeDialogs);
// Invoke "execute_sql"
var routing = _services.GetRequiredService<IRoutingService>();
await routing.InvokeFunction(response.FunctionName, response);
msg.CurrentAgentId = agent.Id;
@ -61,6 +69,11 @@ public class SqlDriverPlanningHook : IPlanningHook
msg.StopCompletion = response.StopCompletion;
}
public async Task OnPlanningCompleted(string planner, RoleDialogModel msg)
{
}
public async Task<string> GetSummaryAdditionalRequirements(string planner, RoleDialogModel message)
{
var settings = _services.GetRequiredService<SqlDriverSetting>();