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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 44 additions & 7 deletions adk/middlewares/summarization/summarization.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,11 @@ type TypedConfig[M adk.MessageType] struct {

// Trigger specifies the conditions that activate summarization.
// Optional. Defaults to triggering when total tokens exceed 160k.
//
// ContextTokens also bounds the default summarization model input. When
// history exceeds that limit, the default input builder keeps the newest
// contiguous messages that fit and evicts older messages. A custom
// GenModelInput retains full control and is not windowed.
Trigger *TriggerCondition

// EmitInternalEvents indicates whether internal events should be emitted during summarization,
Expand Down Expand Up @@ -372,7 +377,7 @@ func (m *TypedMiddleware[M]) shouldSummarize(ctx context.Context, input *TypedTo

func (m *TypedMiddleware[M]) getTriggerContextTokens() int {
const defaultTriggerContextTokens = 160000
if m.cfg.Trigger != nil {
if m.cfg.Trigger != nil && m.cfg.Trigger.ContextTokens > 0 {
return m.cfg.Trigger.ContextTokens
}
return defaultTriggerContextTokens
Expand Down Expand Up @@ -433,12 +438,7 @@ func defaultTypedTokenCounter[M adk.MessageType](_ context.Context, input *Typed

var incrementTokens int
for _, msg := range input.Messages[incrementStart:] {
switch m := any(msg).(type) {
case *schema.Message:
incrementTokens += estimateMessageTokens(m)
case *schema.AgenticMessage:
incrementTokens += estimateAgenticMessageTokens(m)
}
incrementTokens += estimateTypedMessageTokens(msg)
}

for _, tl := range input.Tools {
Expand Down Expand Up @@ -481,6 +481,17 @@ func estimateTokenBytes(tokens int) int {
return tokens * 4
}

func estimateTypedMessageTokens[M adk.MessageType](msg M) int {
switch m := any(msg).(type) {
case *schema.Message:
return estimateMessageTokens(m)
case *schema.AgenticMessage:
return estimateAgenticMessageTokens(m)
default:
return 0
}
}

func (m *TypedMiddleware[M]) summarize(ctx context.Context, originalMsgs []M) (M, []M, error) {
var zero M
_, contextMsgs := splitSystemAndContextMsgs(originalMsgs)
Expand Down Expand Up @@ -614,6 +625,9 @@ func (m *TypedMiddleware[M]) buildSummarizationModelInput(ctx context.Context, o
return input, nil
}

contextMsgs = windowMessagesByTokenLimit(contextMsgs, m.getTriggerContextTokens()-
estimateTypedMessageTokens(sysInstruction)-estimateTypedMessageTokens(userInstruction))

input := make([]M, 0, len(contextMsgs)+2)
input = append(input, sysInstruction)
input = append(input, contextMsgs...)
Expand All @@ -622,6 +636,29 @@ func (m *TypedMiddleware[M]) buildSummarizationModelInput(ctx context.Context, o
return input, nil
}

// windowMessagesByTokenLimit returns the newest contiguous message window that
// fits within tokenLimit. Keeping a suffix preserves the latest user intent and
// tool-call/result ordering while preventing an already-over-limit history from
// being sent unchanged to the summarization model.
func windowMessagesByTokenLimit[M adk.MessageType](messages []M, tokenLimit int) []M {
if tokenLimit <= 0 || len(messages) == 0 {
return nil
}

total := 0
start := len(messages)
for i := len(messages) - 1; i >= 0; i-- {
tokens := estimateTypedMessageTokens(messages[i])
if total+tokens > tokenLimit {
break
}
total += tokens
start = i
}

return messages[start:]
}

func (m *TypedMiddleware[M]) getModelInstructions() (M, M) {
userInstruction := m.cfg.UserInstruction
if userInstruction == "" {
Expand Down
50 changes: 50 additions & 0 deletions adk/middlewares/summarization/summarization_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -575,6 +575,23 @@ func TestMiddlewareShouldSummarize(t *testing.T) {
assert.False(t, triggered)
})

t.Run("message-only trigger does not enable zero token threshold", func(t *testing.T) {
mw := &TypedMiddleware[*schema.Message]{
cfg: &Config{
Trigger: &TriggerCondition{ContextMessages: 3},
},
}

triggered, err := mw.shouldSummarize(ctx, &TokenCounterInput{
Messages: []adk.Message{
schema.UserMessage("msg1"),
schema.UserMessage("msg2"),
},
})
assert.NoError(t, err)
assert.False(t, triggered)
})

t.Run("returns true when over threshold", func(t *testing.T) {
mw := &TypedMiddleware[*schema.Message]{
cfg: &Config{
Expand Down Expand Up @@ -1022,6 +1039,27 @@ func TestMiddlewareBuildSummarizationModelInput(t *testing.T) {
assert.True(t, found, "should contain context message")
})

t.Run("windows default input to context token threshold", func(t *testing.T) {
mw := &TypedMiddleware[*schema.Message]{
cfg: &Config{
Trigger: &TriggerCondition{},
},
}
systemInstruction, userInstruction := mw.getModelInstructions()
mw.cfg.Trigger.ContextTokens = estimateTypedMessageTokens(systemInstruction) +
estimateTypedMessageTokens(userInstruction) + 3

oldMessage := schema.UserMessage(strings.Repeat("old", 20))
recentMessage := schema.UserMessage("recent")
contextMsgs := []adk.Message{oldMessage, recentMessage}

input, err := mw.buildSummarizationModelInput(ctx, contextMsgs, contextMsgs)
require.NoError(t, err)
require.Len(t, input, 3)
assert.Same(t, recentMessage, input[1])
assert.NotContains(t, input, oldMessage)
})

t.Run("uses GenModelInput", func(t *testing.T) {
expectedInput := []adk.Message{
schema.UserMessage("custom input"),
Expand Down Expand Up @@ -1074,6 +1112,18 @@ func TestMiddlewareBuildSummarizationModelInput(t *testing.T) {
})
}

func TestWindowMessagesByTokenLimit(t *testing.T) {
oldMessage := schema.UserMessage(strings.Repeat("o", 80))
recentMessage := schema.UserMessage("recent")
messages := []adk.Message{oldMessage, recentMessage}

window := windowMessagesByTokenLimit(messages, estimateTypedMessageTokens(recentMessage))
require.Len(t, window, 1)
assert.Same(t, recentMessage, window[0])

assert.Nil(t, windowMessagesByTokenLimit(messages, 0))
}

func TestMiddlewareSummarize(t *testing.T) {
ctx := context.Background()

Expand Down