[BREAKING] Add response filter for store input in *Providers (#4327)

* Add response filter for store input for *Providers

* Apply suggestions from code review

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>

* Address feedback

* Apply suggestions from code review

Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>

---------

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>
This commit is contained in:
westey
2026-03-03 11:02:02 +00:00
committed by GitHub
Unverified
parent 2e9319359b
commit 945933c351
31 changed files with 249 additions and 83 deletions
@@ -92,7 +92,6 @@ namespace SampleApp
private readonly IChatClient _chatClient;
public UserInfoMemory(IChatClient chatClient, Func<AgentSession?, UserInfo>? stateInitializer = null)
: base(null, null)
{
this._sessionState = new ProviderSessionState<UserInfo>(
stateInitializer ?? (_ => new UserInfo()),
@@ -73,7 +73,7 @@ AIAgent agent = azureOpenAIClient
// We also want to maintain that exclusion here.
ChatHistoryProvider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions
{
StorageInputMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory)
StorageInputRequestMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory)
}),
});
@@ -80,7 +80,7 @@ AIAgent agent = azureOpenAIClient
// You may choose to persist the TextSearchProvider messages, if you want the search output to be provided to the model in future interactions as well.
ChatHistoryProvider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions()
{
StorageInputMessageFilter = msgs => msgs.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider)
StorageInputRequestMessageFilter = msgs => msgs.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider)
})
});
@@ -85,7 +85,6 @@ namespace SampleApp
VectorStore vectorStore,
Func<AgentSession?, State>? stateInitializer = null,
string? stateKey = null)
: base(provideOutputMessageFilter: null, storeInputMessageFilter: null)
{
this._sessionState = new ProviderSessionState<State>(
stateInitializer ?? (_ => new State(Guid.NewGuid().ToString("N"))),
@@ -49,11 +49,11 @@ AIAgent agent = new AzureOpenAIClient(
""" },
ChatHistoryProvider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions
{
// Use StorageInputMessageFilter to provide a custom filter for messages stored in chat history.
// Use StorageInputRequestMessageFilter to provide a custom filter for request messages stored in chat history.
// By default the chat history provider will store all messages, except for those that came from chat history in the first place.
// In this case, we want to also exclude messages that came from AI context providers.
// You may want to store these messages, depending on their content and your requirements.
StorageInputMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory)
StorageInputRequestMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.AIContextProvider && m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory)
}),
// Add multiple AI context providers: one that maintains a todo list and one that provides upcoming calendar entries.
// The agent will call each provider in sequence, accumulating context from each.
@@ -33,18 +33,23 @@ public abstract class AIContextProvider
{
private static IEnumerable<ChatMessage> DefaultExternalOnlyFilter(IEnumerable<ChatMessage> messages)
=> messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External);
private static IEnumerable<ChatMessage> DefaultNoopFilter(IEnumerable<ChatMessage> messages)
=> messages;
/// <summary>
/// Initializes a new instance of the <see cref="AIContextProvider"/> class.
/// </summary>
/// <param name="provideInputMessageFilter">An optional filter function to apply to input messages before providing context via <see cref="ProvideAIContextAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to request messages before storing context via <see cref="StoreAIContextAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing context via <see cref="StoreAIContextAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing context via <see cref="StoreAIContextAsync"/>. If not set, defaults to a no-op filter that includes all response messages.</param>
protected AIContextProvider(
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideInputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
{
this.ProvideInputMessageFilter = provideInputMessageFilter ?? DefaultExternalOnlyFilter;
this.StoreInputMessageFilter = storeInputMessageFilter ?? DefaultExternalOnlyFilter;
this.StoreInputRequestMessageFilter = storeInputRequestMessageFilter ?? DefaultExternalOnlyFilter;
this.StoreInputResponseMessageFilter = storeInputResponseMessageFilter ?? DefaultNoopFilter;
}
/// <summary>
@@ -55,7 +60,12 @@ public abstract class AIContextProvider
/// <summary>
/// Gets the filter function to apply to request messages before storing context via <see cref="StoreAIContextAsync"/>.
/// </summary>
protected Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> StoreInputMessageFilter { get; }
protected Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> StoreInputRequestMessageFilter { get; }
/// <summary>
/// Gets the filter function to apply to response messages before storing context via <see cref="StoreAIContextAsync"/>.
/// </summary>
protected Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> StoreInputResponseMessageFilter { get; }
/// <summary>
/// Gets the key used to store the provider state in the <see cref="AgentSession.StateBag"/>.
@@ -245,8 +255,10 @@ public abstract class AIContextProvider
/// </para>
/// <para>
/// The default implementation of this method skips execution for any invocation failures,
/// filters the request messages using the configured store-input message filter
/// filters the request messages using the configured store-input request message filter
/// (which defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages),
/// filters the response messages using the configured store-input response message filter
/// (which defaults to a no-op, so all response messages are processed),
/// and calls <see cref="StoreAIContextAsync"/> to process the invocation results.
/// For most scenarios, overriding <see cref="StoreAIContextAsync"/> is sufficient to process invocation results,
/// while still benefiting from the default error handling and filtering behavior.
@@ -261,7 +273,7 @@ public abstract class AIContextProvider
return default;
}
var subContext = new InvokedContext(context.Agent, context.Session, this.StoreInputMessageFilter(context.RequestMessages), context.ResponseMessages!);
var subContext = new InvokedContext(context.Agent, context.Session, this.StoreInputRequestMessageFilter(context.RequestMessages), this.StoreInputResponseMessageFilter(context.ResponseMessages!));
return this.StoreAIContextAsync(subContext, cancellationToken);
}
@@ -42,21 +42,27 @@ public abstract class ChatHistoryProvider
{
private static IEnumerable<ChatMessage> DefaultExcludeChatHistoryFilter(IEnumerable<ChatMessage> messages)
=> messages.Where(m => m.GetAgentRequestMessageSourceType() != AgentRequestMessageSourceType.ChatHistory);
private static IEnumerable<ChatMessage> DefaultNoopFilter(IEnumerable<ChatMessage> messages)
=> messages;
private readonly Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? _provideOutputMessageFilter;
private readonly Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> _storeInputMessageFilter;
private readonly Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> _storeInputRequestMessageFilter;
private readonly Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>> _storeInputResponseMessageFilter;
/// <summary>
/// Initializes a new instance of the <see cref="ChatHistoryProvider"/> class.
/// </summary>
/// <param name="provideOutputMessageFilter">An optional filter function to apply to messages when retrieving them from the chat history.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing them in the chat history. If not set, defaults to a no-op filter that includes all response messages.</param>
protected ChatHistoryProvider(
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideOutputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
{
this._provideOutputMessageFilter = provideOutputMessageFilter;
this._storeInputMessageFilter = storeInputMessageFilter ?? DefaultExcludeChatHistoryFilter;
this._storeInputRequestMessageFilter = storeInputRequestMessageFilter ?? DefaultExcludeChatHistoryFilter;
this._storeInputResponseMessageFilter = storeInputResponseMessageFilter ?? DefaultNoopFilter;
}
/// <summary>
@@ -216,7 +222,7 @@ public abstract class ChatHistoryProvider
/// To check if the invocation was successful, inspect the <see cref="InvokedContext.InvokeException"/> property.
/// </para>
/// <para>
/// The default implementation of this method, skips execution for any invocation failures, filters messages using the optional storage input message filter
/// The default implementation of this method, skips execution for any invocation failures, filters messages using the optional storage input request and response message filters
/// and calls <see cref="StoreChatHistoryAsync"/> to store new chat history messages.
/// For most scenarios, overriding <see cref="StoreChatHistoryAsync"/> is sufficient to store chat history messages, while still benefiting from the default error handling and filtering behavior.
/// However, for scenarios that require more control over error handling or message filtering, overriding this method allows you to directly control the messages that are stored for the invocation.
@@ -229,7 +235,7 @@ public abstract class ChatHistoryProvider
return default;
}
var subContext = new InvokedContext(context.Agent, context.Session, this._storeInputMessageFilter(context.RequestMessages), context.ResponseMessages!);
var subContext = new InvokedContext(context.Agent, context.Session, this._storeInputRequestMessageFilter(context.RequestMessages), this._storeInputResponseMessageFilter(context.ResponseMessages!));
return this.StoreChatHistoryAsync(subContext, cancellationToken);
}
@@ -38,7 +38,8 @@ public sealed class InMemoryChatHistoryProvider : ChatHistoryProvider
public InMemoryChatHistoryProvider(InMemoryChatHistoryProviderOptions? options = null)
: base(
options?.ProvideOutputMessageFilter,
options?.StorageInputMessageFilter)
options?.StorageInputRequestMessageFilter,
options?.StorageInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState<State>(
options?.StateInitializer ?? (_ => new State()),
@@ -59,7 +59,19 @@ public sealed class InMemoryChatHistoryProviderOptions
/// Depending on your requirements, you could provide a different filter, that also excludes
/// messages from e.g. AI context providers.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputMessageFilter { get; set; }
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputRequestMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to response messages before they are added to storage
/// during <see cref="ChatHistoryProvider.InvokedAsync"/>.
/// </summary>
/// <value>
/// When <see langword="null"/>, no filtering is applied to response messages before they are stored.
/// If you want to avoid persisting certain messages (for example, those with
/// <see cref="AgentRequestMessageSourceType.ChatHistory"/> source type or produced by AI context providers),
/// provide a filter that returns only the messages you want to keep.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputResponseMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to messages produced by this provider
@@ -34,11 +34,13 @@ public abstract class MessageAIContextProvider : AIContextProvider
/// Initializes a new instance of the <see cref="MessageAIContextProvider"/> class.
/// </summary>
/// <param name="provideInputMessageFilter">An optional filter function to apply to input messages before providing messages via <see cref="ProvideMessagesAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to request messages before storing messages via <see cref="AIContextProvider.StoreAIContextAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing messages via <see cref="AIContextProvider.StoreAIContextAsync"/>. If not set, defaults to including only <see cref="AgentRequestMessageSourceType.External"/> messages.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing messages via <see cref="AIContextProvider.StoreAIContextAsync"/>. If not set, defaults to including all response messages (no filtering).</param>
protected MessageAIContextProvider(
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideInputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: base(provideInputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: base(provideInputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
}
@@ -87,7 +87,8 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
/// <param name="ownsClient">Whether this instance owns the CosmosClient and should dispose it.</param>
/// <param name="stateKey">An optional key to use for storing the state in the <see cref="AgentSession.StateBag"/>.</param>
/// <param name="provideOutputMessageFilter">An optional filter function to apply to messages when retrieving them from the chat history.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing them in the chat history. If not set, defaults to storing all response messages.</param>
/// <exception cref="ArgumentNullException">Thrown when <paramref name="cosmosClient"/> or <paramref name="stateInitializer"/> is <see langword="null"/>.</exception>
/// <exception cref="ArgumentException">Thrown when any string parameter is null or whitespace.</exception>
public CosmosChatHistoryProvider(
@@ -98,8 +99,9 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
bool ownsClient = false,
string? stateKey = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideOutputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: base(provideOutputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: base(provideOutputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState<State>(
Throw.IfNull(stateInitializer),
@@ -123,7 +125,8 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
/// <param name="stateInitializer">A delegate that initializes the provider state on the first invocation.</param>
/// <param name="stateKey">An optional key to use for storing the state in the <see cref="AgentSession.StateBag"/>.</param>
/// <param name="provideOutputMessageFilter">An optional filter function to apply to messages when retrieving them from the chat history.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing them in the chat history. If not set, defaults to storing all response messages.</param>
/// <exception cref="ArgumentNullException">Thrown when any required parameter is null.</exception>
/// <exception cref="ArgumentException">Thrown when any string parameter is null or whitespace.</exception>
public CosmosChatHistoryProvider(
@@ -133,8 +136,9 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
Func<AgentSession?, State> stateInitializer,
string? stateKey = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideOutputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: this(new CosmosClient(Throw.IfNullOrWhitespace(connectionString)), databaseId, containerId, stateInitializer, ownsClient: true, stateKey, provideOutputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: this(new CosmosClient(Throw.IfNullOrWhitespace(connectionString)), databaseId, containerId, stateInitializer, ownsClient: true, stateKey, provideOutputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
}
@@ -148,7 +152,8 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
/// <param name="stateInitializer">A delegate that initializes the provider state on the first invocation.</param>
/// <param name="stateKey">An optional key to use for storing the state in the <see cref="AgentSession.StateBag"/>.</param>
/// <param name="provideOutputMessageFilter">An optional filter function to apply to messages when retrieving them from the chat history.</param>
/// <param name="storeInputMessageFilter">An optional filter function to apply to messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputRequestMessageFilter">An optional filter function to apply to request messages before storing them in the chat history. If not set, defaults to excluding messages with source type <see cref="AgentRequestMessageSourceType.ChatHistory"/>.</param>
/// <param name="storeInputResponseMessageFilter">An optional filter function to apply to response messages before storing them in the chat history. If not set, defaults to storing all response messages.</param>
/// <exception cref="ArgumentNullException">Thrown when any required parameter is null.</exception>
/// <exception cref="ArgumentException">Thrown when any string parameter is null or whitespace.</exception>
public CosmosChatHistoryProvider(
@@ -159,8 +164,9 @@ public sealed class CosmosChatHistoryProvider : ChatHistoryProvider, IDisposable
Func<AgentSession?, State> stateInitializer,
string? stateKey = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideOutputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: this(new CosmosClient(Throw.IfNullOrWhitespace(accountEndpoint), Throw.IfNull(tokenCredential)), databaseId, containerId, stateInitializer, ownsClient: true, stateKey, provideOutputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: this(new CosmosClient(Throw.IfNullOrWhitespace(accountEndpoint), Throw.IfNull(tokenCredential)), databaseId, containerId, stateInitializer, ownsClient: true, stateKey, provideOutputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
}
@@ -59,7 +59,7 @@ public sealed class FoundryMemoryProvider : AIContextProvider
Func<AgentSession?, State> stateInitializer,
FoundryMemoryProviderOptions? options = null,
ILoggerFactory? loggerFactory = null)
: base(options?.SearchInputMessageFilter, options?.StorageInputMessageFilter)
: base(options?.SearchInputMessageFilter, options?.StorageInputRequestMessageFilter, options?.StorageInputResponseMessageFilter)
{
Throw.IfNull(client);
Throw.IfNullOrWhitespace(memoryStoreName);
@@ -63,5 +63,14 @@ public sealed class FoundryMemoryProviderOptions
/// When <see langword="null"/>, the provider defaults to including only
/// <see cref="AgentRequestMessageSourceType.External"/> messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputMessageFilter { get; set; }
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputRequestMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to response messages when determining which messages to
/// extract memories from during <see cref="AIContextProvider.InvokedAsync"/>.
/// </summary>
/// <value>
/// When <see langword="null"/>, the provider does not filter response messages and includes all messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputResponseMessageFilter { get; set; }
}
@@ -52,7 +52,7 @@ public sealed class Mem0Provider : MessageAIContextProvider
/// </code>
/// </remarks>
public Mem0Provider(HttpClient httpClient, Func<AgentSession?, State> stateInitializer, Mem0ProviderOptions? options = null, ILoggerFactory? loggerFactory = null)
: base(options?.SearchInputMessageFilter, options?.StorageInputMessageFilter)
: base(options?.SearchInputMessageFilter, options?.StorageInputRequestMessageFilter, options?.StorageInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState<State>(
ValidateStateInitializer(Throw.IfNull(stateInitializer)),
@@ -47,5 +47,14 @@ public sealed class Mem0ProviderOptions
/// When <see langword="null"/>, the provider defaults to including only
/// <see cref="AgentRequestMessageSourceType.External"/> messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputMessageFilter { get; set; }
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputRequestMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to response messages when determining which messages to
/// extract memories from during <see cref="AIContextProvider.InvokedAsync"/>.
/// </summary>
/// <value>
/// When <see langword="null"/>, the provider applies no filtering and includes all response messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputResponseMessageFilter { get; set; }
}
@@ -22,7 +22,6 @@ internal sealed class WorkflowChatHistoryProvider : ChatHistoryProvider
/// and source generated serializers are required, or Native AOT / Trimming is required.
/// </param>
public WorkflowChatHistoryProvider(JsonSerializerOptions? jsonSerializerOptions = null)
: base(provideOutputMessageFilter: null, storeInputMessageFilter: null)
{
this._sessionState = new ProviderSessionState<StoreState>(
_ => new StoreState(),
@@ -88,7 +88,7 @@ public sealed class ChatHistoryMemoryProvider : MessageAIContextProvider, IDispo
Func<AgentSession?, State> stateInitializer,
ChatHistoryMemoryProviderOptions? options = null,
ILoggerFactory? loggerFactory = null)
: base(options?.SearchInputMessageFilter, options?.StorageInputMessageFilter)
: base(options?.SearchInputMessageFilter, options?.StorageInputRequestMessageFilter, options?.StorageInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState<State>(
Throw.IfNull(stateInitializer),
@@ -75,8 +75,16 @@ public sealed class ChatHistoryMemoryProviderOptions
/// When <see langword="null"/>, the provider defaults to including only
/// <see cref="AgentRequestMessageSourceType.External"/> messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputMessageFilter { get; set; }
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputRequestMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to response messages when storing recent chat history
/// during <see cref="AIContextProvider.InvokedAsync"/>.
/// </summary>
/// <value>
/// When <see langword="null"/>, the provider does not apply any filtering and includes all response messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputResponseMessageFilter { get; set; }
/// <summary>
/// Behavior choices for the provider.
/// </summary>
@@ -61,7 +61,7 @@ public sealed class TextSearchProvider : MessageAIContextProvider
Func<string, CancellationToken, Task<IEnumerable<TextSearchResult>>> searchAsync,
TextSearchProviderOptions? options = null,
ILoggerFactory? loggerFactory = null)
: base(options?.SearchInputMessageFilter, options?.StorageInputMessageFilter)
: base(options?.SearchInputMessageFilter, options?.StorageInputRequestMessageFilter, options?.StorageInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState<TextSearchProviderState>(
_ => new TextSearchProviderState(),
@@ -86,7 +86,16 @@ public sealed class TextSearchProviderOptions
/// When <see langword="null"/>, the provider defaults to including only
/// <see cref="AgentRequestMessageSourceType.External"/> messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputMessageFilter { get; set; }
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputRequestMessageFilter { get; set; }
/// <summary>
/// Gets or sets an optional filter function applied to response messages when updating the recent message
/// memory during <see cref="AIContextProvider.InvokedAsync"/>.
/// </summary>
/// <value>
/// When <see langword="null"/>, the provider defaults to including all messages.
/// </value>
public Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? StorageInputResponseMessageFilter { get; set; }
/// <summary>
/// Gets or sets the list of <see cref="ChatRole"/> types to filter recent messages to
@@ -543,7 +543,9 @@ public class AIContextProviderTests
var storedRequest = provider.LastStoredContext!.RequestMessages.ToList();
Assert.Single(storedRequest);
Assert.Equal("External", storedRequest[0].Text);
Assert.Same(responseMessages, provider.LastStoredContext.ResponseMessages);
var storedResponse = provider.LastStoredContext.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Response", storedResponse[0].Text);
}
[Fact]
@@ -565,13 +567,14 @@ public class AIContextProviderTests
{
// Arrange - filter that only keeps System messages
var provider = new TestAIContextProvider(
storeInputMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.System));
storeInputRequestMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.System),
storeInputResponseMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.Assistant));
var messages = new[]
{
new ChatMessage(ChatRole.User, "User msg"),
new ChatMessage(ChatRole.System, "System msg")
};
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, messages, [new ChatMessage(ChatRole.Assistant, "Response")]);
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, messages, [new ChatMessage(ChatRole.Assistant, "Response"), new ChatMessage(ChatRole.Tool, "Response")]);
// Act
await provider.InvokedAsync(context);
@@ -581,6 +584,9 @@ public class AIContextProviderTests
var storedRequest = provider.LastStoredContext!.RequestMessages.ToList();
Assert.Single(storedRequest);
Assert.Equal("System msg", storedRequest[0].Text);
var storedResponse = provider.LastStoredContext.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Response", storedResponse[0].Text);
}
[Fact]
@@ -605,6 +611,87 @@ public class AIContextProviderTests
Assert.Equal("External", storedRequest[0].Text);
}
[Fact]
public async Task InvokedCoreAsync_DefaultResponseFilterPassesAllResponseMessagesAsync()
{
// Arrange
var provider = new TestAIContextProvider();
var requestMessages = new[] { new ChatMessage(ChatRole.User, "Request") };
var externalResponse = new ChatMessage(ChatRole.Assistant, "ExternalResp");
var historyResponse = new ChatMessage(ChatRole.Assistant, "HistoryResp")
.WithAgentRequestMessageSource(AgentRequestMessageSourceType.ChatHistory, "src");
var contextResponse = new ChatMessage(ChatRole.Assistant, "ContextResp")
.WithAgentRequestMessageSource(AgentRequestMessageSourceType.AIContextProvider, "src");
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, [externalResponse, historyResponse, contextResponse]);
// Act
await provider.InvokedAsync(context);
// Assert - default response filter is a noop, so all response messages are kept
Assert.NotNull(provider.LastStoredContext);
var storedResponse = provider.LastStoredContext!.ResponseMessages!.ToList();
Assert.Equal(3, storedResponse.Count);
Assert.Equal("ExternalResp", storedResponse[0].Text);
Assert.Equal("HistoryResp", storedResponse[1].Text);
Assert.Equal("ContextResp", storedResponse[2].Text);
}
[Fact]
public async Task InvokedCoreAsync_UsesCustomResponseFilterAsync()
{
// Arrange - response filter that only keeps Assistant messages with specific text
var provider = new TestAIContextProvider(
storeInputResponseMessageFilter: msgs => msgs.Where(m => m.Text == "Keep"));
var requestMessages = new[] { new ChatMessage(ChatRole.User, "Request") };
var responseMessages = new[]
{
new ChatMessage(ChatRole.Assistant, "Keep"),
new ChatMessage(ChatRole.Assistant, "Drop")
};
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, responseMessages);
// Act
await provider.InvokedAsync(context);
// Assert
Assert.NotNull(provider.LastStoredContext);
var storedResponse = provider.LastStoredContext!.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Keep", storedResponse[0].Text);
}
[Fact]
public async Task InvokedCoreAsync_RequestAndResponseFiltersOperateIndependentlyAsync()
{
// Arrange - different filters for request and response
var provider = new TestAIContextProvider(
storeInputRequestMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.System),
storeInputResponseMessageFilter: msgs => msgs.Where(m => m.Text == "Resp1"));
var requestMessages = new[]
{
new ChatMessage(ChatRole.User, "User"),
new ChatMessage(ChatRole.System, "System")
};
var responseMessages = new[]
{
new ChatMessage(ChatRole.Assistant, "Resp1"),
new ChatMessage(ChatRole.Assistant, "Resp2")
};
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, responseMessages);
// Act
await provider.InvokedAsync(context);
// Assert - request filter kept only System, response filter kept only Resp1
Assert.NotNull(provider.LastStoredContext);
var storedRequest = provider.LastStoredContext!.RequestMessages.ToList();
Assert.Single(storedRequest);
Assert.Equal("System", storedRequest[0].Text);
var storedResponse = provider.LastStoredContext!.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Resp1", storedResponse[0].Text);
}
#endregion
private sealed class TestAIContextProvider : AIContextProvider
@@ -620,8 +707,9 @@ public class AIContextProviderTests
AIContext? provideContext = null,
bool captureFilteredContext = false,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideInputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: base(provideInputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: base(provideInputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
this._provideContext = provideContext;
this._captureFilteredContext = captureFilteredContext;
@@ -439,7 +439,9 @@ public class ChatHistoryProviderTests
var storedRequest = provider.LastStoredContext!.RequestMessages.ToList();
Assert.Single(storedRequest);
Assert.Equal("External", storedRequest[0].Text);
Assert.Same(responseMessages, provider.LastStoredContext.ResponseMessages);
var storedResponse = provider.LastStoredContext.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Response", storedResponse[0].Text);
}
[Fact]
@@ -461,13 +463,14 @@ public class ChatHistoryProviderTests
{
// Arrange - filter that only keeps System messages
var provider = new TestChatHistoryProvider(
storeInputMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.System));
storeInputRequestMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.System),
storeInputResponseMessageFilter: msgs => msgs.Where(m => m.Role == ChatRole.Assistant));
var messages = new[]
{
new ChatMessage(ChatRole.User, "User msg"),
new ChatMessage(ChatRole.System, "System msg")
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages, [new ChatMessage(ChatRole.Assistant, "Response")]);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages, [new ChatMessage(ChatRole.Assistant, "Response"), new ChatMessage(ChatRole.Tool, "Response")]);
// Act
await provider.InvokedAsync(context);
@@ -477,6 +480,9 @@ public class ChatHistoryProviderTests
var storedRequest = provider.LastStoredContext!.RequestMessages.ToList();
Assert.Single(storedRequest);
Assert.Equal("System msg", storedRequest[0].Text);
var storedResponse = provider.LastStoredContext.ResponseMessages!.ToList();
Assert.Single(storedResponse);
Assert.Equal("Response", storedResponse[0].Text);
}
[Fact]
@@ -529,8 +535,9 @@ public class ChatHistoryProviderTests
public TestChatHistoryProvider(
IEnumerable<ChatMessage>? provideMessages = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? provideOutputMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputMessageFilter = null)
: base(provideOutputMessageFilter, storeInputMessageFilter)
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputRequestMessageFilter = null,
Func<IEnumerable<ChatMessage>, IEnumerable<ChatMessage>>? storeInputResponseMessageFilter = null)
: base(provideOutputMessageFilter, storeInputRequestMessageFilter, storeInputResponseMessageFilter)
{
this._provideMessages = provideMessages;
}
@@ -418,7 +418,7 @@ public class InMemoryChatHistoryProviderTests
var session = CreateMockSession();
var provider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions
{
StorageInputMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External)
StorageInputRequestMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External)
});
var requestMessages = new List<ChatMessage>
{
@@ -1004,7 +1004,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
s_testDatabaseId,
TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId),
storeInputMessageFilter: messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External));
storeInputRequestMessageFilter: messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External));
var requestMessages = new[]
{
@@ -530,7 +530,7 @@ public sealed class Mem0ProviderTests : IDisposable
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: new Mem0ProviderOptions
{
StorageInputMessageFilter = messages => messages // No filtering - store everything
StorageInputRequestMessageFilter = messages => messages // No filtering - store everything
});
var requestMessages = new List<ChatMessage>
@@ -119,8 +119,8 @@ public class ChatClientAgentOptionsTests
const string Description = "Test description";
var tools = new List<AITool> { AIFunctionFactory.Create(() => "test") };
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null).Object;
var mockAIContextProvider = new Mock<AIContextProvider>(null, null).Object;
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null).Object;
var mockAIContextProvider = new Mock<AIContextProvider>(null, null, null).Object;
var original = new ChatClientAgentOptions()
{
@@ -161,8 +161,8 @@ public class ChatClientAgentOptionsTests
public void Clone_WithoutProvidingChatOptions_ClonesCorrectly()
{
// Arrange
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null).Object;
var mockAIContextProvider = new Mock<AIContextProvider>(null, null).Object;
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null).Object;
var mockAIContextProvider = new Mock<AIContextProvider>(null, null, null).Object;
var original = new ChatClientAgentOptions
{
@@ -488,7 +488,7 @@ public partial class ChatClientAgentTests
})
.ReturnsAsync(new ChatResponse(responseMessages));
var mockProvider = new Mock<AIContextProvider>(null, null);
var mockProvider = new Mock<AIContextProvider>(null, null, null);
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -559,7 +559,7 @@ public partial class ChatClientAgentTests
It.IsAny<CancellationToken>()))
.Throws(new InvalidOperationException("downstream failure"));
var mockProvider = new Mock<AIContextProvider>(null, null);
var mockProvider = new Mock<AIContextProvider>(null, null, null);
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -617,7 +617,7 @@ public partial class ChatClientAgentTests
})
.ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]));
var mockProvider = new Mock<AIContextProvider>(null, null);
var mockProvider = new Mock<AIContextProvider>(null, null, null);
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -677,7 +677,7 @@ public partial class ChatClientAgentTests
.ReturnsAsync(new ChatResponse(responseMessages));
// Provider 1: adds a system message and a tool
var mockProvider1 = new Mock<AIContextProvider>(null, null);
var mockProvider1 = new Mock<AIContextProvider>(null, null, null);
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
@@ -696,7 +696,7 @@ public partial class ChatClientAgentTests
// Provider 2: adds another system message and verifies it receives accumulated context from provider 1
AIContext? provider2ReceivedContext = null;
var mockProvider2 = new Mock<AIContextProvider>(null, null);
var mockProvider2 = new Mock<AIContextProvider>(null, null, null);
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
@@ -784,7 +784,7 @@ public partial class ChatClientAgentTests
It.IsAny<CancellationToken>()))
.ThrowsAsync(new InvalidOperationException("downstream failure"));
var mockProvider1 = new Mock<AIContextProvider>(null, null);
var mockProvider1 = new Mock<AIContextProvider>(null, null, null);
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
@@ -801,7 +801,7 @@ public partial class ChatClientAgentTests
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
var mockProvider2 = new Mock<AIContextProvider>(null, null);
var mockProvider2 = new Mock<AIContextProvider>(null, null, null);
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
@@ -869,7 +869,7 @@ public partial class ChatClientAgentTests
})
.Returns(ToAsyncEnumerableAsync(responseUpdates));
var mockProvider1 = new Mock<AIContextProvider>(null, null);
var mockProvider1 = new Mock<AIContextProvider>(null, null, null);
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
@@ -886,7 +886,7 @@ public partial class ChatClientAgentTests
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
var mockProvider2 = new Mock<AIContextProvider>(null, null);
var mockProvider2 = new Mock<AIContextProvider>(null, null, null);
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
@@ -1828,7 +1828,7 @@ public partial class ChatClientAgentTests
})
.Returns(ToAsyncEnumerableAsync(responseUpdates));
var mockProvider = new Mock<AIContextProvider>(null, null);
var mockProvider = new Mock<AIContextProvider>(null, null, null);
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -1907,7 +1907,7 @@ public partial class ChatClientAgentTests
It.IsAny<CancellationToken>()))
.Throws(new InvalidOperationException("downstream failure"));
var mockProvider = new Mock<AIContextProvider>(null, null);
var mockProvider = new Mock<AIContextProvider>(null, null, null);
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -338,7 +338,7 @@ public class ChatClientAgent_BackgroundResponsesTests
List<ChatMessage> capturedMessages = [];
// Create a mock chat history provider that would normally provide messages
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null);
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null);
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
@@ -346,7 +346,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.ReturnsAsync([new(ChatRole.User, "Message from chat history provider")]);
// Create a mock AI context provider that would normally provide context
var mockContextProvider = new Mock<AIContextProvider>(null, null);
var mockContextProvider = new Mock<AIContextProvider>(null, null, null);
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
@@ -407,7 +407,7 @@ public class ChatClientAgent_BackgroundResponsesTests
List<ChatMessage> capturedMessages = [];
// Create a mock chat history provider that would normally provide messages
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null);
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null);
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
@@ -415,7 +415,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.ReturnsAsync([new(ChatRole.User, "Message from chat history provider")]);
// Create a mock AI context provider that would normally provide context
var mockContextProvider = new Mock<AIContextProvider>(null, null);
var mockContextProvider = new Mock<AIContextProvider>(null, null, null);
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
@@ -638,7 +638,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.Returns(ToAsyncEnumerableAsync(returnUpdates));
List<ChatMessage> capturedMessagesAddedToProvider = [];
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null);
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null);
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
@@ -647,7 +647,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.Returns(new ValueTask());
AIContextProvider.InvokedContext? capturedInvokedContext = null;
var mockContextProvider = new Mock<AIContextProvider>(null, null);
var mockContextProvider = new Mock<AIContextProvider>(null, null, null);
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
@@ -702,7 +702,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.Returns(ToAsyncEnumerableAsync(Array.Empty<ChatResponseUpdate>()));
List<ChatMessage> capturedMessagesAddedToProvider = [];
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null);
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>(null, null, null);
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
@@ -711,7 +711,7 @@ public class ChatClientAgent_BackgroundResponsesTests
.Returns(new ValueTask());
AIContextProvider.InvokedContext? capturedInvokedContext = null;
var mockContextProvider = new Mock<AIContextProvider>(null, null);
var mockContextProvider = new Mock<AIContextProvider>(null, null, null);
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
@@ -185,7 +185,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]));
Mock<ChatHistoryProvider> mockChatHistoryProvider = new(null, null);
Mock<ChatHistoryProvider> mockChatHistoryProvider = new(null, null, null);
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -240,7 +240,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).Throws(new InvalidOperationException("Test Error"));
Mock<ChatHistoryProvider> mockChatHistoryProvider = new(null, null);
Mock<ChatHistoryProvider> mockChatHistoryProvider = new(null, null, null);
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -429,7 +429,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
It.IsAny<CancellationToken>())).ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]));
// Arrange a chat history provider to override the factory provided one.
Mock<ChatHistoryProvider> mockOverrideChatHistoryProvider = new(null, null);
Mock<ChatHistoryProvider> mockOverrideChatHistoryProvider = new(null, null, null);
mockOverrideChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -442,7 +442,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
// Arrange a chat history provider to provide to the agent at construction time.
// This one shouldn't be used since it is being overridden.
Mock<ChatHistoryProvider> mockAgentOptionsChatHistoryProvider = new(null, null);
Mock<ChatHistoryProvider> mockAgentOptionsChatHistoryProvider = new(null, null, null);
mockAgentOptionsChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -467,7 +467,7 @@ public sealed class TextSearchProviderTests
{
RecentMessageMemoryLimit = 10,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.System],
StorageInputMessageFilter = messages => messages // No filtering - store everything
StorageInputRequestMessageFilter = messages => messages // No filtering - store everything
};
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
@@ -687,7 +687,7 @@ public class ChatHistoryMemoryProviderTests
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
options: new ChatHistoryMemoryProviderOptions
{
StorageInputMessageFilter = messages => messages // No filtering - store everything
StorageInputRequestMessageFilter = messages => messages // No filtering - store everything
});
var requestMessages = new List<ChatMessage>