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
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -421,6 +421,12 @@ This release adds sandbox authentication improvements, compaction-model context

## [Unreleased]

- Harden served-agent safety and network controls:
- A2A, MCP HTTP, and chat now default to restricted tool safety; autonomous execution requires `--safety autonomous`, and `safety: autonomous` in YAML fails startup with guidance to use that flag.
- Non-loopback listeners require authentication or `--insecure-no-auth`; Unix sockets remain exempt. A2A and MCP HTTP use `--auth-token`, while chat continues to use `--api-key`; A2A and chat can explicitly allow browser origins with `--cors-origin`.
- A2A context IDs cannot access sessions created by another serve surface. The session-schema migration records session origins; older binaries reject upgraded databases with a newer-database error.
- `mcp.CreateToolHandler` now requires an explicit safety policy.

## What's New

- Splits the sidebar's Token Usage click target in two: clicking the token/context part (glyph, token count, context `%`, the "compacting…" marker, or the `⚠ capped` marker) opens the `/context` dialog, while clicking the cost part (`$` figure, sub-session count) keeps opening `/cost`; the "Token Usage" section title is no longer clickable
Expand Down
68 changes: 62 additions & 6 deletions cmd/root/a2a.go
Original file line number Diff line number Diff line change
@@ -1,19 +1,33 @@
package root

import (
"errors"
"fmt"
"io"
"net"
"strings"

"github.com/spf13/cobra"

"github.com/docker/docker-agent/pkg/a2a"
"github.com/docker/docker-agent/pkg/cli"
"github.com/docker/docker-agent/pkg/config"
"github.com/docker/docker-agent/pkg/httpsec"
"github.com/docker/docker-agent/pkg/servesafety"
"github.com/docker/docker-agent/pkg/session"
"github.com/docker/docker-agent/pkg/telemetry"
)

type a2aFlags struct {
agentName string
listenAddr string
sessionDB string
runConfig config.RuntimeConfig
agentName string
listenAddr string
sessionDB string
safety string
authToken string
corsOrigin string
insecureNoAuth bool
stdout io.Writer
runConfig config.RuntimeConfig
}

func newA2ACmd() *cobra.Command {
Expand All @@ -32,6 +46,10 @@ func newA2ACmd() *cobra.Command {
cmd.PersistentFlags().StringVarP(&flags.agentName, "agent", "a", "", "Name of the agent to run (defaults to the team's first agent)")
cmd.PersistentFlags().StringVarP(&flags.listenAddr, "listen", "l", "127.0.0.1:8082", "Address to listen on")
cmd.PersistentFlags().StringVarP(&flags.sessionDB, "session-db", "s", "", "Path to the session database (default: <data-dir>/session.db)")
cmd.PersistentFlags().StringVar(&flags.safety, "safety", "", "Tool safety policy (strict, balanced, restricted, autonomous)")
cmd.PersistentFlags().StringVar(&flags.authToken, "auth-token", "", "Bearer token required for all A2A requests")
cmd.PersistentFlags().StringVar(&flags.corsOrigin, "cors-origin", "", "Allowed browser origin(s), comma-separated; empty disables CORS")
cmd.PersistentFlags().BoolVar(&flags.insecureNoAuth, "insecure-no-auth", false, "Allow unauthenticated non-loopback binding (insecure)")
addRuntimeConfigFlags(cmd, &flags.runConfig)

return cmd
Expand All @@ -44,7 +62,22 @@ func (f *a2aFlags) runA2ACommand(cmd *cobra.Command, args []string) (commandErr
telemetry.TrackCommandError(ctx, "serve", append([]string{"a2a"}, args...), commandErr)
}()

out := cli.NewPrinter(cmd.OutOrStdout())
if err := validateSafetyFlag(f.safety); err != nil {
return err
}
if f.corsOrigin != "" {
if _, err := httpsec.ParseOrigins(f.corsOrigin); err != nil {
return fmt.Errorf("invalid --cors-origin: %w", err)
}
}
if !isLoopbackListenAddr(f.listenAddr) && f.authToken == "" && !f.insecureNoAuth {
return errors.New("non-loopback A2A listeners require --auth-token or --insecure-no-auth")
}

out := cli.NewPrinter(f.stdout)
if f.stdout == nil {
out = cli.NewPrinter(cmd.OutOrStdout())
}
agentFilename := args[0]

ln, cleanup, err := newListener(ctx, f.listenAddr)
Expand All @@ -54,5 +87,28 @@ func (f *a2aFlags) runA2ACommand(cmd *cobra.Command, args []string) (commandErr
defer cleanup()

out.Println("Listening on", ln.Addr().String())
return a2a.Run(ctx, agentFilename, f.agentName, sessionDBPath(f.sessionDB), &f.runConfig, ln)
return a2a.Run(ctx, agentFilename, f.agentName, sessionDBPath(f.sessionDB), &f.runConfig, ln, a2a.RunOptions{
CLISafety: session.SafetyPolicy(f.safety),
AuthToken: f.authToken,
CORSOrigin: f.corsOrigin,
OnSafetyPolicy: func(resolved servesafety.Resolved) {
out.Printf("Tool safety policy: %s (source: %s)\n", resolved.Policy, resolved.Source)
},
})
}

func isLoopbackListenAddr(addr string) bool {
if strings.HasPrefix(addr, "unix://") {
return true
}
host, _, err := net.SplitHostPort(addr)
if err != nil {
return false
}
host = strings.Trim(host, "[]")
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
28 changes: 28 additions & 0 deletions cmd/root/a2a_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
package root

import "testing"

func TestIsLoopbackListenAddr(t *testing.T) {
t.Parallel()

for _, tc := range []struct {
addr string
want bool
}{
{"127.0.0.1:8082", true},
{"[::1]:8082", true},
{"localhost:8082", true},
{"unix:///tmp/agent.sock", true},
{"unix://", true},
{":8082", false},
{"0.0.0.0:8082", false},
{"[::]:8082", false},
{"192.168.1.1:8082", false},
} {
t.Run(tc.addr, func(t *testing.T) {
if got := isLoopbackListenAddr(tc.addr); got != tc.want {
t.Errorf("isLoopbackListenAddr(%q) = %v, want %v", tc.addr, got, tc.want)
}
})
}
}
34 changes: 27 additions & 7 deletions cmd/root/chat.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
package root

import (
"errors"
"fmt"
"os"
"time"

Expand All @@ -9,6 +11,8 @@ import (
"github.com/docker/docker-agent/pkg/chatserver"
"github.com/docker/docker-agent/pkg/cli"
"github.com/docker/docker-agent/pkg/config"
"github.com/docker/docker-agent/pkg/servesafety"
"github.com/docker/docker-agent/pkg/session"
"github.com/docker/docker-agent/pkg/telemetry"
)

Expand All @@ -23,6 +27,8 @@ type chatFlags struct {
conversationsMaxItems int
conversationTTL time.Duration
maxIdleRuntimes int
safety string
insecureNoAuth bool
runConfig config.RuntimeConfig
}

Expand All @@ -48,6 +54,8 @@ agent without any custom integration.`,
cmd.Flags().StringVar(&flags.corsOrigin, "cors-origin", "", "Allowed CORS origin (e.g. https://example.com); empty disables CORS entirely")
cmd.Flags().StringVar(&flags.apiKey, "api-key", "", "Required Bearer token clients must present (Authorization: Bearer <token>); empty disables auth")
cmd.Flags().StringVar(&flags.apiKeyEnv, "api-key-env", "", "Read the API key from this environment variable instead of the command line")
cmd.Flags().StringVar(&flags.safety, "safety", "", "Tool safety policy (strict, balanced, restricted, autonomous)")
cmd.Flags().BoolVar(&flags.insecureNoAuth, "insecure-no-auth", false, "Allow unauthenticated non-loopback binding (insecure)")
cmd.Flags().Int64Var(&flags.maxRequestSize, "max-request-size", 1<<20, "Maximum request body size in bytes (default 1 MiB)")
cmd.Flags().DurationVar(&flags.requestTimeout, "request-timeout", 5*time.Minute, "Per-request timeout (covers model + tool calls + streaming)")
cmd.Flags().IntVar(&flags.conversationsMaxItems, "conversations-max", 0, "Cache up to N conversations server-side, keyed by X-Conversation-Id (0 disables; clients must resend full history)")
Expand All @@ -65,6 +73,21 @@ func (f *chatFlags) runChatCommand(cmd *cobra.Command, args []string) (commandEr
telemetry.TrackCommandError(ctx, "serve", append([]string{"chat"}, args...), commandErr)
}()

if err := validateSafetyFlag(f.safety); err != nil {
return err
}

apiKey := f.apiKey
if f.apiKeyEnv != "" {
apiKey = os.Getenv(f.apiKeyEnv)
if apiKey == "" {
return fmt.Errorf("environment variable %q is empty or not set", f.apiKeyEnv)
}
}
if !isLoopbackListenAddr(f.listenAddr) && apiKey == "" && !f.insecureNoAuth {
return errors.New("non-loopback chat listeners require --api-key, --api-key-env, or --insecure-no-auth")
}

out := cli.NewPrinter(cmd.OutOrStdout())
agentFilename := args[0]

Expand All @@ -77,13 +100,6 @@ func (f *chatFlags) runChatCommand(cmd *cobra.Command, args []string) (commandEr
out.Println("Listening on", ln.Addr().String())
out.Println("OpenAI-compatible chat completions endpoint: http://" + ln.Addr().String() + "/v1/chat/completions")

apiKey := f.apiKey
if f.apiKeyEnv != "" {
if v := os.Getenv(f.apiKeyEnv); v != "" {
apiKey = v
}
}

return chatserver.Run(ctx, agentFilename, chatserver.Options{
AgentName: f.agentName,
RunConfig: &f.runConfig,
Expand All @@ -94,5 +110,9 @@ func (f *chatFlags) runChatCommand(cmd *cobra.Command, args []string) (commandEr
ConversationsMaxSessions: f.conversationsMaxItems,
ConversationTTL: f.conversationTTL,
MaxIdleRuntimes: f.maxIdleRuntimes,
CLISafety: session.SafetyPolicy(f.safety),
OnSafetyPolicy: func(resolved servesafety.Resolved) {
out.Printf("Tool safety policy: %s (source: %s)\n", resolved.Policy, resolved.Source)
},
}, ln)
}
28 changes: 28 additions & 0 deletions cmd/root/chat_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
package root

import (
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestChatRejectsUnauthenticatedNonLoopbackBind(t *testing.T) {
t.Parallel()

cmd := newChatCmd()
cmd.SetArgs([]string{"agent.yaml", "--listen", "0.0.0.0:8083"})
err := cmd.ExecuteContext(t.Context())
require.Error(t, err)
assert.Contains(t, err.Error(), "require --api-key, --api-key-env, or --insecure-no-auth")
}

func TestChatRejectsEmptyAPIKeyEnvironmentVariable(t *testing.T) {
t.Setenv("CHAT_API_KEY", "")

cmd := newChatCmd()
cmd.SetArgs([]string{"agent.yaml", "--api-key-env", "CHAT_API_KEY"})
err := cmd.ExecuteContext(t.Context())
require.Error(t, err)
assert.Contains(t, err.Error(), "CHAT_API_KEY")
}
41 changes: 35 additions & 6 deletions cmd/root/mcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,21 +3,27 @@ package root
import (
"context"
"errors"
"fmt"

"github.com/spf13/cobra"

"github.com/docker/docker-agent/pkg/config"
"github.com/docker/docker-agent/pkg/mcp"
"github.com/docker/docker-agent/pkg/runregistry"
"github.com/docker/docker-agent/pkg/servesafety"
"github.com/docker/docker-agent/pkg/session"
"github.com/docker/docker-agent/pkg/telemetry"
)

type mcpFlags struct {
agentName string
http bool
listenAddr string
attach string
runConfig config.RuntimeConfig
agentName string
http bool
listenAddr string
attach string
safety string
authToken string
insecureNoAuth bool
runConfig config.RuntimeConfig
}

func newMCPCmd() *cobra.Command {
Expand All @@ -41,6 +47,9 @@ func newMCPCmd() *cobra.Command {
cmd.PersistentFlags().StringVarP(&flags.listenAddr, "listen", "l", "127.0.0.1:8081", "Address to listen on")
cmd.PersistentFlags().StringVar(&flags.attach, "attach", "", "Attach to a running TUI run by pid, address, or session id (or empty for the most recent)")
cmd.PersistentFlags().Lookup("attach").NoOptDefVal = "latest"
cmd.PersistentFlags().StringVar(&flags.safety, "safety", "", "Tool safety policy (strict, balanced, restricted, autonomous); only valid with --http")
cmd.PersistentFlags().StringVar(&flags.authToken, "auth-token", "", "Bearer token required for HTTP MCP requests; only valid with --http")
cmd.PersistentFlags().BoolVar(&flags.insecureNoAuth, "insecure-no-auth", false, "Allow unauthenticated non-loopback HTTP binding (insecure); only valid with --http")
cmd.PersistentFlags().StringVar(&flags.runConfig.MCPToolName, "tool-name", "", "Override the MCP tool identifier clients call (defaults to agent name); only valid when exposing a single agent")
cmd.PersistentFlags().DurationVar(&flags.runConfig.MCPKeepAlive, "mcp-keepalive", 0, "Interval between MCP keep-alive pings (e.g. 30s); 0 disables keep-alive")
addRuntimeConfigFlags(cmd, &flags.runConfig)
Expand All @@ -56,9 +65,19 @@ func (f *mcpFlags) runMCPCommand(cmd *cobra.Command, args []string) (commandErr
}()

if f.attach != "" {
if f.http || f.safety != "" || f.authToken != "" || f.insecureNoAuth {
return errors.New("--http-only safety and authentication flags cannot be used with --attach")
}
return f.runAttach(ctx)
}

if !f.http && (f.safety != "" || f.authToken != "" || f.insecureNoAuth) {
return errors.New("--safety, --auth-token, and --insecure-no-auth require --http")
}
if err := validateSafetyFlag(f.safety); err != nil {
return err
}

if len(args) == 0 {
return errors.New("agent file is required (or use --attach)")
}
Expand All @@ -68,13 +87,23 @@ func (f *mcpFlags) runMCPCommand(cmd *cobra.Command, args []string) (commandErr
return mcp.StartMCPServer(ctx, agentFilename, f.agentName, &f.runConfig)
}

if !isLoopbackListenAddr(f.listenAddr) && f.authToken == "" && !f.insecureNoAuth {
return errors.New("non-loopback MCP HTTP listeners require --auth-token or --insecure-no-auth")
}

ln, cleanup, err := newListener(ctx, f.listenAddr)
if err != nil {
return err
}
defer cleanup()

return mcp.StartHTTPServer(ctx, agentFilename, f.agentName, &f.runConfig, ln)
return mcp.StartHTTPServer(ctx, agentFilename, f.agentName, &f.runConfig, ln, mcp.HTTPOptions{
CLISafety: session.SafetyPolicy(f.safety),
AuthToken: f.authToken,
OnSafetyPolicy: func(resolved servesafety.Resolved) {
fmt.Fprintf(cmd.OutOrStdout(), "Tool safety policy: %s (source: %s)\n", resolved.Policy, resolved.Source)
},
})
}

func (f *mcpFlags) runAttach(ctx context.Context) error {
Expand Down
39 changes: 39 additions & 0 deletions cmd/root/mcp_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
package root

import (
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestMCPHTTPSafetyAndAuthenticationFlagsRequireHTTP(t *testing.T) {
t.Parallel()

for _, args := range [][]string{
{"agent.yaml", "--safety", "restricted"},
{"agent.yaml", "--auth-token", "secret"},
{"agent.yaml", "--insecure-no-auth"},
{"--attach", "--safety", "restricted"},
} {
cmd := newMCPCmd()
cmd.SetArgs(args)
err := cmd.ExecuteContext(t.Context())
require.Error(t, err)
if len(args) > 0 && args[0] == "--attach" {
assert.Contains(t, err.Error(), "--http-only")
} else {
assert.Contains(t, err.Error(), "require --http")
}
}
}

func TestMCPHTTPRejectsUnauthenticatedNonLoopbackBind(t *testing.T) {
t.Parallel()

cmd := newMCPCmd()
cmd.SetArgs([]string{"agent.yaml", "--http", "--listen", "0.0.0.0:8081"})
err := cmd.ExecuteContext(t.Context())
require.Error(t, err)
assert.Contains(t, err.Error(), "require --auth-token or --insecure-no-auth")
}
Loading
Loading