mirror of
https://github.com/MindWorkAI/AI-Studio.git
synced 2026-10-05 10:49:40 +00:00
Added a retrieval context validation agent (#289)
This commit is contained in:
1 parent
96e6372fcd
commit
bfc9f2ea1d
26 files changed
+829
-135
No files matched your search
@@ -1,7 +1,9 @@
|
||||
using System.Text;
|
||||
|
||||
using AIStudio.Agents;
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider;
|
||||
using AIStudio.Settings;
|
||||
|
||||
namespace AIStudio.Tools.RAG.AugmentationProcesses;
|
||||
|
||||
@@ -22,13 +24,36 @@ public sealed class AugmentationOne : IAugmentationProcess
|
||||
public async Task<ChatThread> ProcessAsync(IProvider provider, IContent lastPrompt, ChatThread chatThread, IReadOnlyList<IRetrievalContext> retrievalContexts, CancellationToken token = default)
|
||||
{
|
||||
var logger = Program.SERVICE_PROVIDER.GetService<ILogger<AugmentationOne>>()!;
|
||||
var settings = Program.SERVICE_PROVIDER.GetService<SettingsManager>()!;
|
||||
|
||||
if(retrievalContexts.Count == 0)
|
||||
{
|
||||
logger.LogWarning("No retrieval contexts were issued. Skipping the augmentation process.");
|
||||
return chatThread;
|
||||
}
|
||||
|
||||
|
||||
var numTotalRetrievalContexts = retrievalContexts.Count;
|
||||
|
||||
// Want the user to validate all retrieval contexts?
|
||||
if (settings.ConfigurationData.AgentRetrievalContextValidation.EnableRetrievalContextValidation)
|
||||
{
|
||||
// Let's get the validation agent & set up its provider:
|
||||
var validationAgent = Program.SERVICE_PROVIDER.GetService<AgentRetrievalContextValidation>()!;
|
||||
validationAgent.SetLLMProvider(provider);
|
||||
|
||||
// Let's validate all retrieval contexts:
|
||||
var validationResults = await validationAgent.ValidateRetrievalContextsAsync(lastPrompt, chatThread, retrievalContexts, token);
|
||||
|
||||
//
|
||||
// Now, filter the retrieval contexts to the most relevant ones:
|
||||
//
|
||||
var targetWindow = validationResults.DetermineTargetWindow(TargetWindowStrategy.TOP10_BETTER_THAN_GUESSING);
|
||||
var threshold = validationResults.GetConfidenceThreshold(targetWindow);
|
||||
|
||||
// Filter the retrieval contexts:
|
||||
retrievalContexts = validationResults.Where(x => x.RetrievalContext is not null && x.Confidence >= threshold).Select(x => x.RetrievalContext!).ToList();
|
||||
}
|
||||
|
||||
logger.LogInformation($"Starting the augmentation process over {numTotalRetrievalContexts:###,###,###,###} retrieval contexts.");
|
||||
|
||||
//
|
||||
@@ -38,63 +63,8 @@ public sealed class AugmentationOne : IAugmentationProcess
|
||||
sb.AppendLine("The following useful information will help you in processing the user prompt:");
|
||||
sb.AppendLine();
|
||||
|
||||
var index = 0;
|
||||
foreach(var retrievalContext in retrievalContexts)
|
||||
{
|
||||
index++;
|
||||
sb.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
|
||||
sb.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
|
||||
sb.AppendLine($"Content category: {retrievalContext.Category}");
|
||||
sb.AppendLine($"Content type: {retrievalContext.Type}");
|
||||
sb.AppendLine($"Content path: {retrievalContext.Path}");
|
||||
|
||||
if(retrievalContext.Links.Count > 0)
|
||||
{
|
||||
sb.AppendLine("Additional links:");
|
||||
foreach(var link in retrievalContext.Links)
|
||||
sb.AppendLine($"- {link}");
|
||||
}
|
||||
|
||||
switch(retrievalContext)
|
||||
{
|
||||
case RetrievalTextContext textContext:
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Matched text content:");
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(textContext.MatchedText);
|
||||
sb.AppendLine("````");
|
||||
|
||||
if(textContext.SurroundingContent.Count > 0)
|
||||
{
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Surrounding text content:");
|
||||
foreach(var surrounding in textContext.SurroundingContent)
|
||||
{
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(surrounding);
|
||||
sb.AppendLine("````");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
break;
|
||||
|
||||
case RetrievalImageContext imageContext:
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Matched image content as base64-encoded data:");
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(await imageContext.AsBase64(token));
|
||||
sb.AppendLine("````");
|
||||
break;
|
||||
|
||||
default:
|
||||
logger.LogWarning($"The retrieval content type '{retrievalContext.Type}' of data source '{retrievalContext.DataSourceName}' at location '{retrievalContext.Path}' is not supported yet.");
|
||||
break;
|
||||
}
|
||||
|
||||
sb.AppendLine();
|
||||
}
|
||||
// Let's convert all retrieval contexts to Markdown:
|
||||
await retrievalContexts.AsMarkdown(sb, token);
|
||||
|
||||
//
|
||||
// Append the entire augmentation to the chat thread,
|
||||
|
||||
+3
-35
@@ -70,41 +70,9 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
|
||||
|
||||
if (aiSelectedDataSources.Count > 3)
|
||||
{
|
||||
//
|
||||
// We have more than 3 data sources. Let's filter by confidence.
|
||||
// In order to do that, we must identify the lower and upper
|
||||
// bounds of the confidence interval:
|
||||
//
|
||||
var confidenceValues = aiSelectedDataSources.Select(x => x.Confidence).ToList();
|
||||
var lowerBound = confidenceValues.Min();
|
||||
var upperBound = confidenceValues.Max();
|
||||
|
||||
//
|
||||
// Next, we search for a threshold so that we have between 2 and 3
|
||||
// data sources. When not possible, we take all data sources.
|
||||
//
|
||||
var threshold = 0.0f;
|
||||
|
||||
// Check the case where the confidence values are too close:
|
||||
if (upperBound - lowerBound >= 0.01)
|
||||
{
|
||||
var previousThreshold = 0.0f;
|
||||
for (var i = 0; i < 10; i++)
|
||||
{
|
||||
threshold = lowerBound + (upperBound - lowerBound) * i / 10;
|
||||
var numMatches = aiSelectedDataSources.Count(x => x.Confidence >= threshold);
|
||||
if (numMatches <= 1)
|
||||
{
|
||||
threshold = previousThreshold;
|
||||
break;
|
||||
}
|
||||
|
||||
if (numMatches is <= 3 and >= 2)
|
||||
break;
|
||||
|
||||
previousThreshold = threshold;
|
||||
}
|
||||
}
|
||||
// We have more than 3 data sources. Let's filter by confidence:
|
||||
var targetWindow = aiSelectedDataSources.DetermineTargetWindow(TargetWindowStrategy.A_FEW_GOOD_ONES);
|
||||
var threshold = aiSelectedDataSources.GetConfidenceThreshold(targetWindow);
|
||||
|
||||
//
|
||||
// Filter the data sources by the threshold:
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
using System.Text;
|
||||
|
||||
using AIStudio.Chat;
|
||||
|
||||
namespace AIStudio.Tools.RAG;
|
||||
|
||||
public static class IRetrievalContextExtensions
|
||||
{
|
||||
private static readonly ILogger<IRetrievalContext> LOGGER = Program.SERVICE_PROVIDER.GetService<ILogger<IRetrievalContext>>()!;
|
||||
|
||||
public static async Task<string> AsMarkdown(this IReadOnlyList<IRetrievalContext> retrievalContexts, StringBuilder? sb = null, CancellationToken token = default)
|
||||
{
|
||||
sb ??= new StringBuilder();
|
||||
var index = 0;
|
||||
|
||||
foreach(var retrievalContext in retrievalContexts)
|
||||
{
|
||||
index++;
|
||||
await retrievalContext.AsMarkdown(sb, index, retrievalContexts.Count, token);
|
||||
}
|
||||
|
||||
return sb.ToString();
|
||||
}
|
||||
|
||||
public static async Task<string> AsMarkdown(this IRetrievalContext retrievalContext, StringBuilder? sb = null, int index = -1, int numTotalRetrievalContexts = -1, CancellationToken token = default)
|
||||
{
|
||||
sb ??= new StringBuilder();
|
||||
switch (index)
|
||||
{
|
||||
case > 0 when numTotalRetrievalContexts is -1:
|
||||
sb.AppendLine($"# Retrieval context {index}");
|
||||
break;
|
||||
|
||||
case > 0 when numTotalRetrievalContexts > 0:
|
||||
sb.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
|
||||
break;
|
||||
|
||||
default:
|
||||
sb.AppendLine("# Retrieval context");
|
||||
break;
|
||||
}
|
||||
|
||||
sb.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
|
||||
sb.AppendLine($"Content category: {retrievalContext.Category}");
|
||||
sb.AppendLine($"Content type: {retrievalContext.Type}");
|
||||
sb.AppendLine($"Content path: {retrievalContext.Path}");
|
||||
|
||||
if(retrievalContext.Links.Count > 0)
|
||||
{
|
||||
sb.AppendLine("Additional links:");
|
||||
foreach(var link in retrievalContext.Links)
|
||||
sb.AppendLine($"- {link}");
|
||||
}
|
||||
|
||||
switch(retrievalContext)
|
||||
{
|
||||
case RetrievalTextContext textContext:
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Matched text content:");
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(textContext.MatchedText);
|
||||
sb.AppendLine("````");
|
||||
|
||||
if(textContext.SurroundingContent.Count > 0)
|
||||
{
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Surrounding text content:");
|
||||
foreach(var surrounding in textContext.SurroundingContent)
|
||||
{
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(surrounding);
|
||||
sb.AppendLine("````");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
break;
|
||||
|
||||
case RetrievalImageContext imageContext:
|
||||
sb.AppendLine();
|
||||
sb.AppendLine("Matched image content as base64-encoded data:");
|
||||
sb.AppendLine("````");
|
||||
sb.AppendLine(await imageContext.AsBase64(token));
|
||||
sb.AppendLine("````");
|
||||
break;
|
||||
|
||||
default:
|
||||
LOGGER.LogWarning($"The retrieval content type '{retrievalContext.Type}' of data source '{retrievalContext.DataSourceName}' at location '{retrievalContext.Path}' is not supported yet.");
|
||||
break;
|
||||
}
|
||||
|
||||
sb.AppendLine();
|
||||
return sb.ToString();
|
||||
}
|
||||
}
|
||||
@@ -9,7 +9,7 @@ public static class RetrievalContentCategoryExtensions
|
||||
/// </summary>
|
||||
/// <param name="contentType">The content type yielded by the ERI server.</param>
|
||||
/// <returns>The corresponding retrieval content category.</returns>
|
||||
public static RetrievalContentCategory ToRetrievalContentCategory(ContentType contentType) => contentType switch
|
||||
public static RetrievalContentCategory ToRetrievalContentCategory(this ContentType contentType) => contentType switch
|
||||
{
|
||||
ContentType.NONE => RetrievalContentCategory.NONE,
|
||||
ContentType.UNKNOWN => RetrievalContentCategory.UNKNOWN,
|
||||
|
||||
Reference in new issue
Block a user