diff --git a/internal/proxy/toolmediation.go b/internal/proxy/toolmediation.go index 24753c3..5742858 100644 --- a/internal/proxy/toolmediation.go +++ b/internal/proxy/toolmediation.go @@ -135,6 +135,7 @@ type managedToolDuplicateTracker struct { type managedToolSeen struct { Round int Count int + Method string Result []byte Status string StatusCode int @@ -1268,7 +1269,7 @@ func (t *managedToolDuplicateTracker) StoreOpenAIResult(agentCtx *agentctx.Agent if !ok { return } - t.store(resolved.CanonicalName, call.ArgumentsRaw, call.Arguments, round, outcome) + t.store(resolved.CanonicalName, resolved.Manifest.Execution.Method, call.ArgumentsRaw, call.Arguments, round, outcome) } func (t *managedToolDuplicateTracker) ObserveAnthropic(agentCtx *agentctx.AgentContext, call anthropicToolUse, round int) *managedToolDuplicate { @@ -1290,7 +1291,7 @@ func (t *managedToolDuplicateTracker) StoreAnthropicResult(agentCtx *agentctx.Ag if !ok { return } - t.store(resolved.CanonicalName, call.ArgumentsRaw, call.Arguments, round, outcome) + t.store(resolved.CanonicalName, resolved.Manifest.Execution.Method, call.ArgumentsRaw, call.Arguments, round, outcome) } func (t *managedToolDuplicateTracker) observe(canonicalName, service, method string, rawArgs json.RawMessage, args map[string]any, round int) *managedToolDuplicate { @@ -1304,7 +1305,7 @@ func (t *managedToolDuplicateTracker) observe(canonicalName, service, method str } seen, ok := t.seen[signature] if !ok { - t.seen[signature] = managedToolSeen{Round: round, Count: 1} + t.seen[signature] = managedToolSeen{Round: round, Count: 1, Method: strings.ToUpper(strings.TrimSpace(method))} return nil } seen.Count++ @@ -1328,7 +1329,13 @@ func (t *managedToolDuplicateTracker) observe(canonicalName, service, method str } } -func (t *managedToolDuplicateTracker) store(canonicalName string, rawArgs json.RawMessage, args map[string]any, round int, outcome managedToolOutcome) { +func (t *managedToolDuplicateTracker) store(canonicalName, method string, rawArgs json.RawMessage, args map[string]any, round int, outcome managedToolOutcome) { + if managedToolMethodMayMutate(method) { + _, failed := managedToolFailureClass(outcome) + if !failed || managedToolFailureMayHaveMutated(method, outcome.Trace.StatusCode) { + t.invalidateCachedReads() + } + } signature, _, ok := managedToolDuplicateSignature(canonicalName, rawArgs, args) if !ok { return @@ -1337,6 +1344,7 @@ func (t *managedToolDuplicateTracker) store(canonicalName string, rawArgs json.R if !ok { seen = managedToolSeen{Round: round, Count: 1} } + seen.Method = strings.ToUpper(strings.TrimSpace(method)) seen.Result = append([]byte(nil), outcome.RawJSON...) seen.Status = outcome.Trace.Status seen.StatusCode = outcome.Trace.StatusCode @@ -1357,6 +1365,19 @@ func (t *managedToolDuplicateTracker) store(canonicalName string, rawArgs json.R t.seen[signature] = seen } +func (t *managedToolDuplicateTracker) invalidateCachedReads() { + if t == nil { + return + } + for signature, seen := range t.seen { + if !managedToolMethodMayMutate(seen.Method) { + delete(t.seen, signature) + } + } + t.lastSignature = "" + t.duplicateStreak = 0 +} + func managedToolFailureClass(outcome managedToolOutcome) (string, bool) { status := strings.TrimSpace(outcome.Trace.Status) if status == "" { @@ -1393,6 +1414,15 @@ func managedToolFailureMayHaveMutated(method string, statusCode int) bool { return statusCode == 0 || statusCode == http.StatusRequestTimeout || statusCode >= http.StatusInternalServerError } +func managedToolMethodMayMutate(method string) bool { + switch strings.ToUpper(strings.TrimSpace(method)) { + case http.MethodGet, http.MethodHead, http.MethodOptions: + return false + default: + return true + } +} + func managedToolDuplicateSignature(canonicalName string, rawArgs json.RawMessage, args map[string]any) (string, json.RawMessage, bool) { canonicalArgs, ok := canonicalManagedToolArguments(rawArgs, args) if !ok { diff --git a/internal/proxy/toolmediation_test.go b/internal/proxy/toolmediation_test.go index 2b9b014..6ca1c2d 100644 --- a/internal/proxy/toolmediation_test.go +++ b/internal/proxy/toolmediation_test.go @@ -371,6 +371,90 @@ func TestManagedToolDuplicateTrackerCanonicalizesArguments(t *testing.T) { } } +func TestManagedToolDuplicateTrackerInvalidatesReadsAfterMutation(t *testing.T) { + agentCtx := &agentctx.AgentContext{Tools: &agentctx.ToolManifest{Tools: []agentctx.ToolManifestEntry{ + { + Name: "trading-api.get_trigger", + Execution: agentctx.ToolExecution{Service: "trading-api", Method: http.MethodGet}, + }, + { + Name: "trading-api.update_trigger", + Execution: agentctx.ToolExecution{Service: "trading-api", Method: http.MethodPatch}, + }, + }}} + read := openAIToolCall{ + Name: "trading-api.get_trigger", + Arguments: map[string]any{"trigger_id": float64(732)}, + ArgumentsRaw: json.RawMessage(`{"trigger_id":732}`), + } + write := openAIToolCall{ + Name: "trading-api.update_trigger", + Arguments: map[string]any{"trigger_id": float64(732), "label": "repaired"}, + ArgumentsRaw: json.RawMessage(`{"trigger_id":732,"label":"repaired"}`), + } + tracker := newManagedToolDuplicateTracker() + + if duplicate := tracker.ObserveOpenAI(agentCtx, read, 1); duplicate != nil { + t.Fatalf("first read should execute: %+v", duplicate) + } + oldResult := []byte(`{"id":732,"label":"original"}`) + tracker.StoreOpenAIResult(agentCtx, read, 1, managedToolOutcome{ + RawJSON: oldResult, + Trace: sessionhistory.ToolCallTrace{Status: "ok", StatusCode: http.StatusOK}, + }) + if duplicate := tracker.ObserveOpenAI(agentCtx, write, 2); duplicate != nil { + t.Fatalf("first mutation should execute: %+v", duplicate) + } + tracker.StoreOpenAIResult(agentCtx, write, 2, managedToolOutcome{ + RawJSON: []byte(`{"id":732,"label":"repaired"}`), + Trace: sessionhistory.ToolCallTrace{Status: "ok", StatusCode: http.StatusOK}, + }) + + if duplicate := tracker.ObserveOpenAI(agentCtx, read, 3); duplicate != nil { + t.Fatalf("read after mutation must execute instead of replaying stale state: %+v", duplicate) + } + newResult := []byte(`{"id":732,"label":"repaired"}`) + tracker.StoreOpenAIResult(agentCtx, read, 3, managedToolOutcome{ + RawJSON: newResult, + Trace: sessionhistory.ToolCallTrace{Status: "ok", StatusCode: http.StatusOK}, + }) + duplicate := tracker.ObserveOpenAI(agentCtx, read, 4) + if duplicate == nil || !bytes.Equal(duplicate.CachedResult, newResult) { + t.Fatalf("subsequent duplicate should replay the post-mutation read: %+v", duplicate) + } +} + +func TestManagedToolDuplicateTrackerKeepsReadsAfterRejectedMutation(t *testing.T) { + agentCtx := &agentctx.AgentContext{Tools: &agentctx.ToolManifest{Tools: []agentctx.ToolManifestEntry{ + {Name: "trading-api.get_trigger", Execution: agentctx.ToolExecution{Service: "trading-api", Method: http.MethodGet}}, + {Name: "trading-api.update_trigger", Execution: agentctx.ToolExecution{Service: "trading-api", Method: http.MethodPatch}}, + }}} + read := openAIToolCall{Name: "trading-api.get_trigger", Arguments: map[string]any{"trigger_id": float64(732)}, ArgumentsRaw: json.RawMessage(`{"trigger_id":732}`)} + write := openAIToolCall{Name: "trading-api.update_trigger", Arguments: map[string]any{"trigger_id": float64(732)}, ArgumentsRaw: json.RawMessage(`{"trigger_id":732}`)} + tracker := newManagedToolDuplicateTracker() + _ = tracker.ObserveOpenAI(agentCtx, read, 1) + readResult := []byte(`{"id":732,"label":"original"}`) + tracker.StoreOpenAIResult(agentCtx, read, 1, managedToolOutcome{RawJSON: readResult, Trace: sessionhistory.ToolCallTrace{Status: "ok", StatusCode: http.StatusOK}}) + _ = tracker.ObserveOpenAI(agentCtx, write, 2) + rejected := toolErrorPayload("validation_error", "rejected", http.StatusUnprocessableEntity, nil) + tracker.StoreOpenAIResult(agentCtx, write, 2, managedToolOutcome{RawJSON: rejected, Trace: sessionhistory.ToolCallTrace{Result: rejected, Status: "error", StatusCode: http.StatusUnprocessableEntity}}) + + duplicate := tracker.ObserveOpenAI(agentCtx, read, 3) + if duplicate == nil || !bytes.Equal(duplicate.CachedResult, readResult) { + t.Fatalf("a cleanly rejected mutation should not invalidate prior reads: %+v", duplicate) + } + + ambiguousTracker := newManagedToolDuplicateTracker() + _ = ambiguousTracker.ObserveOpenAI(agentCtx, read, 1) + ambiguousTracker.StoreOpenAIResult(agentCtx, read, 1, managedToolOutcome{RawJSON: readResult, Trace: sessionhistory.ToolCallTrace{Status: "ok", StatusCode: http.StatusOK}}) + _ = ambiguousTracker.ObserveOpenAI(agentCtx, write, 2) + timeout := toolErrorPayload("timeout", "outcome unknown", http.StatusGatewayTimeout, nil) + ambiguousTracker.StoreOpenAIResult(agentCtx, write, 2, managedToolOutcome{RawJSON: timeout, Trace: sessionhistory.ToolCallTrace{Result: timeout, Status: "error", StatusCode: http.StatusGatewayTimeout}}) + if duplicate := ambiguousTracker.ObserveOpenAI(agentCtx, read, 3); duplicate != nil { + t.Fatalf("an ambiguous mutation may have committed, so the next read must execute: %+v", duplicate) + } +} + func TestManagedToolDuplicateTrackerRetriesFailuresBeforeSuppressing(t *testing.T) { agentCtx := &agentctx.AgentContext{Tools: managedToolManifest()} tracker := newManagedToolDuplicateTracker()