-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgoverned_tool.go
More file actions
185 lines (159 loc) · 5.13 KB
/
Copy pathgoverned_tool.go
File metadata and controls
185 lines (159 loc) · 5.13 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
// Copyright 2026 AxonFlow
// SPDX-License-Identifier: MIT
// GovernedTool — framework-agnostic tool governance adapter.
//
// Wraps any Tool implementation with AxonFlow input/output policy enforcement.
// Works transparently with any framework that accepts the Tool interface.
//
// "Your tools run the logic. AxonFlow decides when they're allowed to."
package axonflow
import (
"context"
"encoding/json"
"errors"
"fmt"
)
// Tool is a framework-agnostic interface for any callable tool.
type Tool interface {
Name() string
Description() string
Invoke(ctx context.Context, input any) (any, error)
}
// PolicyViolationError is returned when a tool call is blocked by policy.
type PolicyViolationError struct {
Reason string
}
func (e *PolicyViolationError) Error() string {
return fmt.Sprintf("policy violation: %s", e.Reason)
}
// IsPolicyViolationError checks if an error is a PolicyViolationError.
func IsPolicyViolationError(err error) bool {
var pve *PolicyViolationError
return errors.As(err, &pve)
}
// GovernedToolOptions configures a GovernedTool.
type GovernedToolOptions struct {
// ConnectorTypeFn maps a tool name to a connector type string.
// Defaults to using the tool's Name().
ConnectorTypeFn func(name string) string
// Operation is the operation type passed to MCPCheckInput.
// Defaults to "execute". Use "query" for read-only tools.
Operation string
}
// GovernedTool wraps a Tool with AxonFlow input/output governance.
//
// Every Invoke call runs through two policy checks:
//
// 1. Input check (MCPCheckInput): evaluates tool arguments before execution.
// Blocked calls return PolicyViolationError and the tool never runs.
// 2. Output check (MCPCheckOutput): evaluates tool results after execution.
// Can block (return error), redact (return cleaned data), or allow.
type GovernedTool struct {
wrapped Tool
client *AxonFlowClient
connectorType string
operation string
}
// GovernTool wraps a single tool with governance.
func GovernTool(tool Tool, client *AxonFlowClient, opts *GovernedToolOptions) *GovernedTool {
connectorType := tool.Name()
operation := "execute"
if opts != nil {
if opts.ConnectorTypeFn != nil {
connectorType = opts.ConnectorTypeFn(tool.Name())
}
if opts.Operation != "" {
operation = opts.Operation
}
}
return &GovernedTool{
wrapped: tool,
client: client,
connectorType: connectorType,
operation: operation,
}
}
// GovernTools wraps multiple tools with governance.
func GovernTools(tools []Tool, client *AxonFlowClient, opts *GovernedToolOptions) []*GovernedTool {
governed := make([]*GovernedTool, len(tools))
for i, t := range tools {
governed[i] = GovernTool(t, client, opts)
}
return governed
}
// Name returns the governed tool's name.
func (g *GovernedTool) Name() string { return g.wrapped.Name() }
// Description returns the governed tool's description.
func (g *GovernedTool) Description() string { return g.wrapped.Description() }
// Invoke executes the tool with input/output policy checks.
func (g *GovernedTool) Invoke(ctx context.Context, input any) (any, error) {
// 1. Serialize input
statement, err := serializeContent(input)
if err != nil {
return nil, fmt.Errorf("failed to serialize input: %w", err)
}
// 2. Input policy check
inputCheck, err := g.client.MCPCheckInput(ctx, MCPCheckInputRequest{
ConnectorType: g.connectorType,
Statement: statement,
Operation: g.operation,
})
if err != nil {
return nil, fmt.Errorf("mcp input policy check failed: %w", err)
}
// 3. If blocked, return PolicyViolationError — tool never runs
if !inputCheck.Allowed {
reason := inputCheck.BlockReason
if reason == "" {
reason = "tool call blocked by input policy"
}
return nil, &PolicyViolationError{Reason: reason}
}
// 4. Execute the wrapped tool
result, err := g.wrapped.Invoke(ctx, input)
if err != nil {
return nil, err
}
// 5. Serialize result
serialized, err := serializeContent(result)
if err != nil {
return nil, fmt.Errorf("failed to serialize output: %w", err)
}
// 6. Output policy check
outputCheck, err := g.client.MCPCheckOutput(ctx, MCPCheckOutputRequest{
ConnectorType: g.connectorType,
Message: serialized,
})
if err != nil {
return nil, fmt.Errorf("mcp output policy check failed: %w", err)
}
// 7. If blocked, return PolicyViolationError
if !outputCheck.Allowed {
reason := outputCheck.BlockReason
if reason == "" {
reason = "tool output blocked by policy"
}
return nil, &PolicyViolationError{Reason: reason}
}
// 8. If redacted, return redacted data
if outputCheck.RedactedData != nil {
return outputCheck.RedactedData, nil
}
// 9. Return original result
return result, nil
}
// String returns a string representation.
func (g *GovernedTool) String() string {
return fmt.Sprintf("GovernedTool(name=%s, connectorType=%s)", g.wrapped.Name(), g.connectorType)
}
// serializeContent converts input/output to a string for policy evaluation.
func serializeContent(content any) (string, error) {
if s, ok := content.(string); ok {
return s, nil
}
b, err := json.Marshal(content)
if err != nil {
return fmt.Sprintf("%v", content), nil
}
return string(b), nil
}