// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Linq.Expressions;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Microsoft.Extensions.Logging;
using Microsoft.Extensions.VectorData;
using Microsoft.Shared.Diagnostics;
namespace Microsoft.Agents.AI;
///
/// A context provider that stores all chat history in a vector store and is able to
/// retrieve related chat history later to augment the current conversation.
///
///
///
/// This provider stores chat messages in a vector store and retrieves relevant previous messages
/// to provide as context during agent invocations. It uses the VectorStore and VectorStoreCollection
/// abstractions to work with any compatible vector store implementation.
///
///
/// Messages are stored during the method and retrieved during the
/// method using semantic similarity search.
///
///
/// Behavior is configurable through . When
/// is selected the provider
/// exposes a function tool that the model can invoke to retrieve relevant memories on demand instead of
/// injecting them automatically on each invocation.
///
///
public sealed class ChatHistoryMemoryProvider : MessageAIContextProvider, IDisposable
{
private const string DefaultContextPrompt = "## Memories\nConsider the following memories when answering user questions:";
private const int DefaultMaxResults = 3;
private const string DefaultFunctionToolName = "Search";
private const string DefaultFunctionToolDescription = "Allows searching for related previous chat history to help answer the user question.";
private const string KeyField = "Key";
private const string RoleField = "Role";
private const string MessageIdField = "MessageId";
private const string AuthorNameField = "AuthorName";
private const string ApplicationIdField = "ApplicationId";
private const string AgentIdField = "AgentId";
private const string UserIdField = "UserId";
private const string SessionIdField = "SessionId";
private const string ContentField = "Content";
private const string CreatedAtField = "CreatedAt";
private const string ContentEmbeddingField = "ContentEmbedding";
private readonly ProviderSessionState _sessionState;
private IReadOnlyList? _stateKeys;
#pragma warning disable CA2213 // VectorStore is not owned by this class - caller is responsible for disposal
private readonly VectorStore _vectorStore;
#pragma warning restore CA2213
private readonly VectorStoreCollection> _collection;
private readonly int _maxResults;
private readonly string _contextPrompt;
private readonly bool _enableSensitiveTelemetryData;
private readonly ChatHistoryMemoryProviderOptions.SearchBehavior _searchTime;
private readonly string _toolName;
private readonly string _toolDescription;
private readonly ILogger? _logger;
private bool _collectionInitialized;
private readonly SemaphoreSlim _initializationLock = new(1, 1);
private bool _disposedValue;
///
/// Initializes a new instance of the class.
///
/// The vector store to use for storing and retrieving chat history.
/// The name of the collection for storing chat history in the vector store.
/// The number of dimensions to use for the chat history vector store embeddings.
/// A delegate that initializes the provider state on the first invocation, providing the storage and search scopes.
/// Optional configuration options.
/// Optional logger factory.
/// Thrown when or is .
public ChatHistoryMemoryProvider(
VectorStore vectorStore,
string collectionName,
int vectorDimensions,
Func stateInitializer,
ChatHistoryMemoryProviderOptions? options = null,
ILoggerFactory? loggerFactory = null)
: base(options?.SearchInputMessageFilter, options?.StorageInputRequestMessageFilter, options?.StorageInputResponseMessageFilter)
{
this._sessionState = new ProviderSessionState(
Throw.IfNull(stateInitializer),
options?.StateKey ?? this.GetType().Name,
AgentJsonUtilities.DefaultOptions);
this._vectorStore = Throw.IfNull(vectorStore);
options ??= new ChatHistoryMemoryProviderOptions();
this._maxResults = options.MaxResults.HasValue ? Throw.IfLessThanOrEqual(options.MaxResults.Value, 0) : DefaultMaxResults;
this._contextPrompt = options.ContextPrompt ?? DefaultContextPrompt;
this._enableSensitiveTelemetryData = options.EnableSensitiveTelemetryData;
this._searchTime = options.SearchTime;
this._logger = loggerFactory?.CreateLogger();
this._toolName = options.FunctionToolName ?? DefaultFunctionToolName;
this._toolDescription = options.FunctionToolDescription ?? DefaultFunctionToolDescription;
// Create a definition so that we can use the dimensions provided at runtime.
var definition = new VectorStoreCollectionDefinition
{
Properties =
[
new VectorStoreKeyProperty(KeyField, typeof(Guid)),
new VectorStoreDataProperty(RoleField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(MessageIdField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(AuthorNameField, typeof(string)),
new VectorStoreDataProperty(ApplicationIdField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(AgentIdField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(UserIdField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(SessionIdField, typeof(string)) { IsIndexed = true },
new VectorStoreDataProperty(ContentField, typeof(string)) { IsFullTextIndexed = true },
new VectorStoreDataProperty(CreatedAtField, typeof(string)) { IsIndexed = true },
new VectorStoreVectorProperty(ContentEmbeddingField, typeof(string), Throw.IfLessThan(vectorDimensions, 1))
]
};
this._collection = this._vectorStore.GetDynamicCollection(Throw.IfNullOrWhitespace(collectionName), definition);
}
///
public override IReadOnlyList StateKeys => this._stateKeys ??= [this._sessionState.StateKey];
///
protected override async ValueTask ProvideAIContextAsync(AIContextProvider.InvokingContext context, CancellationToken cancellationToken = default)
{
_ = Throw.IfNull(context);
var state = this._sessionState.GetOrInitializeState(context.Session);
var searchScope = state.SearchScope;
if (this._searchTime == ChatHistoryMemoryProviderOptions.SearchBehavior.OnDemandFunctionCalling)
{
Task InlineSearchAsync(string userQuestion, CancellationToken ct)
=> this.SearchTextAsync(userQuestion, searchScope, ct);
// Create on-demand search tool (only used when behavior is OnDemandFunctionCalling)
AITool[] tools =
[
AIFunctionFactory.Create(
InlineSearchAsync,
name: this._toolName,
description: this._toolDescription)
];
// Expose search tool for on-demand invocation by the model
return new AIContext
{
Tools = tools
};
}
return new AIContext
{
Messages = await this.ProvideMessagesAsync(
new InvokingContext(context.Agent, context.Session, context.AIContext.Messages ?? []),
cancellationToken).ConfigureAwait(false)
};
}
///
protected override ValueTask> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
// This code path is invoked using InvokingAsync on MessageAIContextProvider, which does not support tools and instructions,
// and OnDemandFunctionCalling requires tools.
if (this._searchTime != ChatHistoryMemoryProviderOptions.SearchBehavior.BeforeAIInvoke)
{
throw new InvalidOperationException($"Using the {nameof(ChatHistoryMemoryProvider)} as a {nameof(MessageAIContextProvider)} is not supported when {nameof(ChatHistoryMemoryProviderOptions.SearchTime)} is set to {ChatHistoryMemoryProviderOptions.SearchBehavior.OnDemandFunctionCalling}.");
}
return base.InvokingCoreAsync(context, cancellationToken);
}
///
protected override async ValueTask> ProvideMessagesAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
_ = Throw.IfNull(context);
var state = this._sessionState.GetOrInitializeState(context.Session);
var searchScope = state.SearchScope;
try
{
// Get the text from the current request messages
var requestText = string.Join("\n",
(context.RequestMessages ?? [])
.Where(m => m != null && !string.IsNullOrWhiteSpace(m.Text))
.Select(m => m.Text));
if (string.IsNullOrWhiteSpace(requestText))
{
return [];
}
// Search for relevant chat history
var contextText = await this.SearchTextAsync(requestText, searchScope, cancellationToken).ConfigureAwait(false);
if (string.IsNullOrWhiteSpace(contextText))
{
return [];
}
return [new ChatMessage(ChatRole.User, contextText)];
}
catch (Exception ex)
{
if (this._logger?.IsEnabled(LogLevel.Error) is true)
{
this._logger.LogError(
ex,
"ChatHistoryMemoryProvider: Failed to search for chat history due to error. ApplicationId: '{ApplicationId}', AgentId: '{AgentId}', SessionId: '{SessionId}', UserId: '{UserId}'.",
searchScope.ApplicationId,
searchScope.AgentId,
searchScope.SessionId,
this.SanitizeLogData(searchScope.UserId));
}
return [];
}
}
///
protected override async ValueTask StoreAIContextAsync(InvokedContext context, CancellationToken cancellationToken = default)
{
_ = Throw.IfNull(context);
var state = this._sessionState.GetOrInitializeState(context.Session);
var storageScope = state.StorageScope;
try
{
// Ensure the collection is initialized
var collection = await this.EnsureCollectionExistsAsync(cancellationToken).ConfigureAwait(false);
List> itemsToStore = context.RequestMessages
.Concat(context.ResponseMessages ?? [])
.Select(message => new Dictionary
{
[KeyField] = Guid.NewGuid(),
[RoleField] = message.Role.ToString(),
[MessageIdField] = message.MessageId,
[AuthorNameField] = message.AuthorName,
[ApplicationIdField] = storageScope.ApplicationId,
[AgentIdField] = storageScope.AgentId,
[UserIdField] = storageScope.UserId,
[SessionIdField] = storageScope.SessionId,
[ContentField] = message.Text,
[CreatedAtField] = message.CreatedAt?.ToString("O") ?? DateTimeOffset.UtcNow.ToString("O"),
[ContentEmbeddingField] = message.Text,
})
.ToList();
if (itemsToStore.Count > 0)
{
await collection.UpsertAsync(itemsToStore, cancellationToken).ConfigureAwait(false);
}
}
catch (Exception ex)
{
if (this._logger?.IsEnabled(LogLevel.Error) is true)
{
this._logger.LogError(
ex,
"ChatHistoryMemoryProvider: Failed to add messages to chat history vector store due to error. ApplicationId: '{ApplicationId}', AgentId: '{AgentId}', SessionId: '{SessionId}', UserId: '{UserId}'.",
storageScope.ApplicationId,
storageScope.AgentId,
storageScope.SessionId,
this.SanitizeLogData(storageScope.UserId));
}
}
}
///
/// Function callable by the AI model (when enabled) to perform an ad-hoc chat history search.
///
/// The query text.
/// The scope to filter search results with.
/// Cancellation token.
/// Formatted search results (may be empty).
private async Task SearchTextAsync(string userQuestion, ChatHistoryMemoryProviderScope searchScope, CancellationToken cancellationToken = default)
{
if (string.IsNullOrWhiteSpace(userQuestion))
{
return string.Empty;
}
var results = await this.SearchChatHistoryAsync(userQuestion, searchScope, this._maxResults, cancellationToken).ConfigureAwait(false);
if (!results.Any())
{
return string.Empty;
}
// Format the results as a single context message
var outputResultsText = string.Join("\n", results.Select(x => (string?)x[ContentField]).Where(c => !string.IsNullOrWhiteSpace(c)));
if (string.IsNullOrWhiteSpace(outputResultsText))
{
return string.Empty;
}
var formatted = $"{this._contextPrompt}\n{outputResultsText}";
if (this._logger?.IsEnabled(LogLevel.Trace) is true)
{
this._logger.LogTrace(
"ChatHistoryMemoryProvider: Search Results\nInput:{Input}\nOutput:{MessageText}\n ApplicationId: '{ApplicationId}', AgentId: '{AgentId}', SessionId: '{SessionId}', UserId: '{UserId}'.",
this.SanitizeLogData(userQuestion),
this.SanitizeLogData(formatted),
searchScope.ApplicationId,
searchScope.AgentId,
searchScope.SessionId,
this.SanitizeLogData(searchScope.UserId));
}
return formatted;
}
///
/// Searches for relevant chat history items based on the provided query text.
///
/// The text to search for.
/// The scope to filter search results with.
/// The maximum number of results to return.
/// The cancellation token.
/// A list of relevant chat history items.
private async Task>> SearchChatHistoryAsync(
string queryText,
ChatHistoryMemoryProviderScope searchScope,
int top,
CancellationToken cancellationToken = default)
{
if (string.IsNullOrWhiteSpace(queryText))
{
return [];
}
var collection = await this.EnsureCollectionExistsAsync(cancellationToken).ConfigureAwait(false);
string? applicationId = searchScope.ApplicationId;
string? agentId = searchScope.AgentId;
string? userId = searchScope.UserId;
string? sessionId = searchScope.SessionId;
// Build a combined filter using a single shared parameter to avoid expression tree
// scoping issues when multiple filters are combined with AndAlso.
ParameterExpression parameter = Expression.Parameter(typeof(Dictionary), "x");
Expression? filterBody = null;
if (applicationId != null)
{
filterBody = RebindFilterBody(x => (string?)x[ApplicationIdField] == applicationId, parameter);
}
if (agentId != null)
{
Expression body = RebindFilterBody(x => (string?)x[AgentIdField] == agentId, parameter);
filterBody = filterBody == null ? body : Expression.AndAlso(filterBody, body);
}
if (userId != null)
{
Expression body = RebindFilterBody(x => (string?)x[UserIdField] == userId, parameter);
filterBody = filterBody == null ? body : Expression.AndAlso(filterBody, body);
}
if (sessionId != null)
{
Expression body = RebindFilterBody(x => (string?)x[SessionIdField] == sessionId, parameter);
filterBody = filterBody == null ? body : Expression.AndAlso(filterBody, body);
}
Expression, bool>>? filter = filterBody != null
? Expression.Lambda, bool>>(filterBody, parameter)
: null;
// Use search to find relevant messages
var searchResults = collection.SearchAsync(
queryText,
top,
options: new()
{
Filter = filter
},
cancellationToken: cancellationToken);
var results = new List>();
await foreach (var result in searchResults.WithCancellation(cancellationToken).ConfigureAwait(false))
{
results.Add(result.Record);
}
if (this._logger?.IsEnabled(LogLevel.Information) is true)
{
this._logger.LogInformation(
"ChatHistoryMemoryProvider: Retrieved {Count} search results. ApplicationId: '{ApplicationId}', AgentId: '{AgentId}', SessionId: '{SessionId}', UserId: '{UserId}'.",
results.Count,
searchScope.ApplicationId,
searchScope.AgentId,
searchScope.SessionId,
this.SanitizeLogData(searchScope.UserId));
}
return results;
}
///
/// Ensures the collection exists in the vector store, creating it if necessary.
///
/// The cancellation token.
/// The vector store collection.
private async Task>> EnsureCollectionExistsAsync(
CancellationToken cancellationToken = default)
{
if (this._collectionInitialized)
{
return this._collection;
}
await this._initializationLock.WaitAsync(cancellationToken).ConfigureAwait(false);
try
{
if (this._collectionInitialized)
{
return this._collection;
}
await this._collection.EnsureCollectionExistsAsync(cancellationToken).ConfigureAwait(false);
this._collectionInitialized = true;
return this._collection;
}
finally
{
this._initializationLock.Release();
}
}
///
private void Dispose(bool disposing)
{
if (!this._disposedValue)
{
if (disposing)
{
this._initializationLock.Dispose();
this._collection?.Dispose();
}
this._disposedValue = true;
}
}
///
public void Dispose()
{
// Do not change this code. Put cleanup code in 'Dispose(bool disposing)' method
this.Dispose(disposing: true);
GC.SuppressFinalize(this);
}
private string? SanitizeLogData(string? data) => this._enableSensitiveTelemetryData ? data : "";
///
/// Rebinds a filter expression's body to use the specified shared parameter,
/// replacing the original lambda parameter so that multiple filters can be safely
/// combined with .
///
private static Expression RebindFilterBody(
Expression, bool>> filter,
ParameterExpression sharedParameter)
{
return new ParameterReplacer(filter.Parameters[0], sharedParameter).Visit(filter.Body);
}
///
/// An that replaces one with another.
///
private sealed class ParameterReplacer(ParameterExpression original, ParameterExpression replacement) : ExpressionVisitor
{
protected override Expression VisitParameter(ParameterExpression node)
=> node == original ? replacement : base.VisitParameter(node);
}
///
/// Represents the state of a stored in the .
///
public sealed class State
{
///
/// Initializes a new instance of the class with the specified storage and search scopes.
///
/// The scope to use when storing chat history messages.
/// The scope to use when searching for relevant chat history messages. If null, the storage scope will be used for searching as well.
public State(ChatHistoryMemoryProviderScope storageScope, ChatHistoryMemoryProviderScope? searchScope = null)
{
this.StorageScope = Throw.IfNull(storageScope);
this.SearchScope = searchScope ?? storageScope;
}
///
/// Gets or sets the scope used when storing chat history messages.
///
public ChatHistoryMemoryProviderScope StorageScope { get; }
///
/// Gets or sets the scope used when searching chat history messages.
///
public ChatHistoryMemoryProviderScope SearchScope { get; }
}
}