diff --git a/agent-schema.json b/agent-schema.json index fb51300643..95372ea56f 100644 --- a/agent-schema.json +++ b/agent-schema.json @@ -2452,6 +2452,12 @@ "description": "HTTP timeout in seconds (valid for type 'fetch', 'api', and 'openapi'). Defaults to 30 seconds when omitted.", "minimum": 1 }, + "max_output_bytes": { + "type": "integer", + "description": "Maximum OpenAPI text output in bytes (only valid for type 'openapi'). Defaults to 30000 when omitted; 0 disables the text cutoff.", + "minimum": 0, + "default": 30000 + }, "escape_html": { "type": "boolean", "description": "Restore legacy HTML escaping of <, > and & in multi-URL JSON results (only valid for type 'fetch'). Defaults to false to reduce token usage. Decoded content and single-URL results are unchanged.", @@ -2576,6 +2582,21 @@ } }, "additionalProperties": false, + "if": { + "required": [ + "max_output_bytes" + ] + }, + "then": { + "properties": { + "type": { + "const": "openapi" + } + }, + "required": [ + "type" + ] + }, "anyOf": [ { "allOf": [ diff --git a/docs/tools/openapi/index.md b/docs/tools/openapi/index.md index cd67ad5fd4..c98a703c1b 100644 --- a/docs/tools/openapi/index.md +++ b/docs/tools/openapi/index.md @@ -62,9 +62,10 @@ When Docker Desktop is running, eligible public destinations use its PAC proxy b | Property | Type | Required | Description | | ------------------- | ----------------- | -------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `url` | string | ✓ | URL of the OpenAPI specification (JSON format). Supports `${env.VAR}` interpolation. | +| `url` | string | ✓ | URL of the OpenAPI specification (JSON or YAML format). Supports `${env.VAR}` interpolation. | | `headers` | map[string]string | ✗ | Custom HTTP headers sent with every request — both the spec fetch and every generated tool call. Values support `${env.VAR}` and `${headers.NAME}` placeholders (the latter forwards a header from the caller's incoming request when docker agent is exposed as a server). | | `timeout` | int | ✗ | HTTP client timeout in seconds (default: `30`). Applies to both the spec fetch and the generated tools' requests. | +| `max_output_bytes` | integer | ✗ | Maximum returned response text in bytes. Omit for 30,000; `0` disables this cutoff while retaining the 1 MiB HTTP read cap. | | `allow_private_ips` | boolean | ✗ | Opt in to dialling **non-public** IP addresses (loopback, RFC1918, link-local — including the cloud-metadata endpoint at `169.254.169.254` — multicast and the unspecified address). Set to `true` only when the spec or its servers legitimately target internal services. By default such addresses are refused at dial time, after DNS resolution, so DNS rebinding cannot bypass the check. | ## How it works @@ -75,6 +76,24 @@ When Docker Desktop is running, eligible public destinations use its PAC proxy b 4. Read-only operations (GET, HEAD, OPTIONS) are annotated accordingly. 5. Responses are returned as text; errors include the HTTP status code. +## Returned text size + +Generated tools return at most 30,000 bytes of response text by default. Set +`max_output_bytes` to a larger positive limit, or `0` to disable this text cutoff: + +```yaml +toolsets: + - type: openapi + url: https://raw.githubusercontent.com/PokeAPI/pokeapi/master/openapi.yml + tools: [pokemon_retrieve] + max_output_bytes: 0 +``` + +The separate **1 MiB HTTP response read cap** still applies. Larger outputs also +remain subject to agent-level `max_tool_result_tokens`, context limits and provider +limits. Fetching more text can increase latency and token cost; a field-filtering +adapter may be preferable when most of a response is irrelevant. + ## Limits - The OpenAPI spec must be **10 MB or less**. diff --git a/examples/openapi-pokemon.yaml b/examples/openapi-pokemon.yaml new file mode 100644 index 0000000000..2195363955 --- /dev/null +++ b/examples/openapi-pokemon.yaml @@ -0,0 +1,13 @@ +agents: + root: + model: openai/gpt-4.1-mini + description: Pokédex lookup with complete API responses + instruction: | + Look up species using pokemon_retrieve with id set to their name. + Report types and sum the six stats[].base_stat values. + Use returned facts, not remembered values. + toolsets: + - type: openapi + url: https://raw.githubusercontent.com/PokeAPI/pokeapi/master/openapi.yml + tools: [pokemon_retrieve] + max_output_bytes: 0 diff --git a/pkg/config/latest/openapi_output_test.go b/pkg/config/latest/openapi_output_test.go new file mode 100644 index 0000000000..ed93acb255 --- /dev/null +++ b/pkg/config/latest/openapi_output_test.go @@ -0,0 +1,99 @@ +package latest + +import ( + "encoding/json" + "testing" + + "github.com/goccy/go-yaml" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestToolsetMaxOutputBytesValidation(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + toolset Toolset + wantErr string + }{ + {name: "omitted", toolset: Toolset{Type: "openapi", URL: "https://api.example.com/spec.yaml"}}, + {name: "disabled", toolset: Toolset{Type: "openapi", URL: "https://api.example.com/spec.yaml", MaxOutputBytes: new(0)}}, + {name: "positive", toolset: Toolset{Type: "openapi", URL: "https://api.example.com/spec.yaml", MaxOutputBytes: new(1024)}}, + {name: "negative", toolset: Toolset{Type: "openapi", URL: "https://api.example.com/spec.yaml", MaxOutputBytes: new(-1)}, wantErr: "max_output_bytes must not be negative"}, + {name: "shell zero", toolset: Toolset{Type: "shell", MaxOutputBytes: new(0)}, wantErr: "max_output_bytes can only be used with type 'openapi'"}, + {name: "fetch positive", toolset: Toolset{Type: "fetch", MaxOutputBytes: new(1024)}, wantErr: "max_output_bytes can only be used with type 'openapi'"}, + {name: "missing type", toolset: Toolset{MaxOutputBytes: new(0)}, wantErr: "max_output_bytes can only be used with type 'openapi'"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + inline := Config{Agents: Agents{{Name: "root", Toolsets: []Toolset{tt.toolset}}}} + named := Config{Toolsets: map[string]Toolset{"api": tt.toolset}} + data, err := yaml.Marshal(tt.toolset) + require.NoError(t, err) + var parsed Toolset + parseErr := yaml.Unmarshal(data, &parsed) + + if tt.wantErr != "" { + require.EqualError(t, tt.toolset.validate(), tt.wantErr) + require.ErrorContains(t, inline.Validate(), tt.wantErr) + require.ErrorContains(t, named.Validate(), "toolsets.api: "+tt.wantErr) + require.ErrorContains(t, parseErr, tt.wantErr) + return + } + require.NoError(t, tt.toolset.validate()) + require.NoError(t, inline.Validate()) + require.NoError(t, named.Validate()) + require.NoError(t, parseErr) + assert.Equal(t, tt.toolset.MaxOutputBytes, parsed.MaxOutputBytes) + }) + } +} + +func TestToolsetMaxOutputBytesRoundTrip(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value *int + }{ + {name: "omitted"}, + {name: "disabled", value: new(0)}, + {name: "positive", value: new(1024)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + toolset := Toolset{Type: "openapi", URL: "https://api.example.com/spec.yaml", MaxOutputBytes: tt.value} + + jsonData, err := json.Marshal(toolset) + require.NoError(t, err) + var jsonFields map[string]any + require.NoError(t, json.Unmarshal(jsonData, &jsonFields)) + var fromJSON Toolset + require.NoError(t, json.Unmarshal(jsonData, &fromJSON)) + assert.Equal(t, tt.value, fromJSON.MaxOutputBytes) + + yamlData, err := yaml.Marshal(toolset) + require.NoError(t, err) + var yamlFields map[string]any + require.NoError(t, yaml.Unmarshal(yamlData, &yamlFields)) + var fromYAML Toolset + require.NoError(t, yaml.Unmarshal(yamlData, &fromYAML)) + assert.Equal(t, tt.value, fromYAML.MaxOutputBytes) + + if tt.value == nil { + assert.NotContains(t, jsonFields, "max_output_bytes") + assert.NotContains(t, yamlFields, "max_output_bytes") + return + } + assert.EqualValues(t, *tt.value, jsonFields["max_output_bytes"]) + assert.EqualValues(t, *tt.value, yamlFields["max_output_bytes"]) + }) + } +} diff --git a/pkg/config/latest/types.go b/pkg/config/latest/types.go index 997b22c99e..e00b197ebf 100644 --- a/pkg/config/latest/types.go +++ b/pkg/config/latest/types.go @@ -1692,6 +1692,9 @@ type Toolset struct { // Defaults to 30 seconds when omitted. Timeout int `json:"timeout,omitempty"` + // MaxOutputBytes caps OpenAPI text output in bytes; nil defaults to 30000, 0 disables the cutoff. + MaxOutputBytes *int `json:"max_output_bytes,omitempty" yaml:"max_output_bytes,omitempty"` + // EscapeHTML restores legacy HTML escaping in fetch's multi-URL JSON results. // Defaults to false; single-URL results are unaffected. EscapeHTML *bool `json:"escape_html,omitempty" yaml:"escape_html,omitempty"` diff --git a/pkg/config/latest/validate.go b/pkg/config/latest/validate.go index fa3479d804..e902e7dc19 100644 --- a/pkg/config/latest/validate.go +++ b/pkg/config/latest/validate.go @@ -322,6 +322,9 @@ func (t *Toolset) validate() error { if err := validateNonEmptyEntries("blocked_servers", t.BlockedServers); err != nil { return err } + if t.MaxOutputBytes != nil && t.Type != "openapi" { + return errors.New("max_output_bytes can only be used with type 'openapi'") + } if t.EscapeHTML != nil && t.Type != "fetch" { return errors.New("escape_html can only be used with type 'fetch'") } @@ -459,6 +462,9 @@ func (t *Toolset) validate() error { if t.URL == "" { return errors.New("openapi toolset requires a url to be set") } + if t.MaxOutputBytes != nil && *t.MaxOutputBytes < 0 { + return errors.New("max_output_bytes must not be negative") + } case "open_url": if t.URL == "" { return errors.New("open_url toolset requires a url to be set") diff --git a/pkg/config/schema_test.go b/pkg/config/schema_test.go index da3205971a..d51f18639a 100644 --- a/pkg/config/schema_test.go +++ b/pkg/config/schema_test.go @@ -274,6 +274,110 @@ agents: } } +func TestJsonSchemaOpenAPIMaxOutputBytes(t *testing.T) { + t.Parallel() + + schemaBytes, err := os.ReadFile(schemaFile) + require.NoError(t, err) + schema, err := gojsonschema.NewSchema(gojsonschema.NewBytesLoader(schemaBytes)) + require.NoError(t, err) + + tests := []struct { + name string + toolType string + value any + valid bool + }{ + {name: "omitted", toolType: "openapi", valid: true}, + {name: "disabled", toolType: "openapi", value: 0, valid: true}, + {name: "positive", toolType: "openapi", value: 1024, valid: true}, + {name: "negative", toolType: "openapi", value: -1}, + {name: "fractional", toolType: "openapi", value: 1.5}, + {name: "string", toolType: "openapi", value: "1024"}, + {name: "shell zero", toolType: "shell", value: 0}, + {name: "fetch positive", toolType: "fetch", value: 1024}, + {name: "missing type", value: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + toolset := map[string]any{} + if tt.toolType != "" { + toolset["type"] = tt.toolType + } + if tt.toolType == "openapi" { + toolset["url"] = "https://api.example.com/spec.yaml" + } + if tt.value != nil { + toolset["max_output_bytes"] = tt.value + } + configs := map[string]any{ + "inline": map[string]any{ + "agents": map[string]any{"root": map[string]any{"toolsets": []any{toolset}}}, + }, + "named": map[string]any{ + "agents": map[string]any{"root": map[string]any{"use_toolsets": []any{"api"}}}, + "toolsets": map[string]any{"api": toolset}, + }, + } + for name, config := range configs { + data, err := json.Marshal(config) + require.NoError(t, err) + result, err := schema.Validate(gojsonschema.NewBytesLoader(data)) + require.NoError(t, err) + assert.Equal(t, tt.valid, result.Valid(), "%s: %v", name, result.Errors()) + } + }) + } +} + +func TestLoadOpenAPIMaxOutputBytes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + version string + field string + value *int + wantErr string + }{ + {name: "latest omitted"}, + {name: "legacy omitted", version: "15"}, + {name: "disabled", field: "max_output_bytes: 0", value: new(0)}, + {name: "positive", field: "max_output_bytes: 1024", value: new(1024)}, + {name: "negative", field: "max_output_bytes: -1", wantErr: "max_output_bytes must not be negative"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + data := fmt.Appendf(nil, `agents: + root: + model: openai/gpt-4o + toolsets: + - type: openapi + url: https://api.example.com/spec.yaml + %s +`, tt.field) + if tt.version != "" { + data = fmt.Appendf(data, "version: %q\n", tt.version) + } + cfg, err := Load(t.Context(), NewBytesSource("openapi.yaml", data)) + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + return + } + require.NoError(t, err) + require.Len(t, cfg.Agents, 1) + require.Len(t, cfg.Agents[0].Toolsets, 1) + assert.Equal(t, tt.value, cfg.Agents[0].Toolsets[0].MaxOutputBytes) + }) + } +} + // TestSchemaMatchesGoTypes verifies that every JSON-tagged field in the Go // config structs has a corresponding property in agent-schema.json (and // vice-versa). This prevents the schema from silently drifting out of sync diff --git a/pkg/tools/builtin/openapi/limit.go b/pkg/tools/builtin/openapi/limit.go index 964c0796db..bc2246972e 100644 --- a/pkg/tools/builtin/openapi/limit.go +++ b/pkg/tools/builtin/openapi/limit.go @@ -1,10 +1,22 @@ package openapi +import ( + "fmt" + "unicode/utf8" +) + const maxOutputSize = 30000 -func limitOutput(output string) string { - if len(output) > maxOutputSize { - return output[:maxOutputSize] + "\n\n[Output truncated: exceeded 30,000 character limit]" +func limitOutput(output string, limit int) string { + if limit <= 0 || len(output) <= limit { + return output + } + end := limit + for end > 0 && !utf8.RuneStart(output[end]) { + end-- + } + if limit == maxOutputSize { + return output[:end] + "\n\n[Output truncated: exceeded 30,000 character limit]" } - return output + return output[:end] + fmt.Sprintf("\n\n[Output truncated: exceeded %d byte limit]", limit) } diff --git a/pkg/tools/builtin/openapi/limit_test.go b/pkg/tools/builtin/openapi/limit_test.go new file mode 100644 index 0000000000..57947ac8c2 --- /dev/null +++ b/pkg/tools/builtin/openapi/limit_test.go @@ -0,0 +1,131 @@ +package openapi + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "unicode/utf8" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/config" + "github.com/docker/docker-agent/pkg/config/latest" +) + +func TestOutputLimits(t *testing.T) { + t.Parallel() + for _, test := range []struct { + name string + input string + limit int + want string + }{ + {"below limit", "abc", 4, "abc"}, + {"exact limit", "abcd", 4, "abcd"}, + {"disabled", strings.Repeat("a", maxOutputSize+1), 0, strings.Repeat("a", maxOutputSize+1)}, + {"custom limit", "abcde", 4, "abcd\n\n[Output truncated: exceeded 4 byte limit]"}, + {"UTF-8 boundary", "aéz", 2, "a\n\n[Output truncated: exceeded 2 byte limit]"}, + {"UTF-8 first rune", "éz", 1, "\n\n[Output truncated: exceeded 1 byte limit]"}, + } { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + output := limitOutput(test.input, test.limit) + assert.Equal(t, test.want, output) + assert.True(t, utf8.ValidString(output)) + }) + } + assert.Equal(t, strings.Repeat("a", maxOutputSize)+"\n\n[Output truncated: exceeded 30,000 character limit]", + limitOutput(strings.Repeat("a", maxOutputSize+1), maxOutputSize)) +} + +func TestConfiguredResponseLimits(t *testing.T) { + t.Parallel() + body := `{"padding":"` + strings.Repeat("x", maxOutputSize+100) + `","types":["electric"],"stats":[35,55,40,50,50,90]}` + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/openapi.json" { + _, _ = w.Write([]byte(petStoreSpec)) + return + } + if r.URL.Path == "/oversized" { + _, _ = w.Write([]byte(strings.Repeat("x", (1<<20)+1))) + return + } + if r.URL.Query().Get("error") != "" { + w.WriteHeader(http.StatusBadRequest) + } + _, _ = w.Write([]byte(body)) + })) + t.Cleanup(server.Close) + + for _, test := range []struct { + name string + limit *int + want string + }{ + {"default", nil, limitOutput(body, maxOutputSize)}, + {"custom", new(100), limitOutput(body, 100)}, + {"disabled", new(0), body}, + {"larger", new(len(body)), body}, + } { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + toolset, err := CreateToolSet(t.Context(), latest.Toolset{ + Type: "openapi", URL: server.URL + "/openapi.json", AllowPrivateIPs: new(true), + MaxOutputBytes: test.limit, + }, &config.RuntimeConfig{}) + require.NoError(t, err) + tools, err := toolset.Tools(t.Context()) + require.NoError(t, err) + result := callTool(t, toolByName(t, tools, "listPets"), `{}`) + assert.False(t, result.IsError) + assert.Equal(t, test.want, result.Output) + result = callTool(t, toolByName(t, tools, "listPets"), `{"error":"yes"}`) + assert.True(t, result.IsError) + assert.Equal(t, "HTTP 400: "+test.want, result.Output) + if test.limit != nil && *test.limit == 0 { + var parsed map[string]any + require.NoError(t, json.Unmarshal([]byte(test.want), &parsed)) + assert.Contains(t, parsed, "types") + assert.Contains(t, parsed, "stats") + } + }) + } +} + +func TestDisabledOutputCutoffRetainsHTTPBodyCap(t *testing.T) { + t.Parallel() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(strings.Repeat("x", (1<<20)+1))) + })) + t.Cleanup(server.Close) + handler := &openAPIHandler{ + baseURL: server.URL, path: "/", method: http.MethodGet, + maxOutputBytes: 0, allowPrivateIPs: true, + } + result, err := handler.callTool(t.Context(), nil) + require.NoError(t, err) + assert.True(t, strings.HasPrefix(result.Output, "[WARNING: Response truncated at 1MB limit]\n")) + assert.Len(t, result.Output, len("[WARNING: Response truncated at 1MB limit]\n")+(1<<20)) +} + +func TestHTTPBodyCapPreservesUTF8Boundary(t *testing.T) { + t.Parallel() + body := `{"padding":"` + strings.Repeat("x", (1<<20)-len(`{"padding":"`)-1) + "é" + `"}` + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(body)) + })) + t.Cleanup(server.Close) + for _, limit := range []int{0, 1 << 20} { + handler := &openAPIHandler{ + baseURL: server.URL, path: "/", method: http.MethodGet, + maxOutputBytes: limit, allowPrivateIPs: true, + } + result, err := handler.callTool(t.Context(), nil) + require.NoError(t, err) + assert.True(t, utf8.ValidString(result.Output)) + assert.Equal(t, "[WARNING: Response truncated at 1MB limit]\n"+body[:(1<<20)-1], result.Output) + } +} diff --git a/pkg/tools/builtin/openapi/openapi.go b/pkg/tools/builtin/openapi/openapi.go index 239b5c8833..7def9a22b3 100644 --- a/pkg/tools/builtin/openapi/openapi.go +++ b/pkg/tools/builtin/openapi/openapi.go @@ -12,6 +12,7 @@ import ( "net/url" "strings" "time" + "unicode/utf8" "github.com/pb33f/libopenapi" "github.com/pb33f/libopenapi/datamodel/high/base" @@ -36,6 +37,9 @@ func CreateToolSet(ctx context.Context, toolset latest.Toolset, runConfig *confi if toolset.Timeout > 0 { opts = append(opts, WithTimeout(time.Duration(toolset.Timeout)*time.Second)) } + if toolset.MaxOutputBytes != nil { + opts = append(opts, WithMaxOutputBytes(*toolset.MaxOutputBytes)) + } if toolset.AllowPrivateIPsEnabled() { opts = append(opts, WithAllowPrivateIPs(true)) } @@ -50,6 +54,7 @@ type ToolSet struct { headers map[string]string timeout time.Duration + maxOutputBytes int allowPrivateIPs bool expander *js.Expander } @@ -70,6 +75,12 @@ func WithTimeout(d time.Duration) Option { return func(t *ToolSet) { t.timeout = d } } +// WithMaxOutputBytes limits the returned text; zero disables this cutoff. +// The HTTP response still has a separate 1 MiB read limit. +func WithMaxOutputBytes(limit int) Option { + return func(t *ToolSet) { t.maxOutputBytes = limit } +} + // WithAllowPrivateIPs disables SSRF dial-time protection on both the spec // fetch and the generated tools' HTTP calls. Operators opt in via // `allow_private_ips: true` when the spec or its servers legitimately @@ -85,9 +96,10 @@ func WithExpander(expander *js.Expander) Option { // New creates a new OpenAPI toolset from the given spec URL. func New(specURL string, headers map[string]string, opts ...Option) *ToolSet { t := &ToolSet{ - specURL: specURL, - headers: headers, - timeout: httpclient.DefaultToolHTTPTimeout, + specURL: specURL, + headers: headers, + timeout: httpclient.DefaultToolHTTPTimeout, + maxOutputBytes: maxOutputSize, } for _, opt := range opts { opt(t) @@ -256,6 +268,7 @@ func (t *ToolSet) operationToTool(baseURL, path, method string, op *v3.Operation method: method, headers: t.headers, timeout: t.timeout, + maxOutputBytes: t.maxOutputBytes, allowPrivateIPs: t.allowPrivateIPs, expander: t.expander, }).callTool), @@ -447,6 +460,7 @@ type openAPIHandler struct { headers map[string]string timeout time.Duration + maxOutputBytes int allowPrivateIPs bool expander *js.Expander } @@ -495,8 +509,18 @@ func (h *openAPIHandler) callTool(ctx context.Context, params openAPICallArgs) ( return nil, fmt.Errorf("failed to read response: %w", err) } - output := limitOutput(string(body)) - if len(body) >= 1<<20 { + bodyTruncated := len(body) >= 1<<20 + if bodyTruncated { + start := len(body) - 1 + for start > 0 && !utf8.RuneStart(body[start]) { + start-- + } + if !utf8.FullRune(body[start:]) { + body = body[:start] + } + } + output := limitOutput(string(body), h.maxOutputBytes) + if bodyTruncated { output = "[WARNING: Response truncated at 1MB limit]\n" + output }