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
74 changes: 67 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,52 @@ 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
}
// A token boundary can land between an assistant tool call and its result.
// Do not send an orphaned leading result to the summarization model: providers
// generally require every tool result to follow its corresponding tool call.
for start < len(messages) && isToolResultMessage(messages[start]) {
start++
}

return messages[start:]
}

func isToolResultMessage[M adk.MessageType](msg M) bool {
switch m := any(msg).(type) {
case *schema.Message:
return m != nil && m.Role == schema.Tool
case *schema.AgenticMessage:
if m == nil || m.Role != schema.AgenticRoleTypeUser {
return false
}
for _, block := range m.ContentBlocks {
if block != nil && block.Type == schema.ContentBlockTypeFunctionToolResult {
return true
}
}
}
return false
}

func (m *TypedMiddleware[M]) getModelInstructions() (M, M) {
userInstruction := m.cfg.UserInstruction
if userInstruction == "" {
Expand Down
88 changes: 88 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,56 @@ 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))

t.Run("drops an orphaned message tool result at the boundary", func(t *testing.T) {
toolCall := schema.AssistantMessage("", []schema.ToolCall{{
ID: "call_1",
Function: schema.FunctionCall{
Name: "lookup",
},
}})
toolResult := schema.ToolMessage("result", "call_1")
messages := []adk.Message{toolCall, toolResult, recentMessage}
limit := estimateTypedMessageTokens(toolResult) + estimateTypedMessageTokens(recentMessage)

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

t.Run("drops an orphaned agentic tool result at the boundary", func(t *testing.T) {
toolCall := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolCall{CallID: "call_1", Name: "lookup"}),
},
}
toolResult := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolResult{CallID: "call_1", Name: "lookup"}),
},
}
recentMessage := schema.UserAgenticMessage("recent")
messages := []*schema.AgenticMessage{toolCall, toolResult, recentMessage}
limit := estimateTypedMessageTokens(toolResult) + estimateTypedMessageTokens(recentMessage)

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

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

Expand Down