diff --git a/README.md b/README.md index e4c0fce..366f831 100644 --- a/README.md +++ b/README.md @@ -51,7 +51,11 @@ go get github.com/hung12ct/gopheragent load only when a skill is used. - **Deterministic ReAct loop** — dependency-aware parallel tool scheduling with `` refs, anti-loop detection, token-budget-aware pruning. -- **Streaming & HITL** — SSE streaming, human approvals, plan mode, self-critique. +- **Streaming & HITL** — SSE streaming, human approvals, plan mode, and + self-critique that keeps the best-scoring pass instead of the last one. +- **Honest terminals** — a `context_trace` event says exactly which messages + pruning rewrote and why; a `degraded` terminal reports "the artifact landed, + the bookkeeping did not" instead of forcing a turn into success or failure. - **Custom tools** — one interface, schema derived from a Go struct; a middleware chain for logging, timing, rate limiting, and tracing. - **Multi-provider** — OpenAI, Anthropic, Gemini, Vertex, and OpenAI-compatible diff --git a/pkg/agent/constants.go b/pkg/agent/constants.go index 9bd8f67..101e034 100644 --- a/pkg/agent/constants.go +++ b/pkg/agent/constants.go @@ -18,4 +18,12 @@ const ( // budgetWarnRatio is the fraction of MaxTokenBudget at which // aggressive context pruning kicks in (truncate tool arguments). budgetWarnRatio = 0.85 + + // defaultProtectedEnds is how many trailing messages the routine + // depth prune leaves untouched on every LLM call. + defaultProtectedEnds = 3 + + // emergencyProtectedEnds is the shallower protection used once the + // estimate has blown past MaxTokenBudget outright. + emergencyProtectedEnds = 1 ) diff --git a/pkg/agent/context_trace.go b/pkg/agent/context_trace.go new file mode 100644 index 0000000..c33294c --- /dev/null +++ b/pkg/agent/context_trace.go @@ -0,0 +1,109 @@ +package agent + +import ( + "context" + + "github.com/hung12ct/gopheragent/pkg/history" +) + +// ContextChangeReason classifies why the loop rewrote a message before +// handing the conversation to the provider. It is a closed enum — the +// pruner has exactly three ways to shrink a message. +type ContextChangeReason string + +const ( + // ContextChangeSoftTrim: a long tool/assistant message had its middle + // cut out by the depth-based prune; head and tail survive. + ContextChangeSoftTrim ContextChangeReason = "soft-trim" + // ContextChangeOutlierDiscarded: the message blew past the outlier + // ceiling and its entire payload was replaced by a system notice. This + // is the only reason that loses content outright. + ContextChangeOutlierDiscarded ContextChangeReason = "outlier-discarded" + // ContextChangeArgsTruncated: a tool result was clipped to its first + // few hundred runes because the run crossed the token-budget warn + // threshold. + ContextChangeArgsTruncated ContextChangeReason = "args-truncated" +) + +// ContextPolicy names which enforceTokenBudget path produced a trace, so a +// host can tell routine depth pruning apart from budget pressure. +type ContextPolicy string + +const ( + // ContextPolicyDefault is the unbudgeted path — MaxTokenBudget is + // unset, so only the standard depth prune ran. + ContextPolicyDefault ContextPolicy = "default" + // ContextPolicyBudgetWarn: the estimate crossed budgetWarnRatio of + // MaxTokenBudget and tool arguments were truncated. + ContextPolicyBudgetWarn ContextPolicy = "budget-warn" + // ContextPolicyBudgetEmergency: the estimate exceeded MaxTokenBudget + // outright and the aggressive shallow prune ran. + ContextPolicyBudgetEmergency ContextPolicy = "budget-emergency" +) + +// ContextRef identifies one message the pruner rewrote for a single LLM +// call. Index is the position in the pre-prune slice; because no pruning +// path reorders or removes messages, it also indexes the post-prune +// slice — and, since the loop prunes a transient copy and persists the +// conversation at full fidelity, the same position in the slice returned +// by Sessions.History. That is what makes a trace joinable back to the +// stored transcript. +// +// ToolCallID and CorrelationID are the stable handles for naming a tool +// result across events — CorrelationID is the agent-generated ID also +// carried by ToolCallEvent, so a host can join a trimmed message back to +// the call that produced it. Both are empty for assistant messages. +// +// EstTokensBefore/After use the same 4-chars/token heuristic as +// MaxTokenBudget enforcement, so they are comparable with the event's +// totals but are not exact provider counts. +type ContextRef struct { + Index int `json:"index"` + Role string `json:"role"` + ToolCallID string `json:"tool_call_id,omitempty"` + CorrelationID string `json:"correlation_id,omitempty"` + Reason ContextChangeReason `json:"reason"` + EstTokensBefore int `json:"est_tokens_before"` + EstTokensAfter int `json:"est_tokens_after"` +} + +// contextRefFor builds a trace entry for msg at index i. estAfter is the +// post-rewrite estimate; the caller supplies it because only the rewriting +// branch knows the replacement content. +func contextRefFor(i int, msg history.Message, reason ContextChangeReason, estBefore, estAfter int) ContextRef { + return ContextRef{ + Index: i, + Role: msg.Role, + ToolCallID: msg.ToolCallID, + CorrelationID: msg.CorrelationID, + Reason: reason, + EstTokensBefore: estBefore, + EstTokensAfter: estAfter, + } +} + +// emitContextTrace fires a ContextTraceEvent describing what the pruner +// changed on the way into one LLM call. No-op when nothing was rewritten, +// which is the common case — a turn whose messages all fit emits nothing, +// and the two estimateTokens sweeps only run when there is something to +// report. +func (al *AgentLoop) emitContextTrace( + ctx context.Context, + sessionKey string, + streamChan chan<- StreamEvent, + iteration int, + policy ContextPolicy, + before, after []history.Message, + changes []ContextRef, +) { + if len(changes) == 0 { + return + } + al.emit(ctx, sessionKey, streamChan, Event(ContextTraceEvent{ + Policy: policy, + Iteration: iteration, + Changes: changes, + EstTokensBefore: estimateTokens(before), + EstTokensAfter: estimateTokens(after), + })) +} diff --git a/pkg/agent/context_trace_test.go b/pkg/agent/context_trace_test.go new file mode 100644 index 0000000..ef0ee75 --- /dev/null +++ b/pkg/agent/context_trace_test.go @@ -0,0 +1,158 @@ +package agent + +import ( + "context" + "strings" + "testing" + + "github.com/hung12ct/gopheragent/pkg/history" +) + +// collectTrace drains ch and returns every ContextTraceEvent payload seen. +func collectTrace(ch chan StreamEvent) []ContextTraceEvent { + close(ch) + var out []ContextTraceEvent + for ev := range ch { + if p, ok := ev.Payload.(ContextTraceEvent); ok { + out = append(out, p) + } + } + return out +} + +func TestEnforceTokenBudget_UnbudgetedPathEmitsTrace(t *testing.T) { + al := &AgentLoop{} + ch := make(chan StreamEvent, 16) + msgs := []history.Message{ + {Role: "tool", Content: strings.Repeat("x", softTrimThreshold+100), ToolCallID: "call-1", CorrelationID: "corr-1"}, + {Role: "user", Content: "hi"}, + {Role: "assistant", Content: "ok"}, + {Role: "user", Content: "again"}, + } + + al.enforceTokenBudget(context.Background(), "s", ch, 2, msgs) + + traces := collectTrace(ch) + if len(traces) != 1 { + t.Fatalf("want 1 trace on the unbudgeted path, got %d", len(traces)) + } + tr := traces[0] + if tr.Policy != ContextPolicyDefault { + t.Fatalf("policy = %q, want %q", tr.Policy, ContextPolicyDefault) + } + if tr.Iteration != 2 { + t.Fatalf("iteration = %d, want 2", tr.Iteration) + } + if len(tr.Changes) != 1 { + t.Fatalf("want 1 change, got %d", len(tr.Changes)) + } + c := tr.Changes[0] + if c.Index != 0 || c.Role != "tool" || c.Reason != ContextChangeSoftTrim { + t.Fatalf("change = %+v, want index 0 / tool / soft-trim", c) + } + if c.ToolCallID != "call-1" || c.CorrelationID != "corr-1" { + t.Fatalf("change lost its correlation handles: %+v", c) + } + if c.EstTokensAfter >= c.EstTokensBefore { + t.Fatalf("est tokens should shrink, got before=%d after=%d", c.EstTokensBefore, c.EstTokensAfter) + } + if tr.EstTokensAfter >= tr.EstTokensBefore { + t.Fatalf("totals should shrink, got before=%d after=%d", tr.EstTokensBefore, tr.EstTokensAfter) + } +} + +func TestEnforceTokenBudget_NoChangeEmitsNothing(t *testing.T) { + al := &AgentLoop{} + ch := make(chan StreamEvent, 16) + msgs := []history.Message{ + {Role: "user", Content: "hi"}, + {Role: "assistant", Content: "short"}, + {Role: "user", Content: "again"}, + } + + al.enforceTokenBudget(context.Background(), "s", ch, 0, msgs) + + if traces := collectTrace(ch); len(traces) != 0 { + t.Fatalf("want no trace when nothing was pruned, got %d", len(traces)) + } +} + +func TestEnforceTokenBudget_EmergencyPolicyTagged(t *testing.T) { + // MaxTokenBudget=1 puts the estimate past the ceiling outright, which + // skips the warn window (it only covers thresh < est <= budget) and + // goes straight to the shallow emergency prune. + al := &AgentLoop{MaxTokenBudget: 1} + ch := make(chan StreamEvent, 16) + msgs := []history.Message{ + {Role: "tool", Content: strings.Repeat("y", softTrimThreshold+100), ToolCallID: "call-1"}, + {Role: "user", Content: "hi"}, + } + + al.enforceTokenBudget(context.Background(), "s", ch, 0, msgs) + + traces := collectTrace(ch) + if len(traces) != 1 { + t.Fatalf("want 1 trace, got %d", len(traces)) + } + if traces[0].Policy != ContextPolicyBudgetEmergency { + t.Fatalf("policy = %q, want %q", traces[0].Policy, ContextPolicyBudgetEmergency) + } + if len(traces[0].Changes) != 1 || traces[0].Changes[0].Reason != ContextChangeSoftTrim { + t.Fatalf("changes = %+v, want a single soft-trim", traces[0].Changes) + } +} + +func TestEnforceTokenBudget_WarnPolicyTagsArgsTruncation(t *testing.T) { + // Size the budget so the estimate lands inside the warn window: + // thresh (0.85 * 1200 = 1020) < est (~1100) <= 1200. + al := &AgentLoop{MaxTokenBudget: 1200} + ch := make(chan StreamEvent, 16) + msgs := []history.Message{ + {Role: "tool", Content: strings.Repeat("y", 4400), ToolCallID: "call-1", CorrelationID: "corr-1"}, + {Role: "user", Content: "hi"}, + } + + al.enforceTokenBudget(context.Background(), "s", ch, 0, msgs) + + traces := collectTrace(ch) + if len(traces) != 1 { + t.Fatalf("want 1 trace, got %d", len(traces)) + } + if traces[0].Policy != ContextPolicyBudgetWarn { + t.Fatalf("policy = %q, want %q", traces[0].Policy, ContextPolicyBudgetWarn) + } + if len(traces[0].Changes) != 1 { + t.Fatalf("changes = %+v, want exactly one", traces[0].Changes) + } + c := traces[0].Changes[0] + if c.Reason != ContextChangeArgsTruncated || c.CorrelationID != "corr-1" { + t.Fatalf("change = %+v, want args-truncated on corr-1", c) + } +} + +func TestPruneContextMessages_OutlierReasonRecorded(t *testing.T) { + msgs := []history.Message{ + {Role: "tool", Content: strings.Repeat("z", outlierTrimThreshold+10)}, + {Role: "user", Content: "hi"}, + } + + _, changes := pruneContextMessages(msgs, 1) + + if len(changes) != 1 { + t.Fatalf("want 1 change, got %d", len(changes)) + } + if changes[0].Reason != ContextChangeOutlierDiscarded { + t.Fatalf("reason = %q, want %q", changes[0].Reason, ContextChangeOutlierDiscarded) + } +} + +func TestPruneContextMessages_NoChangeAllocatesNoTrace(t *testing.T) { + msgs := []history.Message{ + {Role: "tool", Content: "small"}, + {Role: "user", Content: "hi"}, + } + + if _, changes := pruneContextMessages(msgs, 1); changes != nil { + t.Fatalf("want nil trace when nothing changed, got %+v", changes) + } +} diff --git a/pkg/agent/degraded.go b/pkg/agent/degraded.go new file mode 100644 index 0000000..2824d40 --- /dev/null +++ b/pkg/agent/degraded.go @@ -0,0 +1,186 @@ +package agent + +import ( + "context" + "fmt" + "strings" + "sync" + + "github.com/hung12ct/gopheragent/pkg/tools" +) + +// ToolDegradation is one tool call that half-succeeded, as reported by a +// tools.Degradation and stamped with the tool that raised it. See +// tools.Degradation for the field semantics. +type ToolDegradation struct { + Tool string `json:"tool"` + Reason string `json:"reason"` + Artifacts []string `json:"artifacts,omitempty"` + Unreliable []string `json:"unreliable,omitempty"` +} + +// DegradedError is the error form of a partial-success terminal, for +// adopters that classify turns by error rather than by event type. It +// matches errors.Is(err, ErrDegraded). +// +// It is never returned from Run — a degraded turn produced a real answer, +// and surfacing it as a returned error would push callers to discard +// work that landed. It reaches adopters through DegradedEvent.Err. +type DegradedError struct { + Units []ToolDegradation +} + +func (e *DegradedError) Error() string { + names := make([]string, 0, len(e.Units)) + for _, u := range e.Units { + names = append(names, u.Tool) + } + return fmt.Sprintf("agent: turn completed with degraded state from %s", strings.Join(names, ", ")) +} + +// Is reports ErrDegraded so errors.Is matches without unwrapping. +func (e *DegradedError) Is(target error) bool { return target == ErrDegraded } + +// degradedKey is the ctx-value key for the per-Run degradation accumulator. +type degradedKey struct{} + +// degradedAcc collects the degradations raised by tools across every +// iteration of one Run. The mutex is load-bearing: tool waves execute in +// parallel goroutines, so several tools can degrade at once. +type degradedAcc struct { + mu sync.Mutex + units []ToolDegradation +} + +func (a *degradedAcc) add(u ToolDegradation) { + a.mu.Lock() + a.units = append(a.units, u) + a.mu.Unlock() +} + +// drain returns the accumulated degradations and clears the accumulator, +// so whichever terminal path fires first reports them and later paths +// (the Run-level deferred sweep) do not emit a duplicate. +func (a *degradedAcc) drain() []ToolDegradation { + a.mu.Lock() + defer a.mu.Unlock() + units := a.units + a.units = nil + return units +} + +func degradedAccFromContext(ctx context.Context) *degradedAcc { + v, _ := ctx.Value(degradedKey{}).(*degradedAcc) + return v +} + +// installDegradationAccumulator stashes a fresh accumulator on ctx and +// returns it alongside a sweep callback that emits any degradation no +// terminal path has claimed yet. +// +// Unlike installRunCostAccumulator, which skips the ctx allocation +// entirely when PriceTable is nil, this one always allocates: whether a +// tool will degrade is not knowable up front and there is no config knob +// to gate on. The cost is one small struct, one ctx value, and one +// closure per Run — deliberate, not an oversight. +// +// Caller pattern: +// +// ctx, sweepDegraded := al.installDegradationAccumulator(ctx, sessionKey, streamChan) +// defer sweepDegraded() +// +// The sweep exists so a Run that degrades and then dies on MaxIters or a +// fatal LLM error still reports which state went unreliable — that is +// exactly the run whose bookkeeping most needs repairing. +func (al *AgentLoop) installDegradationAccumulator(ctx context.Context, sessionKey string, streamChan chan<- StreamEvent) (context.Context, func()) { + ctx = context.WithValue(ctx, degradedKey{}, °radedAcc{}) + return ctx, func() { + al.emitDegradedIfAny(ctx, sessionKey, streamChan) + } +} + +// recordDegradation files a tool's partial-success report against the +// Run's accumulator. No-op when called outside a Run that installed one +// (sub-agent tools invoked directly in tests, for example). +func recordDegradation(ctx context.Context, toolName string, d *tools.Degradation) { + // Check d first: the speculative path calls this on every successful + // execution, and the overwhelming majority pass nil. Testing the + // pointer before walking the ctx chain keeps the common case free. + if d == nil { + return + } + acc := degradedAccFromContext(ctx) + if acc == nil { + return + } + acc.add(ToolDegradation{ + Tool: toolName, + Reason: d.Reason, + Artifacts: d.Artifacts, + Unreliable: d.Unreliable, + }) +} + +// applyDegradation annotates a half-succeeded tool result so the model +// does not redo the half that landed, and files the report against the +// Run's accumulator. Returns result unchanged when the call did not +// degrade, or when it degraded into an outright error — there the error +// is the story and the failure path already tells it. +// +// This is the single filing point for every consumed tool result, +// speculated or not, and it deliberately runs AFTER the OnToolResult +// hook chain: that hook can recover an error into a success or convert a +// success into an error, so deciding earlier would report a degradation +// the loop then hands to the model as a hard failure, or drop one the +// hook just recovered. The only other filer is +// reportOrphanedSpeculation, which covers entries this path can never +// see because they are discarded before being awaited. +func applyDegradation(ctx context.Context, toolName, result string, d *tools.Degradation, execErr error) string { + if d == nil || execErr != nil { + return result + } + recordDegradation(ctx, toolName, d) + return result + degradationNote(d) +} + +// emitDegradedIfAny drains the accumulator and emits a DegradedEvent when +// anything degraded this Run. Called immediately before DoneEvent on the +// final-answer path, and again from the Run-level defer to catch turns +// that ended on a cap or a fatal error. +func (al *AgentLoop) emitDegradedIfAny(ctx context.Context, sessionKey string, streamChan chan<- StreamEvent) { + acc := degradedAccFromContext(ctx) + if acc == nil { + return + } + units := acc.drain() + if len(units) == 0 { + return + } + al.emit(ctx, sessionKey, streamChan, Event(DegradedEvent{ + Units: units, + Err: &DegradedError{Units: units}, + })) +} + +// degradationNote renders the model-facing partial-success annotation +// appended to a degraded tool's result. It tells the model not to redo +// the work that landed, which is the failure mode a bare error would +// cause. +func degradationNote(d *tools.Degradation) string { + reason := d.Reason + if reason == "" { + reason = "the tool reported that part of its work did not complete" + } + var b strings.Builder + b.WriteString("\n\n[System: partial success] ") + b.WriteString(reason) + if len(d.Artifacts) > 0 { + b.WriteString("\nLanded and must NOT be retried or discarded: ") + b.WriteString(strings.Join(d.Artifacts, ", ")) + } + if len(d.Unreliable) > 0 { + b.WriteString("\nNow unreliable — treat as suspect and repair before relying on it: ") + b.WriteString(strings.Join(d.Unreliable, ", ")) + } + return b.String() +} diff --git a/pkg/agent/degraded_test.go b/pkg/agent/degraded_test.go new file mode 100644 index 0000000..8a836e1 --- /dev/null +++ b/pkg/agent/degraded_test.go @@ -0,0 +1,306 @@ +package agent + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "github.com/hung12ct/gopheragent/pkg/cache" + "github.com/hung12ct/gopheragent/pkg/tools" +) + +// halfSuccessTool writes its "artifact" fine and always fails the derived +// bookkeeping, which is the exact shape DegradedEvent exists for. +type halfSuccessTool struct { + name string + err error + cacheable bool +} + +func (t *halfSuccessTool) Descriptor() tools.ToolDescriptor { + return tools.ToolDescriptor{ + Name: t.name, + Description: "writes a report and updates the index", + Parameters: tools.ToolSchema{Type: "object"}, + Cacheable: t.cacheable, + } +} + +func (t *halfSuccessTool) Execute(context.Context, string) (tools.Result, error) { + // The degradation rides along even on the error branch: a tool that + // wrote its artifact and then failed hard still left the artifact + // behind. Whether that gets reported is the loop's decision, and it + // depends on the post-hook error state — which is exactly what the + // speculative and non-speculative paths must agree on. + res := tools.Result{ + Text: "report written to /reports/q3.md", + Degraded: &tools.Degradation{ + Reason: "report written but the search index update failed", + Artifacts: []string{"/reports/q3.md"}, + Unreliable: []string{"search_index"}, + }, + } + if t.err != nil { + return res, t.err + } + return res, nil +} + +// drainEvents runs one streaming turn and returns every event emitted. +func drainEvents(t *testing.T, loop *AgentLoop, sessionKey, msg string) []StreamEvent { + t.Helper() + var evs []StreamEvent + for ev := range loop.RunText(context.Background(), sessionKey, msg) { + evs = append(evs, ev) + } + return evs +} + +func TestDegraded_EmittedBeforeDoneOnFinalAnswer(t *testing.T) { + provider := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {Content: "the report is ready"}, + }} + loop, _ := setup(provider, &halfSuccessTool{name: "write_report"}) + + evs := drainEvents(t, loop, "s1", "write the q3 report") + + degradedAt, doneAt := -1, -1 + var payload DegradedEvent + for i, ev := range evs { + switch p := ev.Payload.(type) { + case DegradedEvent: + degradedAt, payload = i, p + case DoneEvent: + doneAt = i + } + } + if degradedAt < 0 { + t.Fatal("no DegradedEvent emitted for a tool that reported partial success") + } + if doneAt < 0 { + t.Fatal("DoneEvent must still fire — a degraded turn produced a real answer") + } + if degradedAt > doneAt { + t.Fatalf("DegradedEvent must precede DoneEvent, got indices %d and %d", degradedAt, doneAt) + } + if len(payload.Units) != 1 { + t.Fatalf("units = %+v, want exactly one", payload.Units) + } + u := payload.Units[0] + if u.Tool != "write_report" || len(u.Artifacts) != 1 || len(u.Unreliable) != 1 { + t.Fatalf("unit lost its detail: %+v", u) + } + if !errors.Is(payload.Err, ErrDegraded) { + t.Fatalf("Err = %v, want errors.Is(..., ErrDegraded)", payload.Err) + } +} + +func TestDegraded_NoteReachesTheModel(t *testing.T) { + provider := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {Content: "done"}, + }} + loop, sm := setup(provider, &halfSuccessTool{name: "write_report"}) + + if _, err := loop.RunIteration(context.Background(), "s1", "write it"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + msgs, err := sm.History(context.Background(), "s1") + if err != nil { + t.Fatalf("history: %v", err) + } + var toolContent string + for _, m := range msgs { + if m.Role == "tool" { + toolContent = m.Content + } + } + if !strings.Contains(toolContent, "[System: partial success]") { + t.Fatalf("tool result should carry the partial-success note, got %q", toolContent) + } + if !strings.Contains(toolContent, "/reports/q3.md") || !strings.Contains(toolContent, "search_index") { + t.Fatalf("note should name the artifact and the unreliable state, got %q", toolContent) + } +} + +func TestDegraded_SilentWhenNoToolDegrades(t *testing.T) { + provider := &scriptProvider{turns: []LLMResult{{Content: "nothing to do"}}} + loop, _ := setup(provider) + + for _, ev := range drainEvents(t, loop, "s1", "hi") { + if _, ok := ev.Payload.(DegradedEvent); ok { + t.Fatal("DegradedEvent must not fire on a clean turn") + } + } +} + +func TestDegraded_ToolErrorSuppressesDegradation(t *testing.T) { + // A tool that errors is not degraded — it failed, and the error path + // already tells that story. + provider := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {Content: "could not write it"}, + }} + loop, _ := setup(provider, &halfSuccessTool{name: "write_report", err: errors.New("disk full")}) + + for _, ev := range drainEvents(t, loop, "s1", "write it") { + if _, ok := ev.Payload.(DegradedEvent); ok { + t.Fatal("a failed tool call must not report a partial success") + } + } +} + +func TestDegraded_ReportedOnMaxItersTerminal(t *testing.T) { + // The Run degrades and then never reaches a final answer. The + // deferred sweep must still report the unreliable state. + loop, _ := setup(&scriptProvider{turns: []LLMResult{}}, &halfSuccessTool{name: "write_report"}) + loop.MaxIters = 2 + loop.LLM = &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {ToolCalls: []PendingToolCall{{ID: "c2", Name: "write_report", ArgsJSON: `{"n":2}`}}}, + }} + + var seen bool + for _, ev := range drainEvents(t, loop, "s1", "write it") { + if _, ok := ev.Payload.(DegradedEvent); ok { + seen = true + } + } + if !seen { + t.Fatal("a Run that degraded and then hit MaxIters must still report it") + } +} + +func TestDegradedAcc_DrainIsIdempotent(t *testing.T) { + acc := °radedAcc{} + acc.add(ToolDegradation{Tool: "a"}) + if got := acc.drain(); len(got) != 1 { + t.Fatalf("first drain = %+v, want one unit", got) + } + if got := acc.drain(); len(got) != 0 { + t.Fatalf("second drain = %+v, want empty so terminals cannot double-emit", got) + } +} + +func TestDegraded_ParallelWaveCollectsEveryUnit(t *testing.T) { + // Three half-success tools in one wave exercise degradedAcc's mutex. + // Run under -race; an unguarded append would be caught here. + provider := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{ + {ID: "c1", Name: "w1", ArgsJSON: `{}`}, + {ID: "c2", Name: "w2", ArgsJSON: `{}`}, + {ID: "c3", Name: "w3", ArgsJSON: `{}`}, + }}, + {Content: "all three ran"}, + }} + loop, _ := setup(provider, + &halfSuccessTool{name: "w1"}, &halfSuccessTool{name: "w2"}, &halfSuccessTool{name: "w3"}) + + var units []ToolDegradation + for _, ev := range drainEvents(t, loop, "s1", "run all three") { + if p, ok := ev.Payload.(DegradedEvent); ok { + units = append(units, p.Units...) + } + } + if len(units) != 3 { + t.Fatalf("want 3 degradations from a 3-tool wave, got %d: %+v", len(units), units) + } + seen := map[string]bool{} + for _, u := range units { + seen[u.Tool] = true + } + for _, name := range []string{"w1", "w2", "w3"} { + if !seen[name] { + t.Fatalf("missing degradation from %s: %+v", name, units) + } + } +} + +func TestDegraded_NotCachedSoAReplayCannotFakeSuccess(t *testing.T) { + // A cacheable tool that degrades must not populate the cache: a hit + // would replay the partial-success note to a later turn's model while + // the host sees no DegradedEvent at all. + tool := &halfSuccessTool{name: "write_report", cacheable: true} + newRun := func() (*AgentLoop, *cache.SearchCache) { + c := cache.NewSearchCache(10, time.Minute) + provider := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {Content: "done"}, + }} + loop, _ := setup(provider, tool) + loop.Cache = c + return loop, c + } + + loop, shared := newRun() + if evs := degradedUnits(drainEvents(t, loop, "s1", "write it")); len(evs) != 1 { + t.Fatalf("first run: want 1 degradation, got %d", len(evs)) + } + + // Second run reuses the same cache; the tool must execute again. + provider2 := &scriptProvider{turns: []LLMResult{ + {ToolCalls: []PendingToolCall{{ID: "c1", Name: "write_report", ArgsJSON: `{}`}}}, + {Content: "done"}, + }} + loop2, _ := setup(provider2, tool) + loop2.Cache = shared + if evs := degradedUnits(drainEvents(t, loop2, "s2", "write it")); len(evs) != 1 { + t.Fatalf("second run served a degraded result from cache — host saw %d degradations, want 1", len(evs)) + } +} + +// degradedUnits flattens every DegradedEvent payload in evs. +func degradedUnits(evs []StreamEvent) []ToolDegradation { + var out []ToolDegradation + for _, ev := range evs { + if p, ok := ev.Payload.(DegradedEvent); ok { + out = append(out, p.Units...) + } + } + return out +} + +// --- speculative execution --- + +func TestDegraded_SpeculatedResultFiledExactlyOnce(t *testing.T) { + // The tool runs in the drainer's speculation goroutine and its result + // is consumed by the wave executor. Exactly one filer must claim it: + // double-filing would show the operator two failures for one call. + loop, _ := setup( + &streamingToolCallReadyProvider{toolName: "write_report", argsJSON: `{}`}, + &halfSuccessTool{name: "write_report"}, + ) + loop.SpeculativeTools = true + + units := degradedUnits(drainEvents(t, loop, "s1", "write it")) + if len(units) != 1 { + t.Fatalf("want exactly 1 degradation from a speculated call, got %d: %+v", len(units), units) + } + if units[0].Tool != "write_report" { + t.Fatalf("unit = %+v, want write_report", units[0]) + } +} + +func TestDegraded_SpeculatedResultRecoveredByHookIsStillFiled(t *testing.T) { + // OnToolResult can recover an errored call into a success. The + // degradation must be judged against the POST-hook state: deciding at + // execution time reports nothing here, so the model is told the work + // half-landed while the host sees a clean turn. + loop, _ := setup( + &streamingToolCallReadyProvider{toolName: "write_report", argsJSON: `{}`}, + &halfSuccessTool{name: "write_report", err: errors.New("index update failed")}, + ) + loop.SpeculativeTools = true + loop.OnToolResult = func(_ context.Context, _, _, _, _ string, _ any, _ error) (string, error) { + return "report written, index repair queued", nil // error in -> recovered + } + + units := degradedUnits(drainEvents(t, loop, "s1", "write it")) + if len(units) != 1 { + t.Fatalf("hook recovered the call, so the degradation must be reported; got %d: %+v", len(units), units) + } +} diff --git a/pkg/agent/errors.go b/pkg/agent/errors.go index 0072695..54358b0 100644 --- a/pkg/agent/errors.go +++ b/pkg/agent/errors.go @@ -21,6 +21,12 @@ var ( // ErrLoopDetected is returned when the anti-loop detector terminates the cycle. ErrLoopDetected = errors.New("agent: infinite loop detected") + // ErrDegraded classifies a turn that produced a real answer while some + // tool's derived bookkeeping failed — neither a clean success nor a + // failure worth retrying. It is never returned from Run (the answer is + // good); adopters match it via errors.Is on DegradedEvent.Err. + ErrDegraded = errors.New("agent: turn completed with degraded state") + // ErrToolNotFound is returned when the LLM requests a tool that is not registered. ErrToolNotFound = errors.New("agent: tool not found") diff --git a/pkg/agent/event_types.go b/pkg/agent/event_types.go index d8c41de..9393afa 100644 --- a/pkg/agent/event_types.go +++ b/pkg/agent/event_types.go @@ -77,6 +77,15 @@ const ( // computed dollar cost under the configured PriceTable. Skipped // when no PriceTable is configured (zero-cost when unused). EventTypeRunCost StreamEventType = "run_cost" + // EventTypeContextTrace records what the pruner rewrote on the way + // into one LLM call. Emitted only when a prune actually changed + // something, so a turn whose context always fits stays silent. + EventTypeContextTrace StreamEventType = "context_trace" + // EventTypeDegraded reports that the turn finished with some tool's + // derived bookkeeping left inconsistent, even though the turn itself + // produced a real answer. Precedes the terminal event, never replaces + // it — a consumer that ignores it sees exactly today's behavior. + EventTypeDegraded StreamEventType = "degraded" ) // LimitKind enumerates the cap categories surfaced via LimitExhaustedEvent. @@ -227,9 +236,22 @@ func (DoneEvent) eventType() StreamEventType { return EventTypeDone } // ReflectedEvent delivers a post-critique canonical answer. Round indicates // which self-critique pass produced it (1-indexed); consumers typically keep // the last seen payload as the authoritative response. +// +// The event fires only for a round the loop actually adopted, so the +// last-seen-wins contract holds whether or not AgentLoop.Scorer is set. +// Score is the adopted round's rank under that Scorer, and is nil when no +// Scorer is configured — a rejected round emits a thought event naming +// its score instead. +// +// It is a pointer because zero is a legitimate score: a 0–100 rubric can +// return 0, and the Scorer docs offer negated latency as a valid unit +// where 0 is the best possible value. A bare float64 with omitempty would +// erase that score from the wire and make it indistinguishable from +// "unscored" in Go. type ReflectedEvent struct { - Text string `json:"text"` - Round int `json:"round"` + Text string `json:"text"` + Round int `json:"round"` + Score *float64 `json:"score,omitempty"` } func (ReflectedEvent) isEventPayload() {} @@ -381,6 +403,77 @@ type RunCostEvent struct { func (RunCostEvent) isEventPayload() {} func (RunCostEvent) eventType() StreamEventType { return EventTypeRunCost } +// ContextTraceEvent is the typed payload of EventTypeContextTrace. It is +// the answer to "why did the agent forget what I told it earlier" — +// context pruning runs before every LLM call and, without this, leaves no +// artifact behind. +// +// Emitted once per LLM call, and only when the pruner actually rewrote a +// message: a turn whose context comfortably fits produces no events at +// all. Changes lists every rewritten message with its reason, grouped by +// the pass that produced it — argument truncation first, then depth +// pruning — so Index is ascending within a group but not across the +// whole slice. A given Index appears at most once under the current +// thresholds +// (argument truncation clips to well under the soft-trim threshold, and +// only tool messages are eligible for both) — treat that as an +// observation, not a guarantee, and do not key a map on Index. +// +// Because the loop prunes a transient copy and leaves session history at +// full fidelity, the same long tool result is re-trimmed and re-reported +// on every iteration of a turn. Consumers that log these should expect +// near-duplicates across a multi-iteration run. +// +// EstTokensBefore/After are whole-conversation estimates that include +// tool-call arguments, so they will not equal the sum of the per-Change +// numbers, which cover message content only. Both use the 4-chars/token +// heuristic MaxTokenBudget enforcement uses — good for spotting how much +// a prune saved, not for billing. +// +// Iteration is the 0-indexed loop iteration the prune fed. +type ContextTraceEvent struct { + Policy ContextPolicy `json:"policy"` + Iteration int `json:"iteration"` + Changes []ContextRef `json:"changes"` + EstTokensBefore int `json:"est_tokens_before"` + EstTokensAfter int `json:"est_tokens_after"` +} + +func (ContextTraceEvent) isEventPayload() {} +func (ContextTraceEvent) eventType() StreamEventType { return EventTypeContextTrace } + +// DegradedEvent is the typed payload of EventTypeDegraded — the terminal +// state that had no representation before: the expensive artifact landed, +// the derived bookkeeping did not. Reporting it as Done would hide a real +// inconsistency; reporting it as an error would invite a retry that +// duplicates work that already succeeded. +// +// Fires only when at least one tool returned a tools.Degradation, and +// annotates the turn's terminal rather than replacing it, so existing +// consumers are unaffected. Position depends on how the turn ended: +// +// - Final answer: emitted immediately BEFORE DoneEvent. +// - Cap or fatal error (MaxIters, tool-call cap, LLM failure): emitted +// by a deferred Run-level sweep, so it arrives AFTER the +// LimitExhaustedEvent / ErrorEvent frame. +// +// The second case matters: a consumer that tears down on the first error +// frame will miss it. Drain the stream to completion if you need the +// degradation record from failed turns — which is exactly the turn whose +// unreliable state most needs repairing. +// +// Units lists every degradation of the Run in the order the tools raised +// them. Err carries the same information as a *DegradedError for adopters +// that classify by errors.Is(err, ErrDegraded); it is reconstructed from +// Units when the event is decoded from the wire. +type DegradedEvent struct { + Units []ToolDegradation `json:"units"` + Err error `json:"-"` +} + +func (DegradedEvent) isEventPayload() {} +func (DegradedEvent) eventType() StreamEventType { return EventTypeDegraded } + // HITLTimedOutEvent is the typed payload of EventTypeHITLTimedOut. Mirrors // HITLDeniedEvent and additionally carries the configured Timeout so a UI // can show "approval expired after 2m" without reaching back into agent @@ -519,6 +612,17 @@ func decodePayload(t StreamEventType, raw []byte) EventPayload { var p RunCostEvent _ = json.Unmarshal(raw, &p) return p + case EventTypeContextTrace: + var p ContextTraceEvent + _ = json.Unmarshal(raw, &p) + return p + case EventTypeDegraded: + var p DegradedEvent + _ = json.Unmarshal(raw, &p) + if p.Err == nil && len(p.Units) > 0 { + p.Err = &DegradedError{Units: p.Units} + } + return p default: return UnknownEvent{OriginalType: t, RawJSON: string(raw)} } @@ -553,6 +657,8 @@ type EventVisitor interface { VisitMemoryLoaded(MemoryLoadedEvent) VisitMemoryConsolidated(MemoryConsolidatedEvent) VisitRunCost(RunCostEvent) + VisitContextTrace(ContextTraceEvent) + VisitDegraded(DegradedEvent) VisitUnknown(UnknownEvent) } @@ -602,6 +708,10 @@ func (ev StreamEvent) Visit(v EventVisitor) { v.VisitMemoryConsolidated(p) case RunCostEvent: v.VisitRunCost(p) + case ContextTraceEvent: + v.VisitContextTrace(p) + case DegradedEvent: + v.VisitDegraded(p) case UnknownEvent: v.VisitUnknown(p) default: diff --git a/pkg/agent/event_types_test.go b/pkg/agent/event_types_test.go index 7a4930f..8820b12 100644 --- a/pkg/agent/event_types_test.go +++ b/pkg/agent/event_types_test.go @@ -77,6 +77,14 @@ func TestStreamEvent_RoundTripsThroughJSON(t *testing.T) { Event(RegeneratedEvent{PreviousAssistantIndex: 3, TruncatedAt: 2}), Event(ContinuedEvent{ContinuedFromIndex: 7}), Event(TaskListEvent{Tasks: []TaskListItem{{ID: "t1", Title: "one", Status: "pending"}}}), + Event(ContextTraceEvent{ + Policy: ContextPolicyBudgetWarn, + Iteration: 2, + Changes: []ContextRef{{Index: 1, Role: "tool", Reason: ContextChangeArgsTruncated}}, + EstTokensBefore: 900, + EstTokensAfter: 120, + }), + Event(DegradedEvent{Units: []ToolDegradation{{Tool: "write_report", Reason: "index failed"}}}), } for _, ev := range cases { data, err := json.Marshal(ev) @@ -171,8 +179,10 @@ func (r *recordingVisitor) VisitMemoryLoaded(MemoryLoadedEvent) { r.visited func (r *recordingVisitor) VisitMemoryConsolidated(MemoryConsolidatedEvent) { r.visited = "memory_consolidated" } -func (r *recordingVisitor) VisitRunCost(RunCostEvent) { r.visited = "run_cost" } -func (r *recordingVisitor) VisitUnknown(UnknownEvent) { r.visited = "unknown" } +func (r *recordingVisitor) VisitRunCost(RunCostEvent) { r.visited = "run_cost" } +func (r *recordingVisitor) VisitContextTrace(ContextTraceEvent) { r.visited = "context_trace" } +func (r *recordingVisitor) VisitDegraded(DegradedEvent) { r.visited = "degraded" } +func (r *recordingVisitor) VisitUnknown(UnknownEvent) { r.visited = "unknown" } func TestVisit_DispatchesToMatchingMethod(t *testing.T) { cases := []struct { @@ -207,3 +217,55 @@ func TestVisit_DispatchesToMatchingMethod(t *testing.T) { } } } + +func TestDegradedEvent_ErrReconstructedFromWire(t *testing.T) { + // Err has json:"-", so a consumer decoding an SSE frame relies on + // decodePayload rebuilding it — that is the whole point of the field. + wire := []byte(`{"type":"degraded","payload":{"units":[{"tool":"write_report","reason":"index update failed","unreliable":["search_index"]}]}}`) + var got StreamEvent + if err := json.Unmarshal(wire, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + p, ok := got.Payload.(DegradedEvent) + if !ok { + t.Fatalf("expected DegradedEvent, got %T", got.Payload) + } + if len(p.Units) != 1 || p.Units[0].Tool != "write_report" { + t.Fatalf("units lost in transit: %+v", p.Units) + } + if !errors.Is(p.Err, ErrDegraded) { + t.Fatalf("Err = %v, want errors.Is(..., ErrDegraded) after decode", p.Err) + } +} + +func TestReflectedEvent_ZeroScoreSurvivesTheWire(t *testing.T) { + // A 0 score is legitimate (0-100 rubric, negated latency). A bare + // float64 with omitempty would erase it and make a scored round + // indistinguishable from an unscored one. + zero := 0.0 + data, err := json.Marshal(Event(ReflectedEvent{Text: "x", Round: 1, Score: &zero})) + if err != nil { + t.Fatalf("marshal: %v", err) + } + var got StreamEvent + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + p := got.Payload.(ReflectedEvent) + if p.Score == nil { + t.Fatalf("zero score erased by the wire format: %s", data) + } + if *p.Score != 0 { + t.Fatalf("score = %v, want 0", *p.Score) + } + + // An unscored round stays absent, so the two remain distinguishable. + unscored, _ := json.Marshal(Event(ReflectedEvent{Text: "x", Round: 1})) + var got2 StreamEvent + if err := json.Unmarshal(unscored, &got2); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if got2.Payload.(ReflectedEvent).Score != nil { + t.Fatal("an unscored round must decode with a nil Score") + } +} diff --git a/pkg/agent/llm_call.go b/pkg/agent/llm_call.go index 58adaa2..9132997 100644 --- a/pkg/agent/llm_call.go +++ b/pkg/agent/llm_call.go @@ -56,29 +56,13 @@ func (al *AgentLoop) callLLMWithRetry(ctx context.Context, st *iterationState, m func (al *AgentLoop) handleFinalAnswer(ctx context.Context, st *iterationState, msgs []history.Message, finalContent string) { msgs = append(msgs, history.Message{Role: "assistant", Content: finalContent}) if al.Reflect > 0 && finalContent != "" { - for r := 1; r <= al.Reflect; r++ { - al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ - Message: fmt.Sprintf("Self-critique pass %d/%d...", r, al.Reflect), - })) - revised, rerr := al.reflectOnce(ctx, st.sessionKey, msgs, r, st.streamChan) - if rerr != nil { - al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ - Message: fmt.Sprintf("Self-critique aborted: %v", rerr), - })) - break - } - if revised == "" || revised == finalContent { - continue - } - finalContent = revised - msgs[len(msgs)-1].Content = finalContent - al.emit(ctx, st.sessionKey, st.streamChan, Event(ReflectedEvent{ - Text: finalContent, - Round: r, - })) - } + msgs[len(msgs)-1].Content = al.runReflectRounds(ctx, st, msgs, finalContent) } al.saveSession(ctx, st.sessionKey, msgs) + // Degraded precedes Done rather than replacing it: the answer is real + // and consumers that gate on Done must still see it. No-op unless a + // tool reported a partial success this Run. + al.emitDegradedIfAny(ctx, st.sessionKey, st.streamChan) al.emit(ctx, st.sessionKey, st.streamChan, Event(DoneEvent{})) } @@ -115,6 +99,10 @@ func (al *AgentLoop) callLLM(ctx context.Context, st *iterationState, msgs []his st.specMu.Lock() for k, sm := range st.specMap { sm.cancel() + // This entry will never be awaited, so the wave executor's + // post-hook filing cannot happen for it. Report any degradation + // now — the tool already ran and its side effects are real. + reportOrphanedSpeculation(ctx, sm) delete(st.specMap, k) } st.specMu.Unlock() diff --git a/pkg/agent/loop_execute.go b/pkg/agent/loop_execute.go index c785f74..1b2ee47 100644 --- a/pkg/agent/loop_execute.go +++ b/pkg/agent/loop_execute.go @@ -210,12 +210,13 @@ func (al *AgentLoop) executeToolCall(ctx context.Context, st *iterationState, ws var toolResult string var structured any var execErr error + var degraded *tools.Degradation if speculated { al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{Message: fmt.Sprintf("Reusing speculative result for %s.", tCall.Name)})) - toolResult, structured, execErr = awaitSpeculative(toolCtx, sm) + toolResult, structured, degraded, execErr = awaitSpeculative(toolCtx, sm) } else { res, err := tool.Execute(toolCtx, tCall.ArgsJSON) - toolResult, structured, execErr = res.Text, res.Structured, err + toolResult, structured, degraded, execErr = res.Text, res.Structured, res.Degraded, err } if al.OnToolResult != nil { rewritten, hookErr := al.OnToolResult(toolCtx, callID, tCall.Name, tCall.ArgsJSON, toolResult, structured, execErr) @@ -231,6 +232,8 @@ func (al *AgentLoop) executeToolCall(ctx context.Context, st *iterationState, ws execErr = nil } } + toolResult = applyDegradation(ctx, tCall.Name, toolResult, degraded, execErr) + content := toolResult isToolErr := execErr != nil if isToolErr { @@ -248,7 +251,12 @@ func (al *AgentLoop) executeToolCall(ctx context.Context, st *iterationState, ws al.emit(ctx, st.sessionKey, st.streamChan, Event(ContentEvent{Text: "\n\n" + content + "\n\n"})) } - if cacheOK && !isToolErr { + // Never cache a degraded result. Its text carries a partial-success + // note naming artifacts that landed *this* run, and a cache hit + // short-circuits before recordDegradation — so replaying it would + // tell a later turn's model that work already exists while the host + // sees a clean turn with no DegradedEvent. + if cacheOK && !isToolErr && degraded == nil { al.Cache.Put(cacheKey, content) } diff --git a/pkg/agent/loop_iteration.go b/pkg/agent/loop_iteration.go index 664c537..c89a39a 100644 --- a/pkg/agent/loop_iteration.go +++ b/pkg/agent/loop_iteration.go @@ -45,7 +45,7 @@ func (al *AgentLoop) runIteration(ctx context.Context, sessionKey string, stream // *msgs stays at full fidelity so SetHistory / saveSession persist the // untouched conversation — without this distinction every prune would // shrink the on-disk history forever. - msgsForLLM := al.enforceTokenBudget(ctx, sessionKey, streamChan, *msgs) + msgsForLLM := al.enforceTokenBudget(ctx, sessionKey, streamChan, iteration, *msgs) al.emitSoftLandingNudge(ctx, sessionKey, streamChan, iteration) specMap := newSpeculativeMap() diff --git a/pkg/agent/loop_iteration_helpers.go b/pkg/agent/loop_iteration_helpers.go index dc0e224..90e9f14 100644 --- a/pkg/agent/loop_iteration_helpers.go +++ b/pkg/agent/loop_iteration_helpers.go @@ -14,26 +14,45 @@ import ( // SetHistory or saveSession; the input msgs slice is the source of truth // for what gets stored. When MaxTokenBudget is zero, falls back to the // standard PruneContextMessages with default depth. -func (al *AgentLoop) enforceTokenBudget(ctx context.Context, sessionKey string, streamChan chan<- StreamEvent, msgs []history.Message) []history.Message { +// +// Every path reports what it rewrote through a ContextTraceEvent — +// including the unbudgeted default path, which used to prune in complete +// silence. The event is suppressed when nothing changed, so a run whose +// context always fits still emits nothing. +func (al *AgentLoop) enforceTokenBudget(ctx context.Context, sessionKey string, streamChan chan<- StreamEvent, iteration int, msgs []history.Message) []history.Message { + orig := msgs if al.MaxTokenBudget <= 0 { - return pruneContextMessages(msgs, 3) + pruned, changes := pruneContextMessages(msgs, defaultProtectedEnds) + al.emitContextTrace(ctx, sessionKey, streamChan, iteration, ContextPolicyDefault, orig, pruned, changes) + return pruned } estToks := estimateTokens(msgs) thresh := int(float64(al.MaxTokenBudget) * budgetWarnRatio) + policy := ContextPolicyDefault + var changes []ContextRef if estToks > thresh && estToks <= al.MaxTokenBudget { al.emit(ctx, sessionKey, streamChan, Event(ThoughtEvent{Message: fmt.Sprintf("Token budget near threshold (~%d >= %d). Truncating tool arguments.", estToks, thresh)})) - msgs = truncateToolArguments(msgs) + var truncated []ContextRef + msgs, truncated = truncateToolArguments(msgs) + changes = append(changes, truncated...) + policy = ContextPolicyBudgetWarn } - if estimateTokens(msgs) > al.MaxTokenBudget { + depth := defaultProtectedEnds + if postTrim := estimateTokens(msgs); postTrim > al.MaxTokenBudget { al.emit(ctx, sessionKey, streamChan, Event(ThoughtEvent{ - Message: fmt.Sprintf("Token budget exceeded (~%d est. tokens). Applying emergency context pruning.", estimateTokens(msgs)), + Message: fmt.Sprintf("Token budget exceeded (~%d est. tokens). Applying emergency context pruning.", postTrim), })) - return pruneContextMessages(msgs, 1) + depth = emergencyProtectedEnds + policy = ContextPolicyBudgetEmergency } - return pruneContextMessages(msgs, 3) + + pruned, prunedRefs := pruneContextMessages(msgs, depth) + changes = append(changes, prunedRefs...) + al.emitContextTrace(ctx, sessionKey, streamChan, iteration, policy, orig, pruned, changes) + return pruned } // emitSoftLandingNudge fires the user-visible thought event marking the diff --git a/pkg/agent/loop_stream.go b/pkg/agent/loop_stream.go index 414ae94..ba1a25b 100644 --- a/pkg/agent/loop_stream.go +++ b/pkg/agent/loop_stream.go @@ -321,7 +321,9 @@ type AgentLoop struct { // final answer with no pending tool calls. Each pass appends a synthetic // critique prompt and asks the model to revise its answer; the final // round's text becomes the canonical response saved to history and - // emitted to callers. 0 (default) disables reflection entirely. + // emitted to callers — unless Scorer is set, in which case the + // best-scoring round wins instead. 0 (default) disables reflection + // entirely. // // Reflection is opt-in because it multiplies latency and token cost by // (1 + N). It targets correctness-critical tasks — SQL generation, code @@ -337,6 +339,26 @@ type AgentLoop struct { // misses task-specific pitfalls. ReflectPrompt string + // Scorer ranks candidate answers so a self-critique pass keeps the + // best round rather than the last one. nil (default) preserves the + // historical last-wins behavior: every non-empty, textually different + // revision is accepted, so a critique round can make the answer worse + // and the earlier text is unrecoverable. + // + // With a Scorer set, the model's original answer is scored as round 0 + // and each revision must beat the best score so far to be adopted; a + // revision that scores lower is discarded and the next round critiques + // the best answer instead. The kept round's score rides along on + // ReflectedEvent so a UI can show why. + // + // Costs one Score call per round plus one for the original. See Scorer + // for the latency and token-spend caveats. + // + // Only the self-critique path consumes this today, so setting it with + // Reflect == 0 is inert. The interface is deliberately generic — a + // best-of-K runner is the intended second consumer. + Scorer Scorer + // ThinkingBudget turns on extended reasoning for providers that support // it. It is a token hint, not a hard cap — see WithThinkingBudget for the // per-provider mapping. Set to 0 (default) to keep requests on the normal @@ -798,6 +820,14 @@ func (al *AgentLoop) runLogicLoop(ctx context.Context, sessionKey string, userMs ctx, emitCost = al.installRunCostAccumulator(ctx, sessionKey, streamChan) defer emitCost() + // Per-Run degradation accumulator. The final-answer path drains it + // before DoneEvent; this deferred sweep catches Runs that degraded and + // then ended on a cap or fatal error, where the unreliable state most + // needs reporting. drain() makes the two mutually exclusive. + var sweepDegraded func() + ctx, sweepDegraded = al.installDegradationAccumulator(ctx, sessionKey, streamChan) + defer sweepDegraded() + // Load memory notes once per Run and stash on ctx; buildMsgsForLLM // reads the cached value on every iteration so notes show up on every // LLM call without re-hitting the Store. Persisted history is left diff --git a/pkg/agent/pruning.go b/pkg/agent/pruning.go index 10a3057..c1cee41 100644 --- a/pkg/agent/pruning.go +++ b/pkg/agent/pruning.go @@ -55,21 +55,29 @@ func runeSlice(s string, start, end int) string { // It performs a soft trim by cutting the middle of excessively long tool responses // but strictly protects the last 'protectedEnds' messages from any modification. // All slicing is rune-safe to avoid corrupting multi-byte UTF-8 (CJK, emoji). -func pruneContextMessages(msgs []history.Message, protectedEnds int) []history.Message { +// +// The second return value records every message this call rewrote, in +// index order, so enforceTokenBudget can turn the decision into a +// ContextTraceEvent. It is nil when nothing was trimmed — the function +// stays pure and allocates no trace on the healthy path. +func pruneContextMessages(msgs []history.Message, protectedEnds int) ([]history.Message, []ContextRef) { if len(msgs) == 0 { - return msgs + return msgs, nil } result := make([]history.Message, 0, len(msgs)) + var changes []ContextRef protectStartIdx := max(len(msgs)-protectedEnds, 0) for i, msg := range msgs { if (msg.Role == "tool" || msg.Role == "assistant") && i < protectStartIdx { runeLen := utf8.RuneCountInString(msg.Content) + estBefore := len(msg.Content) / 4 if runeLen > outlierTrimThreshold { msg.Content = fmt.Sprintf("\n[System: Outlier Payload Truncated] The tool returned %d characters which exceeds the safety threshold of %d. Payload was completely discarded to avoid context explosion. Retry with tighter parameters.\n", runeLen, outlierTrimThreshold) + changes = append(changes, contextRefFor(i, msg, ContextChangeOutlierDiscarded, estBefore, len(msg.Content)/4)) result = append(result, msg) continue } @@ -80,12 +88,13 @@ func pruneContextMessages(msgs []history.Message, protectedEnds int) []history.M omitted := runeLen - retentionHead - retentionTail msg.Content = fmt.Sprintf("%s\n\n... [%d chars truncated] ...\n\n%s", head, omitted, tail) + changes = append(changes, contextRefFor(i, msg, ContextChangeSoftTrim, estBefore, len(msg.Content)/4)) } } result = append(result, msg) } - return result + return result, changes } // hasDanglingToolCalls walks msgs once and returns true if any assistant @@ -195,19 +204,25 @@ func patchDanglingToolCalls(msgs []history.Message) []history.Message { // truncateToolArguments forcefully truncates tool outputs to prevent context // window overflow when running tight on token budget. Non-tool messages pass // through unchanged. Truncation is rune-safe. -func truncateToolArguments(msgs []history.Message) []history.Message { +// +// Like pruneContextMessages, the second return value names the messages +// this call clipped and is nil when none were. +func truncateToolArguments(msgs []history.Message) ([]history.Message, []ContextRef) { if len(msgs) == 0 { - return msgs + return msgs, nil } result := make([]history.Message, 0, len(msgs)) - for _, msg := range msgs { + var changes []ContextRef + for i, msg := range msgs { if msg.Role == "tool" { runeLen := utf8.RuneCountInString(msg.Content) if runeLen > toolArgTruncateLen { + estBefore := len(msg.Content) / 4 msg.Content = runeSlice(msg.Content, 0, toolArgTruncateLen) + "\n... (output truncated by system to save tokens)" + changes = append(changes, contextRefFor(i, msg, ContextChangeArgsTruncated, estBefore, len(msg.Content)/4)) } } result = append(result, msg) } - return result + return result, changes } diff --git a/pkg/agent/pruning_test.go b/pkg/agent/pruning_test.go index c4e67c1..79501d4 100644 --- a/pkg/agent/pruning_test.go +++ b/pkg/agent/pruning_test.go @@ -48,7 +48,7 @@ func TestPruneContextMessages_ShortMessagesUntouched(t *testing.T) { {Role: "user", Content: "hello"}, {Role: "tool", Content: "short result"}, } - pruned := pruneContextMessages(msgs, 1) + pruned, _ := pruneContextMessages(msgs, 1) if pruned[2].Content != "short result" { t.Fatal("short tool message should not be pruned") } @@ -61,7 +61,7 @@ func TestPruneContextMessages_SoftTrim(t *testing.T) { {Role: "tool", Content: longContent}, {Role: "user", Content: "latest"}, } - pruned := pruneContextMessages(msgs, 1) + pruned, _ := pruneContextMessages(msgs, 1) if len(pruned[1].Content) >= len(longContent) { t.Fatal("expected soft trimmed content to be shorter") } @@ -77,7 +77,7 @@ func TestPruneContextMessages_OutlierGuard(t *testing.T) { {Role: "tool", Content: hugeContent}, {Role: "user", Content: "latest"}, } - pruned := pruneContextMessages(msgs, 1) + pruned, _ := pruneContextMessages(msgs, 1) if !strings.Contains(pruned[1].Content, "Outlier Payload Truncated") { t.Fatal("expected outlier truncation") } @@ -90,7 +90,7 @@ func TestPruneContextMessages_ProtectsRecentMessages(t *testing.T) { {Role: "tool", Content: longContent}, // protected (within last 3) {Role: "user", Content: "latest"}, } - pruned := pruneContextMessages(msgs, 3) // protect all 3 + pruned, _ := pruneContextMessages(msgs, 3) // protect all 3 if pruned[1].Content != longContent { t.Fatal("protected message should not be pruned") } @@ -103,7 +103,7 @@ func TestPruneContextMessages_SystemAndUserNeverPruned(t *testing.T) { {Role: "user", Content: longContent}, {Role: "assistant", Content: "short"}, } - pruned := pruneContextMessages(msgs, 1) + pruned, _ := pruneContextMessages(msgs, 1) if pruned[0].Content != longContent { t.Fatal("system message should never be pruned") } @@ -223,7 +223,7 @@ func TestPatchDanglingToolCalls_DanglingBeforeUserTurn(t *testing.T) { func TestTruncateToolArguments_ShortContent(t *testing.T) { msgs := []history.Message{{Role: "tool", Content: "short"}} - got := truncateToolArguments(msgs) + got, _ := truncateToolArguments(msgs) if got[0].Content != "short" { t.Fatalf("short content should not be truncated, got %q", got[0].Content) } @@ -232,7 +232,7 @@ func TestTruncateToolArguments_ShortContent(t *testing.T) { func TestTruncateToolArguments_LongContent(t *testing.T) { long := strings.Repeat("x", toolArgTruncateLen+500) msgs := []history.Message{{Role: "tool", Content: long}} - got := truncateToolArguments(msgs) + got, _ := truncateToolArguments(msgs) if len(got[0].Content) >= len(long) { t.Fatalf("expected truncation, got len %d >= %d", len(got[0].Content), len(long)) } @@ -248,7 +248,7 @@ func TestTruncateToolArguments_NonToolMessagesUntouched(t *testing.T) { {Role: "assistant", Content: long}, {Role: "system", Content: long}, } - got := truncateToolArguments(msgs) + got, _ := truncateToolArguments(msgs) for i, m := range got { if m.Content != long { t.Fatalf("non-tool message #%d should not be truncated", i) @@ -260,7 +260,7 @@ func TestTruncateToolArguments_UTF8Safe(t *testing.T) { // Each CJK char is 3 bytes but 1 rune. Make a content well over the limit. cjk := strings.Repeat("漢字", toolArgTruncateLen) msgs := []history.Message{{Role: "tool", Content: cjk}} - got := truncateToolArguments(msgs) + got, _ := truncateToolArguments(msgs) // Result must be valid UTF-8 (no corruption from byte-level cut). if !isValidUTF8(got[0].Content) { t.Fatalf("truncation corrupted UTF-8: %q", got[0].Content) diff --git a/pkg/agent/reflect.go b/pkg/agent/reflect.go index be1c1ef..3b73d1d 100644 --- a/pkg/agent/reflect.go +++ b/pkg/agent/reflect.go @@ -18,6 +18,91 @@ const defaultReflectPrompt = "Review your previous answer against the user's ori "If the current answer is already correct, repeat it verbatim — no commentary, no preface, no mention of this review. " + "Output only the final answer text." +// runReflectRounds drives the configured self-critique passes and returns +// the answer to keep. msgs must already end with the assistant's current +// answer; the last message's Content is rewritten in place for each round +// so the next critique sees the answer it is reviewing, and restored when +// a round is rejected. +// +// Without a Scorer this is last-wins: any non-empty, textually different +// revision is adopted, which is what the loop has always done. With a +// Scorer it is best-wins: a revision must strictly beat the best score so +// far, so ties and regressions keep the earlier answer and a critique +// pass can no longer degrade the response. +func (al *AgentLoop) runReflectRounds(ctx context.Context, st *iterationState, msgs []history.Message, finalContent string) string { + best := finalContent + last := len(msgs) - 1 + bestScore, haveBest := al.scoreCandidate(ctx, st, msgs, best, 0) + + for r := 1; r <= al.Reflect; r++ { + al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ + Message: fmt.Sprintf("Self-critique pass %d/%d...", r, al.Reflect), + })) + revised, rerr := al.reflectOnce(ctx, st.sessionKey, msgs, r, st.streamChan) + if rerr != nil { + al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ + Message: fmt.Sprintf("Self-critique aborted: %v", rerr), + })) + break + } + if revised == "" || revised == best { + continue + } + + msgs[last].Content = revised + var adopted *float64 + if al.Scorer != nil { + score, ok := al.scoreCandidate(ctx, st, msgs, revised, r) + // An unranked revision cannot be shown to be an improvement, + // so it loses to the incumbent the same way a lower-scoring + // one does. + if !ok || (haveBest && score <= bestScore) { + msgs[last].Content = best + al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ + Message: fmt.Sprintf("Self-critique pass %d discarded: %s.", r, rejectionDetail(score, ok, bestScore)), + })) + continue + } + bestScore, haveBest, adopted = score, true, &score + } + + best = revised + al.emit(ctx, st.sessionKey, st.streamChan, Event(ReflectedEvent{ + Text: best, + Round: r, + Score: adopted, + })) + } + return best +} + +// scoreCandidate ranks one candidate answer. Reports ok=false when no +// Scorer is configured or the Scorer failed — a scorer error degrades the +// round to unranked rather than failing the turn. +func (al *AgentLoop) scoreCandidate(ctx context.Context, st *iterationState, msgs []history.Message, answer string, round int) (float64, bool) { + if al.Scorer == nil { + return 0, false + } + score, err := al.Scorer.Score(ctx, RunResult{Answer: answer, Messages: msgs, Round: round}) + if err != nil { + al.emit(ctx, st.sessionKey, st.streamChan, Event(ThoughtEvent{ + Message: fmt.Sprintf("Scorer failed on round %d: %v", round, err), + })) + return 0, false + } + return score, true +} + +// rejectionDetail renders why a round lost, for the thought event. A +// round only reaches here unranked or beaten by an existing best, so +// those are the only two cases. +func rejectionDetail(score float64, scored bool, best float64) string { + if !scored { + return "the scorer could not rank it" + } + return fmt.Sprintf("scored %.4g, not better than the kept %.4g", score, best) +} + // reflectOnce runs a single self-critique round against the conversation in // baseMsgs (which must already include the assistant's current answer). // diff --git a/pkg/agent/reflect_test.go b/pkg/agent/reflect_test.go index 53b7fed..9969b8b 100644 --- a/pkg/agent/reflect_test.go +++ b/pkg/agent/reflect_test.go @@ -21,12 +21,28 @@ type reflectProvider struct { idx int capturedTools []bool // whether each call received a non-nil tool registry lastMsgLens []int + // critiqued records the assistant answer each critique round was + // asked to review, so a test can prove a rejected round was rolled + // back before the next round saw the conversation. + critiqued []string +} + +// lastAssistant returns the content of the final assistant message, or "" +// when there is none (the initial answer turn). +func lastAssistant(msgs []history.Message) string { + for i := len(msgs) - 1; i >= 0; i-- { + if msgs[i].Role == "assistant" { + return msgs[i].Content + } + } + return "" } func (p *reflectProvider) GenerateStream(_ context.Context, msgs []history.Message, tl *tools.Registry, ch chan<- StreamEvent) (LLMResult, error) { p.mu.Lock() p.capturedTools = append(p.capturedTools, tl != nil) p.lastMsgLens = append(p.lastMsgLens, len(msgs)) + p.critiqued = append(p.critiqued, lastAssistant(msgs)) var text string if p.idx < len(p.turns) { text = p.turns[p.idx] @@ -241,3 +257,146 @@ func (p *erroringReflectProvider) GenerateStream(_ context.Context, _ []history. ch <- Event(ContentEvent{Text: p.draft}) return LLMResult{Content: p.draft}, nil } + +// --- Scorer: keep the best round, not the last one --- + +// scoreByAnswer ranks candidates from a lookup table, so a test can make +// a later round strictly worse than an earlier one. +func scoreByAnswer(table map[string]float64) Scorer { + return ScorerFunc(func(_ context.Context, r RunResult) (float64, error) { + return table[r.Answer], nil + }) +} + +func TestReflect_ScorerKeepsBestRoundNotLast(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "good", "worse"}} + loop, sm := setup(provider) + loop.Reflect = 2 + loop.Scorer = scoreByAnswer(map[string]float64{ + "original": 10, + "good": 90, + "worse": 20, + }) + + got, err := loop.RunIteration(context.Background(), "s1", "q") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "good" { + t.Fatalf("answer = %q, want the best-scoring round %q", got, "good") + } + + msgs, err := sm.History(context.Background(), "s1") + if err != nil { + t.Fatalf("history: %v", err) + } + final := msgs[len(msgs)-1] + if final.Role != "assistant" || final.Content != "good" { + t.Fatalf("persisted answer = %+v, want the kept round", final) + } +} + +func TestReflect_ScorerRejectsRegressionAndRecritiquesBest(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "worse", "best"}} + loop, _ := setup(provider) + loop.Reflect = 2 + loop.Scorer = scoreByAnswer(map[string]float64{ + "original": 50, + "worse": 10, + "best": 99, + }) + + got, err := loop.RunIteration(context.Background(), "s1", "q") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + // Round 1 regressed and was discarded; round 2 beat the original. + if got != "best" { + t.Fatalf("answer = %q, want %q", got, "best") + } + // The rollback is the actual claim: round 2 must critique the + // incumbent, not the revision that was just thrown away. Asserting + // only the final answer passes even with the rollback deleted, + // because "best" outscores the incumbent either way. + provider.mu.Lock() + critiqued := append([]string(nil), provider.critiqued...) + provider.mu.Unlock() + if len(critiqued) != 3 { + t.Fatalf("want 3 provider calls (answer + 2 critiques), got %d: %q", len(critiqued), critiqued) + } + if critiqued[1] != "original" { + t.Fatalf("round 1 critiqued %q, want the original answer", critiqued[1]) + } + if critiqued[2] != "original" { + t.Fatalf("round 2 critiqued %q — the rejected revision was not rolled back", critiqued[2]) + } +} + +func TestReflect_ScorerTieKeepsIncumbent(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "cosmetic"}} + loop, _ := setup(provider) + loop.Reflect = 1 + loop.Scorer = scoreByAnswer(map[string]float64{"original": 42, "cosmetic": 42}) + + got, err := loop.RunIteration(context.Background(), "s1", "q") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "original" { + t.Fatalf("answer = %q — a tie must not displace the incumbent", got) + } +} + +func TestReflect_ScorerErrorDiscardsRevision(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "unrankable"}} + loop, _ := setup(provider) + loop.Reflect = 1 + loop.Scorer = ScorerFunc(func(_ context.Context, r RunResult) (float64, error) { + if r.Round == 0 { + return 5, nil + } + return 0, errors.New("judge unavailable") + }) + + got, err := loop.RunIteration(context.Background(), "s1", "q") + if err != nil { + t.Fatalf("scorer failure must not fail the turn, got %v", err) + } + if got != "original" { + t.Fatalf("answer = %q — an unrankable revision cannot be shown to be better", got) + } +} + +func TestReflect_NoScorerKeepsLastWinsBehavior(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "second", "third"}} + loop, _ := setup(provider) + loop.Reflect = 2 + + got, err := loop.RunIteration(context.Background(), "s1", "q") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "third" { + t.Fatalf("answer = %q, want last-wins %q without a Scorer", got, "third") + } +} + +func TestReflect_AdoptedRoundCarriesItsScore(t *testing.T) { + provider := &reflectProvider{turns: []string{"original", "better"}} + loop, _ := setup(provider) + loop.Reflect = 1 + loop.Scorer = scoreByAnswer(map[string]float64{"original": 1, "better": 7}) + + var reflected []ReflectedEvent + for ev := range loop.RunText(context.Background(), "s1", "q") { + if p, ok := ev.Payload.(ReflectedEvent); ok { + reflected = append(reflected, p) + } + } + if len(reflected) != 1 { + t.Fatalf("want exactly one ReflectedEvent for the adopted round, got %d", len(reflected)) + } + if reflected[0].Score == nil || *reflected[0].Score != 7 || reflected[0].Round != 1 { + t.Fatalf("event = %+v, want round 1 scored 7", reflected[0]) + } +} diff --git a/pkg/agent/regenerate.go b/pkg/agent/regenerate.go index e398e74..f323c04 100644 --- a/pkg/agent/regenerate.go +++ b/pkg/agent/regenerate.go @@ -108,6 +108,12 @@ func (al *AgentLoop) continueLogicLoop(ctx context.Context, sessionKey string, s ctx, emitCost = al.installRunCostAccumulator(ctx, sessionKey, streamChan) defer emitCost() + // Mirror runLogicLoop so a resumed Run reports partial-success state + // the same way a fresh one does. + var sweepDegraded func() + ctx, sweepDegraded = al.installDegradationAccumulator(ctx, sessionKey, streamChan) + defer sweepDegraded() + msgs, err := al.Sessions.History(ctx, sessionKey) if err != nil { al.emit(ctx, sessionKey, streamChan, errEvent(fmt.Errorf("agent: continue: load history: %w", err))) diff --git a/pkg/agent/scorer.go b/pkg/agent/scorer.go new file mode 100644 index 0000000..1a887f1 --- /dev/null +++ b/pkg/agent/scorer.go @@ -0,0 +1,47 @@ +package agent + +import ( + "context" + + "github.com/hung12ct/gopheragent/pkg/history" +) + +// RunResult is one candidate answer handed to a Scorer for ranking. It is +// deliberately narrow: everything a scorer needs to judge a candidate, +// nothing that ties it to the axis the candidate came from. +// +// Answer is the candidate final-answer text. Messages is the conversation +// it terminates, including the candidate itself as the last assistant +// message — read-only; a scorer must not mutate it. Round is the +// 1-indexed refinement pass for sequential scoring (self-critique), and 0 +// for the model's original, unrevised answer. +type RunResult struct { + Answer string + Messages []history.Message + Round int +} + +// Scorer ranks candidate answers so the loop can keep the best one rather +// than the last one. Higher scores win; the unit is the implementation's +// business (a 0–100 rubric, a compile-pass count, a negated latency). +// +// Score is called once per candidate on the loop goroutine, so a slow +// implementation adds directly to turn latency. It receives the turn's +// ctx and must honor cancellation. Returning an error drops that +// candidate from consideration without failing the turn — the loop keeps +// the best-scoring candidate it did manage to rank. +// +// Implementations that call an LLM to judge multiply the turn's token +// spend; that spend is invisible to BudgetTracker unless the scorer +// itself accounts for it. +type Scorer interface { + Score(ctx context.Context, r RunResult) (float64, error) +} + +// ScorerFunc adapts a plain function to Scorer. +type ScorerFunc func(ctx context.Context, r RunResult) (float64, error) + +// Score implements Scorer. +func (f ScorerFunc) Score(ctx context.Context, r RunResult) (float64, error) { + return f(ctx, r) +} diff --git a/pkg/agent/speculative.go b/pkg/agent/speculative.go index 505db7a..5d0b368 100644 --- a/pkg/agent/speculative.go +++ b/pkg/agent/speculative.go @@ -4,6 +4,8 @@ import ( "context" "strings" "sync" + + "github.com/hung12ct/gopheragent/pkg/tools" ) // speculativeExec carries the in-flight or completed result of a tool that @@ -20,11 +22,13 @@ import ( // after the drainer signals completion — no concurrent map access. type speculativeExec struct { id string + name string doneCh chan struct{} cancel context.CancelFunc result string structured any err error + degraded *tools.Degradation } // newSpeculativeMap returns an initialized map keyed by tool call ID. @@ -92,6 +96,7 @@ func (al *AgentLoop) spawnSpeculative( specCtx, cancel := context.WithCancel(ctx) sm := &speculativeExec{ id: id, + name: name, doneCh: make(chan struct{}), cancel: cancel, } @@ -121,20 +126,40 @@ func (al *AgentLoop) spawnSpeculative( // sub-agent emitter. The wave executor owns user-visible emissions // for this call when it processes the result. res, err := tool.Execute(specCtx, argsJSON) - sm.result, sm.structured, sm.err = res.Text, res.Structured, err + sm.result, sm.structured, sm.err, sm.degraded = res.Text, res.Structured, err, res.Degraded }() } +// reportOrphanedSpeculation files the degradation of a speculation that +// completed but is being discarded without ever being awaited — a retry +// reset or a stream error drops the entry, and the tool's side effects +// are real regardless. Consumed speculations are NOT filed here; the +// wave executor files those after OnToolResult has had its say, so the +// two paths are mutually exclusive and cannot double-count. +// +// Non-blocking: a speculation still in flight has left no side effect to +// report yet, and blocking here would stall the retry. +func reportOrphanedSpeculation(ctx context.Context, sm *speculativeExec) { + select { + case <-sm.doneCh: + if sm.err == nil { + recordDegradation(ctx, sm.name, sm.degraded) + } + default: + } +} + // awaitSpeculative blocks until the speculative execution completes and -// returns its (result, structured, err). Safe to call from the wave executor -// after the LLM stream has closed; doneCh acts as the happens-before barrier. -// structured is non-nil only when the underlying tool implemented -// tools.StructuredResult. -func awaitSpeculative(ctx context.Context, sm *speculativeExec) (string, any, error) { +// returns its (result, structured, degraded, err). Safe to call from the wave +// executor after the LLM stream has closed; doneCh acts as the happens-before +// barrier. structured is non-nil only when the underlying tool implemented +// tools.StructuredResult; degraded is non-nil only when the tool reported a +// partial success. +func awaitSpeculative(ctx context.Context, sm *speculativeExec) (string, any, *tools.Degradation, error) { select { case <-sm.doneCh: - return sm.result, sm.structured, sm.err + return sm.result, sm.structured, sm.degraded, sm.err case <-ctx.Done(): - return "", nil, ctx.Err() + return "", nil, nil, ctx.Err() } } diff --git a/pkg/agent/speculative_test.go b/pkg/agent/speculative_test.go index d6b3f6c..a4b9636 100644 --- a/pkg/agent/speculative_test.go +++ b/pkg/agent/speculative_test.go @@ -90,7 +90,7 @@ func TestSpawnSpeculative_CachesResult(t *testing.T) { t.Fatal("speculative entry should be registered before the goroutine runs") } - result, _, err := awaitSpeculative(context.Background(), sm) + result, _, _, err := awaitSpeculative(context.Background(), sm) if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -124,7 +124,7 @@ func TestSpawnSpeculative_CancelAbortsTool(t *testing.T) { sm.cancel() // Tool must abort with ctx.Canceled rather than completing. - _, _, err := awaitSpeculative(context.Background(), sm) + _, _, _, err := awaitSpeculative(context.Background(), sm) if !errors.Is(err, context.Canceled) { t.Fatalf("expected ctx.Canceled after sm.cancel(), got %v", err) } @@ -162,7 +162,7 @@ func TestSpawnSpeculative_MissingToolSurfacesError(t *testing.T) { mu.Lock() sm := store["c1"] mu.Unlock() - _, _, err := awaitSpeculative(context.Background(), sm) + _, _, _, err := awaitSpeculative(context.Background(), sm) if err == nil { t.Fatal("expected ToolNotFoundError from speculative execution") } @@ -178,7 +178,7 @@ func TestAwaitSpeculative_ContextCancelReturnsEarly(t *testing.T) { cancel() }() - _, _, err := awaitSpeculative(ctx, sm) + _, _, _, err := awaitSpeculative(ctx, sm) if err == nil { t.Fatal("expected ctx.Err() when caller cancels before speculative done") } diff --git a/pkg/tools/builtin/code_interpreter_test.go b/pkg/tools/builtin/code_interpreter_test.go index 6b41d49..158bbec 100644 --- a/pkg/tools/builtin/code_interpreter_test.go +++ b/pkg/tools/builtin/code_interpreter_test.go @@ -80,7 +80,7 @@ func TestCodeInterpreterTool_TimesOut(t *testing.T) { } _ = json.Unmarshal([]byte(out.Text), &env) if !env.TimedOut { - t.Fatalf("expected timeout flag, envelope: %s", out) + t.Fatalf("expected timeout flag, envelope: %s", out.Text) } } diff --git a/pkg/tools/builtin/memory_test.go b/pkg/tools/builtin/memory_test.go index 435ec0d..b2ee88c 100644 --- a/pkg/tools/builtin/memory_test.go +++ b/pkg/tools/builtin/memory_test.go @@ -49,11 +49,11 @@ func TestMemoryTools_SessionsAreIsolated(t *testing.T) { } out, _ := get.Execute(memCtx("s1"), `{"key":"k"}`) if !strings.Contains(out.Text, `"value":"one"`) { - t.Fatalf("s1 cross-talk: %s", out) + t.Fatalf("s1 cross-talk: %s", out.Text) } out, _ = get.Execute(memCtx("s2"), `{"key":"k"}`) if !strings.Contains(out.Text, `"value":"two"`) { - t.Fatalf("s2 cross-talk: %s", out) + t.Fatalf("s2 cross-talk: %s", out.Text) } } @@ -87,7 +87,7 @@ func TestMemoryTools_Delete(t *testing.T) { } out, _ := get.Execute(ctx, `{"key":"k"}`) if !strings.Contains(out.Text, `"found":false`) { - t.Fatalf("key still present: %s", out) + t.Fatalf("key still present: %s", out.Text) } } diff --git a/pkg/tools/builtin/show_media_test.go b/pkg/tools/builtin/show_media_test.go index 9860654..ddae837 100644 --- a/pkg/tools/builtin/show_media_test.go +++ b/pkg/tools/builtin/show_media_test.go @@ -30,7 +30,7 @@ func TestShowMedia_AcceptsDirectImageURL(t *testing.T) { t.Fatalf("unexpected error: %v", err) } if !strings.Contains(out.Text, "![hello](https://example.com/photo.jpg)") { - t.Fatalf("unexpected output: %q", out) + t.Fatalf("unexpected output: %q", out.Text) } } @@ -71,7 +71,7 @@ func TestShowMedia_AcceptsVideoURL(t *testing.T) { t.Fatalf("unexpected error: %v", err) } if !strings.Contains(out.Text, "