.NET: Add Rag AIContext Provider (#1630)

* Add Rag AIContext Provider

* Fix issues

* Improve options naming based on PR feedback.

* Move Rag Provider to Data namespace

* Add Raw Representation to RagSearchResult

* Renaming RagProvider to TextSearchProvider
This commit is contained in:
westey
2025-10-28 14:16:29 +00:00
committed by GitHub
Unverified
parent b70030daec
commit a816408cd4
4 changed files with 932 additions and 0 deletions
@@ -4,6 +4,7 @@ using System.Diagnostics.CodeAnalysis;
using System.Text.Encodings.Web;
using System.Text.Json;
using System.Text.Json.Serialization;
using Microsoft.Agents.AI.Data;
namespace Microsoft.Agents.AI;
@@ -62,6 +63,7 @@ internal static partial class AgentJsonUtilities
// Agent abstraction types
[JsonSerializable(typeof(ChatClientAgentThread.ThreadState))]
[JsonSerializable(typeof(TextSearchProvider.TextSearchProviderState))]
[ExcludeFromCodeCoverage]
internal sealed partial class JsonContext : JsonSerializerContext;
@@ -0,0 +1,311 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Microsoft.Extensions.Logging;
using Microsoft.Shared.Diagnostics;
namespace Microsoft.Agents.AI.Data;
/// <summary>
/// A text search context provider that performs a search over external knowledge
/// and injects the formatted results into the AI invocation context, or exposes a search tool for on-demand use.
/// This provider can be used to enable Retrieval Augmented Generation (RAG) on an agent.
/// </summary>
/// <remarks>
/// <para>
/// The provider supports two behaviors controlled via <see cref="TextSearchProviderOptions.SearchTime"/>:
/// <list type="bullet">
/// <item><description><see cref="TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke"/> Automatically performs a search prior to every AI invocation and injects results as additional instructions.</description></item>
/// <item><description><see cref="TextSearchProviderOptions.TextSearchBehavior.OnDemandFunctionCalling"/> Exposes a function tool that the model may invoke to retrieve contextual information when needed.</description></item>
/// </list>
/// </para>
/// <para>
/// When <see cref="TextSearchProviderOptions.RecentMessageMemoryLimit"/> is greater than zero the provider will retain the most recent
/// user and assistant messages (up to the configured limit) across invocations and prepend them (in chronological order)
/// to the current request messages when forming the search input. This can improve search relevance by providing
/// multi-turn context to the retrieval layer without permanently altering the conversation history.
/// </para>
/// </remarks>
public sealed class TextSearchProvider : AIContextProvider
{
private const string DefaultPluginSearchFunctionName = "Search";
private const string DefaultPluginSearchFunctionDescription = "Allows searching for additional information to help answer the user question.";
private const string DefaultContextPrompt = "## Additional Context\nConsider the following information from source documents when responding to the user:";
private const string DefaultCitationsPrompt = "Include citations to the source document with document name and link if document name and link is available.";
private readonly Func<string, CancellationToken, Task<IEnumerable<TextSearchSearchResult>>> _searchAsync;
private readonly ILogger<TextSearchProvider>? _logger;
private readonly AITool[] _tools;
private readonly Queue<string> _recentMessagesText;
private readonly TextSearchProviderOptions _options;
/// <summary>
/// Initializes a new instance of the <see cref="TextSearchProvider"/> class.
/// </summary>
/// <param name="searchAsync">Delegate that executes the search logic. Must not be <see langword="null"/>.</param>
/// <param name="options">Optional configuration options.</param>
/// <param name="loggerFactory">Optional logger factory.</param>
/// <exception cref="ArgumentNullException">Thrown when <paramref name="searchAsync"/> is <see langword="null"/>.</exception>
public TextSearchProvider(Func<string, CancellationToken, Task<IEnumerable<TextSearchSearchResult>>> searchAsync, TextSearchProviderOptions? options = null, ILoggerFactory? loggerFactory = null)
{
this._searchAsync = searchAsync ?? throw new ArgumentNullException(nameof(searchAsync));
this._options = options ?? new();
Throw.IfLessThan(this._options.RecentMessageMemoryLimit, 0);
this._logger = loggerFactory?.CreateLogger<TextSearchProvider>();
this._recentMessagesText = new();
// Create the on-demand search tool (only used if behavior is OnDemandFunctionCalling)
this._tools =
[
AIFunctionFactory.Create(
this.SearchAsync,
name: this._options.FunctionToolName ?? DefaultPluginSearchFunctionName,
description: this._options.FunctionToolDescription ?? DefaultPluginSearchFunctionDescription)
];
}
/// <summary>
/// Initializes a new instance of the <see cref="TextSearchProvider"/> class from previously serialized state.
/// </summary>
/// <param name="searchAsync">Delegate that executes the search logic. Must not be <see langword="null"/>.</param>
/// <param name="serializedState">A <see cref="JsonElement"/> representing the serialized provider state.</param>
/// <param name="jsonSerializerOptions">Optional serializer options (unused - source generated context is used).</param>
/// <param name="options">Optional configuration options.</param>
/// <param name="loggerFactory">Optional logger factory.</param>
/// <exception cref="ArgumentNullException">Thrown when <paramref name="searchAsync"/> is <see langword="null"/>.</exception>
/// <remarks>
/// Only overridden prompts (function name, function description, context prompt, citations prompt) are restored.
/// If a value was not persisted or matches the defaults it will fall back to the built-in defaults.
/// Custom <see cref="TextSearchProviderOptions.ContextFormatter"/> delegates are not serialized.
/// </remarks>
public TextSearchProvider(Func<string, CancellationToken, Task<IEnumerable<TextSearchSearchResult>>> searchAsync, JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, TextSearchProviderOptions? options = null, ILoggerFactory? loggerFactory = null)
{
this._searchAsync = searchAsync ?? throw new ArgumentNullException(nameof(searchAsync));
this._options = options ?? new();
Throw.IfLessThan(this._options.RecentMessageMemoryLimit, 0);
this._logger = loggerFactory?.CreateLogger<TextSearchProvider>();
List<string>? restoredMessages = null;
var state = serializedState.Deserialize(AgentJsonUtilities.DefaultOptions.GetTypeInfo(typeof(TextSearchProviderState))) as TextSearchProviderState;
if (state?.RecentMessagesText is { Count: > 0 })
{
restoredMessages = state.RecentMessagesText;
}
// Restore recent messages respecting the limit (may truncate if limit changed afterwards).
this._recentMessagesText = restoredMessages is null ? new() : new(restoredMessages.Take(this._options.RecentMessageMemoryLimit));
// Create the on-demand search tool (only used if behavior is OnDemandFunctionCalling)
this._tools =
[
AIFunctionFactory.Create(
this.SearchAsync,
name: this._options.FunctionToolName ?? DefaultPluginSearchFunctionName,
description: this._options.FunctionToolDescription ?? DefaultPluginSearchFunctionDescription)
];
}
/// <inheritdoc />
public override async ValueTask<AIContext> InvokingAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
if (this._options.SearchTime != TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke)
{
// Expose the search tool for on-demand invocation.
return new AIContext { Tools = this._tools }; // No automatic instructions injection.
}
// Aggregate text from memory + current request messages.
var sbInput = new StringBuilder();
foreach (var messageText in this._recentMessagesText)
{
if (sbInput.Length > 0)
{
sbInput.Append('\n');
}
sbInput.Append(messageText);
}
foreach (var message in context.RequestMessages)
{
if (!string.IsNullOrWhiteSpace(message?.Text))
{
if (sbInput.Length > 0)
{
sbInput.Append('\n');
}
sbInput.Append(message.Text);
}
}
string input = sbInput.ToString();
// Search
var results = await this._searchAsync(input, cancellationToken).ConfigureAwait(false);
IList<TextSearchSearchResult> materialized = results as IList<TextSearchSearchResult> ?? results.ToList();
if (materialized.Count == 0)
{
this._logger?.LogWarning("TextSearchProvider: No search results found.");
return new AIContext();
}
// Format search results
string formatted = this.FormatResults(materialized);
this._logger?.LogInformation("TextSearchProvider: Retrieved {Count} search results.", materialized.Count);
this._logger?.LogTrace("TextSearchProvider Input:{Input}\nContext Instructions:{Instructions}", input, formatted);
return new AIContext
{
Messages = [new ChatMessage(ChatRole.User, formatted)]
};
}
/// <inheritdoc />
public override ValueTask InvokedAsync(InvokedContext context, CancellationToken cancellationToken = default)
{
int limit = this._options.RecentMessageMemoryLimit;
if (limit <= 0)
{
return default; // Memory disabled.
}
var messagesText = context.RequestMessages
.Concat(context.ResponseMessages ?? [])
.Where(m => (m.Role == ChatRole.User || m.Role == ChatRole.Assistant) && !string.IsNullOrWhiteSpace(m.Text))
.Select(m => m.Text)
.ToList();
if (messagesText.Count > limit)
{
// If the current request/response exceeds the limit, only keep the most recent messages from it.
messagesText = messagesText.Skip(messagesText.Count - limit).ToList();
}
foreach (var message in messagesText)
{
this._recentMessagesText.Enqueue(message);
}
while (this._recentMessagesText.Count > limit)
{
this._recentMessagesText.Dequeue();
}
return default;
}
/// <summary>
/// Serializes the current provider state to a <see cref="JsonElement"/> containing any overridden prompts or descriptions.
/// </summary>
/// <param name="jsonSerializerOptions">Optional serializer options (ignored, source generated context is used).</param>
/// <returns>A <see cref="JsonElement"/> with overridden values, or default if nothing was overridden.</returns>
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
{
// Only persist values that differ from defaults plus recent memory configuration & messages.
TextSearchProviderState state = new();
if (this._options.RecentMessageMemoryLimit > 0 && this._recentMessagesText.Count > 0)
{
state.RecentMessagesText = this._recentMessagesText.Take(this._options.RecentMessageMemoryLimit).ToList();
}
return JsonSerializer.SerializeToElement(state, AgentJsonUtilities.DefaultOptions.GetTypeInfo(typeof(TextSearchProviderState)));
}
/// <summary>
/// Function callable by the AI model (when enabled) to perform an ad-hoc search.
/// </summary>
/// <param name="userQuestion">The query text.</param>
/// <param name="cancellationToken">Cancellation token.</param>
/// <returns>Formatted search results.</returns>
internal async Task<string> SearchAsync(string userQuestion, CancellationToken cancellationToken = default)
{
var results = await this._searchAsync(userQuestion, cancellationToken).ConfigureAwait(false);
IList<TextSearchSearchResult> materialized = results as IList<TextSearchSearchResult> ?? results.ToList();
string formatted = this.FormatResults(materialized);
this._logger?.LogInformation("TextSearchProvider: Retrieved {Count} search results.", materialized.Count);
this._logger?.LogTrace("TextSearchProvider Input:{UserQuestion}\nContext Instructions:{Instructions}", userQuestion, formatted);
return formatted;
}
/// <summary>
/// Formats search results into an instructions string for model consumption.
/// </summary>
/// <param name="results">The results.</param>
/// <returns>Formatted string (may be empty).</returns>
private string FormatResults(IList<TextSearchSearchResult> results)
{
if (this._options.ContextFormatter is not null)
{
return this._options.ContextFormatter(results) ?? string.Empty;
}
if (results.Count == 0)
{
return string.Empty; // No extra context.
}
var sb = new StringBuilder();
sb.AppendLine(this._options.ContextPrompt ?? DefaultContextPrompt);
for (int i = 0; i < results.Count; i++)
{
var result = results[i];
if (!string.IsNullOrWhiteSpace(result.Name))
{
sb.AppendLine($"SourceDocName: {result.Name}");
}
if (!string.IsNullOrWhiteSpace(result.Link))
{
sb.AppendLine($"SourceDocLink: {result.Link}");
}
sb.AppendLine($"Contents: {result.Value}");
sb.AppendLine("----");
}
sb.AppendLine(this._options.CitationsPrompt ?? DefaultCitationsPrompt);
sb.AppendLine();
return sb.ToString();
}
/// <summary>
/// Represents a single retrieved text search result.
/// </summary>
public sealed class TextSearchSearchResult
{
/// <summary>
/// Gets or sets the display name of the source document (optional).
/// </summary>
public string? Name { get; set; }
/// <summary>
/// Gets or sets a link/URL to the source document (optional).
/// </summary>
public string? Link { get; set; }
/// <summary>
/// Gets or sets the textual content of the retrieved chunk.
/// </summary>
public string Value { get; set; } = string.Empty;
/// <summary>
/// Gets or sets the raw representation of the search result from the data source.
/// </summary>
/// <remarks>
/// If a <see cref="TextSearchSearchResult"/> is created to represent some underlying object from another object
/// model, this property can be used to store that original object. This can be useful for debugging or
/// for enabling the <see cref="TextSearchProviderOptions.ContextFormatter"/> to access the underlying object model if needed.
/// </remarks>
public object? RawRepresentation { get; set; }
}
internal sealed class TextSearchProviderState
{
public List<string>? RecentMessagesText { get; set; }
}
}
@@ -0,0 +1,77 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using Microsoft.Extensions.AI;
namespace Microsoft.Agents.AI.Data;
/// <summary>
/// Options controlling the behavior of <see cref="TextSearchProvider"/>.
/// </summary>
public sealed class TextSearchProviderOptions
{
/// <summary>
/// Gets or sets a value indicating when the search should be executed.
/// </summary>
/// <value><see cref="TextSearchBehavior.BeforeAIInvoke"/> by default.</value>
public TextSearchBehavior SearchTime { get; set; } = TextSearchBehavior.BeforeAIInvoke;
/// <summary>
/// Gets or sets the name of the exposed search tool when operating in on-demand mode.
/// </summary>
/// <value>Defaults to "Search".</value>
public string? FunctionToolName { get; set; }
/// <summary>
/// Gets or sets the description of the exposed search tool when operating in on-demand mode.
/// </summary>
/// <value>Defaults to "Allows searching for additional information to help answer the user question.".</value>
public string? FunctionToolDescription { get; set; }
/// <summary>
/// Gets or sets the context prompt prefixed to automatically injected results.
/// </summary>
public string? ContextPrompt { get; set; }
/// <summary>
/// Gets or sets the instruction appended after automatically injected results to request citations.
/// </summary>
public string? CitationsPrompt { get; set; }
/// <summary>
/// Optional delegate to fully customize formatting of the result list.
/// </summary>
/// <remarks>
/// If provided, <see cref="ContextPrompt"/> and <see cref="CitationsPrompt"/> are ignored.
/// </remarks>
public Func<IList<TextSearchProvider.TextSearchSearchResult>, string>? ContextFormatter { get; set; }
/// <summary>
/// Gets or sets the number of recent conversation messages (both user and assistant) to keep in memory
/// and include when constructing the search input for <see cref="TextSearchBehavior.BeforeAIInvoke"/> searches.
/// </summary>
/// <value>
/// The maximum number of most recent messages to retain. A value of <c>0</c> (default) disables memory and
/// only the current request's messages are used for search input. The value is a count of individual
/// messages, not turns. Only messages with role <see cref="ChatRole.User"/> or
/// <see cref="ChatRole.Assistant"/> are retained.
/// </value>
public int RecentMessageMemoryLimit { get; set; }
/// <summary>
/// Behavior choices for the provider.
/// </summary>
public enum TextSearchBehavior
{
/// <summary>
/// Execute search prior to each invocation and inject results as instructions.
/// </summary>
BeforeAIInvoke,
/// <summary>
/// Expose a function tool to perform search on-demand via function/tool calling.
/// </summary>
OnDemandFunctionCalling
}
}
@@ -0,0 +1,542 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Agents.AI.Data;
using Microsoft.Extensions.AI;
using Microsoft.Extensions.Logging;
using Moq;
namespace Microsoft.Agents.AI.UnitTests.Data;
/// <summary>
/// Contains unit tests for <see cref="TextSearchProvider"/>.
/// </summary>
public sealed class TextSearchProviderTests
{
private readonly Mock<ILogger<TextSearchProvider>> _loggerMock;
private readonly Mock<ILoggerFactory> _loggerFactoryMock;
public TextSearchProviderTests()
{
this._loggerMock = new();
this._loggerFactoryMock = new();
this._loggerFactoryMock
.Setup(f => f.CreateLogger(It.IsAny<string>()))
.Returns(this._loggerMock.Object);
this._loggerFactoryMock
.Setup(f => f.CreateLogger(typeof(TextSearchProvider).FullName!))
.Returns(this._loggerMock.Object);
}
[Theory]
[InlineData(null, null, true)]
[InlineData("Custom context prompt", "Custom citations prompt", false)]
public async Task InvokingAsync_ShouldInjectFormattedResultsAsync(string? overrideContextPrompt, string? overrideCitationsPrompt, bool withLogging)
{
// Arrange
List<TextSearchProvider.TextSearchSearchResult> results =
[
new() { Name = "Doc1", Link = "http://example.com/doc1", Value = "Content of Doc1" },
new() { Name = "Doc2", Link = "http://example.com/doc2", Value = "Content of Doc2" }
];
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>(results);
}
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
ContextPrompt = overrideContextPrompt,
CitationsPrompt = overrideCitationsPrompt
};
var provider = new TextSearchProvider(SearchDelegateAsync, options, withLogging ? this._loggerFactoryMock.Object : null);
var invokingContext = new AIContextProvider.InvokingContext(new[]
{
new ChatMessage(ChatRole.User, "Sample user question?"),
new ChatMessage(ChatRole.User, "Additional part")
});
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Equal("Sample user question?\nAdditional part", capturedInput);
Assert.Null(aiContext.Instructions); // TextSearchProvider uses a user message for context injection.
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
var message = aiContext.Messages!.Single();
Assert.Equal(ChatRole.User, message.Role);
string text = message.Text!;
if (overrideContextPrompt is null)
{
Assert.Contains("## Additional Context", text);
Assert.Contains("Consider the following information from source documents when responding to the user:", text);
}
else
{
Assert.Contains(overrideContextPrompt, text);
}
Assert.Contains("SourceDocName: Doc1", text);
Assert.Contains("SourceDocLink: http://example.com/doc1", text);
Assert.Contains("Contents: Content of Doc1", text);
Assert.Contains("SourceDocName: Doc2", text);
Assert.Contains("SourceDocLink: http://example.com/doc2", text);
Assert.Contains("Contents: Content of Doc2", text);
if (overrideCitationsPrompt is null)
{
Assert.Contains("Include citations to the source document with document name and link if document name and link is available.", text);
}
else
{
Assert.Contains(overrideCitationsPrompt, text);
}
if (withLogging)
{
this._loggerMock.Verify(
l => l.Log(
LogLevel.Information,
It.IsAny<EventId>(),
It.Is<It.IsAnyType>((v, t) => v.ToString()!.Contains("TextSearchProvider: Retrieved 2 search results.")),
It.IsAny<Exception?>(),
It.IsAny<Func<It.IsAnyType, Exception?, string>>()),
Times.AtLeastOnce);
this._loggerMock.Verify(
l => l.Log(
LogLevel.Trace,
It.IsAny<EventId>(),
It.Is<It.IsAnyType>((v, t) => v.ToString()!.Contains("TextSearchProvider Input:Sample user question?\nAdditional part\nContext Instructions")),
It.IsAny<Exception?>(),
It.IsAny<Func<It.IsAnyType, Exception?, string>>()),
Times.AtLeastOnce);
}
}
[Theory]
[InlineData(null, null, "Search", "Allows searching for additional information to help answer the user question.")]
[InlineData("CustomSearch", "CustomDescription", "CustomSearch", "CustomDescription")]
public async Task InvokingAsync_OnDemand_ShouldExposeSearchToolAsync(string? overrideName, string? overrideDescription, string expectedName, string expectedDescription)
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.OnDemandFunctionCalling,
FunctionToolName = overrideName,
FunctionToolDescription = overrideDescription
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(new[] { new ChatMessage(ChatRole.User, "Q?") });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages); // No automatic injection.
Assert.NotNull(aiContext.Tools);
Assert.Single(aiContext.Tools);
var tool = aiContext.Tools.Single();
Assert.Equal(expectedName, tool.Name);
Assert.Equal(expectedDescription, tool.Description);
}
[Theory]
[InlineData(null, null)]
[InlineData("Custom context prompt", "Custom citations prompt")]
public async Task SearchAsync_ShouldReturnFormattedResultsAsync(string? overrideContextPrompt, string? overrideCitationsPrompt)
{
// Arrange
List<TextSearchProvider.TextSearchSearchResult> results =
[
new() { Name = "Doc1", Link = "http://example.com/doc1", Value = "Content of Doc1" },
new() { Name = "Doc2", Link = "http://example.com/doc2", Value = "Content of Doc2" }
];
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>(results);
}
var options = new TextSearchProviderOptions
{
ContextPrompt = overrideContextPrompt,
CitationsPrompt = overrideCitationsPrompt
};
var provider = new TextSearchProvider(SearchDelegateAsync, options);
// Act
var formatted = await provider.SearchAsync("Sample user question?", CancellationToken.None);
// Assert
if (overrideContextPrompt is null)
{
Assert.Contains("## Additional Context", formatted);
Assert.Contains("Consider the following information from source documents when responding to the user:", formatted);
}
else
{
Assert.Contains(overrideContextPrompt, formatted);
}
Assert.Contains("SourceDocName: Doc1", formatted);
Assert.Contains("SourceDocLink: http://example.com/doc1", formatted);
Assert.Contains("Contents: Content of Doc1", formatted);
Assert.Contains("SourceDocName: Doc2", formatted);
Assert.Contains("SourceDocLink: http://example.com/doc2", formatted);
Assert.Contains("Contents: Content of Doc2", formatted);
if (overrideCitationsPrompt is null)
{
Assert.Contains("Include citations to the source document with document name and link if document name and link is available.", formatted);
}
else
{
Assert.Contains(overrideCitationsPrompt, formatted);
}
}
[Fact]
public async Task InvokingAsync_ShouldUseContextFormatterWhenProvidedAsync()
{
// Arrange
List<TextSearchProvider.TextSearchSearchResult> results =
[
new() { Name = "Doc1", Link = "http://example.com/doc1", Value = "Content of Doc1" },
new() { Name = "Doc2", Link = "http://example.com/doc2", Value = "Content of Doc2" }
];
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>(results);
}
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
ContextFormatter = r => $"Custom formatted context with {r.Count} results."
};
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(new[] { new ChatMessage(ChatRole.User, "Q?") });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
Assert.Equal("Custom formatted context with 2 results.", aiContext.Messages![0].Text);
}
[Fact]
public async Task InvokingAsync_WithRawRepresentations_ContextFormatterCanAccessAsync()
{
// Arrange
var payload1 = new RawPayload { Id = "R1" };
var payload2 = new RawPayload { Id = "R2" };
List<TextSearchProvider.TextSearchSearchResult> results =
[
new() { Name = "Doc1", Value = "Content 1", RawRepresentation = payload1 },
new() { Name = "Doc2", Value = "Content 2", RawRepresentation = payload2 }
];
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>(results);
}
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
ContextFormatter = r => string.Join(",", r.Select(x => ((RawPayload)x.RawRepresentation!).Id))
};
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(new[] { new ChatMessage(ChatRole.User, "Q?") });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
Assert.Equal("R1,R2", aiContext.Messages![0].Text);
}
[Fact]
public async Task InvokingAsync_WithNoResults_ShouldReturnEmptyContextAsync()
{
// Arrange
var options = new TextSearchProviderOptions { SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke };
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(new[] { new ChatMessage(ChatRole.User, "Q?") });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages);
Assert.Null(aiContext.Instructions);
Assert.Null(aiContext.Tools);
}
#region Recent Message Memory Tests
[Fact]
public async Task InvokingAsync_WithRecentMessageMemory_ShouldIncludeStoredMessagesInSearchInputAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
};
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]); // No results needed.
}
var provider = new TextSearchProvider(SearchDelegateAsync, options);
// Populate memory with more messages than the limit (A,B,C,D) -> should retain B,C,D
var initialMessages = new[]
{
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B"),
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
};
await provider.InvokedAsync(new(initialMessages));
var invokingContext = new AIContextProvider.InvokingContext(new[]
{
new ChatMessage(ChatRole.User, "E")
});
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Equal("B\nC\nD\nE", capturedInput); // Memory first (truncated) then current request.
}
[Fact]
public async Task InvokingAsync_WithAccumulatedMemoryAcrossInvocations_ShouldIncludeAllUpToLimitAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 5
};
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync, options);
// First memory update (A,B)
await provider.InvokedAsync(new(new[]
{
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B"),
}));
// Second memory update (C,D,E)
await provider.InvokedAsync(new(new[]
{
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
new ChatMessage(ChatRole.User, "E"),
}));
var invokingContext = new AIContextProvider.InvokingContext(new[] { new ChatMessage(ChatRole.User, "F") });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Equal("A\nB\nC\nD\nE\nF", capturedInput); // All retained (limit 5) + current request message.
}
#endregion
#region Serialization Tests
[Fact]
public void Serialize_WithNoRecentMessages_ShouldReturnEmptyState()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
// Act
var state = provider.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, state.ValueKind);
Assert.False(state.TryGetProperty("recentMessagesText", out _));
}
[Fact]
public async Task Serialize_WithRecentMessages_ShouldPersistMessagesUpToLimitAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var messages = new[]
{
new ChatMessage(ChatRole.User, "M1"),
new ChatMessage(ChatRole.Assistant, "M2"),
new ChatMessage(ChatRole.User, "M3"),
};
// Act
await provider.InvokedAsync(new(messages)); // Populate recent memory.
var state = provider.Serialize();
// Assert
Assert.True(state.TryGetProperty("recentMessagesText", out var recentProperty));
Assert.Equal(JsonValueKind.Array, recentProperty.ValueKind);
var list = recentProperty.EnumerateArray().Select(e => e.GetString()).ToList();
Assert.Equal(3, list.Count);
Assert.Equal(["M1", "M2", "M3"], list);
}
[Fact]
public async Task SerializeAndDeserialize_RoundtripRestoresMessagesAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 4
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var messages = new[]
{
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B"),
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
};
await provider.InvokedAsync(new(messages));
// Act
var state = provider.Serialize();
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]);
}
var roundTrippedProvider = new TextSearchProvider(SearchDelegate2Async, state, options: new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 4
});
var emptyMessages = Array.Empty<ChatMessage>();
await roundTrippedProvider.InvokingAsync(new(emptyMessages), CancellationToken.None); // Trigger search to read memory.
// Assert
Assert.NotNull(capturedInput);
Assert.Equal("A\nB\nC\nD", capturedInput);
}
[Fact]
public async Task Deserialize_WithChangedLowerLimit_ShouldTruncateToNewLimitAsync()
{
// Arrange
var initialProvider = new TextSearchProvider(this.NoResultSearchAsync, new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 5
});
var messages = new[]
{
new ChatMessage(ChatRole.User, "L1"),
new ChatMessage(ChatRole.Assistant, "L2"),
new ChatMessage(ChatRole.User, "L3"),
new ChatMessage(ChatRole.Assistant, "L4"),
new ChatMessage(ChatRole.User, "L5"),
};
await initialProvider.InvokedAsync(new(messages));
var state = initialProvider.Serialize();
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]);
}
// Act
var restoredProvider = new TextSearchProvider(SearchDelegate2Async, state, options: new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3 // Lower limit
});
await restoredProvider.InvokingAsync(new(Array.Empty<ChatMessage>()), CancellationToken.None);
// Assert
Assert.NotNull(capturedInput);
Assert.Equal("L1\nL2\nL3", capturedInput);
}
[Fact]
public async Task Deserialize_WithEmptyState_ShouldHaveNoMessagesAsync()
{
// Arrange
var emptyState = JsonSerializer.Deserialize("{}", TestJsonSerializerContext.Default.JsonElement);
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]);
}
// Act
var provider = new TextSearchProvider(SearchDelegate2Async, emptyState, options: new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
});
var emptyMessages = Array.Empty<ChatMessage>();
await provider.InvokingAsync(new(emptyMessages), CancellationToken.None);
// Assert
Assert.NotNull(capturedInput);
Assert.Equal(string.Empty, capturedInput); // No recent messages serialized => empty input.
}
#endregion
private Task<IEnumerable<TextSearchProvider.TextSearchSearchResult>> NoResultSearchAsync(string input, CancellationToken ct)
{
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchSearchResult>>([]);
}
private sealed class RawPayload
{
public string Id { get; set; } = string.Empty;
}
}