Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -239,7 +239,7 @@ protected override async ValueTask<IEnumerable<ChatMessage>> 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)
{
Expand Down Expand Up @@ -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)
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,58 @@ public async Task InvokedAsync_DoesNotThrow_WhenUpsertThrowsAsync()
Times.Once);
}

[Fact]
public async Task InvokedAsync_WhenCallerCancels_PropagatesCancellationAsync()
{
// Arrange
using var cts = new CancellationTokenSource();
cts.Cancel();
Comment on lines +276 to +277

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
await Assert.ThrowsAnyAsync<OperationCanceledException>(
() => provider.InvokedAsync(invokedContext, cts.Token).AsTask());
}

[Fact]
public async Task InvokedAsync_WhenProviderCancelsWithoutCallerCancellation_DoesNotThrowAsync()
{
// Arrange
this._vectorStoreCollectionMock
.Setup(c => c.UpsertAsync(It.IsAny<IEnumerable<Dictionary<string, object?>>>(), It.IsAny<CancellationToken>()))
.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<EventId>(),
It.Is<It.IsAnyType>((v, t) => v.ToString()!.Contains("ChatHistoryMemoryProvider: Failed to add messages to chat history vector store due to error")),
It.IsAny<Exception?>(),
It.IsAny<Func<It.IsAnyType, Exception?, string>>()),
Times.Once);
}

[Theory]
[InlineData(false, false, false, 0)]
[InlineData(false, false, true, 0)]
Expand Down Expand Up @@ -793,6 +845,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<string>(),
It.IsAny<int>(),
It.IsAny<VectorSearchOptions<Dictionary<string, object?>>>(),
It.IsAny<CancellationToken>()))
.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<OperationCanceledException>(
() => provider.InvokingAsync(invokingContext, cts.Token).AsTask());
}

[Fact]
public async Task InvokingAsync_WhenProviderCancelsWithoutCallerCancellation_DoesNotThrowAsync()
{
// Arrange
this._vectorStoreCollectionMock
.Setup(c => c.SearchAsync(
It.IsAny<string>(),
It.IsAny<int>(),
It.IsAny<VectorSearchOptions<Dictionary<string, object?>>>(),
It.IsAny<CancellationToken>()))
.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
Expand Down
Loading