Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 34 additions & 4 deletions internal/proxy/toolmediation.go
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,7 @@ type managedToolDuplicateTracker struct {
type managedToolSeen struct {
Round int
Count int
Method string
Result []byte
Status string
StatusCode int
Expand Down Expand Up @@ -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 {
Expand All @@ -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 {
Expand All @@ -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++
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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 == "" {
Expand Down Expand Up @@ -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 {
Expand Down
84 changes: 84 additions & 0 deletions internal/proxy/toolmediation_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down