Skip to content
Merged
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
45 changes: 28 additions & 17 deletions internal/agent/compaction.go
Original file line number Diff line number Diff line change
Expand Up @@ -381,6 +381,8 @@
// but its token COST must still be counted so usage reports and budgets include it.
onUsage func(Usage)
task *taskState
planner *contextPlanner
trace *trace.Recorder

// calibrationRatio scales the raw byte/4 token estimate toward the provider's
// real prompt-token count. ApproxTextTokens over-counts code-heavy content by
Expand Down Expand Up @@ -427,6 +429,8 @@
preserveLast: options.CompactionPreserveLast,
onUsage: options.OnUsage,
task: task,
planner: newContextPlanner(contextPlannerConfig{contextWindow: options.ContextWindow}),
trace: options.Trace,
}
return state
}
Expand Down Expand Up @@ -482,7 +486,7 @@

compacted, err := Compact(messages, CompactionOptions{
PreserveLast: state.preserveLast,
Summarize: summarizeClosure(ctx, provider, state.onUsage),
Summarize: state.summarizeClosure(ctx, provider),
taskState: state.task.snapshotForCompaction(messages),
})
if err != nil {
Expand Down Expand Up @@ -533,7 +537,7 @@

result, compactErr := Compact(messages, CompactionOptions{
PreserveLast: state.preserveLast,
Summarize: summarizeClosure(ctx, provider, state.onUsage),
Summarize: state.summarizeClosure(ctx, provider),
taskState: state.task.snapshotForCompaction(messages),
})
if compactErr != nil {
Expand Down Expand Up @@ -567,9 +571,12 @@
// provider call. The summary stream intentionally does NOT forward OnText (so
// compaction stays invisible on the user-facing surface), but it DOES forward
// OnUsage so the summarizer's token cost is still counted by usage/budgeting.
func summarizeClosure(ctx context.Context, provider Provider, onUsage func(Usage)) func([]zeroruntime.Message) (string, error) {
func (state *compactionState) summarizeClosure(ctx context.Context, provider Provider) func([]zeroruntime.Message) (string, error) {
if state.planner == nil {
state.planner = newContextPlanner(contextPlannerConfig{})
}
return func(toSummarize []zeroruntime.Message) (string, error) {
return summarizeWithFallback(ctx, provider, toSummarize, onUsage)
return summarizeWithPlanner(ctx, provider, toSummarize, state.onUsage, state.planner, state.trace)
}
}

Expand All @@ -580,8 +587,12 @@
// working when the elided middle is bigger than the summarizer's own context.
// Non-context-limit errors (and a single message that still won't fit) surface
// to the caller unchanged.
func summarizeWithFallback(ctx context.Context, provider Provider, messages []zeroruntime.Message, onUsage func(Usage)) (string, error) {

Check failure on line 590 in internal/agent/compaction.go

View workflow job for this annotation

GitHub Actions / Security & code health

unreachable func: summarizeWithFallback

Check failure on line 590 in internal/agent/compaction.go

View workflow job for this annotation

GitHub Actions / Smoke (windows-latest)

unreachable func: summarizeWithFallback
summary, err := summarizeMessagesOnce(ctx, provider, messages, onUsage)
return summarizeWithPlanner(ctx, provider, messages, onUsage, newContextPlanner(contextPlannerConfig{}), nil)
}

func summarizeWithPlanner(ctx context.Context, provider Provider, messages []zeroruntime.Message, onUsage func(Usage), planner *contextPlanner, recorder *trace.Recorder) (string, error) {
summary, err := summarizeMessagesOnce(ctx, provider, messages, onUsage, planner, recorder)
if err == nil {
return summary, nil
}
Expand All @@ -590,11 +601,11 @@
}

mid := len(messages) / 2
left, leftErr := summarizeWithFallback(ctx, provider, messages[:mid], onUsage)
left, leftErr := summarizeWithPlanner(ctx, provider, messages[:mid], onUsage, planner, recorder)
if leftErr != nil {
return "", leftErr
}
right, rightErr := summarizeWithFallback(ctx, provider, messages[mid:], onUsage)
right, rightErr := summarizeWithPlanner(ctx, provider, messages[mid:], onUsage, planner, recorder)
if rightErr != nil {
return "", rightErr
}
Expand All @@ -607,7 +618,7 @@
combined := strings.TrimSpace(left + "\n\n" + right)
reduced, reduceErr := summarizeMessagesOnce(ctx, provider, []zeroruntime.Message{
{Role: zeroruntime.MessageRoleUser, Content: combined},
}, onUsage)
}, onUsage, planner, recorder)
if reduceErr != nil {
if isContextLimitError(reduceErr.Error()) {
// Even the two combined partials don't fit (extreme): fall back to the
Expand All @@ -623,15 +634,15 @@
}

// summarizeMessagesOnce performs a single tool-less summarization call.
func summarizeMessagesOnce(ctx context.Context, provider Provider, messages []zeroruntime.Message, onUsage func(Usage)) (string, error) {
request := zeroruntime.CompletionRequest{
Messages: []zeroruntime.Message{
{Role: zeroruntime.MessageRoleSystem, Content: CompactionSummaryInstructions},
{Role: zeroruntime.MessageRoleUser, Content: "Summarize this conversation:\n\n" + renderTranscript(messages)},
},
// No tools: this is a plain text summarization call.
}
stream, err := provider.StreamCompletion(ctx, request)
func summarizeMessagesOnce(ctx context.Context, provider Provider, messages []zeroruntime.Message, onUsage func(Usage), planner *contextPlanner, recorder *trace.Recorder) (string, error) {
requestMessages := []zeroruntime.Message{
{Role: zeroruntime.MessageRoleSystem, Content: CompactionSummaryInstructions},
{Role: zeroruntime.MessageRoleUser, Content: "Summarize this conversation:\n\n" + renderTranscript(messages)},
}
parts := systemPromptParts{prompt: CompactionSummaryInstructions, baseInstructions: CompactionSummaryInstructions}
plan := planner.planWithPromptParts(requestMessages, nil, "", parts)
recordContextPlanTrace(recorder, plan)
stream, err := provider.StreamCompletion(ctx, plan.Request)
if err != nil {
return "", err
}
Expand Down
63 changes: 63 additions & 0 deletions internal/agent/compaction_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"testing"

"github.com/Gitlawb/zero/internal/tools"
"github.com/Gitlawb/zero/internal/trace"
"github.com/Gitlawb/zero/internal/zeroruntime"
)

Expand Down Expand Up @@ -428,6 +429,7 @@ type reactiveProvider struct {
summarizeCalls int
turnRequests int
failedOnce bool
connectError bool
bigText string
finalText string
}
Expand All @@ -449,6 +451,9 @@ func (provider *reactiveProvider) StreamCompletion(ctx context.Context, request
return streamEvents(toolTurnWithText(provider.bigText, "1", "read_file", `{"path":"x"}`)), nil
case provider.turnRequests == 2 && !provider.failedOnce:
provider.failedOnce = true
if provider.connectError {
return nil, errors.New("context_length_exceeded")
}
return streamEvents([]zeroruntime.StreamEvent{
{Type: zeroruntime.StreamEventError, Error: "This model's maximum context length is 1000 tokens. Please reduce the length of the messages."},
}), nil
Expand All @@ -457,6 +462,39 @@ func (provider *reactiveProvider) StreamCompletion(ctx context.Context, request
}
}

func TestRunConnectErrorCompactionRecordsReplacementPlan(t *testing.T) {
provider := &reactiveProvider{
connectError: true,
bigText: strings.Repeat("b", 6000),
finalText: "recovered",
}
registry := tools.NewRegistry()
registry.Register(tools.NewScopedReadFileTool(t.TempDir(), nil))
recorder := trace.NewRecorder("connect-compaction", "run-1", "test")
var contextPlans []ContextBreakdown

result, err := Run(context.Background(), strings.Repeat("z", 6000), provider, Options{
Registry: registry,
PermissionMode: PermissionModeUnsafe,
ContextWindow: 10_000_000,
CompactionPreserveLast: 2,
Trace: recorder,
OnContext: func(breakdown ContextBreakdown) {
contextPlans = append(contextPlans, breakdown)
},
})
if err != nil || result.FinalAnswer != "recovered" {
t.Fatalf("connect-error compaction failed: result=%#v err=%v", result, err)
}
if len(contextPlans) != 3 || contextPlans[2].MessageTokens >= contextPlans[1].MessageTokens {
t.Fatalf("replacement request context not reported: %#v", contextPlans)
}
gotTrace := recorder.Finish()
if len(gotTrace.PrefixHashes) != 4 || gotTrace.PrefixHashes[3].InvalidationReason != "unchanged" {
t.Fatalf("replacement request trace not recorded: %#v", gotTrace.PrefixHashes)
}
}

func TestRunReactiveCompactionRecovers(t *testing.T) {
provider := &reactiveProvider{
bigText: strings.Repeat("b", 6000),
Expand All @@ -468,6 +506,8 @@ func TestRunReactiveCompactionRecovers(t *testing.T) {
// bloats the history.
registry := tools.NewRegistry()
registry.Register(tools.NewScopedReadFileTool(t.TempDir(), nil))
recorder := trace.NewRecorder("reactive-compaction", "run-1", "test")
var contextPlans []ContextBreakdown

// ContextWindow large enough that proactive compaction never triggers, so
// only the reactive path can save the run.
Expand All @@ -476,6 +516,10 @@ func TestRunReactiveCompactionRecovers(t *testing.T) {
PermissionMode: PermissionModeUnsafe,
ContextWindow: 10_000_000,
CompactionPreserveLast: 2,
Trace: recorder,
OnContext: func(breakdown ContextBreakdown) {
contextPlans = append(contextPlans, breakdown)
},
})
if err != nil {
t.Fatalf("expected reactive compaction to recover the run, got error: %v", err)
Expand All @@ -486,6 +530,25 @@ func TestRunReactiveCompactionRecovers(t *testing.T) {
if provider.summarizeCalls == 0 {
t.Fatal("expected reactive compaction to invoke the summarizer")
}
if len(contextPlans) != 3 {
t.Fatalf("main request context callbacks = %d, want initial turns plus compacted retry", len(contextPlans))
}
if contextPlans[2].MessageTokens >= contextPlans[1].MessageTokens {
t.Fatalf("retry context was not updated after compaction: before=%d after=%d", contextPlans[1].MessageTokens, contextPlans[2].MessageTokens)
}
gotTrace := recorder.Finish()
if len(gotTrace.PrefixHashes) != 4 {
t.Fatalf("prefix evidence count = %d, want two turns, compaction, and retry: %#v", len(gotTrace.PrefixHashes), gotTrace.PrefixHashes)
}
if gotTrace.PrefixHashes[0].InvalidationReason != "initial" ||
gotTrace.PrefixHashes[1].InvalidationReason != "unchanged" ||
gotTrace.PrefixHashes[2].InvalidationReason != "initial" ||
gotTrace.PrefixHashes[3].InvalidationReason != "unchanged" {
t.Fatalf("unexpected request-specific invalidation evidence: %#v", gotTrace.PrefixHashes)
}
if gotTrace.PrefixHashes[2].CompletePrefixHash == gotTrace.PrefixHashes[1].CompletePrefixHash {
t.Fatalf("compaction fingerprint reused the main request prefix: %#v", gotTrace.PrefixHashes)
}
}

// midStreamReactiveProvider forwards some text BEFORE surfacing a context-limit
Expand Down
36 changes: 36 additions & 0 deletions internal/agent/context_measurement.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,24 @@ type ContextBreakdown struct {
TotalTokens int // SystemTokens + ToolTokens + MessageTokens
ContextWindow int // model context window; 0 when unknown
UsedFraction float64 // TotalTokens / ContextWindow; 0 when window unknown
// Blocks explain why each model-visible category is present without exposing
// its content. The planner currently preserves all three categories.
Blocks []ContextBlock
// CompletePrefixHash identifies the provider-visible system+tool prefix.
CompletePrefixHash string
// PrefixInvalidationReason is "initial", "unchanged", or a stable comma-
// separated list of changed prefix categories.
PrefixInvalidationReason string
}

// ContextBlock is content-free evidence for one model-visible context category.
type ContextBlock struct {
Kind string
Tokens int
CacheClass string
Authority string
Recoverable bool
Reason string
}

// MeasureContext estimates the per-category token footprint of a request: the
Expand All @@ -41,6 +59,24 @@ func MeasureContext(messages []zeroruntime.Message, tools []zeroruntime.ToolDefi
ContextWindow: contextWindow,
}
breakdown.TotalTokens = breakdown.SystemTokens + breakdown.ToolTokens + breakdown.MessageTokens
if breakdown.SystemTokens > 0 {
breakdown.Blocks = append(breakdown.Blocks, ContextBlock{
Kind: "system", Tokens: breakdown.SystemTokens, CacheClass: "stable_prefix",
Authority: "required", Reason: "instructions, workspace policy, and safety policy",
})
}
if breakdown.ToolTokens > 0 {
breakdown.Blocks = append(breakdown.Blocks, ContextBlock{
Kind: "tools", Tokens: breakdown.ToolTokens, CacheClass: "stable_prefix",
Authority: "capability", Reason: "currently permitted and selected tool definitions",
})
}
if breakdown.MessageTokens > 0 {
breakdown.Blocks = append(breakdown.Blocks, ContextBlock{
Kind: "conversation", Tokens: breakdown.MessageTokens, CacheClass: "dynamic_tail",
Authority: "transcript", Reason: "valid provider replay history and current task context",
})
}
if contextWindow > 0 {
breakdown.UsedFraction = float64(breakdown.TotalTokens) / float64(contextWindow)
}
Expand Down
3 changes: 3 additions & 0 deletions internal/agent/context_measurement_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,9 @@ func TestMeasureContextSplitsByCategory(t *testing.T) {
if breakdown.UsedFraction != wantFraction {
t.Fatalf("UsedFraction = %v, want %v", breakdown.UsedFraction, wantFraction)
}
if len(breakdown.Blocks) != 3 {
t.Fatalf("Blocks = %#v, want system/tools/conversation", breakdown.Blocks)
}
}

func TestMeasureContextUnknownWindowHasZeroFraction(t *testing.T) {
Expand Down
Loading
Loading