// 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; } } }