Compare commits

...
5 changed files with 270 additions and 38 deletions
@@ -69,13 +69,11 @@ public sealed class ChatReducerCompactionStrategy : CompactionStrategy
return false;
}
// Rebuild the index from the reduced messages
CompactionMessageIndex rebuilt = CompactionMessageIndex.Create(reducedMessages, index.Tokenizer);
index.Groups.Clear();
foreach (CompactionMessageGroup group in rebuilt.Groups)
{
index.Groups.Add(group);
}
// Rebuild the index from the reduced messages.
// Use Update() rather than directly manipulating Groups so that
// cached metrics (IncludedGroupCount, token counts, etc.) are
// properly invalidated.
index.Update(reducedMessages);
return true;
}
@@ -1,5 +1,6 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Diagnostics.CodeAnalysis;
using System.Text.Json.Serialization;
@@ -39,9 +40,15 @@ public sealed class CompactionMessageGroup
/// </remarks>
public static readonly string SummaryPropertyKey = "_is_summary";
private readonly Action _exclusionChangedCallback;
/// <summary>
/// Initializes a new instance of the <see cref="CompactionMessageGroup"/> class.
/// </summary>
/// <param name="exclusionChangedCallback">
/// A callback invoked when <see cref="IsExcluded"/> changes value.
/// Used internally by <see cref="CompactionMessageIndex"/> to invalidate cached aggregates.
/// </param>
/// <param name="kind">The kind of message group.</param>
/// <param name="messages">The messages in this group. The list is captured as a read-only snapshot.</param>
/// <param name="byteCount">The total UTF-8 byte count of the text content in the messages.</param>
@@ -50,8 +57,15 @@ public sealed class CompactionMessageGroup
/// The user turn this group belongs to, or <see langword="null"/> for <see cref="CompactionGroupKind.System"/>.
/// </param>
[JsonConstructor]
internal CompactionMessageGroup(CompactionGroupKind kind, IReadOnlyList<ChatMessage> messages, int byteCount, int tokenCount, int? turnIndex = null)
internal CompactionMessageGroup(
Action exclusionChangedCallback,
CompactionGroupKind kind,
IReadOnlyList<ChatMessage> messages,
int byteCount,
int tokenCount,
int? turnIndex = null)
{
this._exclusionChangedCallback = exclusionChangedCallback;
this.Kind = kind;
this.Messages = messages;
this.MessageCount = messages.Count;
@@ -107,7 +121,18 @@ public sealed class CompactionMessageGroup
/// Excluded groups are preserved in the collection for diagnostics or storage purposes
/// but are not included when calling <see cref="CompactionMessageIndex.GetIncludedMessages"/>.
/// </remarks>
public bool IsExcluded { get; set; }
public bool IsExcluded
{
get;
set
{
if (field != value)
{
field = value;
this._exclusionChangedCallback.Invoke();
}
}
}
/// <summary>
/// Gets or sets an optional reason explaining why this group was excluded.
@@ -27,6 +27,20 @@ public sealed class CompactionMessageIndex
private int _currentTurn;
private ChatMessage? _lastProcessedMessage;
// Cached values for derived properties — invalidated whenever groups are added/removed
// or a group's IsExcluded state changes.
private int? _cachedTotalMessageCount;
private int? _cachedTotalByteCount;
private int? _cachedTotalTokenCount;
private int? _cachedIncludedGroupCount;
private int? _cachedIncludedMessageCount;
private int? _cachedIncludedByteCount;
private int? _cachedIncludedTokenCount;
private int? _cachedTotalTurnCount;
private int? _cachedIncludedTurnCount;
private int? _cachedIncludedNonSystemGroupCount;
private int? _cachedRawMessageCount;
/// <summary>
/// Gets the list of message groups in this collection.
/// </summary>
@@ -123,6 +137,7 @@ public sealed class CompactionMessageIndex
this.Groups.Clear();
this._currentTurn = 0;
this._lastProcessedMessage = null;
this.InvalidateCache();
return;
}
@@ -184,13 +199,13 @@ public sealed class CompactionMessageIndex
if (message.Role == ChatRole.System)
{
// System messages are not part of any turn
this.Groups.Add(CreateGroup(CompactionGroupKind.System, [message], this.Tokenizer, turnIndex: null));
this.AddGroup(CompactionGroupKind.System, [message], turnIndex: null);
index++;
}
else if (message.Role == ChatRole.User)
{
this._currentTurn++;
this.Groups.Add(CreateGroup(CompactionGroupKind.User, [message], this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.User, [message], this._currentTurn);
index++;
}
else if (message.Role == ChatRole.Assistant && HasToolCalls(message))
@@ -207,11 +222,11 @@ public sealed class CompactionMessageIndex
index++;
}
this.Groups.Add(CreateGroup(CompactionGroupKind.ToolCall, groupMessages, this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.ToolCall, groupMessages, this._currentTurn);
}
else if (message.Role == ChatRole.Assistant && IsSummaryMessage(message))
{
this.Groups.Add(CreateGroup(CompactionGroupKind.Summary, [message], this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.Summary, [message], this._currentTurn);
index++;
}
else if (message.Role == ChatRole.Assistant && HasOnlyReasoning(message))
@@ -247,17 +262,17 @@ public sealed class CompactionMessageIndex
index++;
}
this.Groups.Add(CreateGroup(CompactionGroupKind.ToolCall, groupMessages, this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.ToolCall, groupMessages, this._currentTurn);
}
else
{
this.Groups.Add(CreateGroup(CompactionGroupKind.AssistantText, [message], this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.AssistantText, [message], this._currentTurn);
index++;
}
}
else
{
this.Groups.Add(CreateGroup(CompactionGroupKind.AssistantText, [message], this.Tokenizer, this._currentTurn));
this.AddGroup(CompactionGroupKind.AssistantText, [message], this._currentTurn);
index++;
}
}
@@ -266,6 +281,8 @@ public sealed class CompactionMessageIndex
{
this._lastProcessedMessage = messages[^1];
}
this.InvalidateCache();
}
/// <summary>
@@ -279,8 +296,9 @@ public sealed class CompactionMessageIndex
/// <returns>The newly created <see cref="CompactionMessageGroup"/>.</returns>
public CompactionMessageGroup InsertGroup(int index, CompactionGroupKind kind, IReadOnlyList<ChatMessage> messages, int? turnIndex = null)
{
CompactionMessageGroup group = CreateGroup(kind, messages, this.Tokenizer, turnIndex);
CompactionMessageGroup group = this.CreateGroup(kind, messages, this.Tokenizer, turnIndex); // %%% DERIVE TURNINDEX
this.Groups.Insert(index, group);
this.InvalidateCache();
return group;
}
@@ -294,8 +312,9 @@ public sealed class CompactionMessageIndex
/// <returns>The newly created <see cref="CompactionMessageGroup"/>.</returns>
public CompactionMessageGroup AddGroup(CompactionGroupKind kind, IReadOnlyList<ChatMessage> messages, int? turnIndex = null)
{
CompactionMessageGroup group = CreateGroup(kind, messages, this.Tokenizer, turnIndex);
CompactionMessageGroup group = this.CreateGroup(kind, messages, this.Tokenizer, turnIndex); // %%% DERIVE TURNINDEX
this.Groups.Add(group);
this.InvalidateCache();
return group;
}
@@ -320,57 +339,57 @@ public sealed class CompactionMessageIndex
/// <summary>
/// Gets the total number of messages across all groups, including excluded ones.
/// </summary>
public int TotalMessageCount => this.Groups.Sum(group => group.MessageCount);
public int TotalMessageCount => this._cachedTotalMessageCount ??= this.Groups.Sum(group => group.MessageCount);
/// <summary>
/// Gets the total UTF-8 byte count across all groups, including excluded ones.
/// </summary>
public int TotalByteCount => this.Groups.Sum(group => group.ByteCount);
public int TotalByteCount => this._cachedTotalByteCount ??= this.Groups.Sum(group => group.ByteCount);
/// <summary>
/// Gets the total token count across all groups, including excluded ones.
/// </summary>
public int TotalTokenCount => this.Groups.Sum(group => group.TokenCount);
public int TotalTokenCount => this._cachedTotalTokenCount ??= this.Groups.Sum(group => group.TokenCount);
/// <summary>
/// Gets the total number of groups that are not excluded.
/// </summary>
public int IncludedGroupCount => this.Groups.Count(group => !group.IsExcluded);
public int IncludedGroupCount => this._cachedIncludedGroupCount ??= this.Groups.Count(group => !group.IsExcluded);
/// <summary>
/// Gets the total number of messages across all included (non-excluded) groups.
/// </summary>
public int IncludedMessageCount => this.Groups.Where(group => !group.IsExcluded).Sum(group => group.MessageCount);
public int IncludedMessageCount => this._cachedIncludedMessageCount ??= this.Groups.Where(group => !group.IsExcluded).Sum(group => group.MessageCount);
/// <summary>
/// Gets the total UTF-8 byte count across all included (non-excluded) groups.
/// </summary>
public int IncludedByteCount => this.Groups.Where(group => !group.IsExcluded).Sum(group => group.ByteCount);
public int IncludedByteCount => this._cachedIncludedByteCount ??= this.Groups.Where(group => !group.IsExcluded).Sum(group => group.ByteCount);
/// <summary>
/// Gets the total token count across all included (non-excluded) groups.
/// </summary>
public int IncludedTokenCount => this.Groups.Where(group => !group.IsExcluded).Sum(group => group.TokenCount);
public int IncludedTokenCount => this._cachedIncludedTokenCount ??= this.Groups.Where(group => !group.IsExcluded).Sum(group => group.TokenCount);
/// <summary>
/// Gets the total number of user turns across all groups (including those with excluded groups).
/// </summary>
public int TotalTurnCount => this.Groups.Select(group => group.TurnIndex).Distinct().Count(turnIndex => turnIndex is not null && turnIndex > 0);
public int TotalTurnCount => this._cachedTotalTurnCount ??= this.Groups.Select(group => group.TurnIndex).Distinct().Count(turnIndex => turnIndex is not null && turnIndex > 0);
/// <summary>
/// Gets the number of user turns that have at least one non-excluded group.
/// </summary>
public int IncludedTurnCount => this.Groups.Where(group => !group.IsExcluded && group.TurnIndex is not null && group.TurnIndex > 0).Select(group => group.TurnIndex).Distinct().Count();
public int IncludedTurnCount => this._cachedIncludedTurnCount ??= this.Groups.Where(group => !group.IsExcluded && group.TurnIndex is not null && group.TurnIndex > 0).Select(group => group.TurnIndex).Distinct().Count();
/// <summary>
/// Gets the total number of groups across all included (non-excluded) groups that are not <see cref="CompactionGroupKind.System"/>.
/// </summary>
public int IncludedNonSystemGroupCount => this.Groups.Count(group => !group.IsExcluded && group.Kind != CompactionGroupKind.System);
public int IncludedNonSystemGroupCount => this._cachedIncludedNonSystemGroupCount ??= this.Groups.Count(group => !group.IsExcluded && group.Kind != CompactionGroupKind.System);
/// <summary>
/// Gets the total number of original messages (that are not summaries).
/// </summary>
public int RawMessageCount => this.Groups.Where(group => group.Kind != CompactionGroupKind.Summary).Sum(group => group.MessageCount);
public int RawMessageCount => this._cachedRawMessageCount ??= this.Groups.Where(group => group.Kind != CompactionGroupKind.Summary).Sum(group => group.MessageCount);
/// <summary>
/// Returns all groups that belong to the specified user turn.
@@ -379,6 +398,21 @@ public sealed class CompactionMessageIndex
/// <returns>The groups belonging to the turn, in order.</returns>
public IEnumerable<CompactionMessageGroup> GetTurnGroups(int turnIndex) => this.Groups.Where(group => group.TurnIndex == turnIndex);
private void InvalidateCache()
{
this._cachedTotalMessageCount = null;
this._cachedTotalByteCount = null;
this._cachedTotalTokenCount = null;
this._cachedIncludedGroupCount = null;
this._cachedIncludedMessageCount = null;
this._cachedIncludedByteCount = null;
this._cachedIncludedTokenCount = null;
this._cachedTotalTurnCount = null;
this._cachedIncludedTurnCount = null;
this._cachedIncludedNonSystemGroupCount = null;
this._cachedRawMessageCount = null;
}
/// <summary>
/// Computes the UTF-8 byte count for a set of messages across all content types.
/// </summary>
@@ -497,14 +531,14 @@ public sealed class CompactionMessageIndex
private static int GetStringByteCount(string? value) =>
value is { Length: > 0 } ? Encoding.UTF8.GetByteCount(value) : 0;
private static CompactionMessageGroup CreateGroup(CompactionGroupKind kind, IReadOnlyList<ChatMessage> messages, Tokenizer? tokenizer, int? turnIndex)
private CompactionMessageGroup CreateGroup(CompactionGroupKind kind, IReadOnlyList<ChatMessage> messages, Tokenizer? tokenizer, int? turnIndex)
{
int byteCount = ComputeByteCount(messages);
int tokenCount = tokenizer is not null
? ComputeTokenCount(messages, tokenizer)
: byteCount / 4;
return new CompactionMessageGroup(kind, messages, byteCount, tokenCount, turnIndex);
return new CompactionMessageGroup(this.InvalidateCache, kind, messages, byteCount, tokenCount, turnIndex);
}
private static bool HasToolCalls(ChatMessage message)
@@ -263,7 +263,7 @@ public class CompactionMessageIndexTests
public void MessageGroupStoresPassedCounts()
{
// Arrange & Act
CompactionMessageGroup group = new(CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Hello")], byteCount: 5, tokenCount: 2);
CompactionMessageGroup group = new(InvalidateCallback, CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Hello")], byteCount: 5, tokenCount: 2);
// Assert
Assert.Equal(1, group.MessageCount);
@@ -276,7 +276,7 @@ public class CompactionMessageIndexTests
{
// Arrange
IReadOnlyList<ChatMessage> messages = [new ChatMessage(ChatRole.User, "Hello")];
CompactionMessageGroup group = new(CompactionGroupKind.User, messages, byteCount: 5, tokenCount: 1);
CompactionMessageGroup group = new(InvalidateCallback, CompactionGroupKind.User, messages, byteCount: 5, tokenCount: 1);
// Assert — Messages is IReadOnlyList, not IList
Assert.IsType<IReadOnlyList<ChatMessage>>(group.Messages, exactMatch: false);
@@ -752,9 +752,9 @@ public class CompactionMessageIndexTests
public void ConstructorWithGroupsRestoresTurnIndex()
{
// Arrange — pre-existing groups with turn indices
CompactionMessageGroup group1 = new(CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Q1")], 2, 1, turnIndex: 1);
CompactionMessageGroup group2 = new(CompactionGroupKind.AssistantText, [new ChatMessage(ChatRole.Assistant, "A1")], 2, 1, turnIndex: 1);
CompactionMessageGroup group3 = new(CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Q2")], 2, 1, turnIndex: 2);
CompactionMessageGroup group1 = new(InvalidateCallback, CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Q1")], 2, 1, turnIndex: 1);
CompactionMessageGroup group2 = new(InvalidateCallback, CompactionGroupKind.AssistantText, [new ChatMessage(ChatRole.Assistant, "A1")], 2, 1, turnIndex: 1);
CompactionMessageGroup group3 = new(InvalidateCallback, CompactionGroupKind.User, [new ChatMessage(ChatRole.User, "Q2")], 2, 1, turnIndex: 2);
List<CompactionMessageGroup> groups = [group1, group2, group3];
// Act — constructor should restore _currentTurn from the last group's TurnIndex
@@ -789,7 +789,7 @@ public class CompactionMessageIndexTests
public void ConstructorWithGroupsWithoutTurnIndexSkipsRestore()
{
// Arrange — groups without turn indices (system messages)
CompactionMessageGroup systemGroup = new(CompactionGroupKind.System, [new ChatMessage(ChatRole.System, "Be helpful")], 10, 3, turnIndex: null);
CompactionMessageGroup systemGroup = new(InvalidateCallback, CompactionGroupKind.System, [new ChatMessage(ChatRole.System, "Be helpful")], 10, 3, turnIndex: null);
List<CompactionMessageGroup> groups = [systemGroup];
// Act — constructor won't find a TurnIndex to restore
@@ -1474,4 +1474,179 @@ public class CompactionMessageIndexTests
Assert.Equal(CompactionGroupKind.ToolCall, index.Groups[2].Kind);
Assert.Equal(3, index.Groups[2].MessageCount); // reasoning + toolCall + toolResult
}
// -----------------------------------------------------------------------
// Cache invalidation tests
// -----------------------------------------------------------------------
[Fact]
public void CachedMetricsAreInvalidatedWhenIsExcludedChanges()
{
// Arrange — two groups, read Included* properties (populates cache)
CompactionMessageIndex index = CompactionMessageIndex.Create(
[
new ChatMessage(ChatRole.User, "AAAA"), // 4 bytes, 1 token
new ChatMessage(ChatRole.Assistant, "BBBB"), // 4 bytes, 1 token
]);
// Prime the cache by reading properties
Assert.Equal(2, index.IncludedGroupCount);
Assert.Equal(2, index.IncludedMessageCount);
Assert.Equal(8, index.IncludedByteCount);
Assert.Equal(2, index.IncludedTokenCount);
Assert.Equal(1, index.IncludedTurnCount);
Assert.Equal(2, index.IncludedNonSystemGroupCount);
// Act — exclude both groups in turn 1; cache must be invalidated
index.Groups[0].IsExcluded = true;
index.Groups[1].IsExcluded = true;
// Assert — Included* properties now reflect the exclusions
Assert.Equal(0, index.IncludedGroupCount);
Assert.Equal(0, index.IncludedMessageCount);
Assert.Equal(0, index.IncludedByteCount);
Assert.Equal(0, index.IncludedTokenCount);
Assert.Equal(0, index.IncludedTurnCount); // turn 1 is fully excluded
Assert.Equal(0, index.IncludedNonSystemGroupCount);
}
[Fact]
public void CachedMetricsAreInvalidatedWhenGroupAddedViaUpdate()
{
// Arrange — start with two messages, read RawMessageCount (populates cache)
List<ChatMessage> messages =
[
new ChatMessage(ChatRole.User, "Q1"),
new ChatMessage(ChatRole.Assistant, "A1"),
];
CompactionMessageIndex index = CompactionMessageIndex.Create(messages);
// Prime the cache
Assert.Equal(2, index.RawMessageCount);
Assert.Equal(2, index.IncludedGroupCount);
// Act — add two more messages
messages.Add(new ChatMessage(ChatRole.User, "Q2"));
messages.Add(new ChatMessage(ChatRole.Assistant, "A2"));
index.Update(messages);
// Assert — cached values updated
Assert.Equal(4, index.RawMessageCount);
Assert.Equal(4, index.IncludedGroupCount);
}
[Fact]
public void CachedMetricsAreInvalidatedWhenGroupAddedViaAddGroup()
{
// Arrange — prime the cache
CompactionMessageIndex index = CompactionMessageIndex.Create(
[
new ChatMessage(ChatRole.User, "Hello"),
]);
Assert.Equal(1, index.IncludedGroupCount);
Assert.Equal(1, index.TotalMessageCount);
// Act — add a group manually
index.AddGroup(CompactionGroupKind.AssistantText, [new ChatMessage(ChatRole.Assistant, "Hi")], turnIndex: 1);
// Assert — cache invalidated and recomputed
Assert.Equal(2, index.IncludedGroupCount);
Assert.Equal(2, index.TotalMessageCount);
}
[Fact]
public void CachedMetricsAreInvalidatedWhenGroupInsertedViaInsertGroup()
{
// Arrange — prime the cache
CompactionMessageIndex index = CompactionMessageIndex.Create(
[
new ChatMessage(ChatRole.User, "Hello"),
]);
Assert.Equal(1, index.IncludedGroupCount);
Assert.Equal(1, index.TotalMessageCount);
// Act — insert a group
ChatMessage summaryMsg = new(ChatRole.Assistant, "[Summary]");
(summaryMsg.AdditionalProperties ??= [])[CompactionMessageGroup.SummaryPropertyKey] = true;
index.InsertGroup(0, CompactionGroupKind.Summary, [summaryMsg]);
// Assert — cache invalidated and recomputed
Assert.Equal(2, index.IncludedGroupCount);
Assert.Equal(2, index.TotalMessageCount);
Assert.Equal(1, index.RawMessageCount); // Summary excluded from RawMessageCount
}
[Fact]
public void CachedMetricsAreInvalidatedWhenIndexRebuiltByUpdate()
{
// Arrange — populate and prime cache
List<ChatMessage> messages =
[
new ChatMessage(ChatRole.User, "Q1"),
new ChatMessage(ChatRole.Assistant, "A1"),
];
CompactionMessageIndex index = CompactionMessageIndex.Create(messages);
index.Groups[0].IsExcluded = true;
// Prime the cache after exclusion
Assert.Equal(1, index.IncludedGroupCount);
// Act — update with a completely different (shorter) list — forces full rebuild
List<ChatMessage> newMessages =
[
new ChatMessage(ChatRole.User, "NewQ"),
];
index.Update(newMessages);
// Assert — full rebuild, previously excluded group is gone, cache is correct
Assert.Equal(1, index.IncludedGroupCount);
Assert.Equal(1, index.RawMessageCount);
}
[Fact]
public void CachedMetricsAreInvalidatedWhenUpdateClearsAllMessages()
{
// Arrange — prime the cache
CompactionMessageIndex index = CompactionMessageIndex.Create(
[
new ChatMessage(ChatRole.User, "Q1"),
new ChatMessage(ChatRole.Assistant, "A1"),
]);
Assert.Equal(2, index.IncludedGroupCount);
Assert.Equal(2, index.RawMessageCount);
// Act — update with empty list
index.Update([]);
// Assert — all cached metrics reset
Assert.Equal(0, index.IncludedGroupCount);
Assert.Equal(0, index.TotalGroupCount);
Assert.Equal(0, index.RawMessageCount);
Assert.Equal(0, index.IncludedTokenCount);
}
[Fact]
public void CachedTotalMetricsUnchangedWhenOnlyExcludedStateChanges()
{
// Arrange — prime the cache
CompactionMessageIndex index = CompactionMessageIndex.Create(
[
new ChatMessage(ChatRole.User, "AAAA"), // 4 bytes
new ChatMessage(ChatRole.Assistant, "BBBB"), // 4 bytes
]);
Assert.Equal(2, index.TotalGroupCount);
Assert.Equal(2, index.TotalMessageCount);
Assert.Equal(8, index.TotalByteCount);
// Act — exclude a group
index.Groups[0].IsExcluded = true;
// Assert — Total* properties still include both groups
Assert.Equal(2, index.TotalGroupCount);
Assert.Equal(2, index.TotalMessageCount);
Assert.Equal(8, index.TotalByteCount);
Assert.Equal(2, index.TotalTokenCount);
}
private static void InvalidateCallback() { }
}
@@ -356,7 +356,7 @@ public sealed class CompactionProviderTests
Assert.Empty(state.MessageGroups);
// Act
state.MessageGroups = [new CompactionMessageGroup(CompactionGroupKind.User, [], 0, 0, 0)];
state.MessageGroups = [new CompactionMessageGroup(() => { }, CompactionGroupKind.User, [], 0, 0, 0)];
// Assert
Assert.Single(state.MessageGroups);