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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
218 changes: 170 additions & 48 deletions agent/harness/toolapproval/toolapproval.go
Original file line number Diff line number Diff line change
Expand Up @@ -84,21 +84,22 @@ func (r Rule) matches(toolName string, arguments map[string]string) bool {

// state is persisted in the session across turns.
type state struct {
Rules []Rule `json:"rules,omitempty"`
CollectedApprovalResponses []*message.ToolApprovalResponseContent `json:"collectedResponses,omitempty"`
QueuedApprovalRequests []*message.ToolApprovalRequestContent `json:"queuedRequests,omitempty"`
Rules []Rule `json:"rules,omitempty"`
CollectedApprovalResponses []*message.ToolApprovalResponseContent `json:"collectedResponses,omitempty"`
QueuedApprovalRequests []*message.ToolApprovalRequestContent `json:"queuedRequests,omitempty"`
SurfacedApprovalRequests map[string]*message.ToolApprovalRequestContent `json:"surfacedRequests,omitempty"`
}

func loadState(opts []agent.Option) state {
session, ok := agent.GetOption(opts, agent.WithSession)
if !ok {
return state{}
return normalizedState(state{})
}
var s state
if found, _ := session.Get(stateKey, &s); found {
return s
return normalizedState(s)
}
return state{}
return normalizedState(state{})
}

func saveState(opts []agent.Option, s state) {
Expand Down Expand Up @@ -128,6 +129,15 @@ type Config struct {
// without prompting the caller. Returning an error fails the current run.
AutoApprovalRules []AutoApprovalRule

// DisableApprovalResponseBinding disables rebinding inbound approval responses
// to the tool approval requests previously surfaced by this middleware.
//
// When false (the default), only approval responses tied to a surfaced or
// history-carried approval request are honored, and the recorded request's tool
// call is injected downstream so an approved call matches what was surfaced for
// approval. When true, inbound approval responses are forwarded unchanged.
DisableApprovalResponseBinding bool

// MaxAutoApprovalIterations is the safety cap for how many times the inner
// agent is re-invoked in a single run when every surfaced approval request
// is auto-approved. When nil, DefaultMaxAutoApprovalIterations is used.
Expand All @@ -149,7 +159,7 @@ func run(cfg Config, next agent.RunFunc, ctx context.Context, messages []*messag
st := loadState(opts)

// Step 1: Process inbound approval responses from the caller.
messages, st = prepareInbound(messages, st)
messages, st = prepareInbound(messages, st, !cfg.DisableApprovalResponseBinding)

// Step 2: If we have queued requests from a previous turn, drain any
// that are now auto-approvable and surface the next one.
Expand All @@ -160,6 +170,9 @@ func run(cfg Config, next agent.RunFunc, ctx context.Context, messages []*messag
if len(st.QueuedApprovalRequests) > 0 {
next := st.QueuedApprovalRequests[0]
st.QueuedApprovalRequests = st.QueuedApprovalRequests[1:]
if !cfg.DisableApprovalResponseBinding {
recordSurfacedApprovalRequests(&st, next)
}
saveState(opts, st)
yield(&agent.ResponseUpdate{
Role: message.RoleAssistant,
Expand Down Expand Up @@ -245,6 +258,9 @@ func run(cfg Config, next agent.RunFunc, ctx context.Context, messages []*messag
first := needsApproval[0]
st.QueuedApprovalRequests = append(st.QueuedApprovalRequests, needsApproval[1:]...)
st.CollectedApprovalResponses = append(st.CollectedApprovalResponses, autoApproved...)
if !cfg.DisableApprovalResponseBinding {
recordSurfacedApprovalRequests(&st, first)
}

// Non-approval updates were already yielded during streaming.
if !yield(&agent.ResponseUpdate{
Expand All @@ -268,30 +284,40 @@ func run(cfg Config, next agent.RunFunc, ctx context.Context, messages []*messag

// prepareInbound processes caller messages, extracting approval responses
// and any "always approve" flags into standing rules.
func prepareInbound(messages []*message.Message, st state) ([]*message.Message, state) {
func prepareInbound(messages []*message.Message, st state, bindApprovalResponses bool) ([]*message.Message, state) {
knownRequests := make(map[string]*message.ToolApprovalRequestContent)
if bindApprovalResponses {
knownRequests = knownApprovalRequests(messages, st)
}

var cleaned []*message.Message
for i, msg := range messages {
var hasApproval bool
var remaining []message.Content
for _, c := range msg.Contents {
if collectInboundApproval(&st, c) {
switch resp := c.(type) {
case *message.AlwaysApproveToolApprovalResponseContent:
hasApproval = true
bound := bindApprovalResponse(resp.InnerResponse, &st, knownRequests, bindApprovalResponses)
if addApprovalRuleFromResponse(&st, resp, bound) {
st.CollectedApprovalResponses = append(st.CollectedApprovalResponses, bound)
}
case *message.ToolApprovalResponseContent:
hasApproval = true
if bound := bindApprovalResponse(resp, &st, knownRequests, bindApprovalResponses); bound != nil {
st.CollectedApprovalResponses = append(st.CollectedApprovalResponses, bound)
}
default:
if c != nil {
remaining = append(remaining, c)
}
}
}
if hasApproval {
if cleaned == nil {
cleaned = make([]*message.Message, 0, len(messages))
cleaned = append(cleaned, messages[:i]...)
}
// Strip approval contents from the message, keep the rest.
var remaining []message.Content
for _, c := range msg.Contents {
if isInboundApprovalContent(c) {
continue
}
if c != nil {
remaining = append(remaining, c)
}
}
if len(remaining) > 0 {
clone := msg.Clone()
clone.Contents = remaining
Expand All @@ -307,47 +333,101 @@ func prepareInbound(messages []*message.Message, st state) ([]*message.Message,
return messages, st
}

func collectInboundApproval(st *state, c message.Content) bool {
switch resp := c.(type) {
case *message.AlwaysApproveToolApprovalResponseContent:
collectAlwaysApproveResponse(st, resp)
return true
case *message.ToolApprovalResponseContent:
st.CollectedApprovalResponses = append(st.CollectedApprovalResponses, resp)
return true
default:
return false
func normalizedState(s state) state {
if s.SurfacedApprovalRequests == nil {
s.SurfacedApprovalRequests = make(map[string]*message.ToolApprovalRequestContent)
}
return s
}

func collectAlwaysApproveResponse(st *state, resp *message.AlwaysApproveToolApprovalResponseContent) {
if resp.InnerResponse == nil {
return
func knownApprovalRequests(messages []*message.Message, st state) map[string]*message.ToolApprovalRequestContent {
known := make(map[string]*message.ToolApprovalRequestContent, len(st.SurfacedApprovalRequests))
for requestID, req := range st.SurfacedApprovalRequests {
if req != nil {
known[requestID] = req
}
}
if fc, ok := resp.InnerResponse.ToolCall.(*message.FunctionCallContent); ok && fc != nil {
if resp.AlwaysApproveTool {
addRuleIfNotExists(st, Rule{ToolName: fc.Name})
} else if resp.AlwaysApproveToolWithArguments {
args, err := serializeArguments(fc.Arguments)
if err != nil {
return
for _, msg := range messages {
if msg.Role != message.RoleAssistant {
continue
}
for _, c := range msg.Contents {
req, ok := c.(*message.ToolApprovalRequestContent)
if !ok || req == nil || req.RequestID == "" {
continue
}
addRuleIfNotExists(st, Rule{
ToolName: fc.Name,
Arguments: args,
})
known[req.RequestID] = snapshotToolApprovalRequest(req)
}
}
st.CollectedApprovalResponses = append(st.CollectedApprovalResponses, resp.InnerResponse)
return known
}

func isInboundApprovalContent(c message.Content) bool {
switch c.(type) {
case *message.AlwaysApproveToolApprovalResponseContent, *message.ToolApprovalResponseContent:
func bindApprovalResponse(resp *message.ToolApprovalResponseContent, st *state, knownRequests map[string]*message.ToolApprovalRequestContent, bind bool) *message.ToolApprovalResponseContent {
if resp == nil {
return nil
}
if !bind {
return resp
}

matchedRequest, ok := knownRequests[resp.RequestID]
if !ok || matchedRequest == nil {
return nil
}

delete(knownRequests, resp.RequestID)
delete(st.SurfacedApprovalRequests, resp.RequestID)

bound := &message.ToolApprovalResponseContent{
ContentHeader: cloneContentHeader(resp.ContentHeader),
RequestID: resp.RequestID,
Reason: resp.Reason,
Approved: resp.Approved,
ToolCall: cloneToolCallContent(matchedRequest.ToolCall),
}
return bound
}

func addApprovalRuleFromResponse(st *state, resp *message.AlwaysApproveToolApprovalResponseContent, bound *message.ToolApprovalResponseContent) bool {
if resp == nil || bound == nil {
return false
}
if resp.AlwaysApproveToolWithArguments {
if fc, ok := resp.InnerResponse.ToolCall.(*message.FunctionCallContent); ok && fc != nil {
if _, err := serializeArguments(fc.Arguments); err != nil {
return false
}
}
}
fc, ok := bound.ToolCall.(*message.FunctionCallContent)
if !ok || fc == nil {
return true
default:
}
if resp.AlwaysApproveTool {
addRuleIfNotExists(st, Rule{ToolName: fc.Name})
return true
}
if !resp.AlwaysApproveToolWithArguments {
return true
}
args, err := serializeArguments(fc.Arguments)
if err != nil {
return false
}
addRuleIfNotExists(st, Rule{
ToolName: fc.Name,
Arguments: args,
})
return true
}

func recordSurfacedApprovalRequests(st *state, requests ...*message.ToolApprovalRequestContent) {
for _, req := range requests {
if req == nil || req.RequestID == "" {
continue
}
st.SurfacedApprovalRequests[req.RequestID] = snapshotToolApprovalRequest(req)
}
}

// drainAutoApprovable removes queued requests that now match a standing rule,
Expand Down Expand Up @@ -494,6 +574,48 @@ func addRuleIfNotExists(st *state, rule Rule) {
st.Rules = append(st.Rules, rule)
}

func snapshotToolApprovalRequest(req *message.ToolApprovalRequestContent) *message.ToolApprovalRequestContent {
if req == nil {
return nil
}
return &message.ToolApprovalRequestContent{
ContentHeader: cloneContentHeader(req.ContentHeader),
RequestID: req.RequestID,
ToolCall: cloneToolCallContent(req.ToolCall),
}
}

func cloneToolCallContent(content message.ToolCallContent) message.ToolCallContent {
switch content := content.(type) {
case nil:
return nil
case *message.FunctionCallContent:
if content == nil {
return nil
}
cloned := *content
cloned.ContentHeader = cloneContentHeader(content.ContentHeader)
return &cloned
case *message.MCPServerToolCallContent:
if content == nil {
return nil
}
cloned := *content
cloned.ContentHeader = cloneContentHeader(content.ContentHeader)
return &cloned
default:
return content
}
}

func cloneContentHeader(header message.ContentHeader) message.ContentHeader {
return message.ContentHeader{
AdditionalProperties: maps.Clone(header.AdditionalProperties),
Annotations: slices.Clone(header.Annotations),
RawRepresentation: header.RawRepresentation,
}
}

func responseMessage(responses []*message.ToolApprovalResponseContent) *message.Message {
contents := make([]message.Content, len(responses))
for i, r := range responses {
Expand Down
Loading
Loading