merge with latest main

This commit is contained in:
SergeyMenshykh
2026-02-13 15:08:47 +00:00
Unverified
115 changed files with 5641 additions and 4360 deletions
@@ -37,14 +37,14 @@ public class AnthropicChatCompletionFixture : IChatClientAgentFixture
public async Task<List<ChatMessage>> GetChatHistoryAsync(AIAgent agent, AgentSession session)
{
var typedSession = (ChatClientAgentSession)session;
var chatHistoryProvider = agent.GetService<ChatHistoryProvider>();
if (typedSession.ChatHistoryProvider is null)
if (chatHistoryProvider is null)
{
return [];
}
return (await typedSession.ChatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
return (await chatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
}
public Task<ChatClientAgent> CreateChatClientAgentAsync(
@@ -48,12 +48,14 @@ public class AIProjectClientFixture : IChatClientAgentFixture
return await this.GetChatHistoryFromResponsesChainAsync(chatClientSession.ConversationId);
}
if (chatClientSession.ChatHistoryProvider is null)
var chatHistoryProvider = agent.GetService<ChatHistoryProvider>();
if (chatHistoryProvider is null)
{
return [];
}
return (await chatClientSession.ChatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
return (await chatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
}
private async Task<List<ChatMessage>> GetChatHistoryFromResponsesChainAsync(string conversationId)
@@ -21,10 +21,28 @@ public sealed class A2AAgentSessionTests
// Act
JsonElement serialized = originalSession.Serialize();
A2AAgentSession deserializedSession = new(serialized);
A2AAgentSession deserializedSession = A2AAgentSession.Deserialize(serialized);
// Assert
Assert.Equal(originalSession.ContextId, deserializedSession.ContextId);
Assert.Equal(originalSession.TaskId, deserializedSession.TaskId);
}
[Fact]
public void Constructor_RoundTrip_SerializationPreservesStateBag()
{
// Arrange
A2AAgentSession originalSession = new() { ContextId = "ctx-1", TaskId = "task-1" };
originalSession.StateBag.SetValue("testKey", "testValue");
// Act
JsonElement serialized = originalSession.Serialize();
A2AAgentSession deserializedSession = A2AAgentSession.Deserialize(serialized);
// Assert
Assert.Equal("ctx-1", deserializedSession.ContextId);
Assert.Equal("task-1", deserializedSession.TaskId);
Assert.True(deserializedSession.StateBag.TryGetValue<string>("testKey", out var value));
Assert.Equal("testValue", value);
}
}
@@ -3,7 +3,6 @@
using System;
using System.Collections.Generic;
using System.Collections.ObjectModel;
using System.Linq;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
@@ -16,110 +15,6 @@ public class AIContextProviderTests
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
#region InvokingAsync Message Stamping Tests
[Fact]
public async Task InvokingAsync_StampsMessagesWithSourceTypeAndSourceIdAsync()
{
// Arrange
var provider = new TestAIContextProviderWithMessages();
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
AIContext aiContext = await provider.InvokingAsync(context);
// Assert
Assert.NotNull(aiContext.Messages);
ChatMessage message = aiContext.Messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, typedAttribution.SourceType);
Assert.Equal(typeof(TestAIContextProviderWithMessages).FullName, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_WithCustomSourceId_StampsMessagesWithCustomSourceIdAsync()
{
// Arrange
const string CustomSourceId = "CustomContextSource";
var provider = new TestAIContextProviderWithCustomSource(CustomSourceId);
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
AIContext aiContext = await provider.InvokingAsync(context);
// Assert
Assert.NotNull(aiContext.Messages);
ChatMessage message = aiContext.Messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, typedAttribution.SourceType);
Assert.Equal(CustomSourceId, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_DoesNotReStampAlreadyStampedMessagesAsync()
{
// Arrange
var provider = new TestAIContextProviderWithPreStampedMessages();
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
AIContext aiContext = await provider.InvokingAsync(context);
// Assert
Assert.NotNull(aiContext.Messages);
ChatMessage message = aiContext.Messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, typedAttribution.SourceType);
Assert.Equal(typeof(TestAIContextProviderWithPreStampedMessages).FullName, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_StampsMultipleMessagesAsync()
{
// Arrange
var provider = new TestAIContextProviderWithMultipleMessages();
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
AIContext aiContext = await provider.InvokingAsync(context);
// Assert
Assert.NotNull(aiContext.Messages);
List<ChatMessage> messageList = aiContext.Messages.ToList();
Assert.Equal(3, messageList.Count);
foreach (ChatMessage message in messageList)
{
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, typedAttribution.SourceType);
Assert.Equal(typeof(TestAIContextProviderWithMultipleMessages).FullName, typedAttribution.SourceId);
}
}
[Fact]
public async Task InvokingAsync_WithNullMessages_ReturnsContextWithoutStampingAsync()
{
// Arrange
var provider = new TestAIContextProvider();
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
AIContext aiContext = await provider.InvokingAsync(context);
// Assert
Assert.Null(aiContext.Messages);
}
#endregion
#region Basic Tests
[Fact]
@@ -130,25 +25,12 @@ public class AIContextProviderTests
var messages = new ReadOnlyCollection<ChatMessage>([]);
// Act
ValueTask task = provider.InvokedAsync(new(s_mockAgent, s_mockSession, messages));
ValueTask task = provider.InvokedAsync(new(s_mockAgent, s_mockSession, messages, []));
// Assert
Assert.Equal(default, task);
}
[Fact]
public void Serialize_ReturnsEmptyElement()
{
// Arrange
var provider = new TestAIContextProvider();
// Act
var actual = provider.Serialize();
// Assert
Assert.Equal(default, actual);
}
[Fact]
public void InvokingContext_Constructor_ThrowsForNullMessages()
{
@@ -160,7 +42,7 @@ public class AIContextProviderTests
public void InvokedContext_Constructor_ThrowsForNullMessages()
{
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, null!));
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, null!, []));
}
#endregion
@@ -284,39 +166,33 @@ public class AIContextProviderTests
#region InvokingContext Tests
[Fact]
public void InvokingContext_RequestMessages_SetterThrowsForNull()
public void InvokingContext_Constructor_ThrowsForNullAIContext()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, messages);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => context.RequestMessages = null!);
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, null!));
}
[Fact]
public void InvokingContext_RequestMessages_SetterRoundtrips()
public void InvokingContext_AIContext_ConstructorValueRoundtrips()
{
// Arrange
var initialMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var newMessages = new List<ChatMessage> { new(ChatRole.User, "New message") };
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, initialMessages);
var aiContext = new AIContext { Messages = [new ChatMessage(ChatRole.User, "Hello")] };
// Act
context.RequestMessages = newMessages;
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, aiContext);
// Assert
Assert.Same(newMessages, context.RequestMessages);
Assert.Same(aiContext, context.AIContext);
}
[Fact]
public void InvokingContext_Agent_ReturnsConstructorValue()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var aiContext = new AIContext { Messages = [new ChatMessage(ChatRole.User, "Hello")] };
// Act
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, messages);
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, aiContext);
// Assert
Assert.Same(s_mockAgent, context.Agent);
@@ -326,10 +202,10 @@ public class AIContextProviderTests
public void InvokingContext_Session_ReturnsConstructorValue()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var aiContext = new AIContext { Messages = [new ChatMessage(ChatRole.User, "Hello")] };
// Act
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, messages);
var context = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, aiContext);
// Assert
Assert.Same(s_mockSession, context.Session);
@@ -339,10 +215,10 @@ public class AIContextProviderTests
public void InvokingContext_Session_CanBeNull()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var aiContext = new AIContext { Messages = [new ChatMessage(ChatRole.User, "Hello")] };
// Act
var context = new AIContextProvider.InvokingContext(s_mockAgent, null, messages);
var context = new AIContextProvider.InvokingContext(s_mockAgent, null, aiContext);
// Assert
Assert.Null(context.Session);
@@ -352,52 +228,25 @@ public class AIContextProviderTests
public void InvokingContext_Constructor_ThrowsForNullAgent()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var aiContext = new AIContext { Messages = [new ChatMessage(ChatRole.User, "Hello")] };
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokingContext(null!, s_mockSession, messages));
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokingContext(null!, s_mockSession, aiContext));
}
#endregion
#region InvokedContext Tests
[Fact]
public void InvokedContext_RequestMessages_SetterThrowsForNull()
{
// Arrange
var messages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => context.RequestMessages = null!);
}
[Fact]
public void InvokedContext_RequestMessages_SetterRoundtrips()
{
// Arrange
var initialMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var newMessages = new List<ChatMessage> { new(ChatRole.User, "New message") };
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, initialMessages);
// Act
context.RequestMessages = newMessages;
// Assert
Assert.Same(newMessages, context.RequestMessages);
}
[Fact]
public void InvokedContext_ResponseMessages_Roundtrips()
{
// Arrange
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var responseMessages = new List<ChatMessage> { new(ChatRole.Assistant, "Response message") };
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
// Act
context.ResponseMessages = responseMessages;
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, responseMessages);
// Assert
Assert.Same(responseMessages, context.ResponseMessages);
@@ -409,10 +258,9 @@ public class AIContextProviderTests
// Arrange
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
var exception = new InvalidOperationException("Test exception");
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
// Act
context.InvokeException = exception;
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, exception);
// Assert
Assert.Same(exception, context.InvokeException);
@@ -425,7 +273,7 @@ public class AIContextProviderTests
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, []);
// Assert
Assert.Same(s_mockAgent, context.Agent);
@@ -438,7 +286,7 @@ public class AIContextProviderTests
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
var context = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, []);
// Assert
Assert.Same(s_mockSession, context.Session);
@@ -451,7 +299,7 @@ public class AIContextProviderTests
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act
var context = new AIContextProvider.InvokedContext(s_mockAgent, null, requestMessages);
var context = new AIContextProvider.InvokedContext(s_mockAgent, null, requestMessages, []);
// Assert
Assert.Null(context.Session);
@@ -464,7 +312,27 @@ public class AIContextProviderTests
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(null!, s_mockSession, requestMessages));
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(null!, s_mockSession, requestMessages, []));
}
[Fact]
public void InvokedContext_SuccessConstructor_ThrowsForNullResponseMessages()
{
// Arrange
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, (IEnumerable<ChatMessage>)null!));
}
[Fact]
public void InvokedContext_FailureConstructor_ThrowsForNullException()
{
// Arrange
var requestMessages = new ReadOnlyCollection<ChatMessage>([new(ChatRole.User, "Hello")]);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, (Exception)null!));
}
#endregion
@@ -474,55 +342,4 @@ public class AIContextProviderTests
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(new AIContext());
}
private sealed class TestAIContextProviderWithMessages : AIContextProvider
{
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(new AIContext
{
Messages = [new ChatMessage(ChatRole.System, "Context Message")]
});
}
private sealed class TestAIContextProviderWithCustomSource : AIContextProvider
{
public TestAIContextProviderWithCustomSource(string sourceId) : base(sourceId)
{
}
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(new AIContext
{
Messages = [new ChatMessage(ChatRole.System, "Context Message")]
});
}
private sealed class TestAIContextProviderWithPreStampedMessages : AIContextProvider
{
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
var message = new ChatMessage(ChatRole.System, "Pre-stamped Message");
message.AdditionalProperties = new AdditionalPropertiesDictionary
{
[AgentRequestMessageSourceAttribution.AdditionalPropertiesKey] = new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, this.GetType().FullName!)
};
return new(new AIContext
{
Messages = [message]
});
}
}
private sealed class TestAIContextProviderWithMultipleMessages : AIContextProvider
{
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(new AIContext
{
Messages = [
new ChatMessage(ChatRole.System, "Message 1"),
new ChatMessage(ChatRole.User, "Message 2"),
new ChatMessage(ChatRole.Assistant, "Message 3")
]
});
}
}
@@ -1,5 +1,6 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Linq;
using Microsoft.Extensions.AI;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
@@ -33,9 +34,10 @@ public class AIContextTests
};
Assert.NotNull(context.Messages);
Assert.Equal(2, context.Messages.Count);
Assert.Equal("Hello", context.Messages[0].Text);
Assert.Equal("Hi there!", context.Messages[1].Text);
var messages = context.Messages.ToList();
Assert.Equal(2, messages.Count);
Assert.Equal("Hello", messages[0].Text);
Assert.Equal("Hi there!", messages[1].Text);
}
[Fact]
@@ -51,8 +53,9 @@ public class AIContextTests
};
Assert.NotNull(context.Tools);
Assert.Equal(2, context.Tools.Count);
Assert.Equal("Function1", context.Tools[0].Name);
Assert.Equal("Function2", context.Tools[1].Name);
var tools = context.Tools.ToList();
Assert.Equal(2, tools.Count);
Assert.Equal("Function1", tools[0].Name);
Assert.Equal("Function2", tools[1].Name);
}
}
@@ -390,6 +390,49 @@ public sealed class AgentRequestMessageSourceAttributionTests
#endregion
#region ToString Tests
[Fact]
public void ToString_WithSourceId_ReturnsTypeColonId()
{
// Arrange
AgentRequestMessageSourceAttribution attribution = new(AgentRequestMessageSourceType.AIContextProvider, "MyProvider");
// Act
string result = attribution.ToString();
// Assert
Assert.Equal("AIContextProvider:MyProvider", result);
}
[Fact]
public void ToString_WithNullSourceId_ReturnsTypeOnly()
{
// Arrange
AgentRequestMessageSourceAttribution attribution = new(AgentRequestMessageSourceType.ChatHistory, null);
// Act
string result = attribution.ToString();
// Assert
Assert.Equal("ChatHistory", result);
}
[Fact]
public void ToString_Default_ReturnsExternalOnly()
{
// Arrange
AgentRequestMessageSourceAttribution attribution = default;
// Act
string result = attribution.ToString();
// Assert
Assert.Equal("External", result);
}
#endregion
#region Inequality Operator Tests
[Fact]
@@ -414,6 +414,46 @@ public sealed class AgentRequestMessageSourceTypeTests
#endregion
#region ToString Tests
[Fact]
public void ToString_ReturnsValue()
{
// Arrange
AgentRequestMessageSourceType source = new("CustomSource");
// Act
string result = source.ToString();
// Assert
Assert.Equal("CustomSource", result);
}
[Fact]
public void ToString_StaticExternal_ReturnsExternal()
{
// Arrange & Act
string result = AgentRequestMessageSourceType.External.ToString();
// Assert
Assert.Equal("External", result);
}
[Fact]
public void ToString_Default_ReturnsExternal()
{
// Arrange
AgentRequestMessageSourceType source = default;
// Act
string result = source.ToString();
// Assert
Assert.Equal("External", result);
}
#endregion
#region IEquatable Tests
[Fact]
@@ -0,0 +1,840 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Text.Json;
using Microsoft.Agents.AI.Abstractions.UnitTests.Models;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// <summary>
/// Contains tests for the <see cref="AgentSessionStateBag"/> class.
/// </summary>
public sealed class AgentSessionStateBagTests
{
#region Constructor Tests
[Fact]
public void Constructor_Default_CreatesEmptyStateBag()
{
// Act
var stateBag = new AgentSessionStateBag();
// Assert
Assert.False(stateBag.TryGetValue<string>("nonexistent", out _));
}
#endregion
#region SetValue Tests
[Fact]
public void SetValue_WithValidKeyAndValue_StoresValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
stateBag.SetValue("key1", "value1");
// Assert
Assert.True(stateBag.TryGetValue<string>("key1", out var result));
Assert.Equal("value1", result);
}
[Fact]
public void SetValue_WithNullKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentNullException>(() => stateBag.SetValue(null!, "value"));
}
[Fact]
public void SetValue_WithEmptyKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.SetValue("", "value"));
}
[Fact]
public void SetValue_WithWhitespaceKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.SetValue(" ", "value"));
}
[Fact]
public void SetValue_OverwritesExistingValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "originalValue");
// Act
stateBag.SetValue("key1", "newValue");
// Assert
Assert.Equal("newValue", stateBag.GetValue<string>("key1"));
}
#endregion
#region GetValue Tests
[Fact]
public void GetValue_WithExistingKey_ReturnsValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
// Act
var result = stateBag.GetValue<string>("key1");
// Assert
Assert.Equal("value1", result);
}
[Fact]
public void GetValue_WithNonexistentKey_ReturnsNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
var result = stateBag.GetValue<string>("nonexistent");
// Assert
Assert.Null(result);
}
[Fact]
public void GetValue_WithNullKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentNullException>(() => stateBag.GetValue<string>(null!));
}
[Fact]
public void GetValue_WithEmptyKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.GetValue<string>(""));
}
[Fact]
public void GetValue_CachesDeserializedValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
// Act
var result1 = stateBag.GetValue<string>("key1");
var result2 = stateBag.GetValue<string>("key1");
// Assert
Assert.Same(result1, result2);
}
#endregion
#region TryGetValue Tests
[Fact]
public void TryGetValue_WithExistingKey_ReturnsTrueAndValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
// Act
var found = stateBag.TryGetValue<string>("key1", out var result);
// Assert
Assert.True(found);
Assert.Equal("value1", result);
}
[Fact]
public void TryGetValue_WithNonexistentKey_ReturnsFalseAndNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
var found = stateBag.TryGetValue<string>("nonexistent", out var result);
// Assert
Assert.False(found);
Assert.Null(result);
}
[Fact]
public void TryGetValue_WithNullKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentNullException>(() => stateBag.TryGetValue<string>(null!, out _));
}
[Fact]
public void TryGetValue_WithEmptyKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.TryGetValue<string>("", out _));
}
#endregion
#region Null Value Tests
[Fact]
public void SetValue_WithNullValue_StoresNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
stateBag.SetValue<string>("key1", null);
// Assert
Assert.Equal(1, stateBag.Count);
}
[Fact]
public void TryGetValue_WithNullValue_ReturnsTrueAndNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue<string>("key1", null);
// Act
var found = stateBag.TryGetValue<string>("key1", out var result);
// Assert
Assert.True(found);
Assert.Null(result);
}
[Fact]
public void GetValue_WithNullValue_ReturnsNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue<string>("key1", null);
// Act
var result = stateBag.GetValue<string>("key1");
// Assert
Assert.Null(result);
}
[Fact]
public void SetValue_OverwriteWithNull_ReturnsNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
// Act
stateBag.SetValue<string>("key1", null);
// Assert
Assert.True(stateBag.TryGetValue<string>("key1", out var result));
Assert.Null(result);
}
[Fact]
public void SetValue_OverwriteNullWithValue_ReturnsValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue<string>("key1", null);
// Act
stateBag.SetValue("key1", "newValue");
// Assert
Assert.True(stateBag.TryGetValue<string>("key1", out var result));
Assert.Equal("newValue", result);
}
[Fact]
public void SerializeDeserialize_WithNullValue_SerializesAsNull()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue<string>("nullKey", null);
// Act
var json = stateBag.Serialize();
// Assert - null values are serialized as JSON null
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("nullKey", out var nullElement));
Assert.Equal(JsonValueKind.Null, nullElement.ValueKind);
}
#endregion
#region TryRemoveValue Tests
[Fact]
public void TryRemoveValue_ExistingKey_ReturnsTrueAndRemoves()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
// Act
var removed = stateBag.TryRemoveValue("key1");
// Assert
Assert.True(removed);
Assert.Equal(0, stateBag.Count);
Assert.False(stateBag.TryGetValue<string>("key1", out _));
}
[Fact]
public void TryRemoveValue_NonexistentKey_ReturnsFalse()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
var removed = stateBag.TryRemoveValue("nonexistent");
// Assert
Assert.False(removed);
}
[Fact]
public void TryRemoveValue_WithNullKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentNullException>(() => stateBag.TryRemoveValue(null!));
}
[Fact]
public void TryRemoveValue_WithEmptyKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.TryRemoveValue(""));
}
[Fact]
public void TryRemoveValue_WithWhitespaceKey_ThrowsArgumentException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act & Assert
Assert.Throws<ArgumentException>(() => stateBag.TryRemoveValue(" "));
}
[Fact]
public void TryRemoveValue_DoesNotAffectOtherKeys()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "value1");
stateBag.SetValue("key2", "value2");
// Act
stateBag.TryRemoveValue("key1");
// Assert
Assert.Equal(1, stateBag.Count);
Assert.False(stateBag.TryGetValue<string>("key1", out _));
Assert.True(stateBag.TryGetValue<string>("key2", out var value));
Assert.Equal("value2", value);
}
[Fact]
public void TryRemoveValue_ThenSetValue_Works()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "original");
// Act
stateBag.TryRemoveValue("key1");
stateBag.SetValue("key1", "replacement");
// Assert
Assert.True(stateBag.TryGetValue<string>("key1", out var result));
Assert.Equal("replacement", result);
}
#endregion
#region Serialize/Deserialize Tests
[Fact]
public void Serialize_EmptyStateBag_ReturnsEmptyObject()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
var json = stateBag.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
}
[Fact]
public void Serialize_WithStringValue_ReturnsJsonWithValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("stringKey", "stringValue");
// Act
var json = stateBag.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("stringKey", out _));
}
[Fact]
public void Deserialize_FromJsonDocument_ReturnsEmptyStateBag()
{
// Arrange
var emptyJson = JsonDocument.Parse("{}").RootElement;
// Act
var stateBag = AgentSessionStateBag.Deserialize(emptyJson);
// Assert
Assert.False(stateBag.TryGetValue<string>("nonexistent", out _));
}
[Fact]
public void Deserialize_NullElement_ReturnsEmptyStateBag()
{
// Arrange
var nullJson = default(JsonElement);
// Act
var stateBag = AgentSessionStateBag.Deserialize(nullJson);
// Assert
Assert.False(stateBag.TryGetValue<string>("nonexistent", out _));
}
[Fact]
public void SerializeDeserialize_WithStringValue_Roundtrips()
{
// Arrange
var originalStateBag = new AgentSessionStateBag();
originalStateBag.SetValue("stringKey", "stringValue");
// Act
var json = originalStateBag.Serialize();
var restoredStateBag = AgentSessionStateBag.Deserialize(json);
// Assert
Assert.Equal("stringValue", restoredStateBag.GetValue<string>("stringKey"));
}
#endregion
#region Thread Safety Tests
[Fact]
public async System.Threading.Tasks.Task SetValue_MultipleConcurrentWrites_DoesNotThrowAsync()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var tasks = new System.Threading.Tasks.Task[100];
// Act
for (int i = 0; i < 100; i++)
{
int index = i;
tasks[i] = System.Threading.Tasks.Task.Run(() => stateBag.SetValue($"key{index}", $"value{index}"));
}
await System.Threading.Tasks.Task.WhenAll(tasks);
// Assert
for (int i = 0; i < 100; i++)
{
Assert.True(stateBag.TryGetValue<string>($"key{i}", out var value));
Assert.Equal($"value{i}", value);
}
}
[Fact]
public async System.Threading.Tasks.Task ConcurrentWritesAndSerialize_DoesNotThrowAsync()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("shared", "initial");
var tasks = new System.Threading.Tasks.Task[100];
// Act - concurrently write and serialize the same key
for (int i = 0; i < 100; i++)
{
int index = i;
tasks[i] = System.Threading.Tasks.Task.Run(() =>
{
stateBag.SetValue("shared", $"value{index}");
_ = stateBag.Serialize();
});
}
await System.Threading.Tasks.Task.WhenAll(tasks);
// Assert - should have some value and serialize without error
Assert.True(stateBag.TryGetValue<string>("shared", out var result));
Assert.NotNull(result);
var json = stateBag.Serialize();
Assert.Equal(JsonValueKind.Object, json.ValueKind);
}
[Fact]
public async System.Threading.Tasks.Task ConcurrentReadsAndWrites_DoesNotThrowAsync()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key", "initial");
var tasks = new System.Threading.Tasks.Task[200];
// Act - half readers, half writers on the same key
for (int i = 0; i < 200; i++)
{
int index = i;
tasks[i] = (index % 2 == 0)
? System.Threading.Tasks.Task.Run(() => stateBag.GetValue<string>("key"))
: System.Threading.Tasks.Task.Run(() => stateBag.SetValue("key", $"value{index}"));
}
await System.Threading.Tasks.Task.WhenAll(tasks);
// Assert - should have a consistent value
Assert.True(stateBag.TryGetValue<string>("key", out var result));
Assert.NotNull(result);
}
#endregion
#region Complex Object Tests
[Fact]
public void SetValue_WithComplexObject_StoresValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 1, FullName = "Buddy", Species = Species.Bear };
// Act
stateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Assert
Animal? result = stateBag.GetValue<Animal>("animal", TestJsonSerializerContext.Default.Options);
Assert.NotNull(result);
Assert.Equal(1, result.Id);
Assert.Equal("Buddy", result.FullName);
Assert.Equal(Species.Bear, result.Species);
}
[Fact]
public void GetValue_WithComplexObject_CachesDeserializedValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 2, FullName = "Whiskers", Species = Species.Tiger };
stateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Act
Animal? result1 = stateBag.GetValue<Animal>("animal", TestJsonSerializerContext.Default.Options);
Animal? result2 = stateBag.GetValue<Animal>("animal", TestJsonSerializerContext.Default.Options);
// Assert
Assert.Same(result1, result2);
}
[Fact]
public void TryGetValue_WithComplexObject_ReturnsTrueAndValue()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 3, FullName = "Goldie", Species = Species.Walrus };
stateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Act
bool found = stateBag.TryGetValue("animal", out Animal? result, TestJsonSerializerContext.Default.Options);
// Assert
Assert.True(found);
Assert.NotNull(result);
Assert.Equal(3, result.Id);
Assert.Equal("Goldie", result.FullName);
Assert.Equal(Species.Walrus, result.Species);
}
[Fact]
public void SerializeDeserialize_WithComplexObject_Roundtrips()
{
// Arrange
var originalStateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 4, FullName = "Polly", Species = Species.Bear };
originalStateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Act
JsonElement json = originalStateBag.Serialize();
AgentSessionStateBag restoredStateBag = AgentSessionStateBag.Deserialize(json);
// Assert
Animal? restoredAnimal = restoredStateBag.GetValue<Animal>("animal", TestJsonSerializerContext.Default.Options);
Assert.NotNull(restoredAnimal);
Assert.Equal(4, restoredAnimal.Id);
Assert.Equal("Polly", restoredAnimal.FullName);
Assert.Equal(Species.Bear, restoredAnimal.Species);
}
[Fact]
public void Serialize_WithComplexObject_ReturnsJsonWithProperties()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 7, FullName = "Spot", Species = Species.Walrus };
stateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Act
JsonElement json = stateBag.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("animal", out JsonElement animalElement));
Assert.Equal(JsonValueKind.Object, animalElement.ValueKind);
Assert.Equal(7, animalElement.GetProperty("id").GetInt32());
Assert.Equal("Spot", animalElement.GetProperty("fullName").GetString());
Assert.Equal("Walrus", animalElement.GetProperty("species").GetString());
}
#endregion
#region Type Mismatch Tests
[Fact]
public void TryGetValue_WithDifferentTypeAfterSet_ReturnsFalse()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "hello");
// Act
var found = stateBag.TryGetValue<Animal>("key1", out var result, TestJsonSerializerContext.Default.Options);
// Assert
Assert.False(found);
Assert.Null(result);
}
[Fact]
public void GetValue_WithDifferentTypeAfterSet_ThrowsInvalidOperationException()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "hello");
// Act & Assert
Assert.Throws<InvalidOperationException>(() => stateBag.GetValue<Animal>("key1", TestJsonSerializerContext.Default.Options));
}
[Fact]
public void TryGetValue_WithDifferentTypeAfterDeserializedRead_ReturnsFalse()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "hello");
// First read caches the value as string
var cachedValue = stateBag.GetValue<string>("key1");
Assert.Equal("hello", cachedValue);
// Act - request as a different type
var found = stateBag.TryGetValue<Animal>("key1", out var result, TestJsonSerializerContext.Default.Options);
// Assert
Assert.False(found);
Assert.Null(result);
}
[Fact]
public void GetValue_WithDifferentTypeAfterDeserializedRoundtrip_ThrowsInvalidOperationException()
{
// Arrange
var originalStateBag = new AgentSessionStateBag();
originalStateBag.SetValue("key1", "hello");
// Round-trip through serialization
var json = originalStateBag.Serialize();
var restoredStateBag = AgentSessionStateBag.Deserialize(json);
// First read caches the value as string
var cachedValue = restoredStateBag.GetValue<string>("key1");
Assert.Equal("hello", cachedValue);
// Act & Assert - request as a different type
Assert.Throws<InvalidOperationException>(() => restoredStateBag.GetValue<Animal>("key1", TestJsonSerializerContext.Default.Options));
}
[Fact]
public void TryGetValue_ComplexTypeAfterSetString_ReturnsFalse()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("animal", "not an animal");
// Act
var found = stateBag.TryGetValue<Animal>("animal", out var result, TestJsonSerializerContext.Default.Options);
// Assert
Assert.False(found);
Assert.Null(result);
}
[Fact]
public void GetValue_TypeMismatch_ExceptionMessageContainsBothTypeNames()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key1", "hello");
// Act
var exception = Assert.Throws<InvalidOperationException>(() => stateBag.GetValue<Animal>("key1", TestJsonSerializerContext.Default.Options));
// Assert
Assert.Contains(typeof(string).FullName!, exception.Message);
Assert.Contains(typeof(Animal).FullName!, exception.Message);
}
#endregion
#region JsonSerializer Integration Tests
[Fact]
public void JsonSerializerSerialize_EmptyStateBag_ReturnsEmptyObject()
{
// Arrange
var stateBag = new AgentSessionStateBag();
// Act
var json = JsonSerializer.Serialize(stateBag, AgentAbstractionsJsonUtilities.DefaultOptions);
// Assert
Assert.Equal("{}", json);
}
[Fact]
public void JsonSerializerSerialize_WithStringValue_ProducesSameOutputAsSerializeMethod()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("stringKey", "stringValue");
// Act
var jsonFromSerializer = JsonSerializer.Serialize(stateBag, AgentAbstractionsJsonUtilities.DefaultOptions);
var jsonFromMethod = stateBag.Serialize().GetRawText();
// Assert
Assert.Equal(jsonFromMethod, jsonFromSerializer);
}
[Fact]
public void JsonSerializerRoundtrip_WithStringValue_PreservesData()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("greeting", "hello world");
// Act
var json = JsonSerializer.Serialize(stateBag, AgentAbstractionsJsonUtilities.DefaultOptions);
var restored = JsonSerializer.Deserialize<AgentSessionStateBag>(json, AgentAbstractionsJsonUtilities.DefaultOptions);
// Assert
Assert.NotNull(restored);
Assert.Equal("hello world", restored!.GetValue<string>("greeting"));
}
[Fact]
public void JsonSerializerRoundtrip_WithComplexObject_PreservesData()
{
// Arrange
var stateBag = new AgentSessionStateBag();
var animal = new Animal { Id = 10, FullName = "Rex", Species = Species.Tiger };
stateBag.SetValue("animal", animal, TestJsonSerializerContext.Default.Options);
// Act
var json = JsonSerializer.Serialize(stateBag, AgentAbstractionsJsonUtilities.DefaultOptions);
var restored = JsonSerializer.Deserialize<AgentSessionStateBag>(json, AgentAbstractionsJsonUtilities.DefaultOptions);
// Assert
Assert.NotNull(restored);
var restoredAnimal = restored!.GetValue<Animal>("animal", TestJsonSerializerContext.Default.Options);
Assert.NotNull(restoredAnimal);
Assert.Equal(10, restoredAnimal!.Id);
Assert.Equal("Rex", restoredAnimal.FullName);
Assert.Equal(Species.Tiger, restoredAnimal.Species);
}
[Fact]
public void JsonSerializerDeserialize_NullJson_ReturnsNull()
{
// Arrange
const string Json = "null";
// Act
var stateBag = JsonSerializer.Deserialize<AgentSessionStateBag>(Json, AgentAbstractionsJsonUtilities.DefaultOptions);
// Assert
Assert.Null(stateBag);
}
#if NET10_0_OR_GREATER
[Fact]
public void JsonSerializerSerialize_WithUnknownType_Throws()
{
// Arrange
var stateBag = new AgentSessionStateBag();
stateBag.SetValue("key", new { Name = "Test" }); // Anonymous type which cannot be deserialized
// Act & Assert
Assert.Throws<NotSupportedException>(() => JsonSerializer.Serialize(stateBag, AgentAbstractionsJsonUtilities.DefaultOptions));
}
#endif
#endregion
}
@@ -11,6 +11,21 @@ namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// </summary>
public class AgentSessionTests
{
#region StateBag Tests
[Fact]
public void StateBag_Values_Roundtrips()
{
// Arrange
var session = new TestAgentSession();
// Act & Assert
session.StateBag.SetValue("key1", "value1");
Assert.Equal("value1", session.StateBag.GetValue<string>("key1"));
}
#endregion
#region GetService Method Tests
/// <summary>
@@ -1,141 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Collections.Generic;
using System.Linq;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
using Moq.Protected;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// <summary>
/// Contains tests for the <see cref="ChatHistoryProviderExtensions"/> class.
/// </summary>
public sealed class ChatHistoryProviderExtensionsTests
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
[Fact]
public void WithMessageFilters_ReturnsChatHistoryProviderMessageFilter()
{
// Arrange
Mock<ChatHistoryProvider> providerMock = new();
// Act
ChatHistoryProvider result = providerMock.Object.WithMessageFilters(
invokingMessagesFilter: msgs => msgs,
invokedMessagesFilter: ctx => ctx);
// Assert
Assert.IsType<ChatHistoryProviderMessageFilter>(result);
}
[Fact]
public async Task WithMessageFilters_InvokingFilter_IsAppliedAsync()
{
// Arrange
Mock<ChatHistoryProvider> providerMock = new();
List<ChatMessage> innerMessages = [new(ChatRole.User, "Hello"), new(ChatRole.Assistant, "Hi")];
ChatHistoryProvider.InvokingContext context = new(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Test")]);
providerMock
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(innerMessages);
ChatHistoryProvider filtered = providerMock.Object.WithMessageFilters(
invokingMessagesFilter: msgs => msgs.Where(m => m.Role == ChatRole.User));
// Act
List<ChatMessage> result = (await filtered.InvokingAsync(context, CancellationToken.None)).ToList();
// Assert
Assert.Single(result);
Assert.Equal(ChatRole.User, result[0].Role);
}
[Fact]
public async Task WithMessageFilters_InvokedFilter_IsAppliedAsync()
{
// Arrange
Mock<ChatHistoryProvider> providerMock = new();
List<ChatMessage> requestMessages =
[
new(ChatRole.System, "System") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "TestSource") } } },
new(ChatRole.User, "Hello")
];
ChatHistoryProvider.InvokedContext context = new(s_mockAgent, s_mockSession, requestMessages)
{
ResponseMessages = [new ChatMessage(ChatRole.Assistant, "Response")]
};
ChatHistoryProvider.InvokedContext? capturedContext = null;
providerMock
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Callback<ChatHistoryProvider.InvokedContext, CancellationToken>((ctx, _) => capturedContext = ctx)
.Returns(default(ValueTask));
ChatHistoryProvider filtered = providerMock.Object.WithMessageFilters(
invokedMessagesFilter: ctx =>
{
ctx.ResponseMessages = null;
return ctx;
});
// Act
await filtered.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.NotNull(capturedContext);
Assert.Null(capturedContext.ResponseMessages);
}
[Fact]
public void WithAIContextProviderMessageRemoval_ReturnsChatHistoryProviderMessageFilter()
{
// Arrange
Mock<ChatHistoryProvider> providerMock = new();
// Act
ChatHistoryProvider result = providerMock.Object.WithAIContextProviderMessageRemoval();
// Assert
Assert.IsType<ChatHistoryProviderMessageFilter>(result);
}
[Fact]
public async Task WithAIContextProviderMessageRemoval_RemovesAIContextProviderMessagesAsync()
{
// Arrange
Mock<ChatHistoryProvider> providerMock = new();
List<ChatMessage> requestMessages =
[
new(ChatRole.System, "System") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "TestSource") } } },
new(ChatRole.User, "Hello"),
new(ChatRole.System, "Context") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "TestContextSource") } } }
];
ChatHistoryProvider.InvokedContext context = new(s_mockAgent, s_mockSession, requestMessages);
ChatHistoryProvider.InvokedContext? capturedContext = null;
providerMock
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Callback<ChatHistoryProvider.InvokedContext, CancellationToken>((ctx, _) => capturedContext = ctx)
.Returns(default(ValueTask));
ChatHistoryProvider filtered = providerMock.Object.WithAIContextProviderMessageRemoval();
// Act
await filtered.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.NotNull(capturedContext);
Assert.Equal(2, capturedContext.RequestMessages.Count());
Assert.Contains("System", capturedContext.RequestMessages.Select(x => x.Text));
Assert.Contains("Hello", capturedContext.RequestMessages.Select(x => x.Text));
}
}
@@ -1,221 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
using Moq.Protected;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// <summary>
/// Contains tests for the <see cref="ChatHistoryProviderMessageFilter"/> class.
/// </summary>
public sealed class ChatHistoryProviderMessageFilterTests
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
[Fact]
public void Constructor_WithNullInnerProvider_ThrowsArgumentNullException()
{
// Arrange, Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProviderMessageFilter(null!));
}
[Fact]
public void Constructor_WithOnlyInnerProvider_Throws()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
// Act & Assert
Assert.Throws<ArgumentException>(() => new ChatHistoryProviderMessageFilter(innerProviderMock.Object));
}
[Fact]
public void Constructor_WithAllParameters_CreatesInstance()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
IEnumerable<ChatMessage> InvokingFilter(IEnumerable<ChatMessage> msgs) => msgs;
ChatHistoryProvider.InvokedContext InvokedFilter(ChatHistoryProvider.InvokedContext ctx) => ctx;
// Act
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, InvokingFilter, InvokedFilter);
// Assert
Assert.NotNull(filter);
}
[Fact]
public async Task InvokingAsync_WithNoOpFilters_ReturnsInnerProviderMessagesAsync()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
var expectedMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello"),
new(ChatRole.Assistant, "Hi there!")
};
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Test")]);
innerProviderMock
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(expectedMessages);
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, x => x, x => x);
// Act
var result = (await filter.InvokingAsync(context, CancellationToken.None)).ToList();
// Assert
Assert.Equal(2, result.Count);
Assert.Equal("Hello", result[0].Text);
Assert.Equal("Hi there!", result[1].Text);
innerProviderMock
.Protected()
.Verify<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
}
[Fact]
public async Task InvokingAsync_WithInvokingFilter_AppliesFilterAsync()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
var innerMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello"),
new(ChatRole.Assistant, "Hi there!"),
new(ChatRole.User, "How are you?")
};
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Test")]);
innerProviderMock
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(innerMessages);
// Filter to only user messages
IEnumerable<ChatMessage> InvokingFilter(IEnumerable<ChatMessage> msgs) => msgs.Where(m => m.Role == ChatRole.User);
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, InvokingFilter);
// Act
var result = (await filter.InvokingAsync(context, CancellationToken.None)).ToList();
// Assert
Assert.Equal(2, result.Count);
Assert.All(result, msg => Assert.Equal(ChatRole.User, msg.Role));
innerProviderMock
.Protected()
.Verify<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
}
[Fact]
public async Task InvokingAsync_WithInvokingFilter_CanModifyMessagesAsync()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
var innerMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello"),
new(ChatRole.Assistant, "Hi there!")
};
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Test")]);
innerProviderMock
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(innerMessages);
// Filter that transforms messages
IEnumerable<ChatMessage> InvokingFilter(IEnumerable<ChatMessage> msgs) =>
msgs.Select(m => new ChatMessage(m.Role, $"[FILTERED] {m.Text}"));
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, InvokingFilter);
// Act
var result = (await filter.InvokingAsync(context, CancellationToken.None)).ToList();
// Assert
Assert.Equal(2, result.Count);
Assert.Equal("[FILTERED] Hello", result[0].Text);
Assert.Equal("[FILTERED] Hi there!", result[1].Text);
}
[Fact]
public async Task InvokedAsync_WithInvokedFilter_AppliesFilterAsync()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
List<ChatMessage> requestMessages =
[
new(ChatRole.System, "System") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "TestSource") } } },
new(ChatRole.User, "Hello"),
];
var responseMessages = new List<ChatMessage> { new(ChatRole.Assistant, "Response") };
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages)
{
ResponseMessages = responseMessages
};
ChatHistoryProvider.InvokedContext? capturedContext = null;
innerProviderMock
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Callback<ChatHistoryProvider.InvokedContext, CancellationToken>((ctx, ct) => capturedContext = ctx)
.Returns(default(ValueTask));
// Filter that modifies the context
ChatHistoryProvider.InvokedContext InvokedFilter(ChatHistoryProvider.InvokedContext ctx)
{
var modifiedRequestMessages = ctx.RequestMessages.Where(x => x.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External).Select(m => new ChatMessage(m.Role, $"[FILTERED] {m.Text}")).ToList();
return new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, modifiedRequestMessages)
{
ResponseMessages = ctx.ResponseMessages,
InvokeException = ctx.InvokeException
};
}
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, invokedMessagesFilter: InvokedFilter);
// Act
await filter.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.NotNull(capturedContext);
Assert.Single(capturedContext.RequestMessages);
Assert.Equal("[FILTERED] Hello", capturedContext.RequestMessages.First().Text);
innerProviderMock
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>());
}
[Fact]
public void Serialize_DelegatesToInnerProvider()
{
// Arrange
var innerProviderMock = new Mock<ChatHistoryProvider>();
var expectedJson = JsonSerializer.SerializeToElement("data", TestJsonSerializerContext.Default.String);
innerProviderMock
.Setup(s => s.Serialize(It.IsAny<JsonSerializerOptions>()))
.Returns(expectedJson);
var filter = new ChatHistoryProviderMessageFilter(innerProviderMock.Object, x => x, x => x);
// Act
var result = filter.Serialize();
// Assert
Assert.Equal(expectedJson.GetRawText(), result.GetRawText());
innerProviderMock.Verify(s => s.Serialize(null), Times.Once);
}
}
@@ -3,7 +3,6 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
@@ -19,92 +18,6 @@ public class ChatHistoryProviderTests
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
#region InvokingAsync Message Stamping Tests
[Fact]
public async Task InvokingAsync_StampsMessagesWithSourceTypeAndSourceIdAsync()
{
// Arrange
var provider = new TestChatHistoryProvider();
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
IEnumerable<ChatMessage> messages = await provider.InvokingAsync(context);
// Assert
ChatMessage message = messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, typedAttribution.SourceType);
Assert.Equal(typeof(TestChatHistoryProvider).FullName, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_WithCustomSourceId_StampsMessagesWithCustomSourceIdAsync()
{
// Arrange
const string CustomSourceId = "CustomHistorySource";
var provider = new TestChatHistoryProviderWithCustomSource(CustomSourceId);
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
IEnumerable<ChatMessage> messages = await provider.InvokingAsync(context);
// Assert
ChatMessage message = messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, typedAttribution.SourceType);
Assert.Equal(CustomSourceId, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_DoesNotReStampAlreadyStampedMessagesAsync()
{
// Arrange
var provider = new TestChatHistoryProviderWithPreStampedMessages();
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
IEnumerable<ChatMessage> messages = await provider.InvokingAsync(context);
// Assert
ChatMessage message = messages.Single();
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, typedAttribution.SourceType);
Assert.Equal(typeof(TestChatHistoryProviderWithPreStampedMessages).FullName, typedAttribution.SourceId);
}
[Fact]
public async Task InvokingAsync_StampsMultipleMessagesAsync()
{
// Arrange
var provider = new TestChatHistoryProviderWithMultipleMessages();
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Request")]);
// Act
IEnumerable<ChatMessage> messages = await provider.InvokingAsync(context);
// Assert
List<ChatMessage> messageList = messages.ToList();
Assert.Equal(3, messageList.Count);
foreach (ChatMessage message in messageList)
{
Assert.NotNull(message.AdditionalProperties);
Assert.True(message.AdditionalProperties.TryGetValue(AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, out object? attribution));
var typedAttribution = Assert.IsType<AgentRequestMessageSourceAttribution>(attribution);
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, typedAttribution.SourceType);
Assert.Equal(typeof(TestChatHistoryProviderWithMultipleMessages).FullName, typedAttribution.SourceId);
}
}
#endregion
#region GetService Method Tests
[Fact]
@@ -259,33 +172,7 @@ public class ChatHistoryProviderTests
public void InvokedContext_Constructor_ThrowsForNullRequestMessages()
{
// Arrange & Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, null!));
}
[Fact]
public void InvokedContext_RequestMessages_SetterThrowsForNull()
{
// Arrange
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => context.RequestMessages = null!);
}
[Fact]
public void InvokedContext_RequestMessages_SetterRoundtrips()
{
// Arrange
var initialMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
var newMessages = new List<ChatMessage> { new(ChatRole.User, "New message") };
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, initialMessages);
// Act
context.RequestMessages = newMessages;
// Assert
Assert.Same(newMessages, context.RequestMessages);
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, null!, []));
}
[Fact]
@@ -294,10 +181,9 @@ public class ChatHistoryProviderTests
// Arrange
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
var responseMessages = new List<ChatMessage> { new(ChatRole.Assistant, "Response message") };
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
// Act
context.ResponseMessages = responseMessages;
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, responseMessages);
// Assert
Assert.Same(responseMessages, context.ResponseMessages);
@@ -309,10 +195,9 @@ public class ChatHistoryProviderTests
// Arrange
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
var exception = new InvalidOperationException("Test exception");
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
// Act
context.InvokeException = exception;
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, exception);
// Assert
Assert.Same(exception, context.InvokeException);
@@ -325,7 +210,7 @@ public class ChatHistoryProviderTests
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, []);
// Assert
Assert.Same(s_mockAgent, context.Agent);
@@ -338,7 +223,7 @@ public class ChatHistoryProviderTests
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, []);
// Assert
Assert.Same(s_mockSession, context.Session);
@@ -351,7 +236,7 @@ public class ChatHistoryProviderTests
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, null, requestMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, null, requestMessages, []);
// Assert
Assert.Null(context.Session);
@@ -364,7 +249,27 @@ public class ChatHistoryProviderTests
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(null!, s_mockSession, requestMessages));
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(null!, s_mockSession, requestMessages, []));
}
[Fact]
public void InvokedContext_SuccessConstructor_ThrowsForNullResponseMessages()
{
// Arrange
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, (IEnumerable<ChatMessage>)null!));
}
[Fact]
public void InvokedContext_FailureConstructor_ThrowsForNullException()
{
// Arrange
var requestMessages = new List<ChatMessage> { new(ChatRole.User, "Hello") };
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages, (Exception)null!));
}
#endregion
@@ -372,63 +277,9 @@ public class ChatHistoryProviderTests
private sealed class TestChatHistoryProvider : ChatHistoryProvider
{
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new([new ChatMessage(ChatRole.User, "Test Message")]);
=> new(new ChatMessage[] { new(ChatRole.User, "Test Message") }.Concat(context.RequestMessages));
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
=> default;
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> default;
}
private sealed class TestChatHistoryProviderWithCustomSource : ChatHistoryProvider
{
public TestChatHistoryProviderWithCustomSource(string sourceId) : base(sourceId)
{
}
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new([new ChatMessage(ChatRole.User, "Test Message")]);
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
=> default;
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> default;
}
private sealed class TestChatHistoryProviderWithPreStampedMessages : ChatHistoryProvider
{
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
var message = new ChatMessage(ChatRole.User, "Pre-stamped Message");
message.AdditionalProperties = new AdditionalPropertiesDictionary
{
[AgentRequestMessageSourceAttribution.AdditionalPropertiesKey] = new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, this.GetType().FullName!)
};
return new([message]);
}
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
=> default;
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> default;
}
private sealed class TestChatHistoryProviderWithMultipleMessages : ChatHistoryProvider
{
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new([
new ChatMessage(ChatRole.User, "Message 1"),
new ChatMessage(ChatRole.Assistant, "Message 2"),
new ChatMessage(ChatRole.User, "Message 3")
]);
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
=> default;
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> default;
}
}
@@ -350,7 +350,7 @@ public sealed class ChatMessageExtensionsTests
ChatMessage message = new(ChatRole.User, "Hello");
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.External, "TestSourceId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.External, "TestSourceId");
// Assert
Assert.NotSame(message, result);
@@ -368,7 +368,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.AIContextProvider, "ProviderSourceId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.AIContextProvider, "ProviderSourceId");
// Assert
Assert.NotSame(message, result);
@@ -389,7 +389,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.ChatHistory, "HistoryId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.ChatHistory, "HistoryId");
// Assert
Assert.Same(message, result);
@@ -408,7 +408,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.AIContextProvider, "SourceId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.AIContextProvider, "SourceId");
// Assert
Assert.NotSame(message, result);
@@ -429,7 +429,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.External, "NewId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.External, "NewId");
// Assert
Assert.NotSame(message, result);
@@ -444,7 +444,7 @@ public sealed class ChatMessageExtensionsTests
ChatMessage message = new(ChatRole.User, "Hello");
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.ChatHistory);
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.ChatHistory);
// Assert
Assert.NotSame(message, result);
@@ -465,7 +465,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.External);
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.External);
// Assert
Assert.Same(message, result);
@@ -478,7 +478,7 @@ public sealed class ChatMessageExtensionsTests
ChatMessage message = new(ChatRole.User, "Hello");
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.AIContextProvider, "ProviderId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.AIContextProvider, "ProviderId");
// Assert
Assert.Null(message.AdditionalProperties);
@@ -499,7 +499,7 @@ public sealed class ChatMessageExtensionsTests
};
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.External, "SourceId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.External, "SourceId");
// Assert
Assert.NotSame(message, result);
@@ -514,7 +514,7 @@ public sealed class ChatMessageExtensionsTests
ChatMessage message = new(ChatRole.Assistant, "Test content");
// Act
ChatMessage result = message.AsAgentRequestMessageSourcedMessage(AgentRequestMessageSourceType.ChatHistory, "HistoryId");
ChatMessage result = message.WithAgentRequestMessageSource(AgentRequestMessageSourceType.ChatHistory, "HistoryId");
// Assert
Assert.Equal(ChatRole.Assistant, result.Role);
@@ -1,155 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using Microsoft.Extensions.AI;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// <summary>
/// Contains tests for <see cref="InMemoryAgentSession"/>.
/// </summary>
public class InMemoryAgentSessionTests
{
#region Constructor and Property Tests
[Fact]
public void Constructor_SetsDefaultChatHistoryProvider()
{
// Arrange & Act
var session = new TestInMemoryAgentSession();
// Assert
Assert.NotNull(session.GetChatHistoryProvider());
Assert.Empty(session.GetChatHistoryProvider());
}
[Fact]
public void Constructor_WithChatHistoryProvider_SetsProperty()
{
// Arrange
InMemoryChatHistoryProvider provider = [new(ChatRole.User, "Hello")];
// Act
var session = new TestInMemoryAgentSession(provider);
// Assert
Assert.Same(provider, session.GetChatHistoryProvider());
Assert.Single(session.GetChatHistoryProvider());
Assert.Equal("Hello", session.GetChatHistoryProvider()[0].Text);
}
[Fact]
public void Constructor_WithMessages_SetsProperty()
{
// Arrange
var messages = new List<ChatMessage> { new(ChatRole.User, "Hi") };
// Act
var session = new TestInMemoryAgentSession(messages);
// Assert
Assert.NotNull(session.GetChatHistoryProvider());
Assert.Single(session.GetChatHistoryProvider());
Assert.Equal("Hi", session.GetChatHistoryProvider()[0].Text);
}
[Fact]
public void Constructor_WithSerializedState_SetsProperty()
{
// Arrange
InMemoryChatHistoryProvider provider = [new(ChatRole.User, "TestMsg")];
var providerState = provider.Serialize();
var sessionStateWrapper = new InMemoryAgentSession.InMemoryAgentSessionState { ChatHistoryProviderState = providerState };
var json = JsonSerializer.SerializeToElement(sessionStateWrapper, TestJsonSerializerContext.Default.InMemoryAgentSessionState);
// Act
var session = new TestInMemoryAgentSession(json);
// Assert
Assert.NotNull(session.GetChatHistoryProvider());
Assert.Single(session.GetChatHistoryProvider());
Assert.Equal("TestMsg", session.GetChatHistoryProvider()[0].Text);
}
[Fact]
public void Constructor_WithInvalidJson_ThrowsArgumentException()
{
// Arrange
var invalidJson = JsonSerializer.SerializeToElement(42, TestJsonSerializerContext.Default.Int32);
// Act & Assert
Assert.Throws<ArgumentException>(() => new TestInMemoryAgentSession(invalidJson));
}
#endregion
#region SerializeAsync Tests
[Fact]
public void Serialize_ReturnsCorrectJson_WhenMessagesExist()
{
// Arrange
var session = new TestInMemoryAgentSession([new(ChatRole.User, "TestContent")]);
// Act
var json = session.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("chatHistoryProviderState", out var providerStateProperty));
Assert.Equal(JsonValueKind.Object, providerStateProperty.ValueKind);
Assert.True(providerStateProperty.TryGetProperty("messages", out var messagesProperty));
Assert.Equal(JsonValueKind.Array, messagesProperty.ValueKind);
var messagesList = messagesProperty.EnumerateArray().ToList();
Assert.Single(messagesList);
}
[Fact]
public void Serialize_ReturnsEmptyMessages_WhenNoMessages()
{
// Arrange
var session = new TestInMemoryAgentSession();
// Act
var json = session.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("chatHistoryProviderState", out var providerStateProperty));
Assert.Equal(JsonValueKind.Object, providerStateProperty.ValueKind);
Assert.True(providerStateProperty.TryGetProperty("messages", out var messagesProperty));
Assert.Equal(JsonValueKind.Array, messagesProperty.ValueKind);
Assert.Empty(messagesProperty.EnumerateArray());
}
#endregion
#region GetService Tests
[Fact]
public void GetService_RequestingChatHistoryProvider_ReturnsChatHistoryProvider()
{
// Arrange
var session = new TestInMemoryAgentSession();
// Act & Assert
Assert.NotNull(session.GetService(typeof(ChatHistoryProvider)));
Assert.Same(session.GetChatHistoryProvider(), session.GetService(typeof(ChatHistoryProvider)));
Assert.Same(session.GetChatHistoryProvider(), session.GetService(typeof(InMemoryChatHistoryProvider)));
}
#endregion
// Sealed test subclass to expose protected members for testing
private sealed class TestInMemoryAgentSession : InMemoryAgentSession
{
public TestInMemoryAgentSession() { }
public TestInMemoryAgentSession(InMemoryChatHistoryProvider? provider) : base(provider) { }
public TestInMemoryAgentSession(IEnumerable<ChatMessage> messages) : base(messages) { }
public TestInMemoryAgentSession(JsonElement serializedSessionState) : base(serializedSessionState) { }
public InMemoryChatHistoryProvider GetChatHistoryProvider() => this.ChatHistoryProvider;
}
}
@@ -3,9 +3,6 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Encodings.Web;
using System.Text.Json;
using System.Text.Json.Serialization.Metadata;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
@@ -19,22 +16,18 @@ namespace Microsoft.Agents.AI.Abstractions.UnitTests;
public class InMemoryChatHistoryProviderTests
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
[Fact]
public void Constructor_Throws_ForNullReducer() =>
// Arrange & Act & Assert
Assert.Throws<ArgumentNullException>(() => new InMemoryChatHistoryProvider(null!));
private static AgentSession CreateMockSession() => new Mock<AgentSession>().Object;
[Fact]
public void Constructor_DefaultsToBeforeMessageRetrieval_ForNotProvidedTriggerEvent()
{
// Arrange & Act
var reducerMock = new Mock<IChatReducer>();
var provider = new InMemoryChatHistoryProvider(reducerMock.Object);
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object });
// Assert
Assert.Equal(InMemoryChatHistoryProvider.ChatReducerTriggerEvent.BeforeMessagesRetrieval, provider.ReducerTriggerEvent);
Assert.Equal(InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.BeforeMessagesRetrieval, provider.ReducerTriggerEvent);
}
[Fact]
@@ -42,20 +35,43 @@ public class InMemoryChatHistoryProviderTests
{
// Arrange & Act
var reducerMock = new Mock<IChatReducer>();
var provider = new InMemoryChatHistoryProvider(reducerMock.Object, InMemoryChatHistoryProvider.ChatReducerTriggerEvent.AfterMessageAdded);
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object, ReducerTriggerEvent = InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.AfterMessageAdded });
// Assert
Assert.Same(reducerMock.Object, provider.ChatReducer);
Assert.Equal(InMemoryChatHistoryProvider.ChatReducerTriggerEvent.AfterMessageAdded, provider.ReducerTriggerEvent);
Assert.Equal(InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.AfterMessageAdded, provider.ReducerTriggerEvent);
}
[Fact]
public void StateKey_ReturnsDefaultKey_WhenNoOptionsProvided()
{
// Arrange & Act
var provider = new InMemoryChatHistoryProvider();
// Assert
Assert.Equal("InMemoryChatHistoryProvider", provider.StateKey);
}
[Fact]
public void StateKey_ReturnsCustomKey_WhenSetViaOptions()
{
// Arrange & Act
var provider = new InMemoryChatHistoryProvider(new() { StateKey = "custom-key" });
// Assert
Assert.Equal("custom-key", provider.StateKey);
}
[Fact]
public async Task InvokedAsyncAddsMessagesAsync()
{
var session = CreateMockSession();
// Arrange
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello"),
new(ChatRole.System, "additional context") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "TestSource") } } },
new(ChatRole.System, "additional context") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "TestSource") } } },
};
var responseMessages = new List<ChatMessage>
{
@@ -67,125 +83,146 @@ public class InMemoryChatHistoryProviderTests
};
var provider = new InMemoryChatHistoryProvider();
provider.Add(providerMessages[0]);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages)
{
ResponseMessages = responseMessages
};
provider.SetMessages(session, providerMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, responseMessages);
await provider.InvokedAsync(context, CancellationToken.None);
Assert.Equal(4, provider.Count);
Assert.Equal("original instructions", provider[0].Text);
Assert.Equal("Hello", provider[1].Text);
Assert.Equal("additional context", provider[2].Text);
Assert.Equal("Hi there!", provider[3].Text);
// Assert
var messages = provider.GetMessages(session);
Assert.Equal(4, messages.Count);
Assert.Equal("original instructions", messages[0].Text);
Assert.Equal("Hello", messages[1].Text);
Assert.Equal("additional context", messages[2].Text);
Assert.Equal("Hi there!", messages[3].Text);
}
[Fact]
public async Task InvokedAsyncWithEmptyDoesNotFailAsync()
{
var session = CreateMockSession();
// Arrange
var provider = new InMemoryChatHistoryProvider();
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, []);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [], []);
await provider.InvokedAsync(context, CancellationToken.None);
Assert.Empty(provider);
// Assert
Assert.Empty(provider.GetMessages(session));
}
[Fact]
public async Task InvokingAsyncReturnsAllMessagesAsync()
{
var provider = new InMemoryChatHistoryProvider
var session = CreateMockSession();
// Arrange
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello"),
};
var provider = new InMemoryChatHistoryProvider();
provider.SetMessages(session,
[
new ChatMessage(ChatRole.User, "Test1"),
new ChatMessage(ChatRole.Assistant, "Test2")
};
]);
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, requestMessages);
var result = (await provider.InvokingAsync(context, CancellationToken.None)).ToList();
Assert.Equal(2, result.Count);
// Assert
Assert.Equal(3, result.Count);
Assert.Contains(result, m => m.Text == "Test1");
Assert.Contains(result, m => m.Text == "Test2");
Assert.Contains(result, m => m.Text == "Hello");
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, result[0].GetAgentRequestMessageSourceType());
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, result[1].GetAgentRequestMessageSourceType());
Assert.Equal(AgentRequestMessageSourceType.External, result[2].GetAgentRequestMessageSourceType());
}
[Fact]
public async Task DeserializeConstructorWithEmptyElementAsync()
public void StateInitializer_IsInvoked_WhenSessionHasNoState()
{
var emptyObject = JsonSerializer.Deserialize("{}", TestJsonSerializerContext.Default.JsonElement);
// Arrange
var initialMessages = new List<ChatMessage>
{
new(ChatRole.User, "Initial message")
};
var provider = new InMemoryChatHistoryProvider(new()
{
StateInitializer = _ => new InMemoryChatHistoryProvider.State { Messages = initialMessages }
});
var newProvider = new InMemoryChatHistoryProvider(emptyObject);
// Act
var messages = provider.GetMessages(CreateMockSession());
Assert.Empty(newProvider);
// Assert
Assert.Single(messages);
Assert.Equal("Initial message", messages[0].Text);
}
[Fact]
public async Task SerializeAndDeserializeConstructorRoundtripsAsync()
public void GetMessages_ReturnsEmptyList_WhenNullSession()
{
var provider = new InMemoryChatHistoryProvider
{
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B")
};
// Arrange
var provider = new InMemoryChatHistoryProvider();
var jsonElement = provider.Serialize();
var newProvider = new InMemoryChatHistoryProvider(jsonElement);
// Act
var messages = provider.GetMessages(null);
Assert.Equal(2, newProvider.Count);
Assert.Equal("A", newProvider[0].Text);
Assert.Equal("B", newProvider[1].Text);
// Assert
Assert.Empty(messages);
}
[Fact]
public async Task SerializeAndDeserializeConstructorRoundtripsWithCustomAIContentAsync()
public void SetMessages_ThrowsForNullMessages()
{
JsonSerializerOptions options = new(TestJsonSerializerContext.Default.Options)
{
TypeInfoResolver = JsonTypeInfoResolver.Combine(AgentAbstractionsJsonUtilities.DefaultOptions.TypeInfoResolver, TestJsonSerializerContext.Default),
Encoder = JavaScriptEncoder.UnsafeRelaxedJsonEscaping,
};
options.AddAIContentType<TestAIContent>(typeDiscriminatorId: "testContent");
// Arrange
var provider = new InMemoryChatHistoryProvider();
var provider = new InMemoryChatHistoryProvider
{
new ChatMessage(ChatRole.User, [new TestAIContent("foo data")]),
};
var jsonElement = provider.Serialize(options);
var newProvider = new InMemoryChatHistoryProvider(jsonElement, options);
Assert.Single(newProvider);
var actualTestAIContent = Assert.IsType<TestAIContent>(newProvider[0].Contents[0]);
Assert.Equal("foo data", actualTestAIContent.TestData);
// Act & Assert
Assert.Throws<ArgumentNullException>(() => provider.SetMessages(CreateMockSession(), null!));
}
[Fact]
public async Task SerializeAndDeserializeWorksWithExperimentalContentTypesAsync()
public void SetMessages_UpdatesState()
{
var provider = new InMemoryChatHistoryProvider
var session = CreateMockSession();
// Arrange
var provider = new InMemoryChatHistoryProvider();
var messages = new List<ChatMessage>
{
new ChatMessage(ChatRole.User, [new FunctionApprovalRequestContent("call123", new FunctionCallContent("call123", "some_func"))]),
new ChatMessage(ChatRole.Assistant, [new FunctionApprovalResponseContent("call123", true, new FunctionCallContent("call123", "some_func"))])
new(ChatRole.User, "Hello"),
new(ChatRole.Assistant, "World")
};
var jsonElement = provider.Serialize();
var newProvider = new InMemoryChatHistoryProvider(jsonElement);
// Act
provider.SetMessages(session, messages);
var retrieved = provider.GetMessages(session);
Assert.Equal(2, newProvider.Count);
Assert.IsType<FunctionApprovalRequestContent>(newProvider[0].Contents[0]);
Assert.IsType<FunctionApprovalResponseContent>(newProvider[1].Contents[0]);
// Assert
Assert.Equal(2, retrieved.Count);
Assert.Equal("Hello", retrieved[0].Text);
Assert.Equal("World", retrieved[1].Text);
}
[Fact]
public async Task InvokedAsyncWithEmptyMessagesDoesNotChangeProviderAsync()
{
var session = CreateMockSession();
// Arrange
var provider = new InMemoryChatHistoryProvider();
var messages = new List<ChatMessage>();
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, messages, []);
await provider.InvokedAsync(context, CancellationToken.None);
Assert.Empty(provider);
// Assert
Assert.Empty(provider.GetMessages(session));
}
[Fact]
@@ -198,308 +235,11 @@ public class InMemoryChatHistoryProviderTests
await Assert.ThrowsAsync<ArgumentNullException>(() => provider.InvokedAsync(null!, CancellationToken.None).AsTask());
}
[Fact]
public void DeserializeContructor_WithNullSerializedState_CreatesEmptyProvider()
{
// Act
var provider = new InMemoryChatHistoryProvider(new JsonElement());
// Assert
Assert.Empty(provider);
}
[Fact]
public async Task DeserializeContructor_WithEmptyMessages_DoesNotAddMessagesAsync()
{
// Arrange
var stateWithEmptyMessages = JsonSerializer.SerializeToElement(
new Dictionary<string, object> { ["messages"] = new List<ChatMessage>() },
TestJsonSerializerContext.Default.IDictionaryStringObject);
// Act
var provider = new InMemoryChatHistoryProvider(stateWithEmptyMessages);
// Assert
Assert.Empty(provider);
}
[Fact]
public async Task DeserializeConstructor_WithNullMessages_DoesNotAddMessagesAsync()
{
// Arrange
var stateWithNullMessages = JsonSerializer.SerializeToElement(
new Dictionary<string, object> { ["messages"] = null! },
TestJsonSerializerContext.Default.DictionaryStringObject);
// Act
var provider = new InMemoryChatHistoryProvider(stateWithNullMessages);
// Assert
Assert.Empty(provider);
}
[Fact]
public async Task DeserializeConstructor_WithValidMessages_AddsMessagesAsync()
{
// Arrange
var messages = new List<ChatMessage>
{
new(ChatRole.User, "User message"),
new(ChatRole.Assistant, "Assistant message")
};
var state = new Dictionary<string, object> { ["messages"] = messages };
var serializedState = JsonSerializer.SerializeToElement(
state,
TestJsonSerializerContext.Default.DictionaryStringObject);
// Act
var provider = new InMemoryChatHistoryProvider(serializedState);
// Assert
Assert.Equal(2, provider.Count);
Assert.Equal("User message", provider[0].Text);
Assert.Equal("Assistant message", provider[1].Text);
}
[Fact]
public void IndexerGet_ReturnsCorrectMessage()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
provider.Add(message2);
// Act & Assert
Assert.Same(message1, provider[0]);
Assert.Same(message2, provider[1]);
}
[Fact]
public void IndexerSet_UpdatesMessage()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var originalMessage = new ChatMessage(ChatRole.User, "Original");
var newMessage = new ChatMessage(ChatRole.User, "Updated");
provider.Add(originalMessage);
// Act
provider[0] = newMessage;
// Assert
Assert.Same(newMessage, provider[0]);
Assert.Equal("Updated", provider[0].Text);
}
[Fact]
public void IsReadOnly_ReturnsFalse()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
// Act & Assert
Assert.False(provider.IsReadOnly);
}
[Fact]
public void IndexOf_ReturnsCorrectIndex()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
var message3 = new ChatMessage(ChatRole.User, "Third");
provider.Add(message1);
provider.Add(message2);
// Act & Assert
Assert.Equal(0, provider.IndexOf(message1));
Assert.Equal(1, provider.IndexOf(message2));
Assert.Equal(-1, provider.IndexOf(message3)); // Not in provider
}
[Fact]
public void Insert_InsertsMessageAtCorrectIndex()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
var insertMessage = new ChatMessage(ChatRole.User, "Inserted");
provider.Add(message1);
provider.Add(message2);
// Act
provider.Insert(1, insertMessage);
// Assert
Assert.Equal(3, provider.Count);
Assert.Same(message1, provider[0]);
Assert.Same(insertMessage, provider[1]);
Assert.Same(message2, provider[2]);
}
[Fact]
public void RemoveAt_RemovesMessageAtIndex()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
var message3 = new ChatMessage(ChatRole.User, "Third");
provider.Add(message1);
provider.Add(message2);
provider.Add(message3);
// Act
provider.RemoveAt(1);
// Assert
Assert.Equal(2, provider.Count);
Assert.Same(message1, provider[0]);
Assert.Same(message3, provider[1]);
}
[Fact]
public void Clear_RemovesAllMessages()
{
// Arrange
var provider = new InMemoryChatHistoryProvider
{
new ChatMessage(ChatRole.User, "First"),
new ChatMessage(ChatRole.Assistant, "Second")
};
// Act
provider.Clear();
// Assert
Assert.Empty(provider);
}
[Fact]
public void Contains_ReturnsTrueForExistingMessage()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
// Act & Assert
Assert.Contains(message1, provider);
Assert.DoesNotContain(message2, provider);
}
[Fact]
public void CopyTo_CopiesMessagesToArray()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
provider.Add(message2);
var array = new ChatMessage[4];
// Act
provider.CopyTo(array, 1);
// Assert
Assert.Null(array[0]);
Assert.Same(message1, array[1]);
Assert.Same(message2, array[2]);
Assert.Null(array[3]);
}
[Fact]
public void Remove_RemovesSpecificMessage()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
var message3 = new ChatMessage(ChatRole.User, "Third");
provider.Add(message1);
provider.Add(message2);
provider.Add(message3);
// Act
var removed = provider.Remove(message2);
// Assert
Assert.True(removed);
Assert.Equal(2, provider.Count);
Assert.Same(message1, provider[0]);
Assert.Same(message3, provider[1]);
}
[Fact]
public void Remove_ReturnsFalseForNonExistentMessage()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
// Act
var removed = provider.Remove(message2);
// Assert
Assert.False(removed);
Assert.Single(provider);
}
[Fact]
public void GetEnumerator_Generic_ReturnsAllMessages()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
provider.Add(message2);
// Act
var messages = new List<ChatMessage>();
messages.AddRange(provider);
// Assert
Assert.Equal(2, messages.Count);
Assert.Same(message1, messages[0]);
Assert.Same(message2, messages[1]);
}
[Fact]
public void GetEnumerator_NonGeneric_ReturnsAllMessages()
{
// Arrange
var provider = new InMemoryChatHistoryProvider();
var message1 = new ChatMessage(ChatRole.User, "First");
var message2 = new ChatMessage(ChatRole.Assistant, "Second");
provider.Add(message1);
provider.Add(message2);
// Act
var messages = new List<ChatMessage>();
var enumerator = ((System.Collections.IEnumerable)provider).GetEnumerator();
while (enumerator.MoveNext())
{
messages.Add((ChatMessage)enumerator.Current);
}
// Assert
Assert.Equal(2, messages.Count);
Assert.Same(message1, messages[0]);
Assert.Same(message2, messages[1]);
}
[Fact]
public async Task AddMessagesAsync_WithReducer_AfterMessageAdded_InvokesReducerAsync()
{
var session = CreateMockSession();
// Arrange
var originalMessages = new List<ChatMessage>
{
@@ -516,21 +256,24 @@ public class InMemoryChatHistoryProviderTests
.Setup(r => r.ReduceAsync(It.Is<List<ChatMessage>>(x => x.SequenceEqual(originalMessages)), It.IsAny<CancellationToken>()))
.ReturnsAsync(reducedMessages);
var provider = new InMemoryChatHistoryProvider(reducerMock.Object, InMemoryChatHistoryProvider.ChatReducerTriggerEvent.AfterMessageAdded);
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object, ReducerTriggerEvent = InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.AfterMessageAdded });
// Act
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, originalMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, originalMessages, []);
await provider.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.Single(provider);
Assert.Equal("Reduced", provider[0].Text);
var messages = provider.GetMessages(session);
Assert.Single(messages);
Assert.Equal("Reduced", messages[0].Text);
reducerMock.Verify(r => r.ReduceAsync(It.Is<List<ChatMessage>>(x => x.SequenceEqual(originalMessages)), It.IsAny<CancellationToken>()), Times.Once);
}
[Fact]
public async Task GetMessagesAsync_WithReducer_BeforeMessagesRetrieval_InvokesReducerAsync()
{
var session = CreateMockSession();
// Arrange
var originalMessages = new List<ChatMessage>
{
@@ -547,15 +290,11 @@ public class InMemoryChatHistoryProviderTests
.Setup(r => r.ReduceAsync(It.Is<List<ChatMessage>>(x => x.SequenceEqual(originalMessages)), It.IsAny<CancellationToken>()))
.ReturnsAsync(reducedMessages);
var provider = new InMemoryChatHistoryProvider(reducerMock.Object, InMemoryChatHistoryProvider.ChatReducerTriggerEvent.BeforeMessagesRetrieval);
// Add messages directly to the provider for this test
foreach (var msg in originalMessages)
{
provider.Add(msg);
}
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object, ReducerTriggerEvent = InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.BeforeMessagesRetrieval });
provider.SetMessages(session, new List<ChatMessage>(originalMessages));
// Act
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, Array.Empty<ChatMessage>());
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, Array.Empty<ChatMessage>());
var result = (await provider.InvokingAsync(invokingContext, CancellationToken.None)).ToList();
// Assert
@@ -567,6 +306,8 @@ public class InMemoryChatHistoryProviderTests
[Fact]
public async Task AddMessagesAsync_WithReducer_ButWrongTrigger_DoesNotInvokeReducerAsync()
{
var session = CreateMockSession();
// Arrange
var originalMessages = new List<ChatMessage>
{
@@ -575,21 +316,24 @@ public class InMemoryChatHistoryProviderTests
var reducerMock = new Mock<IChatReducer>();
var provider = new InMemoryChatHistoryProvider(reducerMock.Object, InMemoryChatHistoryProvider.ChatReducerTriggerEvent.BeforeMessagesRetrieval);
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object, ReducerTriggerEvent = InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.BeforeMessagesRetrieval });
// Act
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, originalMessages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, originalMessages, []);
await provider.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.Single(provider);
Assert.Equal("Hello", provider[0].Text);
var messages = provider.GetMessages(session);
Assert.Single(messages);
Assert.Equal("Hello", messages[0].Text);
reducerMock.Verify(r => r.ReduceAsync(It.IsAny<IEnumerable<ChatMessage>>(), It.IsAny<CancellationToken>()), Times.Never);
}
[Fact]
public async Task GetMessagesAsync_WithReducer_ButWrongTrigger_DoesNotInvokeReducerAsync()
{
var session = CreateMockSession();
// Arrange
var originalMessages = new List<ChatMessage>
{
@@ -598,13 +342,11 @@ public class InMemoryChatHistoryProviderTests
var reducerMock = new Mock<IChatReducer>();
var provider = new InMemoryChatHistoryProvider(reducerMock.Object, InMemoryChatHistoryProvider.ChatReducerTriggerEvent.AfterMessageAdded)
{
originalMessages[0]
};
var provider = new InMemoryChatHistoryProvider(new() { ChatReducer = reducerMock.Object, ReducerTriggerEvent = InMemoryChatHistoryProviderOptions.ChatReducerTriggerEvent.AfterMessageAdded });
provider.SetMessages(session, new List<ChatMessage>(originalMessages));
// Act
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, Array.Empty<ChatMessage>());
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, Array.Empty<ChatMessage>());
var result = (await provider.InvokingAsync(invokingContext, CancellationToken.None)).ToList();
// Assert
@@ -616,27 +358,21 @@ public class InMemoryChatHistoryProviderTests
[Fact]
public async Task InvokedAsync_WithException_DoesNotAddMessagesAsync()
{
var session = CreateMockSession();
// Arrange
var provider = new InMemoryChatHistoryProvider();
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "Hello")
};
var responseMessages = new List<ChatMessage>
{
new(ChatRole.Assistant, "Hi there!")
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages)
{
ResponseMessages = responseMessages,
InvokeException = new InvalidOperationException("Test exception")
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, new InvalidOperationException("Test exception"));
// Act
await provider.InvokedAsync(context, CancellationToken.None);
// Assert
Assert.Empty(provider);
Assert.Empty(provider.GetMessages(session));
}
[Fact]
@@ -649,6 +385,85 @@ public class InMemoryChatHistoryProviderTests
await Assert.ThrowsAsync<ArgumentNullException>(() => provider.InvokingAsync(null!, CancellationToken.None).AsTask());
}
[Fact]
public async Task InvokedAsync_DefaultFilter_ExcludesChatHistoryMessagesAsync()
{
// Arrange
var session = CreateMockSession();
var provider = new InMemoryChatHistoryProvider();
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(context, CancellationToken.None);
// Assert - ChatHistory message excluded, AIContextProvider message included
var messages = provider.GetMessages(session);
Assert.Equal(3, messages.Count);
Assert.Equal("External message", messages[0].Text);
Assert.Equal("From context provider", messages[1].Text);
Assert.Equal("Response", messages[2].Text);
}
[Fact]
public async Task InvokedAsync_CustomFilter_OverridesDefaultAsync()
{
// Arrange
var session = CreateMockSession();
var provider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions
{
StorageInputMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External)
});
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(context, CancellationToken.None);
// Assert - Custom filter keeps only External messages (both ChatHistory and AIContextProvider excluded)
var messages = provider.GetMessages(session);
Assert.Equal(2, messages.Count);
Assert.Equal("External message", messages[0].Text);
Assert.Equal("Response", messages[1].Text);
}
[Fact]
public async Task InvokingAsync_OutputFilter_FiltersOutputMessagesAsync()
{
// Arrange
var session = CreateMockSession();
var provider = new InMemoryChatHistoryProvider(new InMemoryChatHistoryProviderOptions
{
RetrievalOutputMessageFilter = messages => messages.Where(m => m.Role == ChatRole.User)
});
provider.SetMessages(session,
[
new ChatMessage(ChatRole.User, "User message"),
new ChatMessage(ChatRole.Assistant, "Assistant message"),
new ChatMessage(ChatRole.System, "System message")
]);
// Act
var context = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var result = (await provider.InvokingAsync(context, CancellationToken.None)).ToList();
// Assert - Only user messages pass through the output filter
Assert.Single(result);
Assert.Equal("User message", result[0].Text);
}
public class TestAIContent(string testData) : AIContent
{
public string TestData => testData;
@@ -1,119 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Text.Json;
namespace Microsoft.Agents.AI.Abstractions.UnitTests;
/// <summary>
/// Tests for <see cref="ServiceIdAgentSession"/>.
/// </summary>
public class ServiceIdAgentSessionTests
{
#region Constructor and Property Tests
[Fact]
public void Constructor_SetsDefaults()
{
// Arrange & Act
var session = new TestServiceIdAgentSession();
// Assert
Assert.Null(session.GetServiceSessionId());
}
[Fact]
public void Constructor_WithServiceSessionId_SetsProperty()
{
// Arrange & Act
var session = new TestServiceIdAgentSession("service-id-123");
// Assert
Assert.Equal("service-id-123", session.GetServiceSessionId());
}
[Fact]
public void Constructor_WithSerializedId_SetsProperty()
{
// Arrange
var serviceSessionWrapper = new ServiceIdAgentSession.ServiceIdAgentSessionState { ServiceSessionId = "service-id-456" };
var json = JsonSerializer.SerializeToElement(serviceSessionWrapper, TestJsonSerializerContext.Default.ServiceIdAgentSessionState);
// Act
var session = new TestServiceIdAgentSession(json);
// Assert
Assert.Equal("service-id-456", session.GetServiceSessionId());
}
[Fact]
public void Constructor_WithSerializedUndefinedId_SetsProperty()
{
// Arrange
var emptyObject = new EmptyObject();
var json = JsonSerializer.SerializeToElement(emptyObject, TestJsonSerializerContext.Default.EmptyObject);
// Act
var session = new TestServiceIdAgentSession(json);
// Assert
Assert.Null(session.GetServiceSessionId());
}
[Fact]
public void Constructor_WithInvalidJson_ThrowsArgumentException()
{
// Arrange
var invalidJson = JsonSerializer.SerializeToElement(42, TestJsonSerializerContext.Default.Int32);
// Act & Assert
Assert.Throws<ArgumentException>(() => new TestServiceIdAgentSession(invalidJson));
}
#endregion
#region SerializeAsync Tests
[Fact]
public void Serialize_ReturnsCorrectJson_WhenServiceSessionIdIsSet()
{
// Arrange
var session = new TestServiceIdAgentSession("service-id-789");
// Act
var json = session.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("serviceSessionId", out var idProperty));
Assert.Equal("service-id-789", idProperty.GetString());
}
[Fact]
public void Serialize_ReturnsUndefinedServiceSessionId_WhenNotSet()
{
// Arrange
var session = new TestServiceIdAgentSession();
// Act
var json = session.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.False(json.TryGetProperty("serviceSessionId", out _));
}
#endregion
// Sealed test subclass to expose protected members for testing
private sealed class TestServiceIdAgentSession : ServiceIdAgentSession
{
public TestServiceIdAgentSession() { }
public TestServiceIdAgentSession(string serviceSessionId) : base(serviceSessionId) { }
public TestServiceIdAgentSession(JsonElement serializedSessionState) : base(serializedSessionState) { }
public string? GetServiceSessionId() => this.ServiceSessionId;
}
// Helper class to represent empty objects
internal sealed class EmptyObject;
}
@@ -20,8 +20,5 @@ namespace Microsoft.Agents.AI.Abstractions.UnitTests;
[JsonSerializable(typeof(Dictionary<string, object?>))]
[JsonSerializable(typeof(string[]))]
[JsonSerializable(typeof(int))]
[JsonSerializable(typeof(InMemoryAgentSession.InMemoryAgentSessionState))]
[JsonSerializable(typeof(ServiceIdAgentSession.ServiceIdAgentSessionState))]
[JsonSerializable(typeof(ServiceIdAgentSessionTests.EmptyObject))]
[JsonSerializable(typeof(InMemoryChatHistoryProviderTests.TestAIContent))]
internal sealed partial class TestJsonSerializerContext : JsonSerializerContext;
@@ -2310,23 +2310,18 @@ public sealed class AzureAIProjectChatClientExtensionsTests
#region CreateChatClientAgentOptions - Options Preservation Tests
/// <summary>
/// Verify that CreateChatClientAgentOptions preserves AIContextProviderFactory.
/// Verify that CreateChatClientAgentOptions preserves AIContextProviders.
/// </summary>
[Fact]
public async Task GetAIAgentAsync_WithAIContextProviderFactory_PreservesFactoryAsync()
public async Task GetAIAgentAsync_WithAIContextProviders_PreservesProviderAsync()
{
// Arrange
AIProjectClient client = this.CreateTestAgentClient();
bool factoryInvoked = false;
var options = new ChatClientAgentOptions
{
Name = "test-agent",
ChatOptions = new ChatOptions { Instructions = "Test" },
AIContextProviderFactory = (_, _) =>
{
factoryInvoked = true;
return new ValueTask<AIContextProvider>(new TestAIContextProvider());
}
AIContextProviders = [new TestAIContextProvider()]
};
// Act
@@ -2334,15 +2329,13 @@ public sealed class AzureAIProjectChatClientExtensionsTests
// Assert
Assert.NotNull(agent);
// Verify the factory was captured (though not necessarily invoked yet)
Assert.False(factoryInvoked); // Factory is not invoked during creation
}
/// <summary>
/// Verify that CreateChatClientAgentOptions preserves ChatHistoryProviderFactory.
/// Verify that CreateChatClientAgentOptions preserves ChatHistoryProvider.
/// </summary>
[Fact]
public async Task GetAIAgentAsync_WithChatHistoryProviderFactory_PreservesFactoryAsync()
public async Task GetAIAgentAsync_WithChatHistoryProvider_PreservesProviderAsync()
{
// Arrange
AIProjectClient client = this.CreateTestAgentClient();
@@ -2350,7 +2343,7 @@ public sealed class AzureAIProjectChatClientExtensionsTests
{
Name = "test-agent",
ChatOptions = new ChatOptions { Instructions = "Test" },
ChatHistoryProviderFactory = (_, _) => new ValueTask<ChatHistoryProvider>(new TestChatHistoryProvider())
ChatHistoryProvider = new TestChatHistoryProvider()
};
// Act
@@ -3142,7 +3135,7 @@ public sealed class AzureAIProjectChatClientExtensionsTests
{
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
return new ValueTask<AIContext>(new AIContext());
return new ValueTask<AIContext>(context.AIContext);
}
}
@@ -3153,18 +3146,13 @@ public sealed class AzureAIProjectChatClientExtensionsTests
{
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
{
return new ValueTask<IEnumerable<ChatMessage>>(Array.Empty<ChatMessage>());
return new ValueTask<IEnumerable<ChatMessage>>(context.RequestMessages);
}
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
{
return default;
}
public override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
{
return default;
}
}
}
@@ -3,8 +3,6 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Text.Json.Serialization.Metadata;
using System.Threading.Tasks;
using Azure.Core;
using Azure.Identity;
@@ -42,7 +40,8 @@ namespace Microsoft.Agents.AI.CosmosNoSql.UnitTests;
public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
private static readonly AIAgent s_mockAgent = new Moq.Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Moq.Mock<AgentSession>().Object;
private static AgentSession CreateMockSession() => new Moq.Mock<AgentSession>().Object;
// Cosmos DB Emulator connection settings
private const string EmulatorEndpoint = "https://localhost:8081";
@@ -149,6 +148,35 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
#region Constructor Tests
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void StateKey_ReturnsDefaultKey_WhenNoStateKeyProvided()
{
// Arrange & Act
this.SkipIfEmulatorNotAvailable();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State("test-conversation"));
// Assert
Assert.Equal("CosmosChatHistoryProvider", provider.StateKey);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void StateKey_ReturnsCustomKey_WhenSetViaConstructor()
{
// Arrange & Act
this.SkipIfEmulatorNotAvailable();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State("test-conversation"),
stateKey: "custom-key");
// Assert
Assert.Equal("custom-key", provider.StateKey);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithConnectionString_ShouldCreateInstance()
@@ -157,28 +185,11 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
this.SkipIfEmulatorNotAvailable();
// Act
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, "test-conversation");
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State("test-conversation"));
// Assert
Assert.NotNull(provider);
Assert.Equal("test-conversation", provider.ConversationId);
Assert.Equal(s_testDatabaseId, provider.DatabaseId);
Assert.Equal(TestContainerId, provider.ContainerId);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithConnectionStringNoConversationId_ShouldCreateInstance()
{
// Arrange
this.SkipIfEmulatorNotAvailable();
// Act
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId);
// Assert
Assert.NotNull(provider);
Assert.NotNull(provider.ConversationId);
Assert.Equal(s_testDatabaseId, provider.DatabaseId);
Assert.Equal(TestContainerId, provider.ContainerId);
}
@@ -189,18 +200,19 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange & Act & Assert
Assert.Throws<ArgumentNullException>(() =>
new CosmosChatHistoryProvider((string)null!, s_testDatabaseId, TestContainerId, "test-conversation"));
new CosmosChatHistoryProvider((string)null!, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State("test-conversation")));
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithEmptyConversationId_ShouldThrowArgumentException()
public void Constructor_WithNullStateInitializer_ShouldThrowArgumentNullException()
{
// Arrange & Act & Assert
this.SkipIfEmulatorNotAvailable();
Assert.Throws<ArgumentException>(() =>
new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, ""));
Assert.Throws<ArgumentNullException>(() =>
new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, null!));
}
#endregion
@@ -213,14 +225,13 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = Guid.NewGuid().ToString();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversationId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId));
var message = new ChatMessage(ChatRole.User, "Hello, world!");
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [message])
{
ResponseMessages = []
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [message], []);
// Act
await provider.InvokedAsync(context);
@@ -229,7 +240,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(100);
// Assert
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = await provider.InvokingAsync(invokingContext);
var messageList = messages.ToList();
@@ -279,8 +290,10 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = Guid.NewGuid().ToString();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversationId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId));
var requestMessages = new[]
{
new ChatMessage(ChatRole.User, "First message"),
@@ -293,16 +306,13 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
new ChatMessage(ChatRole.Assistant, "Response message")
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages)
{
ResponseMessages = responseMessages
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, responseMessages);
// Act
await provider.InvokedAsync(context);
// Assert
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var retrievedMessages = await provider.InvokingAsync(invokingContext);
var messageList = retrievedMessages.ToList();
Assert.Equal(5, messageList.Count);
@@ -323,10 +333,12 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, Guid.NewGuid().ToString());
var session = CreateMockSession();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(Guid.NewGuid().ToString()));
// Act
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = await provider.InvokingAsync(invokingContext);
// Assert
@@ -339,21 +351,25 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversation1 = Guid.NewGuid().ToString();
var conversation2 = Guid.NewGuid().ToString();
using var store1 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversation1);
using var store2 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversation2);
// Use different stateKey values so the providers don't overwrite each other's state in the shared session
using var store1 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversation1), stateKey: "conv1");
using var store2 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversation2), stateKey: "conv2");
var context1 = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Message for conversation 1")]);
var context2 = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Message for conversation 2")]);
var context1 = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Message for conversation 1")], []);
var context2 = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Message for conversation 2")], []);
await store1.InvokedAsync(context1);
await store2.InvokedAsync(context2);
// Act
var invokingContext1 = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext2 = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext1 = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var invokingContext2 = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages1 = await store1.InvokingAsync(invokingContext1);
var messages2 = await store2.InvokingAsync(invokingContext2);
@@ -365,6 +381,8 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
Assert.Single(messageList2);
Assert.Equal("Message for conversation 1", messageList1[0].Text);
Assert.Equal("Message for conversation 2", messageList2[0].Text);
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, messageList1[0].GetAgentRequestMessageSourceType());
Assert.Equal(AgentRequestMessageSourceType.ChatHistory, messageList2[0].GetAgentRequestMessageSourceType());
}
#endregion
@@ -377,8 +395,10 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = $"test-conversation-{Guid.NewGuid():N}"; // Use unique conversation ID
using var originalStore = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversationId);
using var originalStore = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId));
var messages = new[]
{
@@ -390,18 +410,21 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
};
// Act 1: Add messages
var invokedContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
var invokedContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, messages, []);
await originalStore.InvokedAsync(invokedContext);
// Act 2: Verify messages were added
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var retrievedMessages = await originalStore.InvokingAsync(invokingContext);
var retrievedList = retrievedMessages.ToList();
Assert.Equal(5, retrievedList.Count);
// Act 3: Create new provider instance for same conversation (test persistence)
using var newProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, conversationId);
var persistedMessages = await newProvider.InvokingAsync(invokingContext);
using var newProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId));
var newSession = CreateMockSession();
var newInvokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, newSession, []);
var persistedMessages = await newProvider.InvokingAsync(newInvokingContext);
var persistedList = persistedMessages.ToList();
// Assert final state
@@ -423,7 +446,8 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, Guid.NewGuid().ToString());
var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(Guid.NewGuid().ToString()));
// Act & Assert
provider.Dispose(); // Should not throw
@@ -435,7 +459,8 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, Guid.NewGuid().ToString());
var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(Guid.NewGuid().ToString()));
// Act & Assert
provider.Dispose(); // First call
@@ -454,11 +479,11 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
this.SkipIfEmulatorNotAvailable();
// Act
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, "tenant-123", "user-456", "session-789");
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State("session-789", "tenant-123", "user-456"));
// Assert
Assert.NotNull(provider);
Assert.Equal("session-789", provider.ConversationId);
Assert.Equal(s_testDatabaseId, provider.DatabaseId);
Assert.Equal(HierarchicalTestContainerId, provider.ContainerId);
}
@@ -472,11 +497,11 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
// Act
TokenCredential credential = new DefaultAzureCredential();
using var provider = new CosmosChatHistoryProvider(EmulatorEndpoint, credential, s_testDatabaseId, HierarchicalTestContainerId, "tenant-123", "user-456", "session-789");
using var provider = new CosmosChatHistoryProvider(EmulatorEndpoint, credential, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State("session-789", "tenant-123", "user-456"));
// Assert
Assert.NotNull(provider);
Assert.Equal("session-789", provider.ConversationId);
Assert.Equal(s_testDatabaseId, provider.DatabaseId);
Assert.Equal(HierarchicalTestContainerId, provider.ContainerId);
}
@@ -489,46 +514,31 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
this.SkipIfEmulatorNotAvailable();
using var cosmosClient = new CosmosClient(EmulatorEndpoint, EmulatorKey);
using var provider = new CosmosChatHistoryProvider(cosmosClient, s_testDatabaseId, HierarchicalTestContainerId, "tenant-123", "user-456", "session-789");
using var provider = new CosmosChatHistoryProvider(cosmosClient, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State("session-789", "tenant-123", "user-456"));
// Assert
Assert.NotNull(provider);
Assert.Equal("session-789", provider.ConversationId);
Assert.Equal(s_testDatabaseId, provider.DatabaseId);
Assert.Equal(HierarchicalTestContainerId, provider.ContainerId);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithHierarchicalNullTenantId_ShouldThrowArgumentException()
public void State_WithEmptyConversationId_ShouldThrowArgumentException()
{
// Arrange & Act & Assert
this.SkipIfEmulatorNotAvailable();
Assert.Throws<ArgumentNullException>(() =>
new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, null!, "user-456", "session-789"));
Assert.Throws<ArgumentException>(() =>
new CosmosChatHistoryProvider.State(""));
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithHierarchicalEmptyUserId_ShouldThrowArgumentException()
public void State_WithWhitespaceConversationId_ShouldThrowArgumentException()
{
// Arrange & Act & Assert
this.SkipIfEmulatorNotAvailable();
Assert.Throws<ArgumentException>(() =>
new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, "tenant-123", "", "session-789"));
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public void Constructor_WithHierarchicalWhitespaceSessionId_ShouldThrowArgumentException()
{
// Arrange & Act & Assert
this.SkipIfEmulatorNotAvailable();
Assert.Throws<ArgumentException>(() =>
new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, "tenant-123", "user-456", " "));
new CosmosChatHistoryProvider.State(" "));
}
[SkippableFact]
@@ -537,14 +547,16 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string TenantId = "tenant-123";
const string UserId = "user-456";
const string SessionId = "session-789";
// Test hierarchical partitioning constructor with connection string
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, TenantId, UserId, SessionId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, TenantId, UserId));
var message = new ChatMessage(ChatRole.User, "Hello from hierarchical partitioning!");
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [message]);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [message], []);
// Act
await provider.InvokedAsync(context);
@@ -553,7 +565,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(100);
// Assert
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = await provider.InvokingAsync(invokingContext);
var messageList = messages.ToList();
@@ -589,11 +601,13 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string TenantId = "tenant-batch";
const string UserId = "user-batch";
const string SessionId = "session-batch";
// Test hierarchical partitioning constructor with connection string
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, TenantId, UserId, SessionId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, TenantId, UserId));
var messages = new[]
{
new ChatMessage(ChatRole.User, "First hierarchical message"),
@@ -601,7 +615,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
new ChatMessage(ChatRole.User, "Third hierarchical message")
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, messages, []);
// Act
await provider.InvokedAsync(context);
@@ -610,7 +624,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(100);
// Assert
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var retrievedMessages = await provider.InvokingAsync(invokingContext);
var messageList = retrievedMessages.ToList();
@@ -626,18 +640,22 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string TenantId = "tenant-isolation";
const string UserId1 = "user-1";
const string UserId2 = "user-2";
const string SessionId = "session-isolation";
// Different userIds create different hierarchical partitions, providing proper isolation
using var store1 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, TenantId, UserId1, SessionId);
using var store2 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, TenantId, UserId2, SessionId);
// Use different stateKey values so the providers don't overwrite each other's state in the shared session
using var store1 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, TenantId, UserId1), stateKey: "user1");
using var store2 = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, TenantId, UserId2), stateKey: "user2");
// Add messages to both stores
var context1 = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Message from user 1")]);
var context2 = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Message from user 2")]);
var context1 = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Message from user 1")], []);
var context2 = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Message from user 2")], []);
await store1.InvokedAsync(context1);
await store2.InvokedAsync(context2);
@@ -646,8 +664,8 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(100);
// Act & Assert
var invokingContext1 = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext2 = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext1 = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var invokingContext2 = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages1 = await store1.InvokingAsync(invokingContext1);
var messageList1 = messages1.ToList();
@@ -664,43 +682,37 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
[SkippableFact]
[Trait("Category", "CosmosDB")]
public async Task SerializeDeserialize_WithHierarchicalPartitioning_ShouldPreserveStateAsync()
public async Task StateBag_WithHierarchicalPartitioning_ShouldPreserveStateAcrossProviderInstancesAsync()
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string TenantId = "tenant-serialize";
const string UserId = "user-serialize";
const string SessionId = "session-serialize";
using var originalStore = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, TenantId, UserId, SessionId);
using var originalStore = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, TenantId, UserId));
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Test serialization message")]);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Test serialization message")], []);
await originalStore.InvokedAsync(context);
// Act - Serialize the provider state
var serializedState = originalStore.Serialize();
// Create a new provider from the serialized state
using var cosmosClient = new CosmosClient(EmulatorEndpoint, EmulatorKey);
var serializerOptions = new JsonSerializerOptions
{
TypeInfoResolver = new DefaultJsonTypeInfoResolver()
};
using var deserializedStore = CosmosChatHistoryProvider.CreateFromSerializedState(cosmosClient, serializedState, s_testDatabaseId, HierarchicalTestContainerId, serializerOptions);
// Wait a moment for eventual consistency
await Task.Delay(100);
// Assert - The deserialized provider should have the same functionality
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var messages = await deserializedStore.InvokingAsync(invokingContext);
// Act - Create a new provider that uses a different intializer, but we will use the same session.
using var newStore = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(Guid.NewGuid().ToString()));
// Assert - The new provider should read the same messages from Cosmos DB
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = await newStore.InvokingAsync(invokingContext);
var messageList = messages.ToList();
Assert.Single(messageList);
Assert.Equal("Test serialization message", messageList[0].Text);
Assert.Equal(SessionId, deserializedStore.ConversationId);
Assert.Equal(s_testDatabaseId, deserializedStore.DatabaseId);
Assert.Equal(HierarchicalTestContainerId, deserializedStore.ContainerId);
Assert.Equal(s_testDatabaseId, newStore.DatabaseId);
Assert.Equal(HierarchicalTestContainerId, newStore.ContainerId);
}
[SkippableFact]
@@ -711,13 +723,17 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
this.SkipIfEmulatorNotAvailable();
const string SessionId = "coexist-session";
var session = CreateMockSession();
// Create simple provider using simple partitioning container and hierarchical provider using hierarchical container
using var simpleProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, SessionId);
using var hierarchicalProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId, "tenant-coexist", "user-coexist", SessionId);
// Use different stateKey values so the providers don't overwrite each other's state in the shared session
using var simpleProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId), stateKey: "simple");
using var hierarchicalProvider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, HierarchicalTestContainerId,
_ => new CosmosChatHistoryProvider.State(SessionId, "tenant-coexist", "user-coexist"), stateKey: "hierarchical");
// Add messages to both
var simpleContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Simple partitioning message")]);
var hierarchicalContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Hierarchical partitioning message")]);
var simpleContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Simple partitioning message")], []);
var hierarchicalContext = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, [new ChatMessage(ChatRole.User, "Hierarchical partitioning message")], []);
await simpleProvider.InvokedAsync(simpleContext);
await hierarchicalProvider.InvokedAsync(hierarchicalContext);
@@ -726,7 +742,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(100);
// Act & Assert
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var simpleMessages = await simpleProvider.InvokingAsync(invokingContext);
var simpleMessageList = simpleMessages.ToList();
@@ -747,9 +763,11 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string ConversationId = "max-messages-test";
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, ConversationId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(ConversationId));
// Add 10 messages
var messages = new List<ChatMessage>();
@@ -759,7 +777,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
await Task.Delay(10); // Small delay to ensure different timestamps
}
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, messages, []);
await provider.InvokedAsync(context);
// Wait for eventual consistency
@@ -767,7 +785,7 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
// Act - Set max to 5 and retrieve
provider.MaxMessagesToRetrieve = 5;
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var retrievedMessages = await provider.InvokingAsync(invokingContext);
var messageList = retrievedMessages.ToList();
@@ -786,9 +804,11 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
const string ConversationId = "max-messages-null-test";
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId, ConversationId);
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(ConversationId));
// Add 10 messages
var messages = new List<ChatMessage>();
@@ -797,14 +817,14 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
messages.Add(new ChatMessage(ChatRole.User, $"Message {i}"));
}
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, s_mockSession, messages);
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, messages, []);
await provider.InvokedAsync(context);
// Wait for eventual consistency
await Task.Delay(100);
// Act - No limit set (default null)
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, s_mockSession, []);
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var retrievedMessages = await provider.InvokingAsync(invokingContext);
var messageList = retrievedMessages.ToList();
@@ -815,4 +835,119 @@ public sealed class CosmosChatHistoryProviderTests : IAsyncLifetime, IDisposable
}
#endregion
#region Message Filter Tests
[SkippableFact]
[Trait("Category", "CosmosDB")]
public async Task InvokedAsync_DefaultFilter_ExcludesChatHistoryMessagesFromStorageAsync()
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = Guid.NewGuid().ToString();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId));
var requestMessages = new[]
{
new ChatMessage(ChatRole.User, "External message"),
new ChatMessage(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new ChatMessage(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(context);
// Wait for eventual consistency
await Task.Delay(100);
// Assert - ChatHistory message excluded, External + AIContextProvider + Response stored
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = (await provider.InvokingAsync(invokingContext)).ToList();
Assert.Equal(3, messages.Count);
Assert.Equal("External message", messages[0].Text);
Assert.Equal("From context provider", messages[1].Text);
Assert.Equal("Response", messages[2].Text);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public async Task InvokedAsync_CustomStorageInputFilter_OverridesDefaultAsync()
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = Guid.NewGuid().ToString();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId))
{
// Custom filter: only store External messages (also exclude AIContextProvider)
StorageInputMessageFilter = messages => messages.Where(m => m.GetAgentRequestMessageSourceType() == AgentRequestMessageSourceType.External)
};
var requestMessages = new[]
{
new ChatMessage(ChatRole.User, "External message"),
new ChatMessage(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new ChatMessage(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(context);
// Wait for eventual consistency
await Task.Delay(100);
// Assert - Custom filter: only External + Response stored (both ChatHistory and AIContextProvider excluded)
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = (await provider.InvokingAsync(invokingContext)).ToList();
Assert.Equal(2, messages.Count);
Assert.Equal("External message", messages[0].Text);
Assert.Equal("Response", messages[1].Text);
}
[SkippableFact]
[Trait("Category", "CosmosDB")]
public async Task InvokingAsync_RetrievalOutputFilter_FiltersRetrievedMessagesAsync()
{
// Arrange
this.SkipIfEmulatorNotAvailable();
var session = CreateMockSession();
var conversationId = Guid.NewGuid().ToString();
using var provider = new CosmosChatHistoryProvider(this._connectionString, s_testDatabaseId, TestContainerId,
_ => new CosmosChatHistoryProvider.State(conversationId))
{
// Only return User messages when retrieving
RetrievalOutputMessageFilter = messages => messages.Where(m => m.Role == ChatRole.User)
};
var requestMessages = new[]
{
new ChatMessage(ChatRole.User, "User message"),
new ChatMessage(ChatRole.System, "System message"),
};
var context = new ChatHistoryProvider.InvokedContext(s_mockAgent, session, requestMessages, [new ChatMessage(ChatRole.Assistant, "Assistant response")]);
await provider.InvokedAsync(context);
// Wait for eventual consistency
await Task.Delay(100);
// Act
var invokingContext = new ChatHistoryProvider.InvokingContext(s_mockAgent, session, []);
var messages = (await provider.InvokingAsync(invokingContext)).ToList();
// Assert - Only User messages returned (System and Assistant filtered by RetrievalOutputMessageFilter)
Assert.Single(messages);
Assert.Equal("User message", messages[0].Text);
Assert.Equal(ChatRole.User, messages[0].Role);
}
#endregion
}
@@ -15,7 +15,7 @@ public sealed class DurableAgentSessionTests
JsonElement serializedSession = session.Serialize();
// Expected format: "{\"sessionId\":\"@dafx-test-agent@<random-key>\"}"
string expectedSerializedSession = $"{{\"sessionId\":\"@dafx-{sessionId.Name}@{sessionId.Key}\"}}";
string expectedSerializedSession = $"{{\"sessionId\":\"@dafx-{sessionId.Name}@{sessionId.Key}\",\"stateBag\":{{}}}}";
Assert.Equal(expectedSerializedSession, serializedSession.ToString());
DurableAgentSession deserializedSession = DurableAgentSession.Deserialize(serializedSession);
@@ -33,11 +33,47 @@ public sealed class DurableAgentSessionTests
string serializedSession = JsonSerializer.Serialize(session, typeof(DurableAgentSession));
// Expected format: "{\"sessionId\":\"@dafx-test-agent@<random-key>\"}"
string expectedSerializedSession = $"{{\"sessionId\":\"@dafx-{sessionId.Name}@{sessionId.Key}\"}}";
string expectedSerializedSession = $"{{\"sessionId\":\"@dafx-{sessionId.Name}@{sessionId.Key}\",\"stateBag\":{{}}}}";
Assert.Equal(expectedSerializedSession, serializedSession);
DurableAgentSession? deserializedSession = JsonSerializer.Deserialize<DurableAgentSession>(serializedSession);
Assert.NotNull(deserializedSession);
Assert.Equal(sessionId, deserializedSession.SessionId);
}
[Fact]
public void BuiltInSerialization_RoundTrip_PreservesStateBag()
{
// Arrange
AgentSessionId sessionId = AgentSessionId.WithRandomKey("test-agent");
DurableAgentSession session = new(sessionId);
session.StateBag.SetValue("durableKey", "durableValue");
// Act
JsonElement serializedSession = session.Serialize();
DurableAgentSession deserializedSession = DurableAgentSession.Deserialize(serializedSession);
// Assert
Assert.Equal(sessionId, deserializedSession.SessionId);
Assert.True(deserializedSession.StateBag.TryGetValue<string>("durableKey", out var value));
Assert.Equal("durableValue", value);
}
[Fact]
public void STJSerialization_RoundTrip_PreservesStateBag()
{
// Arrange
AgentSessionId sessionId = AgentSessionId.WithRandomKey("test-agent");
DurableAgentSession session = new(sessionId);
session.StateBag.SetValue("stjKey", "stjValue");
// Act
string serializedSession = JsonSerializer.Serialize(session, typeof(DurableAgentSession));
DurableAgentSession? deserializedSession = JsonSerializer.Deserialize<DurableAgentSession>(serializedSession);
// Assert
Assert.NotNull(deserializedSession);
Assert.True(deserializedSession.StateBag.TryGetValue<string>("stjKey", out var value));
Assert.Equal("stjValue", value);
}
}
@@ -7,6 +7,7 @@ using System.Linq;
using System.Net.Http;
using System.Runtime.CompilerServices;
using System.Text.Json;
using System.Text.Json.Serialization;
using System.Threading;
using System.Threading.Tasks;
using FluentAssertions;
@@ -281,10 +282,10 @@ internal sealed class FakeChatClientAgent : AIAgent
public override string? Description => "A fake agent for testing";
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession());
new(new FakeAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<FakeAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
=> throw new NotImplementedException();
@@ -326,15 +327,14 @@ internal sealed class FakeChatClientAgent : AIAgent
}
}
private sealed class FakeInMemoryAgentSession : InMemoryAgentSession
private sealed class FakeAgentSession : AgentSession
{
public FakeInMemoryAgentSession()
: base()
public FakeAgentSession()
{
}
public FakeInMemoryAgentSession(JsonElement serializedSession, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSession, jsonSerializerOptions)
[JsonConstructor]
public FakeAgentSession(AgentSessionStateBag stateBag) : base(stateBag)
{
}
}
@@ -348,19 +348,19 @@ internal sealed class FakeMultiMessageAgent : AIAgent
public override string? Description => "A fake agent that sends multiple messages for testing";
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession());
new(new FakeAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<FakeAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
if (session is not FakeInMemoryAgentSession fakeSession)
if (session is not FakeAgentSession fakeSession)
{
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(fakeSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(fakeSession, jsonSerializerOptions));
}
protected override async Task<AgentResponse> RunCoreAsync(
@@ -427,19 +427,16 @@ internal sealed class FakeMultiMessageAgent : AIAgent
}
}
private sealed class FakeInMemoryAgentSession : InMemoryAgentSession
private sealed class FakeAgentSession : AgentSession
{
public FakeInMemoryAgentSession()
: base()
public FakeAgentSession()
{
}
public FakeInMemoryAgentSession(JsonElement serializedSession, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSession, jsonSerializerOptions)
[JsonConstructor]
public FakeAgentSession(AgentSessionStateBag stateBag) : base(stateBag)
{
}
internal new JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> base.Serialize(jsonSerializerOptions);
}
public override object? GetService(Type serviceType, object? serviceKey = null) => null;
@@ -9,6 +9,7 @@ using System.Net.ServerSentEvents;
using System.Runtime.CompilerServices;
using System.Text;
using System.Text.Json;
using System.Text.Json.Serialization;
using System.Threading;
using System.Threading.Tasks;
using FluentAssertions;
@@ -335,35 +336,31 @@ internal sealed class FakeForwardedPropsAgent : AIAgent
}
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession());
new(new FakeAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<FakeAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
if (session is not FakeInMemoryAgentSession fakeSession)
if (session is not FakeAgentSession fakeSession)
{
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(fakeSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(fakeSession, jsonSerializerOptions));
}
private sealed class FakeInMemoryAgentSession : InMemoryAgentSession
private sealed class FakeAgentSession : AgentSession
{
public FakeInMemoryAgentSession()
: base()
public FakeAgentSession()
{
}
public FakeInMemoryAgentSession(JsonElement serializedSession, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSession, jsonSerializerOptions)
[JsonConstructor]
public FakeAgentSession(AgentSessionStateBag stateBag) : base(stateBag)
{
}
internal new JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> base.Serialize(jsonSerializerOptions);
}
public override object? GetService(Type serviceType, object? serviceKey = null) => null;
@@ -7,6 +7,7 @@ using System.Linq;
using System.Net.Http;
using System.Runtime.CompilerServices;
using System.Text.Json;
using System.Text.Json.Serialization;
using System.Threading;
using System.Threading.Tasks;
using FluentAssertions;
@@ -418,35 +419,31 @@ internal sealed class FakeStateAgent : AIAgent
}
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession());
new(new FakeAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new FakeInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<FakeAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
if (session is not FakeInMemoryAgentSession fakeSession)
if (session is not FakeAgentSession fakeSession)
{
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(fakeSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(fakeSession, jsonSerializerOptions));
}
private sealed class FakeInMemoryAgentSession : InMemoryAgentSession
private sealed class FakeAgentSession : AgentSession
{
public FakeInMemoryAgentSession()
: base()
public FakeAgentSession()
{
}
public FakeInMemoryAgentSession(JsonElement serializedSession, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSession, jsonSerializerOptions)
[JsonConstructor]
public FakeAgentSession(AgentSessionStateBag stateBag) : base(stateBag)
{
}
internal new JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> base.Serialize(jsonSerializerOptions);
}
public override object? GetService(Type serviceType, object? serviceKey = null) => null;
@@ -6,6 +6,7 @@ using System.IO;
using System.Linq;
using System.Text;
using System.Text.Json;
using System.Text.Json.Serialization;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Agents.AI.Hosting.AGUI.AspNetCore.Shared;
@@ -426,19 +427,19 @@ public sealed class AGUIEndpointRouteBuilderExtensionsTests
public override string? Description => "Agent that produces multiple text chunks";
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new TestInMemoryAgentSession());
new(new TestAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new TestInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<TestAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
if (session is not TestInMemoryAgentSession testSession)
if (session is not TestAgentSession testSession)
{
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(testSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(testSession, jsonSerializerOptions));
}
protected override Task<AgentResponse> RunCoreAsync(IEnumerable<ChatMessage> messages, AgentSession? session = null, AgentRunOptions? options = null, CancellationToken cancellationToken = default)
@@ -506,20 +507,16 @@ public sealed class AGUIEndpointRouteBuilderExtensionsTests
};
}
private sealed class TestInMemoryAgentSession : InMemoryAgentSession
private sealed class TestAgentSession : AgentSession
{
public TestInMemoryAgentSession()
: base()
public TestAgentSession()
{
}
public TestInMemoryAgentSession(JsonElement serializedSessionState, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSessionState, jsonSerializerOptions, null)
[JsonConstructor]
public TestAgentSession(AgentSessionStateBag stateBag) : base(stateBag)
{
}
internal new JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
=> base.Serialize(jsonSerializerOptions);
}
private sealed class TestAgent : AIAgent
@@ -529,19 +526,19 @@ public sealed class AGUIEndpointRouteBuilderExtensionsTests
public override string? Description => "Test agent";
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new TestInMemoryAgentSession());
new(new TestAgentSession());
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default) =>
new(new TestInMemoryAgentSession(serializedState, jsonSerializerOptions));
new(serializedState.Deserialize<TestAgentSession>(jsonSerializerOptions)!);
protected override ValueTask<JsonElement> SerializeSessionCoreAsync(AgentSession session, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
if (session is not TestInMemoryAgentSession testSession)
if (session is not TestAgentSession testSession)
{
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(testSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(testSession, jsonSerializerOptions));
}
protected override Task<AgentResponse> RunCoreAsync(IEnumerable<ChatMessage> messages, AgentSession? session = null, AgentRunOptions? options = null, CancellationToken cancellationToken = default)
@@ -1,6 +1,8 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Net.Http;
using System.Net.Http.Headers;
using System.Threading;
@@ -19,7 +21,6 @@ public sealed class Mem0ProviderTests : IDisposable
private const string SkipReason = "Requires a Mem0 service configured"; // Set to null to enable.
private static readonly AIAgent s_mockAgent = new Moq.Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Moq.Mock<AgentSession>().Object;
private readonly HttpClient _httpClient;
@@ -49,21 +50,22 @@ public sealed class Mem0ProviderTests : IDisposable
var question = new ChatMessage(ChatRole.User, "What is my name?");
var input = new ChatMessage(ChatRole.User, "Hello, my name is Caoimhe.");
var storageScope = new Mem0ProviderScope { ThreadId = "it-thread-1", UserId = "it-user-1" };
var sut = new Mem0Provider(this._httpClient, storageScope);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
await sut.ClearStoredMemoriesAsync();
var ctxBefore = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
Assert.DoesNotContain("Caoimhe", ctxBefore.Messages?[0].Text ?? string.Empty);
await sut.ClearStoredMemoriesAsync(mockSession);
var ctxBefore = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { question } }));
Assert.DoesNotContain("Caoimhe", ctxBefore.Messages?.LastOrDefault()?.Text ?? string.Empty);
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [input]));
var ctxAfterAdding = await GetContextWithRetryAsync(sut, question);
await sut.ClearStoredMemoriesAsync();
var ctxAfterClearing = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, [input], []));
var ctxAfterAdding = await GetContextWithRetryAsync(sut, mockSession, question);
await sut.ClearStoredMemoriesAsync(mockSession);
var ctxAfterClearing = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { question } }));
// Assert
Assert.Contains("Caoimhe", ctxAfterAdding.Messages?[0].Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterClearing.Messages?[0].Text ?? string.Empty);
Assert.Contains("Caoimhe", ctxAfterAdding.Messages?.LastOrDefault()?.Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterClearing.Messages?.LastOrDefault()?.Text ?? string.Empty);
}
[Fact(Skip = SkipReason)]
@@ -73,21 +75,22 @@ public sealed class Mem0ProviderTests : IDisposable
var question = new ChatMessage(ChatRole.User, "What is your name?");
var assistantIntro = new ChatMessage(ChatRole.Assistant, "Hello, I'm a friendly assistant and my name is Caoimhe.");
var storageScope = new Mem0ProviderScope { AgentId = "it-agent-1" };
var sut = new Mem0Provider(this._httpClient, storageScope);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
await sut.ClearStoredMemoriesAsync();
var ctxBefore = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
Assert.DoesNotContain("Caoimhe", ctxBefore.Messages?[0].Text ?? string.Empty);
await sut.ClearStoredMemoriesAsync(mockSession);
var ctxBefore = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { question } }));
Assert.DoesNotContain("Caoimhe", ctxBefore.Messages?.LastOrDefault()?.Text ?? string.Empty);
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [assistantIntro]));
var ctxAfterAdding = await GetContextWithRetryAsync(sut, question);
await sut.ClearStoredMemoriesAsync();
var ctxAfterClearing = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, [assistantIntro], []));
var ctxAfterAdding = await GetContextWithRetryAsync(sut, mockSession, question);
await sut.ClearStoredMemoriesAsync(mockSession);
var ctxAfterClearing = await sut.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { question } }));
// Assert
Assert.Contains("Caoimhe", ctxAfterAdding.Messages?[0].Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterClearing.Messages?[0].Text ?? string.Empty);
Assert.Contains("Caoimhe", ctxAfterAdding.Messages?.LastOrDefault()?.Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterClearing.Messages?.LastOrDefault()?.Text ?? string.Empty);
}
[Fact(Skip = SkipReason)]
@@ -96,38 +99,42 @@ public sealed class Mem0ProviderTests : IDisposable
// Arrange
var question = new ChatMessage(ChatRole.User, "What is your name?");
var assistantIntro = new ChatMessage(ChatRole.Assistant, "I'm an AI tutor and my name is Caoimhe.");
var sut1 = new Mem0Provider(this._httpClient, new Mem0ProviderScope { AgentId = "it-agent-a" });
var sut2 = new Mem0Provider(this._httpClient, new Mem0ProviderScope { AgentId = "it-agent-b" });
var storageScope1 = new Mem0ProviderScope { AgentId = "it-agent-a" };
var storageScope2 = new Mem0ProviderScope { AgentId = "it-agent-b" };
var mockSession1 = new TestAgentSession();
var mockSession2 = new TestAgentSession();
var sut1 = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope1));
var sut2 = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope2));
await sut1.ClearStoredMemoriesAsync();
await sut2.ClearStoredMemoriesAsync();
await sut1.ClearStoredMemoriesAsync(mockSession1);
await sut2.ClearStoredMemoriesAsync(mockSession2);
var ctxBefore1 = await sut1.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
var ctxBefore2 = await sut2.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]));
Assert.DoesNotContain("Caoimhe", ctxBefore1.Messages?[0].Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxBefore2.Messages?[0].Text ?? string.Empty);
var ctxBefore1 = await sut1.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession1, new AIContext { Messages = new List<ChatMessage> { question } }));
var ctxBefore2 = await sut2.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, mockSession2, new AIContext { Messages = new List<ChatMessage> { question } }));
Assert.DoesNotContain("Caoimhe", ctxBefore1.Messages?.LastOrDefault()?.Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxBefore2.Messages?.LastOrDefault()?.Text ?? string.Empty);
// Act
await sut1.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [assistantIntro]));
var ctxAfterAdding1 = await GetContextWithRetryAsync(sut1, question);
var ctxAfterAdding2 = await GetContextWithRetryAsync(sut2, question);
await sut1.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession1, [assistantIntro], []));
var ctxAfterAdding1 = await GetContextWithRetryAsync(sut1, mockSession1, question);
var ctxAfterAdding2 = await GetContextWithRetryAsync(sut2, mockSession2, question);
// Assert
Assert.Contains("Caoimhe", ctxAfterAdding1.Messages?[0].Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterAdding2.Messages?[0].Text ?? string.Empty);
Assert.Contains("Caoimhe", ctxAfterAdding1.Messages?.LastOrDefault()?.Text ?? string.Empty);
Assert.DoesNotContain("Caoimhe", ctxAfterAdding2.Messages?.LastOrDefault()?.Text ?? string.Empty);
// Cleanup
await sut1.ClearStoredMemoriesAsync();
await sut2.ClearStoredMemoriesAsync();
await sut1.ClearStoredMemoriesAsync(mockSession1);
await sut2.ClearStoredMemoriesAsync(mockSession2);
}
private static async Task<AIContext> GetContextWithRetryAsync(Mem0Provider provider, ChatMessage question, int attempts = 5, int delayMs = 1000)
private static async Task<AIContext> GetContextWithRetryAsync(Mem0Provider provider, AgentSession session, ChatMessage question, int attempts = 5, int delayMs = 1000)
{
AIContext? ctx = null;
for (int i = 0; i < attempts; i++)
{
ctx = await provider.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [question]), CancellationToken.None);
var text = ctx.Messages?[0].Text;
ctx = await provider.InvokingAsync(new AIContextProvider.InvokingContext(s_mockAgent, session, new AIContext { Messages = new List<ChatMessage> { question } }), CancellationToken.None);
var text = ctx.Messages?.LastOrDefault()?.Text;
if (!string.IsNullOrEmpty(text) && text.IndexOf("Caoimhe", StringComparison.OrdinalIgnoreCase) >= 0)
{
break;
@@ -141,4 +148,12 @@ public sealed class Mem0ProviderTests : IDisposable
{
this._httpClient.Dispose();
}
private sealed class TestAgentSession : AgentSession
{
public TestAgentSession()
{
this.StateBag = new AgentSessionStateBag();
}
}
}
@@ -19,7 +19,6 @@ namespace Microsoft.Agents.AI.Mem0.UnitTests;
public sealed class Mem0ProviderTests : IDisposable
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
private readonly Mock<ILogger<Mem0Provider>> _loggerMock;
private readonly Mock<ILoggerFactory> _loggerFactoryMock;
@@ -55,35 +54,39 @@ public sealed class Mem0ProviderTests : IDisposable
using HttpClient client = new();
// Act & Assert
var ex = Assert.Throws<ArgumentException>(() => new Mem0Provider(client, new Mem0ProviderScope() { ThreadId = "tid" }));
var ex = Assert.Throws<ArgumentException>(() => new Mem0Provider(client, _ => new Mem0Provider.State(new Mem0ProviderScope { ThreadId = "tid" })));
Assert.StartsWith("The HttpClient BaseAddress must be set for Mem0 operations.", ex.Message);
}
[Fact]
public void Constructor_Throws_WhenNoStorageScopeValueIsSet()
public void Constructor_Throws_WhenStateInitializerIsNull()
{
// Act & Assert
var ex = Assert.Throws<ArgumentException>(() => new Mem0Provider(this._httpClient, new Mem0ProviderScope()));
Assert.StartsWith("At least one of ApplicationId, AgentId, ThreadId, or UserId must be provided for the storage scope.", ex.Message);
var ex = Assert.Throws<ArgumentNullException>(() => new Mem0Provider(this._httpClient, null!));
Assert.Contains("stateInitializer", ex.Message);
}
[Fact]
public void Constructor_Throws_WhenNoSearchScopeValueIsSet()
public void StateKey_ReturnsDefaultKey_WhenNoOptionsProvided()
{
// Act & Assert
var ex = Assert.Throws<ArgumentException>(() => new Mem0Provider(this._httpClient, new Mem0ProviderScope() { ThreadId = "tid" }, new Mem0ProviderScope()));
Assert.StartsWith("At least one of ApplicationId, AgentId, ThreadId, or UserId must be provided for the search scope.", ex.Message);
// Arrange & Act
var provider = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(new Mem0ProviderScope { ThreadId = "tid" }));
// Assert
Assert.Equal("Mem0Provider", provider.StateKey);
}
[Fact]
public void DeserializingConstructor_Throws_WithEmptyJsonElement()
public void StateKey_ReturnsCustomKey_WhenSetViaOptions()
{
// Arrange
var jsonElement = JsonSerializer.SerializeToElement(new object(), Mem0JsonUtilities.DefaultOptions);
// Arrange & Act
var provider = new Mem0Provider(
this._httpClient,
_ => new Mem0Provider.State(new Mem0ProviderScope { ThreadId = "tid" }),
new Mem0ProviderOptions { StateKey = "custom-key" });
// Act & Assert
var ex = Assert.Throws<InvalidOperationException>(() => new Mem0Provider(this._httpClient, jsonElement));
Assert.StartsWith("The Mem0Provider state did not contain the required scope properties.", ex.Message);
// Assert
Assert.Equal("custom-key", provider.StateKey);
}
[Fact]
@@ -98,8 +101,9 @@ public sealed class Mem0ProviderTests : IDisposable
ThreadId = "session",
UserId = "user"
};
var sut = new Mem0Provider(this._httpClient, storageScope, options: new() { EnableSensitiveTelemetryData = true }, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "What is my name?")]);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: new() { EnableSensitiveTelemetryData = true }, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "What is my name?") } });
// Act
var aiContext = await sut.InvokingAsync(invokingContext);
@@ -114,9 +118,13 @@ public sealed class Mem0ProviderTests : IDisposable
Assert.Equal("What is my name?", doc.RootElement.GetProperty("query").GetString());
Assert.NotNull(aiContext.Messages);
var contextMessage = Assert.Single(aiContext.Messages);
var messages = aiContext.Messages.ToList();
Assert.Equal(2, messages.Count);
Assert.Equal(AgentRequestMessageSourceType.External, messages[0].GetAgentRequestMessageSourceType());
var contextMessage = messages[1];
Assert.Equal(ChatRole.User, contextMessage.Role);
Assert.Contains("Name is Caoimhe", contextMessage.Text);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, contextMessage.GetAgentRequestMessageSourceType());
this._loggerMock.Verify(
l => l.Log(
@@ -162,9 +170,10 @@ public sealed class Mem0ProviderTests : IDisposable
UserId = "user"
};
var options = new Mem0ProviderOptions { EnableSensitiveTelemetryData = enableSensitiveTelemetryData };
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, storageScope, options: options, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Who am I?")]);
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: options, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Who am I?") } });
// Act
await sut.InvokingAsync(invokingContext, CancellationToken.None);
@@ -204,7 +213,8 @@ public sealed class Mem0ProviderTests : IDisposable
this._handler.EnqueueEmptyOk(); // For second CreateMemory
this._handler.EnqueueEmptyOk(); // For third CreateMemory
var storageScope = new Mem0ProviderScope { ApplicationId = "a", AgentId = "b", ThreadId = "c", UserId = "d" };
var sut = new Mem0Provider(this._httpClient, storageScope);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
var requestMessages = new List<ChatMessage>
{
@@ -218,7 +228,7 @@ public sealed class Mem0ProviderTests : IDisposable
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages) { ResponseMessages = responseMessages });
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, responseMessages));
// Assert
var memoryPosts = this._handler.Requests.Where(r => r.RequestMessage.RequestUri!.AbsolutePath == "/v1/memories/" && r.RequestMessage.Method == HttpMethod.Post).ToList();
@@ -235,7 +245,8 @@ public sealed class Mem0ProviderTests : IDisposable
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "a", AgentId = "b", ThreadId = "c", UserId = "d" };
var sut = new Mem0Provider(this._httpClient, storageScope);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
var requestMessages = new List<ChatMessage>
{
@@ -245,7 +256,7 @@ public sealed class Mem0ProviderTests : IDisposable
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages) { ResponseMessages = null, InvokeException = new InvalidOperationException("Request Failed") });
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, new InvalidOperationException("Request Failed")));
// Assert
Assert.Empty(this._handler.Requests);
@@ -256,7 +267,8 @@ public sealed class Mem0ProviderTests : IDisposable
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "a", AgentId = "b", ThreadId = "c", UserId = "d" };
var sut = new Mem0Provider(this._httpClient, storageScope, loggerFactory: this._loggerFactoryMock.Object);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), loggerFactory: this._loggerFactoryMock.Object);
this._handler.EnqueueEmptyInternalServerError();
var requestMessages = new List<ChatMessage>
@@ -271,7 +283,7 @@ public sealed class Mem0ProviderTests : IDisposable
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages) { ResponseMessages = responseMessages });
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, responseMessages));
// Assert
this._loggerMock.Verify(
@@ -310,7 +322,8 @@ public sealed class Mem0ProviderTests : IDisposable
};
var options = new Mem0ProviderOptions { EnableSensitiveTelemetryData = enableSensitiveTelemetryData };
var sut = new Mem0Provider(this._httpClient, storageScope, options: options, loggerFactory: this._loggerFactoryMock.Object);
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: options, loggerFactory: this._loggerFactoryMock.Object);
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "User text")
@@ -321,7 +334,7 @@ public sealed class Mem0ProviderTests : IDisposable
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, requestMessages) { ResponseMessages = responseMessages });
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, responseMessages));
// Assert
Assert.Equal(expectedLogCount, this._loggerMock.Invocations.Count);
@@ -343,73 +356,33 @@ public sealed class Mem0ProviderTests : IDisposable
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "app", AgentId = "agent", ThreadId = "session", UserId = "user" };
var sut = new Mem0Provider(this._httpClient, storageScope);
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
this._handler.EnqueueEmptyOk(); // for DELETE
var mockSession = new TestAgentSession();
// Act
await sut.ClearStoredMemoriesAsync();
await sut.ClearStoredMemoriesAsync(mockSession);
// Assert
var delete = Assert.Single(this._handler.Requests, r => r.RequestMessage.Method == HttpMethod.Delete);
Assert.Equal("https://localhost/v1/memories/?app_id=app&agent_id=agent&run_id=session&user_id=user", delete.RequestMessage.RequestUri!.AbsoluteUri);
}
[Fact]
public void Serialize_RoundTripsScopes()
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "app", AgentId = "agent", ThreadId = "session", UserId = "user" };
var sut = new Mem0Provider(this._httpClient, storageScope, options: new() { ContextPrompt = "Custom:" }, loggerFactory: this._loggerFactoryMock.Object);
// Act
var stateElement = sut.Serialize();
using JsonDocument doc = JsonDocument.Parse(stateElement.GetRawText());
var storageScopeElement = doc.RootElement.GetProperty("storageScope");
Assert.Equal("app", storageScopeElement.GetProperty("applicationId").GetString());
Assert.Equal("agent", storageScopeElement.GetProperty("agentId").GetString());
Assert.Equal("session", storageScopeElement.GetProperty("threadId").GetString());
Assert.Equal("user", storageScopeElement.GetProperty("userId").GetString());
var sut2 = new Mem0Provider(this._httpClient, stateElement);
var stateElement2 = sut2.Serialize();
// Assert
using JsonDocument doc2 = JsonDocument.Parse(stateElement2.GetRawText());
var storageScopeElement2 = doc2.RootElement.GetProperty("storageScope");
Assert.Equal("app", storageScopeElement2.GetProperty("applicationId").GetString());
Assert.Equal("agent", storageScopeElement2.GetProperty("agentId").GetString());
Assert.Equal("session", storageScopeElement2.GetProperty("threadId").GetString());
Assert.Equal("user", storageScopeElement2.GetProperty("userId").GetString());
}
[Fact]
public void Serialize_DoesNotIncludeDefaultContextPrompt()
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "app" };
var sut = new Mem0Provider(this._httpClient, storageScope);
// Act
var stateElement = sut.Serialize();
// Assert
using JsonDocument doc = JsonDocument.Parse(stateElement.GetRawText());
Assert.False(doc.RootElement.TryGetProperty("contextPrompt", out _));
}
[Fact]
public async Task InvokingAsync_ShouldNotThrow_WhenSearchFailsAsync()
{
// Arrange
var storageScope = new Mem0ProviderScope { ApplicationId = "app" };
var provider = new Mem0Provider(this._httpClient, storageScope, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var mockSession = new TestAgentSession();
var provider = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages);
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages);
Assert.Null(aiContext.Tools);
this._loggerMock.Verify(
l => l.Log(
@@ -421,6 +394,159 @@ public sealed class Mem0ProviderTests : IDisposable
Times.Once);
}
[Fact]
public async Task StateInitializer_IsCalledOnceAndStoredInStateBagAsync()
{
// Arrange
this._handler.EnqueueJsonResponse("[]");
this._handler.EnqueueJsonResponse("[]");
var storageScope = new Mem0ProviderScope { ApplicationId = "app" };
var mockSession = new TestAgentSession();
int initializerCallCount = 0;
var sut = new Mem0Provider(this._httpClient, _ =>
{
initializerCallCount++;
return new Mem0Provider.State(storageScope);
});
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
await sut.InvokingAsync(invokingContext, CancellationToken.None);
await sut.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Equal(1, initializerCallCount);
}
[Fact]
public async Task StateKey_CanBeConfiguredViaOptionsAsync()
{
// Arrange
this._handler.EnqueueJsonResponse("[]");
var storageScope = new Mem0ProviderScope { ApplicationId = "app" };
var mockSession = new TestAgentSession();
const string CustomKey = "MyCustomKey";
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: new() { StateKey = CustomKey });
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
await sut.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.True(mockSession.StateBag.TryGetValue<Mem0Provider.State>(CustomKey, out var state, Mem0JsonUtilities.DefaultOptions));
Assert.NotNull(state);
}
[Fact]
public async Task InvokingAsync_DefaultFilter_ExcludesNonExternalMessagesFromSearchAsync()
{
// Arrange
this._handler.EnqueueJsonResponse("[]"); // Empty search results
var storageScope = new Mem0ProviderScope { ApplicationId = "app", AgentId = "agent", ThreadId = "session", UserId = "user" };
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = requestMessages });
// Act
await sut.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Search query should only contain the External message
var searchRequest = Assert.Single(this._handler.Requests, r => r.RequestMessage.Method == HttpMethod.Post);
using JsonDocument doc = JsonDocument.Parse(searchRequest.RequestBody);
Assert.Equal("External message", doc.RootElement.GetProperty("query").GetString());
}
[Fact]
public async Task InvokingAsync_CustomSearchInputFilter_OverridesDefaultAsync()
{
// Arrange
this._handler.EnqueueJsonResponse("[]"); // Empty search results
var storageScope = new Mem0ProviderScope { ApplicationId = "app", AgentId = "agent", ThreadId = "session", UserId = "user" };
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: new Mem0ProviderOptions
{
SearchInputMessageFilter = messages => messages // No filtering
});
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, mockSession, new AIContext { Messages = requestMessages });
// Act
await sut.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Search query should contain all messages (custom identity filter)
var searchRequest = Assert.Single(this._handler.Requests, r => r.RequestMessage.Method == HttpMethod.Post);
using JsonDocument doc = JsonDocument.Parse(searchRequest.RequestBody);
var queryText = doc.RootElement.GetProperty("query").GetString();
Assert.Contains("External message", queryText);
Assert.Contains("From history", queryText);
}
[Fact]
public async Task InvokedAsync_DefaultFilter_ExcludesNonExternalMessagesFromStorageAsync()
{
// Arrange
this._handler.EnqueueEmptyOk(); // For the one message that should be stored
var storageScope = new Mem0ProviderScope { ApplicationId = "a", AgentId = "b", ThreadId = "c", UserId = "d" };
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope));
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, []));
// Assert - Only the External message should be persisted
var memoryPosts = this._handler.Requests.Where(r => r.RequestMessage.RequestUri!.AbsolutePath == "/v1/memories/" && r.RequestMessage.Method == HttpMethod.Post).ToList();
Assert.Single(memoryPosts);
Assert.Contains("External message", memoryPosts[0].RequestBody);
Assert.DoesNotContain(memoryPosts, r => ContainsOrdinal(r.RequestBody, "From history"));
}
[Fact]
public async Task InvokedAsync_CustomStorageInputFilter_OverridesDefaultAsync()
{
// Arrange
this._handler.EnqueueEmptyOk(); // For first CreateMemory
this._handler.EnqueueEmptyOk(); // For second CreateMemory
var storageScope = new Mem0ProviderScope { ApplicationId = "a", AgentId = "b", ThreadId = "c", UserId = "d" };
var mockSession = new TestAgentSession();
var sut = new Mem0Provider(this._httpClient, _ => new Mem0Provider.State(storageScope), options: new Mem0ProviderOptions
{
StorageInputMessageFilter = messages => messages // No filtering - store everything
});
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
// Act
await sut.InvokedAsync(new AIContextProvider.InvokedContext(s_mockAgent, mockSession, requestMessages, []));
// Assert - Both messages should be persisted (identity filter overrides default)
var memoryPosts = this._handler.Requests.Where(r => r.RequestMessage.RequestUri!.AbsolutePath == "/v1/memories/" && r.RequestMessage.Method == HttpMethod.Post).ToList();
Assert.Equal(2, memoryPosts.Count);
}
private static bool ContainsOrdinal(string source, string value) => source.IndexOf(value, StringComparison.Ordinal) >= 0;
public void Dispose()
@@ -465,4 +591,12 @@ public sealed class Mem0ProviderTests : IDisposable
public void EnqueueEmptyInternalServerError() => this._responses.Enqueue(new HttpResponseMessage(System.Net.HttpStatusCode.InternalServerError));
}
private sealed class TestAgentSession : AgentSession
{
public TestAgentSession()
{
this.StateBag = new AgentSessionStateBag();
}
}
}
@@ -1,8 +1,6 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Collections.Generic;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
@@ -23,8 +21,8 @@ public class ChatClientAgentOptionsTests
Assert.Null(options.Name);
Assert.Null(options.Description);
Assert.Null(options.ChatOptions);
Assert.Null(options.ChatHistoryProviderFactory);
Assert.Null(options.AIContextProviderFactory);
Assert.Null(options.ChatHistoryProvider);
Assert.Null(options.AIContextProviders);
}
[Fact]
@@ -36,8 +34,8 @@ public class ChatClientAgentOptionsTests
// Assert
Assert.Null(options.Name);
Assert.Null(options.Description);
Assert.Null(options.AIContextProviderFactory);
Assert.Null(options.ChatHistoryProviderFactory);
Assert.Null(options.AIContextProviders);
Assert.Null(options.ChatHistoryProvider);
Assert.NotNull(options.ChatOptions);
Assert.Null(options.ChatOptions.Instructions);
Assert.Null(options.ChatOptions.Tools);
@@ -117,11 +115,8 @@ public class ChatClientAgentOptionsTests
const string Description = "Test description";
var tools = new List<AITool> { AIFunctionFactory.Create(() => "test") };
static ValueTask<ChatHistoryProvider> ChatHistoryProviderFactoryAsync(
ChatClientAgentOptions.ChatHistoryProviderFactoryContext ctx, CancellationToken ct) => new(new Mock<ChatHistoryProvider>().Object);
static ValueTask<AIContextProvider> AIContextProviderFactoryAsync(
ChatClientAgentOptions.AIContextProviderFactoryContext ctx, CancellationToken ct) => new(new Mock<AIContextProvider>().Object);
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>().Object;
var mockAIContextProvider = new Mock<AIContextProvider>().Object;
var original = new ChatClientAgentOptions()
{
@@ -129,8 +124,8 @@ public class ChatClientAgentOptionsTests
Description = Description,
ChatOptions = new() { Tools = tools },
Id = "test-id",
ChatHistoryProviderFactory = ChatHistoryProviderFactoryAsync,
AIContextProviderFactory = AIContextProviderFactoryAsync
ChatHistoryProvider = mockChatHistoryProvider,
AIContextProviders = [mockAIContextProvider]
};
// Act
@@ -141,8 +136,8 @@ public class ChatClientAgentOptionsTests
Assert.Equal(original.Id, clone.Id);
Assert.Equal(original.Name, clone.Name);
Assert.Equal(original.Description, clone.Description);
Assert.Same(original.ChatHistoryProviderFactory, clone.ChatHistoryProviderFactory);
Assert.Same(original.AIContextProviderFactory, clone.AIContextProviderFactory);
Assert.Same(original.ChatHistoryProvider, clone.ChatHistoryProvider);
Assert.Equal(original.AIContextProviders, clone.AIContextProviders);
// ChatOptions should be cloned, not the same reference
Assert.NotSame(original.ChatOptions, clone.ChatOptions);
@@ -154,11 +149,16 @@ public class ChatClientAgentOptionsTests
public void Clone_WithoutProvidingChatOptions_ClonesCorrectly()
{
// Arrange
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>().Object;
var mockAIContextProvider = new Mock<AIContextProvider>().Object;
var original = new ChatClientAgentOptions
{
Id = "test-id",
Name = "Test name",
Description = "Test description"
Description = "Test description",
ChatHistoryProvider = mockChatHistoryProvider,
AIContextProviders = [mockAIContextProvider]
};
// Act
@@ -170,8 +170,8 @@ public class ChatClientAgentOptionsTests
Assert.Equal(original.Name, clone.Name);
Assert.Equal(original.Description, clone.Description);
Assert.Null(original.ChatOptions);
Assert.Null(clone.ChatHistoryProviderFactory);
Assert.Null(clone.AIContextProviderFactory);
Assert.Same(original.ChatHistoryProvider, clone.ChatHistoryProvider);
Assert.Equal(original.AIContextProviders, clone.AIContextProviders);
}
private static void AssertSameTools(IList<AITool>? expected, IList<AITool>? actual)
@@ -1,12 +1,9 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
#pragma warning disable CA1861 // Avoid constant arrays as arguments
@@ -24,7 +21,6 @@ public class ChatClientAgentSessionTests
// Assert
Assert.Null(session.ConversationId);
Assert.Null(session.ChatHistoryProvider);
}
[Fact]
@@ -39,53 +35,6 @@ public class ChatClientAgentSessionTests
// Assert
Assert.Equal(ConversationId, session.ConversationId);
Assert.Null(session.ChatHistoryProvider);
}
[Fact]
public void SetChatHistoryProviderRoundtrips()
{
// Arrange
var session = new ChatClientAgentSession();
var chatHistoryProvider = new InMemoryChatHistoryProvider();
// Act
session.ChatHistoryProvider = chatHistoryProvider;
// Assert
Assert.Same(chatHistoryProvider, session.ChatHistoryProvider);
Assert.Null(session.ConversationId);
}
[Fact]
public void SetConversationIdThrowsWhenChatHistoryProviderIsSet()
{
// Arrange
var session = new ChatClientAgentSession
{
ChatHistoryProvider = new InMemoryChatHistoryProvider()
};
// Act & Assert
var exception = Assert.Throws<InvalidOperationException>(() => session.ConversationId = "new-session-id");
Assert.Equal("Only the ConversationId or ChatHistoryProvider may be set, but not both and switching from one to another is not supported.", exception.Message);
Assert.NotNull(session.ChatHistoryProvider);
}
[Fact]
public void SetChatHistoryProviderThrowsWhenConversationIdIsSet()
{
// Arrange
var session = new ChatClientAgentSession
{
ConversationId = "existing-session-id"
};
var provider = new InMemoryChatHistoryProvider();
// Act & Assert
var exception = Assert.Throws<InvalidOperationException>(() => session.ChatHistoryProvider = provider);
Assert.Equal("Only the ConversationId or ChatHistoryProvider may be set, but not both and switching from one to another is not supported.", exception.Message);
Assert.NotNull(session.ConversationId);
}
#endregion Constructor and Property Tests
@@ -93,29 +42,33 @@ public class ChatClientAgentSessionTests
#region Deserialize Tests
[Fact]
public async Task VerifyDeserializeWithMessagesAsync()
public void VerifyDeserializeWithMessages()
{
// Arrange
var json = JsonSerializer.Deserialize("""
{
"chatHistoryProviderState": { "messages": [{"authorName": "testAuthor"}] }
"stateBag": {
"InMemoryChatHistoryProvider": {
"messages": [{"authorName": "testAuthor"}]
}
}
}
""", TestJsonSerializerContext.Default.JsonElement);
// Act.
var session = await ChatClientAgentSession.DeserializeAsync(json);
var session = ChatClientAgentSession.Deserialize(json, TestJsonSerializerContext.Default.Options);
// Assert
Assert.Null(session.ConversationId);
var chatHistoryProvider = session.ChatHistoryProvider as InMemoryChatHistoryProvider;
Assert.NotNull(chatHistoryProvider);
Assert.Single(chatHistoryProvider);
Assert.Equal("testAuthor", chatHistoryProvider[0].AuthorName);
var chatHistoryProvider = new InMemoryChatHistoryProvider();
var messages = chatHistoryProvider.GetMessages(session);
Assert.Single(messages);
Assert.Equal("testAuthor", messages[0].AuthorName);
}
[Fact]
public async Task VerifyDeserializeWithIdAsync()
public void VerifyDeserializeWithId()
{
// Arrange
var json = JsonSerializer.Deserialize("""
@@ -125,42 +78,43 @@ public class ChatClientAgentSessionTests
""", TestJsonSerializerContext.Default.JsonElement);
// Act
var session = await ChatClientAgentSession.DeserializeAsync(json);
var session = ChatClientAgentSession.Deserialize(json);
// Assert
Assert.Equal("TestConvId", session.ConversationId);
Assert.Null(session.ChatHistoryProvider);
}
[Fact]
public async Task VerifyDeserializeWithAIContextProviderAsync()
public void VerifyDeserializeWithStateBag()
{
// Arrange
var json = JsonSerializer.Deserialize("""
{
"conversationId": "TestConvId",
"aiContextProviderState": ["CP1"]
"stateBag": {
"dog": {
"name": "Fido"
}
}
}
""", TestJsonSerializerContext.Default.JsonElement);
Mock<AIContextProvider> mockProvider = new();
// Act
var session = await ChatClientAgentSession.DeserializeAsync(json, aiContextProviderFactory: (_, _, _) => new(mockProvider.Object));
var session = ChatClientAgentSession.Deserialize(json);
// Assert
Assert.Null(session.ChatHistoryProvider);
Assert.Same(session.AIContextProvider, mockProvider.Object);
var dog = session.StateBag.GetValue<Animal>("dog", TestJsonSerializerContext.Default.Options);
Assert.NotNull(dog);
Assert.Equal("Fido", dog.Name);
}
[Fact]
public async Task DeserializeWithInvalidJsonThrowsAsync()
public void DeserializeWithInvalidJsonThrows()
{
// Arrange
var invalidJson = JsonSerializer.Deserialize("[42]", TestJsonSerializerContext.Default.JsonElement);
var session = new ChatClientAgentSession();
// Act & Assert
await Assert.ThrowsAsync<ArgumentException>(() => ChatClientAgentSession.DeserializeAsync(invalidJson));
Assert.Throws<ArgumentException>(() => ChatClientAgentSession.Deserialize(invalidJson));
}
#endregion Deserialize Tests
@@ -195,8 +149,9 @@ public class ChatClientAgentSessionTests
public void VerifySessionSerializationWithMessages()
{
// Arrange
InMemoryChatHistoryProvider provider = [new(ChatRole.User, "TestContent") { AuthorName = "TestAuthor" }];
var session = new ChatClientAgentSession { ChatHistoryProvider = provider };
var provider = new InMemoryChatHistoryProvider();
var session = new ChatClientAgentSession();
provider.SetMessages(session, [new(ChatRole.User, "TestContent") { AuthorName = "TestAuthor" }]);
// Act
var json = session.Serialize();
@@ -206,10 +161,12 @@ public class ChatClientAgentSessionTests
Assert.False(json.TryGetProperty("conversationId", out _));
Assert.True(json.TryGetProperty("chatHistoryProviderState", out var chatHistoryProviderStateProperty));
Assert.Equal(JsonValueKind.Object, chatHistoryProviderStateProperty.ValueKind);
Assert.True(chatHistoryProviderStateProperty.TryGetProperty("messages", out var messagesProperty));
// Messages should be stored in the stateBag
Assert.True(json.TryGetProperty("stateBag", out var stateBagProperty));
Assert.Equal(JsonValueKind.Object, stateBagProperty.ValueKind);
Assert.True(stateBagProperty.TryGetProperty("InMemoryChatHistoryProvider", out var providerStateProperty));
Assert.Equal(JsonValueKind.Object, providerStateProperty.ValueKind);
Assert.True(providerStateProperty.TryGetProperty("messages", out var messagesProperty));
Assert.Equal(JsonValueKind.Array, messagesProperty.ValueKind);
Assert.Single(messagesProperty.EnumerateArray());
@@ -224,29 +181,23 @@ public class ChatClientAgentSessionTests
}
[Fact]
public void VerifySessionSerializationWithWithAIContextProvider()
public void VerifySessionSerializationWithWithStateBag()
{
// Arrange
Mock<AIContextProvider> mockProvider = new();
mockProvider
.Setup(m => m.Serialize(It.IsAny<JsonSerializerOptions?>()))
.Returns(JsonSerializer.SerializeToElement(["CP1"], TestJsonSerializerContext.Default.StringArray));
var session = new ChatClientAgentSession
{
AIContextProvider = mockProvider.Object
};
var session = new ChatClientAgentSession();
session.StateBag.SetValue("dog", new Animal { Name = "Fido" }, TestJsonSerializerContext.Default.Options);
// Act
var json = session.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.True(json.TryGetProperty("aiContextProviderState", out var providerStateProperty));
Assert.Equal(JsonValueKind.Array, providerStateProperty.ValueKind);
Assert.Single(providerStateProperty.EnumerateArray());
Assert.Equal("CP1", providerStateProperty.EnumerateArray().First().GetString());
mockProvider.Verify(m => m.Serialize(It.IsAny<JsonSerializerOptions?>()), Times.Once);
Assert.True(json.TryGetProperty("stateBag", out var stateBagProperty));
Assert.Equal(JsonValueKind.Object, stateBagProperty.ValueKind);
Assert.True(stateBagProperty.TryGetProperty("dog", out var dogProperty));
Assert.Equal(JsonValueKind.Object, dogProperty.ValueKind);
Assert.True(dogProperty.TryGetProperty("name", out var nameProperty));
Assert.Equal("Fido", nameProperty.GetString());
}
/// <summary>
@@ -258,17 +209,7 @@ public class ChatClientAgentSessionTests
// Arrange
var session = new ChatClientAgentSession();
JsonSerializerOptions options = new() { PropertyNamingPolicy = JsonNamingPolicy.SnakeCaseLower };
options.TypeInfoResolverChain.Add(AgentAbstractionsJsonUtilities.DefaultOptions.TypeInfoResolver!);
var chatHistoryProviderStateElement = JsonSerializer.SerializeToElement(
new Dictionary<string, object> { ["Key"] = "TestValue" },
TestJsonSerializerContext.Default.DictionaryStringObject);
var chatHistoryProviderMock = new Mock<ChatHistoryProvider>();
chatHistoryProviderMock
.Setup(m => m.Serialize(options))
.Returns(chatHistoryProviderStateElement);
session.ChatHistoryProvider = chatHistoryProviderMock.Object;
options.TypeInfoResolverChain.Add(AgentJsonUtilities.DefaultOptions.TypeInfoResolver!);
// Act
var json = session.Serialize(options);
@@ -276,54 +217,29 @@ public class ChatClientAgentSessionTests
// Assert
Assert.Equal(JsonValueKind.Object, json.ValueKind);
Assert.False(json.TryGetProperty("conversationId", out var idProperty));
Assert.True(json.TryGetProperty("chatHistoryProviderState", out var chatHistoryProviderStateProperty));
Assert.Equal(JsonValueKind.Object, chatHistoryProviderStateProperty.ValueKind);
Assert.True(chatHistoryProviderStateProperty.TryGetProperty("Key", out var keyProperty));
Assert.Equal("TestValue", keyProperty.GetString());
chatHistoryProviderMock.Verify(m => m.Serialize(options), Times.Once);
// [JsonPropertyName] takes precedence over naming policy
Assert.True(json.TryGetProperty("conversationId", out var _));
}
#endregion Serialize Tests
#region GetService Tests
#region StateBag Roundtrip Tests
[Fact]
public void GetService_RequestingAIContextProvider_ReturnsAIContextProvider()
public void VerifyStateBagRoundtrips()
{
// Arrange
var session = new ChatClientAgentSession();
var mockProvider = new Mock<AIContextProvider>();
mockProvider
.Setup(m => m.GetService(It.Is<Type>(x => x == typeof(AIContextProvider)), null))
.Returns(mockProvider.Object);
session.AIContextProvider = mockProvider.Object;
session.StateBag.SetValue("dog", new Animal { Name = "Fido" }, TestJsonSerializerContext.Default.Options);
// Act
var result = session.GetService(typeof(AIContextProvider));
var serializedSession = session.Serialize();
var deserializedSession = ChatClientAgentSession.Deserialize(serializedSession);
// Assert
Assert.NotNull(result);
Assert.Same(mockProvider.Object, result);
}
[Fact]
public void GetService_RequestingChatHistoryProvider_ReturnsChatHistoryProvider()
{
// Arrange
var session = new ChatClientAgentSession();
var chatHistoryProvider = new InMemoryChatHistoryProvider();
session.ChatHistoryProvider = chatHistoryProvider;
// Act
var result = session.GetService(typeof(ChatHistoryProvider));
// Assert
Assert.NotNull(result);
Assert.Same(chatHistoryProvider, result);
var dog = deserializedSession.StateBag.GetValue<Animal>("dog", TestJsonSerializerContext.Default.Options);
Assert.NotNull(dog);
Assert.Equal("Fido", dog.Name);
}
#endregion
@@ -43,6 +43,154 @@ public partial class ChatClientAgentTests
Assert.Equal("FunctionInvokingChatClient", agent.ChatClient.GetType().Name);
}
/// <summary>
/// Verify that the constructor throws when two AIContextProviders use the same StateKey.
/// </summary>
[Fact]
public void Constructor_ThrowsWhenDuplicateAIContextProviderStateKeys()
{
// Arrange
var chatClient = new Mock<IChatClient>().Object;
var provider1 = new TestAIContextProvider("SharedKey");
var provider2 = new TestAIContextProvider("SharedKey");
// Act & Assert
var ex = Assert.Throws<InvalidOperationException>(() =>
new ChatClientAgent(chatClient, options: new()
{
AIContextProviders = [provider1, provider2]
}));
Assert.Contains("SharedKey", ex.Message);
}
/// <summary>
/// Verify that the constructor throws when an AIContextProvider uses the same StateKey as the default InMemoryChatHistoryProvider
/// and no explicit ChatHistoryProvider is configured.
/// </summary>
[Fact]
public void Constructor_ThrowsWhenAIContextProviderStateKeyClashesWithDefaultInMemoryChatHistoryProvider()
{
// Arrange
var chatClient = new Mock<IChatClient>().Object;
var contextProvider = new TestAIContextProvider(nameof(InMemoryChatHistoryProvider));
// Act & Assert
var ex = Assert.Throws<InvalidOperationException>(() =>
new ChatClientAgent(chatClient, options: new()
{
AIContextProviders = [contextProvider]
}));
Assert.Contains(nameof(InMemoryChatHistoryProvider), ex.Message);
}
/// <summary>
/// Verify that the constructor throws when a ChatHistoryProvider uses the same StateKey as an AIContextProvider.
/// </summary>
[Fact]
public void Constructor_ThrowsWhenChatHistoryProviderStateKeyClashesWithAIContextProvider()
{
// Arrange
var chatClient = new Mock<IChatClient>().Object;
var contextProvider = new TestAIContextProvider("SharedKey");
var historyProvider = new TestChatHistoryProvider("SharedKey");
// Act & Assert
var ex = Assert.Throws<InvalidOperationException>(() =>
new ChatClientAgent(chatClient, options: new()
{
AIContextProviders = [contextProvider],
ChatHistoryProvider = historyProvider
}));
Assert.Contains("SharedKey", ex.Message);
Assert.Contains(nameof(ChatHistoryProvider), ex.Message);
}
/// <summary>
/// Verify that the constructor succeeds when all providers use unique StateKeys.
/// </summary>
[Fact]
public void Constructor_SucceedsWithUniqueProviderStateKeys()
{
// Arrange
var chatClient = new Mock<IChatClient>().Object;
var contextProvider1 = new TestAIContextProvider("Key1");
var contextProvider2 = new TestAIContextProvider("Key2");
var historyProvider = new TestChatHistoryProvider("Key3");
// Act & Assert - should not throw
_ = new ChatClientAgent(chatClient, options: new()
{
AIContextProviders = [contextProvider1, contextProvider2],
ChatHistoryProvider = historyProvider
});
}
/// <summary>
/// Verify that RunAsync throws when an override ChatHistoryProvider's StateKey clashes with an AIContextProvider.
/// </summary>
[Fact]
public async Task RunAsync_ThrowsWhenOverrideChatHistoryProviderStateKeyClashesWithAIContextProviderAsync()
{
// Arrange
Mock<IChatClient> mockService = new();
mockService.Setup(
s => s.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]));
var contextProvider = new TestAIContextProvider("SharedKey");
var overrideHistoryProvider = new TestChatHistoryProvider("SharedKey");
ChatClientAgent agent = new(mockService.Object, options: new()
{
AIContextProviders = [contextProvider]
});
// Act & Assert
ChatClientAgentSession? session = await agent.CreateSessionAsync() as ChatClientAgentSession;
AdditionalPropertiesDictionary additionalProperties = new();
additionalProperties.Add<ChatHistoryProvider>(overrideHistoryProvider);
var ex = await Assert.ThrowsAsync<InvalidOperationException>(() =>
agent.RunAsync([new(ChatRole.User, "test")], session, options: new AgentRunOptions { AdditionalProperties = additionalProperties }));
Assert.Contains("SharedKey", ex.Message);
}
/// <summary>
/// Verify that RunAsync succeeds when an override ChatHistoryProvider uses the same StateKey as the default ChatHistoryProvider.
/// </summary>
[Fact]
public async Task RunAsync_SucceedsWhenOverrideChatHistoryProviderSharesKeyWithDefaultAsync()
{
// Arrange
Mock<IChatClient> mockService = new();
mockService.Setup(
s => s.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]));
var defaultHistoryProvider = new TestChatHistoryProvider("SameKey");
var overrideHistoryProvider = new TestChatHistoryProvider("SameKey");
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatHistoryProvider = defaultHistoryProvider
});
// Act & Assert - should not throw
ChatClientAgentSession? session = await agent.CreateSessionAsync() as ChatClientAgentSession;
AdditionalPropertiesDictionary additionalProperties = new();
additionalProperties.Add<ChatHistoryProvider>(overrideHistoryProvider);
await agent.RunAsync([new(ChatRole.User, "test")], session, options: new AgentRunOptions { AdditionalProperties = additionalProperties });
}
#endregion
#region RunAsync Tests
@@ -343,18 +491,19 @@ public partial class ChatClientAgentTests
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(new AIContext
{
Messages = aiContextProviderMessages,
Instructions = "context provider instructions",
Tools = [AIFunctionFactory.Create(() => { }, "context provider function")]
});
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat(aiContextProviderMessages),
Instructions = ctx.AIContext.Instructions + "\ncontext provider instructions",
Tools = (ctx.AIContext.Tools ?? []).Concat(new[] { AIFunctionFactory.Create(() => { }, "context provider function") })
}));
mockProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviderFactory = (_, _) => new(mockProvider.Object), ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviders = [mockProvider.Object], ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
// Act
var session = await agent.CreateSessionAsync() as ChatClientAgentSession;
@@ -373,11 +522,13 @@ public partial class ChatClientAgentTests
Assert.Contains(capturedTools, t => t.Name == "context provider function");
// Verify that the session was updated with the ai context provider, input and response messages
var chatHistoryProvider = Assert.IsType<InMemoryChatHistoryProvider>(session!.ChatHistoryProvider);
Assert.Equal(3, chatHistoryProvider.Count);
Assert.Equal("user message", chatHistoryProvider[0].Text);
Assert.Equal("context provider message", chatHistoryProvider[1].Text);
Assert.Equal("response", chatHistoryProvider[2].Text);
var chatHistoryProvider = agent.ChatHistoryProvider as InMemoryChatHistoryProvider;
Assert.NotNull(chatHistoryProvider);
var messages = chatHistoryProvider.GetMessages(session);
Assert.Equal(3, messages.Count);
Assert.Equal("user message", messages[0].Text);
Assert.Equal("context provider message", messages[1].Text);
Assert.Equal("response", messages[2].Text);
mockProvider
.Protected()
@@ -411,16 +562,17 @@ public partial class ChatClientAgentTests
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(new AIContext
{
Messages = aiContextProviderMessages,
});
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat(aiContextProviderMessages),
}));
mockProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviderFactory = (_, _) => new(mockProvider.Object), ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviders = [mockProvider.Object], ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
// Act
await Assert.ThrowsAsync<InvalidOperationException>(() => agent.RunAsync(requestMessages));
@@ -468,9 +620,15 @@ public partial class ChatClientAgentTests
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(new AIContext());
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Instructions = ctx.AIContext.Instructions,
Messages = ctx.AIContext.Messages,
Tools = ctx.AIContext.Tools
}));
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviderFactory = (_, _) => new(mockProvider.Object), ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
ChatClientAgent agent = new(mockService.Object, options: new() { AIContextProviders = [mockProvider.Object], ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] } });
// Act
await agent.RunAsync([new(ChatRole.User, "user message")]);
@@ -488,6 +646,299 @@ public partial class ChatClientAgentTests
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
}
/// <summary>
/// Verify that RunAsync invokes multiple AIContextProviders in sequence, each receiving the accumulated context.
/// </summary>
[Fact]
public async Task RunAsyncInvokesMultipleAIContextProvidersInOrderAsync()
{
// Arrange
ChatMessage[] requestMessages = [new(ChatRole.User, "user message")];
ChatMessage[] responseMessages = [new(ChatRole.Assistant, "response")];
Mock<IChatClient> mockService = new();
List<ChatMessage> capturedMessages = [];
string capturedInstructions = string.Empty;
List<AITool> capturedTools = [];
mockService
.Setup(s => s.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>()))
.Callback<IEnumerable<ChatMessage>, ChatOptions, CancellationToken>((msgs, opts, ct) =>
{
capturedMessages.AddRange(msgs);
capturedInstructions = opts.Instructions ?? string.Empty;
if (opts.Tools is not null)
{
capturedTools.AddRange(opts.Tools);
}
})
.ReturnsAsync(new ChatResponse(responseMessages));
// Provider 1: adds a system message and a tool
var mockProvider1 = new Mock<AIContextProvider>();
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat([new ChatMessage(ChatRole.System, "provider1 context")]).ToList(),
Instructions = ctx.AIContext.Instructions + "\nprovider1 instructions",
Tools = (ctx.AIContext.Tools ?? []).Concat([AIFunctionFactory.Create(() => { }, "provider1 function")]).ToList()
}));
mockProvider1
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
// Provider 2: adds another system message and verifies it receives accumulated context from provider 1
AIContext? provider2ReceivedContext = null;
var mockProvider2 = new Mock<AIContextProvider>();
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
{
provider2ReceivedContext = ctx.AIContext;
return new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat([new ChatMessage(ChatRole.System, "provider2 context")]).ToList(),
Instructions = ctx.AIContext.Instructions + "\nprovider2 instructions",
Tools = (ctx.AIContext.Tools ?? []).Concat([AIFunctionFactory.Create(() => { }, "provider2 function")]).ToList()
});
});
mockProvider2
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
ChatClientAgent agent = new(mockService.Object, options: new()
{
AIContextProviders = [mockProvider1.Object, mockProvider2.Object],
ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] }
});
// Act
var session = await agent.CreateSessionAsync() as ChatClientAgentSession;
await agent.RunAsync(requestMessages, session);
// Assert
// Provider 2 should have received accumulated context from provider 1
Assert.NotNull(provider2ReceivedContext);
Assert.Contains(provider2ReceivedContext.Messages!, m => m.Text == "provider1 context");
Assert.Contains("provider1 instructions", provider2ReceivedContext.Instructions);
// Final captured messages should contain user message + both provider contexts
Assert.Equal(3, capturedMessages.Count);
Assert.Equal("user message", capturedMessages[0].Text);
Assert.Equal("provider1 context", capturedMessages[1].Text);
Assert.Equal("provider2 context", capturedMessages[2].Text);
// Instructions should be accumulated
Assert.Equal("base instructions\nprovider1 instructions\nprovider2 instructions", capturedInstructions);
// Tools should contain base + both provider tools
Assert.Equal(3, capturedTools.Count);
Assert.Contains(capturedTools, t => t.Name == "base function");
Assert.Contains(capturedTools, t => t.Name == "provider1 function");
Assert.Contains(capturedTools, t => t.Name == "provider2 function");
// Both providers should have been invoked
mockProvider1
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
// Both providers should have been notified of success
mockProvider1
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.ResponseMessages == responseMessages &&
x.InvokeException == null), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.ResponseMessages == responseMessages &&
x.InvokeException == null), ItExpr.IsAny<CancellationToken>());
}
/// <summary>
/// Verify that RunAsync invokes InvokedCoreAsync on all AIContextProviders when the downstream GetResponse call fails.
/// </summary>
[Fact]
public async Task RunAsyncInvokesMultipleAIContextProvidersOnFailureAsync()
{
// Arrange
ChatMessage[] requestMessages = [new(ChatRole.User, "user message")];
Mock<IChatClient> mockService = new();
mockService
.Setup(s => s.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>()))
.ThrowsAsync(new InvalidOperationException("downstream failure"));
var mockProvider1 = new Mock<AIContextProvider>();
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = ctx.AIContext.Messages?.ToList(),
Instructions = ctx.AIContext.Instructions,
Tools = ctx.AIContext.Tools
}));
mockProvider1
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
var mockProvider2 = new Mock<AIContextProvider>();
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = ctx.AIContext.Messages?.ToList(),
Instructions = ctx.AIContext.Instructions,
Tools = ctx.AIContext.Tools
}));
mockProvider2
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
ChatClientAgent agent = new(mockService.Object, options: new()
{
AIContextProviders = [mockProvider1.Object, mockProvider2.Object],
ChatOptions = new() { Instructions = "base instructions" }
});
// Act
await Assert.ThrowsAsync<InvalidOperationException>(() => agent.RunAsync(requestMessages));
// Assert - both providers should have been notified of the failure
mockProvider1
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
mockProvider1
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.InvokeException is InvalidOperationException), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.InvokeException is InvalidOperationException), ItExpr.IsAny<CancellationToken>());
}
/// <summary>
/// Verify that RunStreamingAsync invokes multiple AIContextProviders in sequence.
/// </summary>
[Fact]
public async Task RunStreamingAsyncInvokesMultipleAIContextProvidersAsync()
{
// Arrange
ChatMessage[] requestMessages = [new(ChatRole.User, "user message")];
ChatResponseUpdate[] responseUpdates = [new(ChatRole.Assistant, "response")];
Mock<IChatClient> mockService = new();
List<ChatMessage> capturedMessages = [];
string capturedInstructions = string.Empty;
mockService
.Setup(s => s.GetStreamingResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>()))
.Callback<IEnumerable<ChatMessage>, ChatOptions, CancellationToken>((msgs, opts, ct) =>
{
capturedMessages.AddRange(msgs);
capturedInstructions = opts.Instructions ?? string.Empty;
})
.Returns(ToAsyncEnumerableAsync(responseUpdates));
var mockProvider1 = new Mock<AIContextProvider>();
mockProvider1.SetupGet(p => p.StateKey).Returns("Provider1");
mockProvider1
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat([new ChatMessage(ChatRole.System, "provider1 context")]).ToList(),
Instructions = ctx.AIContext.Instructions + "\nprovider1 instructions",
Tools = ctx.AIContext.Tools
}));
mockProvider1
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
var mockProvider2 = new Mock<AIContextProvider>();
mockProvider2.SetupGet(p => p.StateKey).Returns("Provider2");
mockProvider2
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat([new ChatMessage(ChatRole.System, "provider2 context")]).ToList(),
Instructions = ctx.AIContext.Instructions + "\nprovider2 instructions",
Tools = ctx.AIContext.Tools
}));
mockProvider2
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
ChatClientAgent agent = new(
mockService.Object,
options: new()
{
ChatOptions = new() { Instructions = "base instructions" },
AIContextProviders = [mockProvider1.Object, mockProvider2.Object]
});
// Act
var session = await agent.CreateSessionAsync() as ChatClientAgentSession;
var updates = agent.RunStreamingAsync(requestMessages, session);
_ = await updates.ToAgentResponseAsync();
// Assert
Assert.Equal(3, capturedMessages.Count);
Assert.Equal("user message", capturedMessages[0].Text);
Assert.Equal("provider1 context", capturedMessages[1].Text);
Assert.Equal("provider2 context", capturedMessages[2].Text);
Assert.Equal("base instructions\nprovider1 instructions\nprovider2 instructions", capturedInstructions);
// Both providers should have been invoked and notified
mockProvider1
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask<AIContext>>("InvokingCoreAsync", Times.Once(), ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>());
mockProvider1
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.InvokeException == null), ItExpr.IsAny<CancellationToken>());
mockProvider2
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(), ItExpr.Is<AIContextProvider.InvokedContext>(x =>
x.InvokeException == null), ItExpr.IsAny<CancellationToken>());
}
#endregion
#region Property Override Tests
@@ -1259,12 +1710,10 @@ public partial class ChatClientAgentTests
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).Returns(ToAsyncEnumerableAsync(returnUpdates));
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(new InMemoryChatHistoryProvider());
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = new InMemoryChatHistoryProvider()
});
// Act
@@ -1272,11 +1721,11 @@ public partial class ChatClientAgentTests
await agent.RunStreamingAsync([new(ChatRole.User, "test")], session).ToListAsync();
// Assert
var chatHistoryProvider = Assert.IsType<InMemoryChatHistoryProvider>(session!.ChatHistoryProvider);
Assert.Equal(2, chatHistoryProvider.Count);
Assert.Equal("test", chatHistoryProvider[0].Text);
Assert.Equal("what?", chatHistoryProvider[1].Text);
mockFactory.Verify(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>()), Times.Once);
var chatHistoryProvider = Assert.IsType<InMemoryChatHistoryProvider>(agent.GetService(typeof(ChatHistoryProvider)));
var historyMessages = chatHistoryProvider.GetMessages(session);
Assert.Equal(2, historyMessages.Count);
Assert.Equal("test", historyMessages[0].Text);
Assert.Equal("what?", historyMessages[1].Text);
}
/// <summary>
@@ -1319,10 +1768,10 @@ public partial class ChatClientAgentTests
}
/// <summary>
/// Verify that RunStreamingAsync throws when a <see cref="ChatHistoryProvider"/> factory is provided and the chat client returns a conversation id.
/// Verify that RunStreamingAsync throws when a <see cref="ChatHistoryProvider"/> is provided and the chat client returns a conversation id.
/// </summary>
[Fact]
public async Task RunStreamingAsyncThrowsWhenChatHistoryProviderFactoryProvidedAndConversationIdReturnedByChatClientAsync()
public async Task RunStreamingAsyncThrowsWhenChatHistoryProviderProvidedAndConversationIdReturnedByChatClientAsync()
{
// Arrange
Mock<IChatClient> mockService = new();
@@ -1336,18 +1785,16 @@ public partial class ChatClientAgentTests
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).Returns(ToAsyncEnumerableAsync(returnUpdates));
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(new InMemoryChatHistoryProvider());
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = new InMemoryChatHistoryProvider()
});
// Act & Assert
ChatClientAgentSession? session = await agent.CreateSessionAsync() as ChatClientAgentSession;
var exception = await Assert.ThrowsAsync<InvalidOperationException>(async () => await agent.RunStreamingAsync([new(ChatRole.User, "test")], session).ToListAsync());
Assert.Equal("Only the ConversationId or ChatHistoryProvider may be set, but not both and switching from one to another is not supported.", exception.Message);
Assert.Equal("Only ConversationId or ChatHistoryProvider may be used, but not both. The service returned a conversation id indicating server-side chat history management, but the agent has a ChatHistoryProvider configured.", exception.Message);
}
/// <summary>
@@ -1384,12 +1831,13 @@ public partial class ChatClientAgentTests
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(new AIContext
{
Messages = aiContextProviderMessages,
Instructions = "context provider instructions",
Tools = [AIFunctionFactory.Create(() => { }, "context provider function")]
});
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat(aiContextProviderMessages),
Instructions = ctx.AIContext.Instructions + "\ncontext provider instructions",
Tools = (ctx.AIContext.Tools ?? []).Concat(new[] { AIFunctionFactory.Create(() => { }, "context provider function") })
}));
mockProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
@@ -1400,7 +1848,7 @@ public partial class ChatClientAgentTests
options: new()
{
ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] },
AIContextProviderFactory = (_, _) => new(mockProvider.Object)
AIContextProviders = [mockProvider.Object]
});
// Act
@@ -1421,11 +1869,13 @@ public partial class ChatClientAgentTests
Assert.Contains(capturedTools, t => t.Name == "context provider function");
// Verify that the session was updated with the input, ai context provider, and response messages
var chatHistoryProvider = Assert.IsType<InMemoryChatHistoryProvider>(session!.ChatHistoryProvider);
Assert.Equal(3, chatHistoryProvider.Count);
Assert.Equal("user message", chatHistoryProvider[0].Text);
Assert.Equal("context provider message", chatHistoryProvider[1].Text);
Assert.Equal("response", chatHistoryProvider[2].Text);
var chatHistoryProvider = agent.ChatHistoryProvider as InMemoryChatHistoryProvider;
Assert.NotNull(chatHistoryProvider);
var historyMessages2 = chatHistoryProvider.GetMessages(session);
Assert.Equal(3, historyMessages2.Count);
Assert.Equal("user message", historyMessages2[0].Text);
Assert.Equal("context provider message", historyMessages2[1].Text);
Assert.Equal("response", historyMessages2[2].Text);
mockProvider
.Protected()
@@ -1460,10 +1910,11 @@ public partial class ChatClientAgentTests
mockProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(new AIContext
{
Messages = aiContextProviderMessages,
});
.Returns((AIContextProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<AIContext>(new AIContext
{
Messages = (ctx.AIContext.Messages ?? []).Concat(aiContextProviderMessages),
}));
mockProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
@@ -1474,7 +1925,7 @@ public partial class ChatClientAgentTests
options: new()
{
ChatOptions = new() { Instructions = "base instructions", Tools = [AIFunctionFactory.Create(() => { }, "base function")] },
AIContextProviderFactory = (_, _) => new(mockProvider.Object)
AIContextProviders = [mockProvider.Object]
});
// Act
@@ -1506,4 +1957,27 @@ public partial class ChatClientAgentTests
yield return update;
}
}
[JsonSourceGenerationOptions(UseStringEnumConverter = true, PropertyNamingPolicy = JsonKnownNamingPolicy.CamelCase)]
[JsonSerializable(typeof(Animal))]
private sealed partial class JsonContext2 : JsonSerializerContext;
private sealed class TestAIContextProvider(string stateKey) : AIContextProvider
{
public override string StateKey => stateKey;
protected override ValueTask<AIContext> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(context.AIContext);
}
private sealed class TestChatHistoryProvider(string stateKey) : ChatHistoryProvider
{
public override string StateKey => stateKey;
protected override ValueTask<IEnumerable<ChatMessage>> InvokingCoreAsync(InvokingContext context, CancellationToken cancellationToken = default)
=> new(context.RequestMessages);
protected override ValueTask InvokedCoreAsync(InvokedContext context, CancellationToken cancellationToken = default)
=> default;
}
}
@@ -339,6 +339,7 @@ public class ChatClientAgent_BackgroundResponsesTests
// Create a mock chat history provider that would normally provide messages
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -346,6 +347,7 @@ public class ChatClientAgent_BackgroundResponsesTests
// Create a mock AI context provider that would normally provide context
var mockContextProvider = new Mock<AIContextProvider>();
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -365,14 +367,14 @@ public class ChatClientAgent_BackgroundResponsesTests
capturedMessages.AddRange(msgs))
.ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "continued response")]));
ChatClientAgent agent = new(mockChatClient.Object);
// Create a session with both chat history provider and AI context provider
ChatClientAgentSession? session = new()
ChatClientAgent agent = new(mockChatClient.Object, options: new()
{
ChatHistoryProvider = mockChatHistoryProvider.Object,
AIContextProvider = mockContextProvider.Object
};
AIContextProviders = [mockContextProvider.Object]
});
// Create a session
ChatClientAgentSession? session = new();
AgentRunOptions runOptions = new()
{
@@ -406,6 +408,7 @@ public class ChatClientAgent_BackgroundResponsesTests
// Create a mock chat history provider that would normally provide messages
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -413,6 +416,7 @@ public class ChatClientAgent_BackgroundResponsesTests
// Create a mock AI context provider that would normally provide context
var mockContextProvider = new Mock<AIContextProvider>();
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
.Setup<ValueTask<AIContext>>("InvokingCoreAsync", ItExpr.IsAny<AIContextProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
@@ -432,14 +436,14 @@ public class ChatClientAgent_BackgroundResponsesTests
capturedMessages.AddRange(msgs))
.Returns(ToAsyncEnumerableAsync([new ChatResponseUpdate(role: ChatRole.Assistant, content: "continued response")]));
ChatClientAgent agent = new(mockChatClient.Object);
// Create a session with both chat history provider and AI context provider
ChatClientAgentSession? session = new()
ChatClientAgent agent = new(mockChatClient.Object, options: new()
{
ChatHistoryProvider = mockChatHistoryProvider.Object,
AIContextProvider = mockContextProvider.Object
};
AIContextProviders = [mockContextProvider.Object]
});
// Create a session
ChatClientAgentSession? session = new();
AgentRunOptions runOptions = new()
{
@@ -633,10 +637,9 @@ public class ChatClientAgent_BackgroundResponsesTests
It.IsAny<CancellationToken>()))
.Returns(ToAsyncEnumerableAsync(returnUpdates));
ChatClientAgent agent = new(mockChatClient.Object);
List<ChatMessage> capturedMessagesAddedToProvider = [];
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
@@ -645,17 +648,20 @@ public class ChatClientAgent_BackgroundResponsesTests
AIContextProvider.InvokedContext? capturedInvokedContext = null;
var mockContextProvider = new Mock<AIContextProvider>();
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Callback<AIContextProvider.InvokedContext, CancellationToken>((context, ct) => capturedInvokedContext = context)
.Returns(new ValueTask());
ChatClientAgentSession? session = new()
ChatClientAgent agent = new(mockChatClient.Object, options: new()
{
ChatHistoryProvider = mockChatHistoryProvider.Object,
AIContextProvider = mockContextProvider.Object
};
AIContextProviders = [mockContextProvider.Object]
});
ChatClientAgentSession? session = new();
AgentRunOptions runOptions = new()
{
@@ -695,10 +701,9 @@ public class ChatClientAgent_BackgroundResponsesTests
It.IsAny<CancellationToken>()))
.Returns(ToAsyncEnumerableAsync(Array.Empty<ChatResponseUpdate>()));
ChatClientAgent agent = new(mockChatClient.Object);
List<ChatMessage> capturedMessagesAddedToProvider = [];
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
mockChatHistoryProvider.SetupGet(p => p.StateKey).Returns("ChatHistoryProvider");
mockChatHistoryProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
@@ -707,17 +712,20 @@ public class ChatClientAgent_BackgroundResponsesTests
AIContextProvider.InvokedContext? capturedInvokedContext = null;
var mockContextProvider = new Mock<AIContextProvider>();
mockContextProvider.SetupGet(p => p.StateKey).Returns("Provider1");
mockContextProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<AIContextProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Callback<AIContextProvider.InvokedContext, CancellationToken>((context, ct) => capturedInvokedContext = context)
.Returns(new ValueTask());
ChatClientAgentSession? session = new()
ChatClientAgent agent = new(mockChatClient.Object, options: new()
{
ChatHistoryProvider = mockChatHistoryProvider.Object,
AIContextProvider = mockContextProvider.Object
};
AIContextProviders = [mockContextProvider.Object]
});
ChatClientAgentSession? session = new();
AgentRunOptions runOptions = new()
{
@@ -163,17 +163,19 @@ public class ChatClientAgent_ChatHistoryManagementTests
await agent.RunAsync([new(ChatRole.User, "test")], session);
// Assert
InMemoryChatHistoryProvider chatHistoryProvider = Assert.IsType<InMemoryChatHistoryProvider>(session!.ChatHistoryProvider);
Assert.Equal(2, chatHistoryProvider.Count);
Assert.Equal("test", chatHistoryProvider[0].Text);
Assert.Equal("response", chatHistoryProvider[1].Text);
var inMemoryProvider = agent.ChatHistoryProvider as InMemoryChatHistoryProvider;
Assert.NotNull(inMemoryProvider);
var messages = inMemoryProvider.GetMessages(session!);
Assert.Equal(2, messages.Count);
Assert.Equal("test", messages[0].Text);
Assert.Equal("response", messages[1].Text);
}
/// <summary>
/// Verify that RunAsync uses the ChatHistoryProvider factory when the chat client returns no conversation id.
/// Verify that RunAsync uses the ChatHistoryProvider when the chat client returns no conversation id.
/// </summary>
[Fact]
public async Task RunAsync_UsesChatHistoryProviderFactory_WhenProvidedAndNoConversationIdReturnedByChatClientAsync()
public async Task RunAsync_UsesChatHistoryProvider_WhenProvidedAndNoConversationIdReturnedByChatClientAsync()
{
// Arrange
Mock<IChatClient> mockService = new();
@@ -187,19 +189,17 @@ public class ChatClientAgent_ChatHistoryManagementTests
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync([new ChatMessage(ChatRole.User, "Existing Chat History")]);
.Returns((ChatHistoryProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<IEnumerable<ChatMessage>>(new List<ChatMessage> { new(ChatRole.User, "Existing Chat History") }.Concat(ctx.RequestMessages).ToList()));
mockChatHistoryProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(mockChatHistoryProvider.Object);
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = mockChatHistoryProvider.Object
});
// Act
@@ -207,7 +207,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
await agent.RunAsync([new(ChatRole.User, "test")], session);
// Assert
Assert.IsType<ChatHistoryProvider>(session!.ChatHistoryProvider, exactMatch: false);
Assert.Same(mockChatHistoryProvider.Object, agent.ChatHistoryProvider);
mockService.Verify(
x => x.GetResponseAsync(
It.Is<IEnumerable<ChatMessage>>(msgs => msgs.Count() == 2 && msgs.Any(m => m.Text == "Existing Chat History") && msgs.Any(m => m.Text == "test")),
@@ -222,9 +222,8 @@ public class ChatClientAgent_ChatHistoryManagementTests
mockChatHistoryProvider
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(),
ItExpr.Is<ChatHistoryProvider.InvokedContext>(x => x.RequestMessages.Count() == 1 && x.ResponseMessages!.Count() == 1),
ItExpr.Is<ChatHistoryProvider.InvokedContext>(x => x.RequestMessages.Count() == 2 && x.ResponseMessages!.Count() == 1),
ItExpr.IsAny<CancellationToken>());
mockFactory.Verify(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>()), Times.Once);
}
/// <summary>
@@ -242,14 +241,16 @@ public class ChatClientAgent_ChatHistoryManagementTests
It.IsAny<CancellationToken>())).Throws(new InvalidOperationException("Test Error"));
Mock<ChatHistoryProvider> mockChatHistoryProvider = new();
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(mockChatHistoryProvider.Object);
mockChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.Returns((ChatHistoryProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<IEnumerable<ChatMessage>>(ctx.RequestMessages.ToList()));
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = mockChatHistoryProvider.Object
});
// Act
@@ -257,20 +258,19 @@ public class ChatClientAgent_ChatHistoryManagementTests
await Assert.ThrowsAsync<InvalidOperationException>(() => agent.RunAsync([new(ChatRole.User, "test")], session));
// Assert
Assert.IsType<ChatHistoryProvider>(session!.ChatHistoryProvider, exactMatch: false);
Assert.Same(mockChatHistoryProvider.Object, agent.ChatHistoryProvider);
mockChatHistoryProvider
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(),
ItExpr.Is<ChatHistoryProvider.InvokedContext>(x => x.RequestMessages.Count() == 1 && x.ResponseMessages == null && x.InvokeException!.Message == "Test Error"),
ItExpr.IsAny<CancellationToken>());
mockFactory.Verify(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>()), Times.Once);
}
/// <summary>
/// Verify that RunAsync throws when a ChatHistoryProvider Factory is provided and the chat client returns a conversation id.
/// Verify that RunAsync throws when a ChatHistoryProvider is provided and the chat client returns a conversation id.
/// </summary>
[Fact]
public async Task RunAsync_Throws_WhenChatHistoryProviderFactoryProvidedAndConversationIdReturnedByChatClientAsync()
public async Task RunAsync_Throws_WhenChatHistoryProviderProvidedAndConversationIdReturnedByChatClientAsync()
{
// Arrange
Mock<IChatClient> mockService = new();
@@ -279,18 +279,16 @@ public class ChatClientAgent_ChatHistoryManagementTests
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>())).ReturnsAsync(new ChatResponse([new(ChatRole.Assistant, "response")]) { ConversationId = "ConvId" });
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(new InMemoryChatHistoryProvider());
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = new InMemoryChatHistoryProvider()
});
// Act & Assert
ChatClientAgentSession? session = await agent.CreateSessionAsync() as ChatClientAgentSession;
InvalidOperationException exception = await Assert.ThrowsAsync<InvalidOperationException>(() => agent.RunAsync([new(ChatRole.User, "test")], session));
Assert.Equal("Only the ConversationId or ChatHistoryProvider may be set, but not both and switching from one to another is not supported.", exception.Message);
Assert.Equal("Only ConversationId or ChatHistoryProvider may be used, but not both. The service returned a conversation id indicating server-side chat history management, but the agent has a ChatHistoryProvider configured.", exception.Message);
}
#endregion
@@ -317,31 +315,29 @@ public class ChatClientAgent_ChatHistoryManagementTests
mockOverrideChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ReturnsAsync([new ChatMessage(ChatRole.User, "Existing Chat History")]);
.Returns((ChatHistoryProvider.InvokingContext ctx, CancellationToken _) =>
new ValueTask<IEnumerable<ChatMessage>>(new List<ChatMessage> { new(ChatRole.User, "Existing Chat History") }.Concat(ctx.RequestMessages).ToList()));
mockOverrideChatHistoryProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Returns(new ValueTask());
// Arrange a chat history provider to provide to the agent via a factory at construction time.
// 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> mockFactoryChatHistoryProvider = new();
mockFactoryChatHistoryProvider
Mock<ChatHistoryProvider> mockAgentOptionsChatHistoryProvider = new();
mockAgentOptionsChatHistoryProvider
.Protected()
.Setup<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(), ItExpr.IsAny<CancellationToken>())
.ThrowsAsync(FailException.ForFailure("Base ChatHistoryProvider shouldn't be used."));
mockFactoryChatHistoryProvider
mockAgentOptionsChatHistoryProvider
.Protected()
.Setup<ValueTask>("InvokedCoreAsync", ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(), ItExpr.IsAny<CancellationToken>())
.Throws(FailException.ForFailure("Base ChatHistoryProvider shouldn't be used."));
Mock<Func<ChatClientAgentOptions.ChatHistoryProviderFactoryContext, CancellationToken, ValueTask<ChatHistoryProvider>>> mockFactory = new();
mockFactory.Setup(f => f(It.IsAny<ChatClientAgentOptions.ChatHistoryProviderFactoryContext>(), It.IsAny<CancellationToken>())).ReturnsAsync(mockFactoryChatHistoryProvider.Object);
ChatClientAgent agent = new(mockService.Object, options: new()
{
ChatOptions = new() { Instructions = "test instructions" },
ChatHistoryProviderFactory = mockFactory.Object
ChatHistoryProvider = mockAgentOptionsChatHistoryProvider.Object
});
// Act
@@ -351,7 +347,7 @@ public class ChatClientAgent_ChatHistoryManagementTests
await agent.RunAsync([new(ChatRole.User, "test")], session, options: new AgentRunOptions { AdditionalProperties = additionalProperties });
// Assert
Assert.Same(mockFactoryChatHistoryProvider.Object, session!.ChatHistoryProvider);
Assert.Same(mockAgentOptionsChatHistoryProvider.Object, agent.ChatHistoryProvider);
mockService.Verify(
x => x.GetResponseAsync(
It.Is<IEnumerable<ChatMessage>>(msgs => msgs.Count() == 2 && msgs.Any(m => m.Text == "Existing Chat History") && msgs.Any(m => m.Text == "test")),
@@ -366,15 +362,15 @@ public class ChatClientAgent_ChatHistoryManagementTests
mockOverrideChatHistoryProvider
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Once(),
ItExpr.Is<ChatHistoryProvider.InvokedContext>(x => x.RequestMessages.Count() == 1 && x.ResponseMessages!.Count() == 1),
ItExpr.Is<ChatHistoryProvider.InvokedContext>(x => x.RequestMessages.Count() == 2 && x.ResponseMessages!.Count() == 1),
ItExpr.IsAny<CancellationToken>());
mockFactoryChatHistoryProvider
mockAgentOptionsChatHistoryProvider
.Protected()
.Verify<ValueTask<IEnumerable<ChatMessage>>>("InvokingCoreAsync", Times.Never(),
ItExpr.IsAny<ChatHistoryProvider.InvokingContext>(),
ItExpr.IsAny<CancellationToken>());
mockFactoryChatHistoryProvider
mockAgentOptionsChatHistoryProvider
.Protected()
.Verify<ValueTask>("InvokedCoreAsync", Times.Never(),
ItExpr.IsAny<ChatHistoryProvider.InvokedContext>(),
@@ -11,77 +11,6 @@ namespace Microsoft.Agents.AI.UnitTests;
/// </summary>
public class ChatClientAgent_CreateSessionTests
{
[Fact]
public async Task CreateSession_UsesAIContextProviderFactory_IfProvidedAsync()
{
// Arrange
var mockChatClient = new Mock<IChatClient>();
var mockContextProvider = new Mock<AIContextProvider>();
var factoryCalled = false;
var agent = new ChatClientAgent(mockChatClient.Object, new ChatClientAgentOptions
{
ChatOptions = new() { Instructions = "Test instructions" },
AIContextProviderFactory = (_, _) =>
{
factoryCalled = true;
return new ValueTask<AIContextProvider>(mockContextProvider.Object);
}
});
// Act
var session = await agent.CreateSessionAsync();
// Assert
Assert.True(factoryCalled, "AIContextProviderFactory was not called.");
Assert.IsType<ChatClientAgentSession>(session);
var typedSession = (ChatClientAgentSession)session;
Assert.Same(mockContextProvider.Object, typedSession.AIContextProvider);
}
[Fact]
public async Task CreateSession_UsesChatHistoryProviderFactory_IfProvidedAsync()
{
// Arrange
var mockChatClient = new Mock<IChatClient>();
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
var factoryCalled = false;
var agent = new ChatClientAgent(mockChatClient.Object, new ChatClientAgentOptions
{
ChatOptions = new() { Instructions = "Test instructions" },
ChatHistoryProviderFactory = (_, _) =>
{
factoryCalled = true;
return new ValueTask<ChatHistoryProvider>(mockChatHistoryProvider.Object);
}
});
// Act
var session = await agent.CreateSessionAsync();
// Assert
Assert.True(factoryCalled, "ChatHistoryProviderFactory was not called.");
Assert.IsType<ChatClientAgentSession>(session);
var typedSession = (ChatClientAgentSession)session;
Assert.Same(mockChatHistoryProvider.Object, typedSession.ChatHistoryProvider);
}
[Fact]
public async Task CreateSession_UsesChatHistoryProvider_FromTypedOverloadAsync()
{
// Arrange
var mockChatClient = new Mock<IChatClient>();
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
var agent = new ChatClientAgent(mockChatClient.Object);
// Act
var session = await agent.CreateSessionAsync(mockChatHistoryProvider.Object);
// Assert
Assert.IsType<ChatClientAgentSession>(session);
var typedSession = (ChatClientAgentSession)session;
Assert.Same(mockChatHistoryProvider.Object, typedSession.ChatHistoryProvider);
}
[Fact]
public async Task CreateSession_UsesConversationId_FromTypedOverloadAsync()
{
@@ -1,80 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Text.Json;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
namespace Microsoft.Agents.AI.UnitTests;
/// <summary>
/// Contains unit tests for the ChatClientAgent.DeserializeSession methods.
/// </summary>
public class ChatClientAgent_DeserializeSessionTests
{
[Fact]
public async Task DeserializeSession_UsesAIContextProviderFactory_IfProvidedAsync()
{
// Arrange
var mockChatClient = new Mock<IChatClient>();
var mockContextProvider = new Mock<AIContextProvider>();
var factoryCalled = false;
var agent = new ChatClientAgent(mockChatClient.Object, new ChatClientAgentOptions
{
ChatOptions = new() { Instructions = "Test instructions" },
AIContextProviderFactory = (_, _) =>
{
factoryCalled = true;
return new ValueTask<AIContextProvider>(mockContextProvider.Object);
}
});
var json = JsonSerializer.Deserialize("""
{
"aiContextProviderState": ["CP1"]
}
""", TestJsonSerializerContext.Default.JsonElement);
// Act
var session = await agent.DeserializeSessionAsync(json);
// Assert
Assert.True(factoryCalled, "AIContextProviderFactory was not called.");
Assert.IsType<ChatClientAgentSession>(session);
var typedSession = (ChatClientAgentSession)session;
Assert.Same(mockContextProvider.Object, typedSession.AIContextProvider);
}
[Fact]
public async Task DeserializeSession_UsesChatHistoryProviderFactory_IfProvidedAsync()
{
// Arrange
var mockChatClient = new Mock<IChatClient>();
var mockChatHistoryProvider = new Mock<ChatHistoryProvider>();
var factoryCalled = false;
var agent = new ChatClientAgent(mockChatClient.Object, new ChatClientAgentOptions
{
ChatOptions = new() { Instructions = "Test instructions" },
ChatHistoryProviderFactory = (_, _) =>
{
factoryCalled = true;
return new ValueTask<ChatHistoryProvider>(mockChatHistoryProvider.Object);
}
});
var json = JsonSerializer.Deserialize("""
{
"chatHistoryProviderState": { }
}
""", TestJsonSerializerContext.Default.JsonElement);
// Act
var session = await agent.DeserializeSessionAsync(json);
// Assert
Assert.True(factoryCalled, "ChatHistoryProviderFactory was not called.");
Assert.IsType<ChatClientAgentSession>(session);
var typedSession = (ChatClientAgentSession)session;
Assert.Same(mockChatHistoryProvider.Object, typedSession.ChatHistoryProvider);
}
}
@@ -18,7 +18,6 @@ namespace Microsoft.Agents.AI.UnitTests.Data;
public sealed class TextSearchProviderTests
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
private readonly Mock<ILogger<TextSearchProvider>> _loggerMock;
private readonly Mock<ILoggerFactory> _loggerFactoryMock;
@@ -39,6 +38,28 @@ public sealed class TextSearchProviderTests
.Returns(true);
}
[Fact]
public void StateKey_ReturnsDefaultKey_WhenNoOptionsProvided()
{
// Arrange & Act
var provider = new TextSearchProvider((_, _) => Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]));
// Assert
Assert.Equal("TextSearchProvider", provider.StateKey);
}
[Fact]
public void StateKey_ReturnsCustomKey_WhenSetViaOptions()
{
// Arrange & Act
var provider = new TextSearchProvider(
(_, _) => Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]),
new TextSearchProviderOptions { StateKey = "custom-key" });
// Assert
Assert.Equal("custom-key", provider.StateKey);
}
[Theory]
[InlineData(null, null, true)]
[InlineData("Custom context prompt", "Custom citations prompt", false)]
@@ -64,15 +85,19 @@ public sealed class TextSearchProviderTests
ContextPrompt = overrideContextPrompt,
CitationsPrompt = overrideCitationsPrompt
};
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options, withLogging ? this._loggerFactoryMock.Object : null);
var provider = new TextSearchProvider(SearchDelegateAsync, options, withLogging ? this._loggerFactoryMock.Object : null);
var invokingContext = new AIContextProvider.InvokingContext(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "Sample user question?"),
new ChatMessage(ChatRole.User, "Additional part")
]);
new TestAgentSession(),
new AIContext
{
Messages = new List<ChatMessage>
{
new(ChatRole.User, "Sample user question?"),
new(ChatRole.User, "Additional part")
}
});
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -81,9 +106,15 @@ public sealed class TextSearchProviderTests
Assert.Equal("Sample user question?\nAdditional part", capturedInput);
Assert.Null(aiContext.Instructions); // TextSearchProvider uses a user message for context injection.
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
var message = aiContext.Messages!.Single();
var messages = aiContext.Messages!.ToList();
Assert.Equal(3, messages.Count); // 2 input messages + 1 search result message
Assert.Equal("Sample user question?", messages[0].Text);
Assert.Equal("Additional part", messages[1].Text);
Assert.Equal(AgentRequestMessageSourceType.External, messages[0].GetAgentRequestMessageSourceType());
Assert.Equal(AgentRequestMessageSourceType.External, messages[1].GetAgentRequestMessageSourceType());
var message = messages.Last();
Assert.Equal(ChatRole.User, message.Role);
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, message.GetAgentRequestMessageSourceType());
string text = message.Text!;
if (overrideContextPrompt is null)
@@ -143,17 +174,21 @@ public sealed class TextSearchProviderTests
FunctionToolName = overrideName,
FunctionToolDescription = overrideDescription
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, default, null, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages); // No automatic injection.
Assert.NotNull(aiContext.Messages); // Input messages are preserved.
var messages = aiContext.Messages!.ToList();
Assert.Single(messages);
Assert.Equal("Q?", messages[0].Text);
Assert.NotNull(aiContext.Tools);
Assert.Single(aiContext.Tools);
var tool = aiContext.Tools.Single();
var tools = aiContext.Tools!.ToList();
Assert.Single(tools);
var tool = tools[0];
Assert.Equal(expectedName, tool.Name);
Assert.Equal(expectedDescription, tool.Description);
}
@@ -162,14 +197,17 @@ public sealed class TextSearchProviderTests
public async Task InvokingAsync_ShouldNotThrow_WhenSearchFailsAsync()
{
// Arrange
var provider = new TextSearchProvider(this.FailingSearchAsync, default, null, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var provider = new TextSearchProvider(this.FailingSearchAsync, loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages);
Assert.NotNull(aiContext.Messages); // Input messages are preserved on error.
var messages = aiContext.Messages!.ToList();
Assert.Single(messages);
Assert.Equal("Q?", messages[0].Text);
Assert.Null(aiContext.Tools);
this._loggerMock.Verify(
l => l.Log(
@@ -203,7 +241,7 @@ public sealed class TextSearchProviderTests
ContextPrompt = overrideContextPrompt,
CitationsPrompt = overrideCitationsPrompt
};
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
// Act
var formatted = await provider.SearchAsync("Sample user question?", CancellationToken.None);
@@ -255,16 +293,18 @@ public sealed class TextSearchProviderTests
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
ContextFormatter = r => $"Custom formatted context with {r.Count} results."
};
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
Assert.Equal("Custom formatted context with 2 results.", aiContext.Messages![0].Text);
var messages = aiContext.Messages!.ToList();
Assert.Equal(2, messages.Count); // 1 input message + 1 formatted result message
Assert.Equal("Q?", messages[0].Text);
Assert.Equal("Custom formatted context with 2 results.", messages[1].Text);
}
[Fact]
@@ -289,16 +329,18 @@ public sealed class TextSearchProviderTests
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
ContextFormatter = r => string.Join(",", r.Select(x => ((RawPayload)x.RawRepresentation!).Id))
};
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.NotNull(aiContext.Messages);
Assert.Single(aiContext.Messages!);
Assert.Equal("R1,R2", aiContext.Messages![0].Text);
var messages = aiContext.Messages!.ToList();
Assert.Equal(2, messages.Count); // 1 input message + 1 formatted result message
Assert.Equal("Q?", messages[0].Text);
Assert.Equal("R1,R2", messages[1].Text);
}
[Fact]
@@ -306,18 +348,155 @@ public sealed class TextSearchProviderTests
{
// Arrange
var options = new TextSearchProviderOptions { SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke };
var provider = new TextSearchProvider(this.NoResultSearchAsync, default, null, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "Q?")]);
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Q?") } });
// Act
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
Assert.Null(aiContext.Messages);
Assert.NotNull(aiContext.Messages); // Input messages are preserved when no results found.
var messages = aiContext.Messages!.ToList();
Assert.Single(messages);
Assert.Equal("Q?", messages[0].Text);
Assert.Null(aiContext.Instructions);
Assert.Null(aiContext.Tools);
}
#region Message Filter Tests
[Fact]
public async Task InvokingAsync_DefaultFilter_ExcludesNonExternalMessagesFromSearchInputAsync()
{
// Arrange
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync);
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = requestMessages });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Only external messages should be used for search input
Assert.Equal("External message", capturedInput);
}
[Fact]
public async Task InvokingAsync_CustomSearchInputFilter_OverridesDefaultAsync()
{
// Arrange
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync, new TextSearchProviderOptions
{
SearchInputMessageFilter = messages => messages.Where(m => m.Role == ChatRole.System)
});
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "User message"),
new(ChatRole.System, "System message"),
};
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = requestMessages });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Custom filter keeps only System messages
Assert.Equal("System message", capturedInput);
}
[Fact]
public async Task InvokedAsync_DefaultFilter_ExcludesNonExternalMessagesFromStorageAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
RecentMessageMemoryLimit = 10,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.System]
};
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var session = new TestAgentSession();
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
// Store messages via InvokedAsync
await provider.InvokedAsync(new(s_mockAgent, session, requestMessages, []));
// Now invoke to read stored memory
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, session, new AIContext { Messages = [new ChatMessage(ChatRole.User, "Next")] });
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Only "External message" was stored in memory, so search input = "External message" + "Next"
Assert.Equal("External message\nNext", capturedInput);
}
[Fact]
public async Task InvokedAsync_CustomStorageInputFilter_OverridesDefaultAsync()
{
// Arrange
var options = new TextSearchProviderOptions
{
RecentMessageMemoryLimit = 10,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.System],
StorageInputMessageFilter = messages => messages // No filtering - store everything
};
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegateAsync(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var session = new TestAgentSession();
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
// Store messages via InvokedAsync
await provider.InvokedAsync(new(s_mockAgent, session, requestMessages, []));
// Now invoke to read stored memory
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, session, new AIContext { Messages = [new ChatMessage(ChatRole.User, "Next")] });
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Both messages stored (identity filter), so search input includes all + current
Assert.Equal("External message\nFrom history\nNext", capturedInput);
}
#endregion
#region Recent Message Memory Tests
[Fact]
@@ -335,7 +514,7 @@ public sealed class TextSearchProviderTests
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]); // No results needed.
}
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
// Populate memory with more messages than the limit (A,B,C,D) -> should retain B,C,D
var initialMessages = new[]
@@ -345,14 +524,14 @@ public sealed class TextSearchProviderTests
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
};
await provider.InvokedAsync(new(s_mockAgent, s_mockSession, initialMessages) { InvokeException = new InvalidOperationException("Request Failed") });
var session = new TestAgentSession();
await provider.InvokedAsync(new(s_mockAgent, session, initialMessages, new InvalidOperationException("Request Failed")));
var invokingContext = new AIContextProvider.InvokingContext(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "E")
]);
session,
new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "E") } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -377,7 +556,8 @@ public sealed class TextSearchProviderTests
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]); // No results needed.
}
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var session = new TestAgentSession();
// Populate memory with more messages than the limit (A,B,C,D) -> should retain B,C,D
var initialMessages = new[]
@@ -387,14 +567,12 @@ public sealed class TextSearchProviderTests
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
};
await provider.InvokedAsync(new(s_mockAgent, s_mockSession, initialMessages));
await provider.InvokedAsync(new(s_mockAgent, session, initialMessages, []));
var invokingContext = new AIContextProvider.InvokingContext(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "E")
]);
session,
new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "E") } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -419,28 +597,31 @@ public sealed class TextSearchProviderTests
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var session = new TestAgentSession();
// First memory update (A,B)
await provider.InvokedAsync(new(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B"),
]));
s_mockAgent,
session,
[
new ChatMessage(ChatRole.User, "A"),
new ChatMessage(ChatRole.Assistant, "B"),
],
[]));
// Second memory update (C,D,E)
await provider.InvokedAsync(new(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
new ChatMessage(ChatRole.User, "E"),
]));
s_mockAgent,
session,
[
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
new ChatMessage(ChatRole.User, "E"),
],
[]));
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "F")]);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, session, new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "F") } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -465,7 +646,8 @@ public sealed class TextSearchProviderTests
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]); // No results needed for this test.
}
var provider = new TextSearchProvider(SearchDelegateAsync, default, null, options);
var provider = new TextSearchProvider(SearchDelegateAsync, options);
var session = new TestAgentSession();
// Populate memory with mixed roles; only Assistant messages (A1,A2) should be retained.
var initialMessages = new[]
@@ -475,14 +657,12 @@ public sealed class TextSearchProviderTests
new ChatMessage(ChatRole.User, "U2"),
new ChatMessage(ChatRole.Assistant, "A2"),
};
await provider.InvokedAsync(new(s_mockAgent, s_mockSession, initialMessages));
await provider.InvokedAsync(new(s_mockAgent, session, initialMessages, []));
var invokingContext = new AIContextProvider.InvokingContext(
s_mockAgent,
s_mockSession,
[
new ChatMessage(ChatRole.User, "Question?") // Current request message always appended.
]);
session,
new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "Question?") } }); // Current request message always appended.
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -496,26 +676,7 @@ public sealed class TextSearchProviderTests
#region Serialization Tests
[Fact]
public void Serialize_WithNoRecentMessages_ShouldReturnEmptyState()
{
// Arrange
var options = new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, default, null, options);
// Act
var state = provider.Serialize();
// Assert
Assert.Equal(JsonValueKind.Object, state.ValueKind);
Assert.False(state.TryGetProperty("recentMessagesText", out _));
}
[Fact]
public async Task Serialize_WithRecentMessages_ShouldPersistMessagesUpToLimitAsync()
public async Task InvokedAsync_ShouldPersistMessagesToSessionStateBagAsync()
{
// Arrange
var options = new TextSearchProviderOptions
@@ -524,7 +685,8 @@ public sealed class TextSearchProviderTests
RecentMessageMemoryLimit = 3,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.Assistant]
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, default, null, options);
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var session = new TestAgentSession();
var messages = new[]
{
new ChatMessage(ChatRole.User, "M1"),
@@ -533,11 +695,12 @@ public sealed class TextSearchProviderTests
};
// Act
await provider.InvokedAsync(new(s_mockAgent, s_mockSession, messages)); // Populate recent memory.
var state = provider.Serialize();
await provider.InvokedAsync(new(s_mockAgent, session, messages, [])); // Populate recent memory.
// Assert
Assert.True(state.TryGetProperty("recentMessagesText", out var recentProperty));
// Assert - State should be in the session's StateBag
var stateBagSerialized = session.StateBag.Serialize();
Assert.True(stateBagSerialized.TryGetProperty("TextSearchProvider", out var stateProperty));
Assert.True(stateProperty.TryGetProperty("recentMessagesText", out var recentProperty));
Assert.Equal(JsonValueKind.Array, recentProperty.ValueKind);
var list = recentProperty.EnumerateArray().Select(e => e.GetString()).ToList();
Assert.Equal(3, list.Count);
@@ -545,7 +708,7 @@ public sealed class TextSearchProviderTests
}
[Fact]
public async Task SerializeAndDeserialize_RoundtripRestoresMessagesAsync()
public async Task StateBag_RoundtripRestoresMessagesAsync()
{
// Arrange
var options = new TextSearchProviderOptions
@@ -554,7 +717,8 @@ public sealed class TextSearchProviderTests
RecentMessageMemoryLimit = 4,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.Assistant]
};
var provider = new TextSearchProvider(this.NoResultSearchAsync, default, null, options);
var provider = new TextSearchProvider(this.NoResultSearchAsync, options);
var session = new TestAgentSession();
var messages = new[]
{
new ChatMessage(ChatRole.User, "A"),
@@ -562,23 +726,24 @@ public sealed class TextSearchProviderTests
new ChatMessage(ChatRole.User, "C"),
new ChatMessage(ChatRole.Assistant, "D"),
};
await provider.InvokedAsync(new(s_mockAgent, s_mockSession, messages));
await provider.InvokedAsync(new(s_mockAgent, session, messages, []));
// Act - Serialize and deserialize the StateBag
var serializedStateBag = session.StateBag.Serialize();
var restoredSession = new TestAgentSession(AgentSessionStateBag.Deserialize(serializedStateBag));
// Act
var state = provider.Serialize();
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
var roundTrippedProvider = new TextSearchProvider(SearchDelegate2Async, state, options: new TextSearchProviderOptions
var newProvider = new TextSearchProvider(SearchDelegate2Async, new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 4
});
var emptyMessages = Array.Empty<ChatMessage>();
await roundTrippedProvider.InvokingAsync(new(s_mockAgent, s_mockSession, emptyMessages), CancellationToken.None); // Trigger search to read memory.
await newProvider.InvokingAsync(new(s_mockAgent, restoredSession, new AIContext()), CancellationToken.None); // Trigger search to read memory.
// Assert
Assert.NotNull(capturedInput);
@@ -586,25 +751,10 @@ public sealed class TextSearchProviderTests
}
[Fact]
public async Task Deserialize_WithChangedLowerLimit_ShouldTruncateToNewLimitAsync()
public async Task InvokingAsync_WithEmptyStateBag_ShouldHaveNoMessagesAsync()
{
// Arrange
var initialProvider = new TextSearchProvider(this.NoResultSearchAsync, default, null, new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 5,
RecentMessageRolesIncluded = [ChatRole.User, ChatRole.Assistant]
});
var messages = new[]
{
new ChatMessage(ChatRole.User, "L1"),
new ChatMessage(ChatRole.Assistant, "L2"),
new ChatMessage(ChatRole.User, "L3"),
new ChatMessage(ChatRole.Assistant, "L4"),
new ChatMessage(ChatRole.User, "L5"),
};
await initialProvider.InvokedAsync(new(s_mockAgent, s_mockSession, messages));
var state = initialProvider.Serialize();
var session = new TestAgentSession(); // Fresh session with empty StateBag
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
@@ -614,43 +764,16 @@ public sealed class TextSearchProviderTests
}
// Act
var restoredProvider = new TextSearchProvider(SearchDelegate2Async, state, options: new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3 // Lower limit
});
await restoredProvider.InvokingAsync(new(s_mockAgent, s_mockSession, Array.Empty<ChatMessage>()), CancellationToken.None);
// Assert
Assert.NotNull(capturedInput);
Assert.Equal("L1\nL2\nL3", capturedInput);
}
[Fact]
public async Task Deserialize_WithEmptyState_ShouldHaveNoMessagesAsync()
{
// Arrange
var emptyState = JsonSerializer.Deserialize("{}", TestJsonSerializerContext.Default.JsonElement);
string? capturedInput = null;
Task<IEnumerable<TextSearchProvider.TextSearchResult>> SearchDelegate2Async(string input, CancellationToken ct)
{
capturedInput = input;
return Task.FromResult<IEnumerable<TextSearchProvider.TextSearchResult>>([]);
}
// Act
var provider = new TextSearchProvider(SearchDelegate2Async, emptyState, options: new TextSearchProviderOptions
var provider = new TextSearchProvider(SearchDelegate2Async, new TextSearchProviderOptions
{
SearchTime = TextSearchProviderOptions.TextSearchBehavior.BeforeAIInvoke,
RecentMessageMemoryLimit = 3
});
var emptyMessages = Array.Empty<ChatMessage>();
await provider.InvokingAsync(new(s_mockAgent, s_mockSession, emptyMessages), CancellationToken.None);
await provider.InvokingAsync(new(s_mockAgent, session, new AIContext()), CancellationToken.None);
// Assert
Assert.NotNull(capturedInput);
Assert.Equal(string.Empty, capturedInput); // No recent messages serialized => empty input.
Assert.Equal(string.Empty, capturedInput); // No recent messages in StateBag => empty input.
}
#endregion
@@ -669,4 +792,16 @@ public sealed class TextSearchProviderTests
{
public string Id { get; set; } = string.Empty;
}
private sealed class TestAgentSession : AgentSession
{
public TestAgentSession()
{
}
public TestAgentSession(AgentSessionStateBag stateBag)
{
this.StateBag = stateBag;
}
}
}
@@ -3,7 +3,6 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
@@ -19,7 +18,6 @@ namespace Microsoft.Agents.AI.Memory.UnitTests;
public class ChatHistoryMemoryProviderTests
{
private static readonly AIAgent s_mockAgent = new Mock<AIAgent>().Object;
private static readonly AgentSession s_mockSession = new Mock<AgentSession>().Object;
private readonly Mock<ILogger<ChatHistoryMemoryProvider>> _loggerMock;
private readonly Mock<ILoggerFactory> _loggerFactoryMock;
@@ -57,33 +55,82 @@ public class ChatHistoryMemoryProviderTests
.Returns(this._vectorStoreCollectionMock.Object);
}
[Fact]
public void StateKey_ReturnsDefaultKey_WhenNoOptionsProvided()
{
// Arrange & Act
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }));
// Assert
Assert.Equal("ChatHistoryMemoryProvider", provider.StateKey);
}
[Fact]
public void StateKey_ReturnsCustomKey_WhenSetViaOptions()
{
// Arrange & Act
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
new ChatHistoryMemoryProviderOptions { StateKey = "custom-key" });
// Assert
Assert.Equal("custom-key", provider.StateKey);
}
[Fact]
public void Constructor_Throws_ForNullVectorStore()
{
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(null!, "testcollection", 1, new ChatHistoryMemoryProviderScope() { UserId = "UID" }));
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(
null!,
"testcollection",
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" })));
}
[Fact]
public void Constructor_Throws_ForNullCollectionName()
{
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, null!, 1, new ChatHistoryMemoryProviderScope() { UserId = "UID" }));
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
null!,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" })));
}
[Fact]
public void Constructor_Throws_ForNullStorageScope()
public void Constructor_Throws_ForNullStateInitializer()
{
// Act & Assert
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, "testcollection", 1, null!));
Assert.Throws<ArgumentNullException>(() => new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
"testcollection",
1,
null!));
}
[Fact]
public void Constructor_Throws_ForInvalidVectorDimensions()
{
// Act & Assert
Assert.Throws<ArgumentOutOfRangeException>(() => new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, "testcollection", 0, new ChatHistoryMemoryProviderScope() { UserId = "UID" }));
Assert.Throws<ArgumentOutOfRangeException>(() => new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, "testcollection", -5, new ChatHistoryMemoryProviderScope() { UserId = "UID" }));
Assert.Throws<ArgumentOutOfRangeException>(() => new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
"testcollection",
0,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" })));
Assert.Throws<ArgumentOutOfRangeException>(() => new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
"testcollection",
-5,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" })));
}
#region InvokedAsync Tests
@@ -113,16 +160,17 @@ public class ChatHistoryMemoryProviderTests
UserId = "user1"
};
var provider = new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, TestCollectionName, 1, storeScope);
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(storeScope));
var requestMsgWithValues = new ChatMessage(ChatRole.User, "request text") { MessageId = "req-1", AuthorName = "user1", CreatedAt = new DateTimeOffset(new DateTime(2000, 1, 1), TimeSpan.Zero) };
var requestMsgWithNulls = new ChatMessage(ChatRole.User, "request text nulls");
var responseMsg = new ChatMessage(ChatRole.Assistant, "response text") { MessageId = "resp-1", AuthorName = "assistant" };
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [requestMsgWithValues, requestMsgWithNulls])
{
ResponseMessages = [responseMsg]
};
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsgWithValues, requestMsgWithNulls], [responseMsg]);
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
@@ -175,12 +223,9 @@ public class ChatHistoryMemoryProviderTests
this._vectorStoreMock.Object,
TestCollectionName,
1,
new ChatHistoryMemoryProviderScope() { UserId = "UID" });
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }));
var requestMsg = new ChatMessage(ChatRole.User, "request text") { MessageId = "req-1" };
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [requestMsg])
{
InvokeException = new InvalidOperationException("Invoke failed")
};
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsg], new InvalidOperationException("Invoke failed"));
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
@@ -203,10 +248,10 @@ public class ChatHistoryMemoryProviderTests
this._vectorStoreMock.Object,
TestCollectionName,
1,
new ChatHistoryMemoryProviderScope() { UserId = "UID" },
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
loggerFactory: this._loggerFactoryMock.Object);
var requestMsg = new ChatMessage(ChatRole.User, "request text") { MessageId = "req-1" };
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [requestMsg]);
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsg], []);
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
@@ -252,12 +297,12 @@ public class ChatHistoryMemoryProviderTests
this._vectorStoreMock.Object,
TestCollectionName,
1,
new ChatHistoryMemoryProviderScope { UserId = "user1" },
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "user1" }),
options: options,
loggerFactory: this._loggerFactoryMock.Object);
var requestMsg = new ChatMessage(ChatRole.User, "request text");
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, s_mockSession, [requestMsg]);
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), [requestMsg], []);
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
@@ -326,14 +371,14 @@ public class ChatHistoryMemoryProviderTests
this._vectorStoreMock.Object,
TestCollectionName,
1,
new ChatHistoryMemoryProviderScope() { UserId = "UID" },
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
options: providerOptions);
var requestMsg = new ChatMessage(ChatRole.User, "requesting relevant history");
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [requestMsg]);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { requestMsg } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
var aiContext = await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert
this._vectorStoreCollectionMock.Verify(
@@ -343,6 +388,12 @@ public class ChatHistoryMemoryProviderTests
It.IsAny<VectorSearchOptions<Dictionary<string, object?>>>(),
It.IsAny<CancellationToken>()),
Times.Once);
Assert.NotNull(aiContext.Messages);
var messages = aiContext.Messages.ToList();
Assert.Equal(2, messages.Count);
Assert.Equal(AgentRequestMessageSourceType.External, messages[0].GetAgentRequestMessageSourceType());
Assert.Equal(AgentRequestMessageSourceType.AIContextProvider, messages[1].GetAgentRequestMessageSourceType());
}
[Fact]
@@ -378,10 +429,15 @@ public class ChatHistoryMemoryProviderTests
})
.Returns(ToAsyncEnumerableAsync(new List<VectorSearchResult<Dictionary<string, object?>>>()));
var provider = new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, TestCollectionName, 1, options: providerOptions, storageScope: searchScope, searchScope: searchScope);
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(searchScope, searchScope),
options: providerOptions);
var requestMsg = new ChatMessage(ChatRole.User, "requesting relevant history");
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [requestMsg]);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { requestMsg } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -440,12 +496,11 @@ public class ChatHistoryMemoryProviderTests
this._vectorStoreMock.Object,
TestCollectionName,
1,
storageScope: scope,
searchScope: scope,
_ => new ChatHistoryMemoryProvider.State(scope, scope),
options: options,
loggerFactory: this._loggerFactoryMock.Object);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, s_mockSession, [new ChatMessage(ChatRole.User, "requesting relevant history")]);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = new List<ChatMessage> { new(ChatRole.User, "requesting relevant history") } });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
@@ -479,52 +534,178 @@ public class ChatHistoryMemoryProviderTests
#endregion
#region Serialization Tests
#region Message Filter Tests
[Fact]
public void Serialize_Deserialize_RoundtripsScopes()
public async Task InvokingAsync_DefaultFilter_ExcludesNonExternalMessagesFromSearchAsync()
{
// Arrange
var storageScope = new ChatHistoryMemoryProviderScope
var providerOptions = new ChatHistoryMemoryProviderOptions
{
ApplicationId = "app",
AgentId = "agent",
SessionId = "session",
UserId = "user"
SearchTime = ChatHistoryMemoryProviderOptions.SearchBehavior.BeforeAIInvoke,
};
var searchScope = new ChatHistoryMemoryProviderScope
string? capturedQuery = null;
this._vectorStoreCollectionMock
.Setup(c => c.SearchAsync(
It.IsAny<string>(),
It.IsAny<int>(),
It.IsAny<VectorSearchOptions<Dictionary<string, object?>>>(),
It.IsAny<CancellationToken>()))
.Callback<string, int, VectorSearchOptions<Dictionary<string, object?>>, CancellationToken>((query, _, _, _) => capturedQuery = query)
.Returns(ToAsyncEnumerableAsync(new List<VectorSearchResult<Dictionary<string, object?>>>()));
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
options: providerOptions);
var requestMessages = new List<ChatMessage>
{
ApplicationId = "app2",
AgentId = "agent2",
SessionId = "session2",
UserId = "user2"
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var provider = new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, TestCollectionName, 1, storageScope: storageScope, searchScope: searchScope);
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = requestMessages });
// Act
var stateElement = provider.Serialize();
await provider.InvokingAsync(invokingContext, CancellationToken.None);
using JsonDocument doc = JsonDocument.Parse(stateElement.GetRawText());
var storage = doc.RootElement.GetProperty("storageScope");
Assert.Equal("app", storage.GetProperty("applicationId").GetString());
Assert.Equal("agent", storage.GetProperty("agentId").GetString());
Assert.Equal("session", storage.GetProperty("sessionId").GetString());
Assert.Equal("user", storage.GetProperty("userId").GetString());
// Assert - Only External message used for search query
Assert.Equal("External message", capturedQuery);
}
var search = doc.RootElement.GetProperty("searchScope");
Assert.Equal("app2", search.GetProperty("applicationId").GetString());
Assert.Equal("agent2", search.GetProperty("agentId").GetString());
Assert.Equal("session2", search.GetProperty("sessionId").GetString());
Assert.Equal("user2", search.GetProperty("userId").GetString());
[Fact]
public async Task InvokingAsync_CustomSearchInputFilter_OverridesDefaultAsync()
{
// Arrange
var providerOptions = new ChatHistoryMemoryProviderOptions
{
SearchTime = ChatHistoryMemoryProviderOptions.SearchBehavior.BeforeAIInvoke,
SearchInputMessageFilter = messages => messages // No filtering
};
// Act - deserialize and serialize again
var provider2 = new ChatHistoryMemoryProvider(this._vectorStoreMock.Object, TestCollectionName, 1, serializedState: stateElement);
var stateElement2 = provider2.Serialize();
string? capturedQuery = null;
this._vectorStoreCollectionMock
.Setup(c => c.SearchAsync(
It.IsAny<string>(),
It.IsAny<int>(),
It.IsAny<VectorSearchOptions<Dictionary<string, object?>>>(),
It.IsAny<CancellationToken>()))
.Callback<string, int, VectorSearchOptions<Dictionary<string, object?>>, CancellationToken>((query, _, _, _) => capturedQuery = query)
.Returns(ToAsyncEnumerableAsync(new List<VectorSearchResult<Dictionary<string, object?>>>()));
// Assert - roundtrip the state
Assert.Equal(stateElement.GetRawText(), stateElement2.GetRawText());
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
options: providerOptions);
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
var invokingContext = new AIContextProvider.InvokingContext(s_mockAgent, new TestAgentSession(), new AIContext { Messages = requestMessages });
// Act
await provider.InvokingAsync(invokingContext, CancellationToken.None);
// Assert - Both messages should be included in search query (identity filter)
Assert.NotNull(capturedQuery);
Assert.Contains("External message", capturedQuery);
Assert.Contains("From history", capturedQuery);
}
[Fact]
public async Task InvokedAsync_DefaultFilter_ExcludesNonExternalMessagesFromStorageAsync()
{
// Arrange
var stored = new List<Dictionary<string, object?>>();
this._vectorStoreCollectionMock
.Setup(c => c.UpsertAsync(It.IsAny<IEnumerable<Dictionary<string, object?>>>(), It.IsAny<CancellationToken>()))
.Callback<IEnumerable<Dictionary<string, object?>>, CancellationToken>((items, ct) =>
{
if (items != null)
{
stored.AddRange(items);
}
})
.Returns(Task.CompletedTask);
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }));
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
new(ChatRole.System, "From context provider") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.AIContextProvider, "ContextSource") } } },
};
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
// Assert - Only External message + response stored (ChatHistory and AIContextProvider excluded by default)
Assert.Equal(2, stored.Count);
Assert.Equal("External message", stored[0]["Content"]);
Assert.Equal("Response", stored[1]["Content"]);
}
[Fact]
public async Task InvokedAsync_CustomStorageInputFilter_OverridesDefaultAsync()
{
// Arrange
var stored = new List<Dictionary<string, object?>>();
this._vectorStoreCollectionMock
.Setup(c => c.UpsertAsync(It.IsAny<IEnumerable<Dictionary<string, object?>>>(), It.IsAny<CancellationToken>()))
.Callback<IEnumerable<Dictionary<string, object?>>, CancellationToken>((items, ct) =>
{
if (items != null)
{
stored.AddRange(items);
}
})
.Returns(Task.CompletedTask);
var provider = new ChatHistoryMemoryProvider(
this._vectorStoreMock.Object,
TestCollectionName,
1,
_ => new ChatHistoryMemoryProvider.State(new ChatHistoryMemoryProviderScope { UserId = "UID" }),
options: new ChatHistoryMemoryProviderOptions
{
StorageInputMessageFilter = messages => messages // No filtering - store everything
});
var requestMessages = new List<ChatMessage>
{
new(ChatRole.User, "External message"),
new(ChatRole.System, "From history") { AdditionalProperties = new() { { AgentRequestMessageSourceAttribution.AdditionalPropertiesKey, new AgentRequestMessageSourceAttribution(AgentRequestMessageSourceType.ChatHistory, "HistorySource") } } },
};
var invokedContext = new AIContextProvider.InvokedContext(s_mockAgent, new TestAgentSession(), requestMessages, [new ChatMessage(ChatRole.Assistant, "Response")]);
// Act
await provider.InvokedAsync(invokedContext, CancellationToken.None);
// Assert - All messages stored (identity filter overrides default)
Assert.Equal(3, stored.Count);
Assert.Equal("External message", stored[0]["Content"]);
Assert.Equal("From history", stored[1]["Content"]);
Assert.Equal("Response", stored[2]["Content"]);
}
#endregion
@@ -537,4 +718,16 @@ public class ChatHistoryMemoryProviderTests
yield return update;
}
}
private sealed class TestAgentSession : AgentSession
{
public TestAgentSession()
{
}
public TestAgentSession(AgentSessionStateBag stateBag)
{
this.StateBag = stateBag;
}
}
}
@@ -15,4 +15,5 @@ namespace Microsoft.Agents.AI.UnitTests;
[JsonSerializable(typeof(string[]))]
[JsonSerializable(typeof(Dictionary<string, object?>))]
[JsonSerializable(typeof(ChatClientAgentSessionTests.Animal))]
[JsonSerializable(typeof(ChatClientAgentSession))]
internal sealed partial class TestJsonSerializerContext : JsonSerializerContext;
@@ -161,7 +161,7 @@ public class AgentWorkflowBuilderTests
}
}
private sealed class DoubleEchoAgentSession() : InMemoryAgentSession();
private sealed class DoubleEchoAgentSession() : AgentSession();
[Fact]
public async Task BuildConcurrent_AgentsRunInParallelAsync()
@@ -195,5 +195,5 @@ public class InProcessExecutionTests
/// <summary>
/// Simple session implementation for SimpleTestAgent.
/// </summary>
private sealed class SimpleTestAgentSession : InMemoryAgentSession;
private sealed class SimpleTestAgentSession : AgentSession;
}
@@ -46,5 +46,5 @@ internal sealed class RoleCheckAgent(bool allowOtherAssistantRoles, string? id =
};
}
private sealed class RoleCheckAgentSession : InMemoryAgentSession;
private sealed class RoleCheckAgentSession : AgentSession;
}
@@ -90,4 +90,4 @@ internal sealed class HelloAgent(string id = nameof(HelloAgent)) : AIAgent
}
}
internal sealed class HelloAgentSession() : InMemoryAgentSession();
internal sealed class HelloAgentSession() : AgentSession();
@@ -5,6 +5,7 @@ using System.Collections.Generic;
using System.Linq;
using System.Runtime.CompilerServices;
using System.Text.Json;
using System.Text.Json.Serialization;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
@@ -16,6 +17,8 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
protected override string? IdCore => id;
public override string? Name => name ?? base.Name;
public InMemoryChatHistoryProvider ChatHistoryProvider { get; } = new();
protected override async ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
return serializedState.Deserialize<EchoAgentSession>(jsonSerializerOptions) ?? await this.CreateSessionAsync(cancellationToken);
@@ -28,15 +31,15 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
throw new InvalidOperationException("The provided session is not compatible with the agent. Only sessions created by the agent can be serialized.");
}
return new(typedSession.Serialize(jsonSerializerOptions));
return new(JsonSerializer.SerializeToElement(typedSession, jsonSerializerOptions));
}
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default) =>
new(new EchoAgentSession());
private static ChatMessage UpdateSession(ChatMessage message, InMemoryAgentSession? session = null)
private ChatMessage UpdateSession(ChatMessage message, AgentSession? session = null)
{
session?.ChatHistoryProvider.Add(message);
this.ChatHistoryProvider.GetMessages(session).Add(message);
return message;
}
@@ -45,7 +48,7 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
{
foreach (ChatMessage message in messages)
{
UpdateSession(message, session as InMemoryAgentSession);
this.UpdateSession(message, session);
}
IEnumerable<ChatMessage> echoMessages
@@ -53,14 +56,14 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
where message.Role == ChatRole.User &&
!string.IsNullOrEmpty(message.Text)
select
UpdateSession(new ChatMessage(ChatRole.Assistant, $"{prefix}{message.Text}")
this.UpdateSession(new ChatMessage(ChatRole.Assistant, $"{prefix}{message.Text}")
{
AuthorName = this.Name ?? this.Id,
CreatedAt = DateTimeOffset.Now,
MessageId = Guid.NewGuid().ToString("N")
}, session as InMemoryAgentSession);
}, session);
return echoMessages.Concat(this.GetEpilogueMessages(options).Select(m => UpdateSession(m, session as InMemoryAgentSession)));
return echoMessages.Concat(this.GetEpilogueMessages(options).Select(m => this.UpdateSession(m, session)));
}
protected virtual IEnumerable<ChatMessage> GetEpilogueMessages(AgentRunOptions? options = null)
@@ -99,11 +102,11 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
}
}
private sealed class EchoAgentSession : InMemoryAgentSession
private sealed class EchoAgentSession : AgentSession
{
internal new JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
{
return base.Serialize(jsonSerializerOptions);
}
internal EchoAgentSession() { }
[JsonConstructor]
internal EchoAgentSession(AgentSessionStateBag stateBag) : base(stateBag) { }
}
}
@@ -104,5 +104,5 @@ public class TestReplayAgent(List<ChatMessage>? messages = null, string? id = nu
return candidateMessages;
}
private sealed class ReplayAgentSession() : InMemoryAgentSession();
private sealed class ReplayAgentSession() : AgentSession();
}
@@ -330,7 +330,7 @@ internal sealed class TestRequestAgent(TestAgentRequestType requestType, int unp
}
}
private sealed class TestRequestAgentSession<TRequest, TResponse> : InMemoryAgentSession
private sealed class TestRequestAgentSession<TRequest, TResponse> : AgentSession
where TRequest : AIContent
where TResponse : AIContent
{
@@ -343,19 +343,13 @@ internal sealed class TestRequestAgent(TestAgentRequestType requestType, int unp
public HashSet<string> ServicedRequests { get; } = new();
public HashSet<string> PairedRequests { get; } = new();
private static JsonElement DeserializeAndExtractState(JsonElement serializedState,
out TestRequestAgentSessionState state,
JsonSerializerOptions? jsonSerializerOptions = null)
public TestRequestAgentSession(JsonElement element, JsonSerializerOptions? jsonSerializerOptions = null)
{
state = JsonSerializer.Deserialize<TestRequestAgentSessionState>(serializedState, jsonSerializerOptions)
var state = JsonSerializer.Deserialize<TestRequestAgentSessionState>(element, jsonSerializerOptions)
?? throw new ArgumentException("Unable to deserialize session state.");
return state.SessionState;
}
this.StateBag = AgentSessionStateBag.Deserialize(state.SessionState);
public TestRequestAgentSession(JsonElement element, JsonSerializerOptions? jsonSerializerOptions = null)
: base(DeserializeAndExtractState(element, out TestRequestAgentSessionState state, jsonSerializerOptions))
{
this.UnservicedRequests = state.UnservicedRequests.ToDictionary(
keySelector: item => item.Key,
elementSelector: item => item.Value.As<TRequest>()!);
@@ -364,9 +358,9 @@ internal sealed class TestRequestAgent(TestAgentRequestType requestType, int unp
this.PairedRequests = state.PairedRequests;
}
protected override JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
internal JsonElement Serialize(JsonSerializerOptions? jsonSerializerOptions = null)
{
JsonElement sessionState = base.Serialize(jsonSerializerOptions);
JsonElement sessionState = this.StateBag.Serialize();
Dictionary<string, PortableValue> portableUnservicedRequests =
this.UnservicedRequests.ToDictionary(
@@ -32,18 +32,16 @@ public class WorkflowHostSmokeTests
{
private sealed class AlwaysFailsAIAgent(bool failByThrowing) : AIAgent
{
private sealed class Session : InMemoryAgentSession
private sealed class Session : AgentSession
{
public Session() { }
public Session(JsonElement serializedSession, JsonSerializerOptions? jsonSerializerOptions = null)
: base(serializedSession, jsonSerializerOptions)
{ }
public Session(AgentSessionStateBag stateBag) : base(stateBag) { }
}
protected override ValueTask<AgentSession> DeserializeSessionCoreAsync(JsonElement serializedState, JsonSerializerOptions? jsonSerializerOptions = null, CancellationToken cancellationToken = default)
{
return new(new Session(serializedState, jsonSerializerOptions));
return new(serializedState.Deserialize<Session>(jsonSerializerOptions)!);
}
protected override ValueTask<AgentSession> CreateSessionCoreAsync(CancellationToken cancellationToken = default)
@@ -30,14 +30,14 @@ public class OpenAIChatCompletionFixture : IChatClientAgentFixture
public async Task<List<ChatMessage>> GetChatHistoryAsync(AIAgent agent, AgentSession session)
{
var typedSession = (ChatClientAgentSession)session;
var chatHistoryProvider = agent.GetService<ChatHistoryProvider>();
if (typedSession.ChatHistoryProvider is null)
if (chatHistoryProvider is null)
{
return [];
}
return (await typedSession.ChatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
return (await chatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
}
public Task<ChatClientAgent> CreateChatClientAgentAsync(
@@ -50,12 +50,14 @@ public class OpenAIResponseFixture(bool store) : IChatClientAgentFixture
return [.. previousMessages, responseMessage];
}
if (typedSession.ChatHistoryProvider is null)
var chatHistoryProvider = agent.GetService<ChatHistoryProvider>();
if (chatHistoryProvider is null)
{
return [];
}
return (await typedSession.ChatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
return (await chatHistoryProvider.InvokingAsync(new(agent, session, []))).ToList();
}
private static ChatMessage ConvertToChatMessage(ResponseItem item)