From 8c31ae3b950dc9479cf55c28544269bffb5a463b Mon Sep 17 00:00:00 2001 From: zhangzherui Date: Fri, 24 Jul 2026 16:50:55 +0800 Subject: [PATCH] fix(adk): bound summarization model input --- .../summarization/summarization.go | 51 ++++++++++++++++--- .../summarization/summarization_test.go | 50 ++++++++++++++++++ 2 files changed, 94 insertions(+), 7 deletions(-) diff --git a/adk/middlewares/summarization/summarization.go b/adk/middlewares/summarization/summarization.go index 2d77c2e85..0ea819880 100644 --- a/adk/middlewares/summarization/summarization.go +++ b/adk/middlewares/summarization/summarization.go @@ -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, @@ -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 @@ -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 { @@ -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) @@ -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...) @@ -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 == "" { diff --git a/adk/middlewares/summarization/summarization_test.go b/adk/middlewares/summarization/summarization_test.go index d70f396b0..ba7298c4f 100644 --- a/adk/middlewares/summarization/summarization_test.go +++ b/adk/middlewares/summarization/summarization_test.go @@ -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{ @@ -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"), @@ -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()