diff --git a/dotnet/src/Microsoft.Agents.AI/Memory/ChatHistoryMemoryProvider.cs b/dotnet/src/Microsoft.Agents.AI/Memory/ChatHistoryMemoryProvider.cs index f8fd7abaa9e..530b4bae727 100644 --- a/dotnet/src/Microsoft.Agents.AI/Memory/ChatHistoryMemoryProvider.cs +++ b/dotnet/src/Microsoft.Agents.AI/Memory/ChatHistoryMemoryProvider.cs @@ -239,7 +239,7 @@ protected override async ValueTask> ProvideMessagesAsyn return [new ChatMessage(ChatRole.User, contextText)]; } - catch (Exception ex) + catch (Exception ex) when (ex is not OperationCanceledException || !cancellationToken.IsCancellationRequested) { if (this._logger?.IsEnabled(LogLevel.Error) is true) { @@ -292,7 +292,7 @@ protected override async ValueTask StoreAIContextAsync(InvokedContext context, C await collection.UpsertAsync(itemsToStore, cancellationToken).ConfigureAwait(false); } } - catch (Exception ex) + catch (Exception ex) when (ex is not OperationCanceledException || !cancellationToken.IsCancellationRequested) { if (this._logger?.IsEnabled(LogLevel.Error) is true) { diff --git a/dotnet/tests/Microsoft.Agents.AI.UnitTests/Memory/ChatHistoryMemoryProviderTests.cs b/dotnet/tests/Microsoft.Agents.AI.UnitTests/Memory/ChatHistoryMemoryProviderTests.cs index 43cabebaed7..567594846b8 100644 --- a/dotnet/tests/Microsoft.Agents.AI.UnitTests/Memory/ChatHistoryMemoryProviderTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.UnitTests/Memory/ChatHistoryMemoryProviderTests.cs @@ -269,6 +269,66 @@ public async Task InvokedAsync_DoesNotThrow_WhenUpsertThrowsAsync() Times.Once); } + [Fact] + public async Task InvokedAsync_WhenCallerCancels_PropagatesCancellationAsync() + { + // Arrange + using var cts = new CancellationTokenSource(); + + this._vectorStoreCollectionMock + .Setup(c => c.UpsertAsync(It.IsAny>>(), cts.Token)) + .Callback(() => cts.Cancel()) + .ThrowsAsync(new OperationCanceledException(cts.Token)); + + var provider = new ChatHistoryMemoryProvider( + this._vectorStoreMock.Object, + TestCollectionName, + 1, + _ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }), + loggerFactory: this._loggerFactoryMock.Object); + var requestMsg = new ChatMessage(ChatRole.User, "request text"); + var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsg], []); + + // Act & Assert + var exception = await Assert.ThrowsAnyAsync( + () => provider.InvokedAsync(invokedContext, cts.Token).AsTask()); + Assert.Equal(cts.Token, exception.CancellationToken); + this._vectorStoreCollectionMock.Verify( + c => c.UpsertAsync(It.IsAny>>(), cts.Token), + Times.Once); + } + + [Fact] + public async Task InvokedAsync_WhenProviderCancelsWithoutCallerCancellation_DoesNotThrowAsync() + { + // Arrange + this._vectorStoreCollectionMock + .Setup(c => c.UpsertAsync(It.IsAny>>(), It.IsAny())) + .ThrowsAsync(new OperationCanceledException("Provider cancelled")); + + var provider = new ChatHistoryMemoryProvider( + this._vectorStoreMock.Object, + TestCollectionName, + 1, + _ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }), + loggerFactory: this._loggerFactoryMock.Object); + var requestMsg = new ChatMessage(ChatRole.User, "request text"); + var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsg], []); + + // Act + await provider.InvokedAsync(invokedContext, CancellationToken.None); + + // Assert + this._loggerMock.Verify( + l => l.Log( + LogLevel.Error, + It.IsAny(), + It.Is((v, t) => v.ToString()!.Contains("ChatHistoryMemoryProvider: Failed to add messages to chat history vector store due to error")), + It.IsAny(), + It.IsAny>()), + Times.Once); + } + [Theory] [InlineData(false, false, false, 0)] [InlineData(false, false, true, 0)] @@ -793,6 +853,79 @@ public async Task InvokedAsync_CustomStorageInputFilter_OverridesDefaultAsync() Assert.Equal("Response", stored[2]["Content"]); } + [Fact] + public async Task InvokingAsync_WhenCallerCancels_PropagatesCancellationAsync() + { + // Arrange + using var cts = new CancellationTokenSource(); + cts.Cancel(); + + this._vectorStoreCollectionMock + .Setup(c => c.SearchAsync( + It.IsAny(), + It.IsAny(), + It.IsAny>>(), + It.IsAny())) + .Throws(new OperationCanceledException(cts.Token)); + + var provider = new ChatHistoryMemoryProvider( + this._vectorStoreMock.Object, + TestCollectionName, + 1, + _ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }), + options: new ChatHistoryMemoryProviderOptions + { + SearchTime = ChatHistoryMemoryProviderOptions.SearchBehavior.BeforeAIInvoke + }, + loggerFactory: this._loggerFactoryMock.Object); + + var invokingContext = new AIContextProvider.InvokingContext( + s_mockAgent, + new TestAgentSession(), + new AIContext { Messages = [new ChatMessage(ChatRole.User, "What was discussed?")] }); + + // Act & Assert + await Assert.ThrowsAnyAsync( + () => provider.InvokingAsync(invokingContext, cts.Token).AsTask()); + } + + [Fact] + public async Task InvokingAsync_WhenProviderCancelsWithoutCallerCancellation_DoesNotThrowAsync() + { + // Arrange + this._vectorStoreCollectionMock + .Setup(c => c.SearchAsync( + It.IsAny(), + It.IsAny(), + It.IsAny>>(), + It.IsAny())) + .Throws(new OperationCanceledException("Provider cancelled")); + + var provider = new ChatHistoryMemoryProvider( + this._vectorStoreMock.Object, + TestCollectionName, + 1, + _ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }), + options: new ChatHistoryMemoryProviderOptions + { + SearchTime = ChatHistoryMemoryProviderOptions.SearchBehavior.BeforeAIInvoke + }, + loggerFactory: this._loggerFactoryMock.Object); + + var invokingContext = new AIContextProvider.InvokingContext( + s_mockAgent, + new TestAgentSession(), + new AIContext { Messages = [new ChatMessage(ChatRole.User, "What was discussed?")] }); + + // Act + var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None); + + // Assert + Assert.NotNull(aiContext.Messages); + Assert.Single(aiContext.Messages); + Assert.Equal("What was discussed?", aiContext.Messages.Single().Text); + } + #endregion #region MessageAIContextProvider.InvokingAsync Tests