diff --git a/cmd/root/dmr.go b/cmd/root/dmr.go new file mode 100644 index 000000000..3dfefeaac --- /dev/null +++ b/cmd/root/dmr.go @@ -0,0 +1,112 @@ +package root + +import ( + "cmp" + "context" + "fmt" + "log/slog" + "net" + "os" + "path/filepath" + "slices" + "strings" + + "github.com/docker/cli/cli/command" + dockerconfig "github.com/docker/cli/cli/config" + "github.com/docker/cli/cli/connhelper/commandconn" + cliflags "github.com/docker/cli/cli/flags" + "github.com/spf13/pflag" + + "github.com/docker/docker-agent/pkg/model/provider/dmr/dmrmodels" +) + +func withDMRDockerConnection(ctx context.Context, dockerCLI command.Cli, flags *pflag.FlagSet) (context.Context, error) { + args, err := dmrDockerArgs(ctx, dockerCLI, flags) + if err != nil { + return nil, err + } + + // Delegate socket, TLS, SSH and named-pipe handling to the same Docker CLI + // that runs model status. No engine connection is opened until HTTP dials. + dialArgs := slices.Concat(args, []string{"system", "dial-stdio"}) + return dmrmodels.ContextWithDockerConnection(ctx, args, func(ctx context.Context) (net.Conn, error) { + return commandconn.New(ctx, "docker", dialArgs...) + }), nil +} + +func dmrDockerArgs(ctx context.Context, dockerCLI command.Cli, flags *pflag.FlagSet) ([]string, error) { + if flags == nil { + flags = pflag.NewFlagSet("docker", pflag.ContinueOnError) + cliflags.NewClientOptions().InstallFlags(flags) + } + + // Resolve paths before --working-dir changes the process directory. + configDir, err := filepath.Abs(cmp.Or(flags.Lookup("config").Value.String(), dockerconfig.Dir())) + if err != nil { + return nil, fmt.Errorf("resolving Docker config directory: %w", err) + } + args := []string{"--config=" + configDir} + for _, name := range []string{"tls", "tlsverify"} { + if flag := flags.Lookup(name); flag.Changed { + args = append(args, "--"+name+"="+flag.Value.String()) + } + } + contextName, _ := flags.GetString("context") + if contextName == "" && !flags.Changed("host") { + switch { + case dockerCLI != nil: + contextName = dockerCLI.CurrentContext() + case os.Getenv("DOCKER_HOST") != "": + contextName = command.DefaultContextName + case os.Getenv("DOCKER_CONTEXT") != "": + contextName = os.Getenv("DOCKER_CONTEXT") + default: + cfg, err := dockerconfig.Load(configDir) + if err != nil { + // Docker itself tolerates an unreadable config for connection selection. + slog.DebugContext(ctx, "Failed to load Docker config", "error", err) + } + contextName = cfg.CurrentContext + } + } + if flags.Changed("host") || (contextName == command.DefaultContextName && os.Getenv("DOCKER_HOST") != "") { + host := os.Getenv("DOCKER_HOST") + if flags.Changed("host") { + host = flags.Lookup("host").Value.String() + } + if socket, ok := strings.CutPrefix(strings.TrimSpace(host), "unix://"); ok && socket != "" { + path, err := filepath.Abs(socket) + if err != nil { + return nil, fmt.Errorf("resolving Docker socket: %w", err) + } + host = "unix://" + path + } + args = append(args, "--host="+host) + } else { + args = append(args, "--context="+cmp.Or(contextName, command.DefaultContextName)) + } + + tls, _ := flags.GetBool("tls") + tlsVerify, _ := flags.GetBool("tlsverify") + for _, name := range []string{"tlscacert", "tlscert", "tlskey"} { + flag := flags.Lookup(name) + if !flag.Changed && !tls && !tlsVerify && !flags.Changed("tlsverify") { + continue + } + path := flag.Value.String() + // Docker ignores missing default client certificates, but not explicit ones. + if !flag.Changed && name != "tlscacert" { + if _, err := os.Stat(path); os.IsNotExist(err) { + path = "" + } + } + if path != "" { + path, err = filepath.Abs(path) + if err != nil { + return nil, fmt.Errorf("resolving Docker --%s: %w", name, err) + } + } + args = append(args, "--"+name+"="+path) + } + return args, nil +} diff --git a/cmd/root/dmr_test.go b/cmd/root/dmr_test.go new file mode 100644 index 000000000..421ed6140 --- /dev/null +++ b/cmd/root/dmr_test.go @@ -0,0 +1,378 @@ +package root + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "testing" + "time" + + "github.com/docker/cli/cli/command" + dockerconfig "github.com/docker/cli/cli/config" + "github.com/docker/cli/cli/context/docker" + "github.com/docker/cli/cli/context/store" + cliflags "github.com/docker/cli/cli/flags" + "github.com/spf13/pflag" + "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" + "github.com/docker/docker-agent/pkg/model/provider/dmr" + "github.com/docker/docker-agent/pkg/model/provider/dmr/dmrmodels" +) + +func TestDMRDockerConnection(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim and Unix socket") + } + if _, err := os.Stat("/.dockerenv"); err == nil { + t.Skip("Desktop engine routing is host-only") + } + + t.Setenv("MODEL_RUNNER_HOST", "") + t.Setenv("DOCKER_HOST", "") + t.Setenv("DOCKER_CONTEXT", "") + configDir := t.TempDir() + oldConfigDir := dockerconfig.Dir() + dockerconfig.SetDir(configDir) + t.Cleanup(func() { dockerconfig.SetDir(oldConfigDir) }) + t.Setenv("DOCKER_CONFIG", configDir) + + // A relative socket path keeps macOS's sockaddr_un path below its limit. + t.Chdir(t.TempDir()) + var lc net.ListenConfig + listener, err := lc.Listen(t.Context(), "unix", "engine.sock") + require.NoError(t, err) + wd, err := os.Getwd() + require.NoError(t, err) + host := "unix://" + filepath.Join(wd, "engine.sock") + server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/exp/vDD4.40/engines/v1/models": + _, _ = io.WriteString(w, `{"data":[{"id":"ai/test"}]}`) + case "/exp/vDD4.40/engines/_configure": + w.WriteHeader(http.StatusAccepted) + case "/exp/vDD4.40/engines/v1/embeddings": + _, _ = io.WriteString(w, `{"data":[{"embedding":[0.1,0.2]}]}`) + default: + t.Errorf("unexpected engine request: %s", r.URL) + w.WriteHeader(http.StatusNotFound) + } + })) + require.NoError(t, server.Listener.Close()) + server.Listener = listener + server.Start() + defer server.Close() + + contexts := store.New(filepath.Join(configDir, "contexts"), command.DefaultContextStoreConfig()) + require.NoError(t, contexts.CreateOrUpdate(store.Metadata{ + Name: "desktop-test", + Endpoints: map[string]any{docker.DockerEndpoint: docker.EndpointMeta{Host: host}}, + })) + require.NoError(t, os.WriteFile(filepath.Join(configDir, "config.json"), []byte(`{"currentContext":"unusable"}`), 0o600)) + + shimDir := t.TempDir() + argsFile := filepath.Join(shimDir, "args") + t.Setenv("DMR_ARGS_FILE", argsFile) + testBinary, err := os.Executable() + require.NoError(t, err) + t.Setenv("DMR_DOCKER_HELPER", "1") + script := "#!/bin/sh\nexec \"" + testBinary + "\" -test.run=^TestDMRDockerHelper$ -- \"$@\"\n" + require.NoError(t, os.WriteFile(filepath.Join(shimDir, "docker"), []byte(script), 0o755)) + t.Setenv("PATH", shimDir+string(os.PathListSeparator)+os.Getenv("PATH")) + + relativeConfig, err := filepath.Rel(wd, configDir) + require.NoError(t, err) + workspace := t.TempDir() + for _, tt := range []struct { + name, envName, envValue string + args []string + }{ + {name: "plugin context overrides environment", envName: "DOCKER_CONTEXT", envValue: "unusable", args: []string{"--context=desktop-test", "--config=" + configDir}}, + {name: "standalone DOCKER_CONTEXT", envName: "DOCKER_CONTEXT", envValue: "desktop-test"}, + {name: "standalone DOCKER_HOST", envName: "DOCKER_HOST", envValue: host}, + {name: "plugin host", args: []string{"--host=" + host}}, + {name: "standalone current context"}, + {name: "empty context pins current selection", args: []string{"--context="}}, + {name: "relative host with working directory", args: []string{"--host=unix://engine.sock"}}, + {name: "relative env host with working directory", envName: "DOCKER_HOST", envValue: "unix://engine.sock"}, + {name: "relative config with working directory", args: []string{"--context=desktop-test", "--config=" + relativeConfig}}, + } { + t.Run(tt.name, func(t *testing.T) { + if tt.envName != "" { + t.Setenv(tt.envName, tt.envValue) + } + if tt.name == "standalone current context" || tt.name == "empty context pins current selection" { + require.NoError(t, os.WriteFile(filepath.Join(configDir, "config.json"), []byte(`{"currentContext":"desktop-test"}`), 0o600)) + } + var dockerCLI command.Cli + var flags *pflag.FlagSet + if tt.args != nil { + opts := cliflags.NewClientOptions() + flags = pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + require.NoError(t, flags.Parse(tt.args)) + opts.SetDefaultOptions(flags) + cli, err := command.NewDockerCli(command.WithOutputStream(io.Discard), command.WithErrorStream(io.Discard)) + require.NoError(t, err) + require.NoError(t, cli.Initialize(opts)) + dockerCLI = cli + } + + ctx, err := withDMRDockerConnection(t.Context(), dockerCLI, flags) + require.NoError(t, err) + if strings.Contains(tt.name, "with working directory") { + t.Chdir(wd) // Restore the process directory after setupWorkingDirectory. + require.NoError(t, setupWorkingDirectory(&config.RuntimeConfig{WorkingDir: workspace})) + } + models, err := dmrmodels.ListModels(ctx) + require.NoError(t, err) + assert.Equal(t, []string{"ai/test"}, models) + if tt.name == "standalone current context" || tt.name == "empty context pins current selection" { + baseURL, transport := dmrmodels.ResolveBaseURL(ctx, nil, "http://model-runner.docker.internal/engines/v1/") + require.NotNil(t, transport) + defer transport.CloseIdleConnections() + _, err := dmrmodels.ListModelsAt(t.Context(), transport, baseURL) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(configDir, "config.json"), []byte(`{"currentContext":"unusable"}`), 0o600)) + transport.CloseIdleConnections() + models, err := dmrmodels.ListModelsAt(t.Context(), transport, baseURL) + require.NoError(t, err) + assert.Equal(t, []string{"ai/test"}, models) + } + gotArgs, err := os.ReadFile(argsFile) + require.NoError(t, err) + for _, arg := range tt.args { + if arg == "--context=" { + arg = "--context=desktop-test" + } + if arg == "--host=unix://engine.sock" { + arg = "--host=" + host + } + if strings.HasPrefix(arg, "--config=") { + arg = "--config=" + configDir + } + assert.Contains(t, strings.Split(string(gotArgs), "\n"), arg) + } + assert.True(t, strings.HasSuffix(string(gotArgs), "model\nstatus\n--json\n")) + + var doctor doctorFlags + withDoctorTestEnv(nil, nil, nil)(&doctor) + doctor.dmrLister = nil + report, err := doctor.buildReport(ctx, "") + require.NoError(t, err) + assert.Equal(t, dmrStatusReachable, report.DMR.Status) + assert.Empty(t, report.Issues) + + client, err := dmr.NewClient(ctx, &latest.ModelConfig{Provider: "dmr", Model: "ai/test"}) + require.NoError(t, err) + _, err = client.CreateBatchEmbedding(t.Context(), []string{"hello"}) + require.NoError(t, err) + }) + } +} + +func TestDMRDockerConnectionOverridesRemainLazy(t *testing.T) { + t.Setenv("DOCKER_CONTEXT", "does-not-exist") + ctx, err := withDMRDockerConnection(t.Context(), nil, nil) + require.NoError(t, err) + baseURL, client := dmrmodels.ResolveBaseURL(ctx, &latest.ModelConfig{BaseURL: "http://explicit/"}, "") + assert.Equal(t, "http://explicit/", baseURL) + assert.Nil(t, client) + + t.Setenv("MODEL_RUNNER_HOST", "http://runner") + baseURL, client = dmrmodels.ResolveBaseURL(ctx, nil, "") + assert.Equal(t, "http://runner/engines/v1/", baseURL) + assert.Nil(t, client) +} + +func TestDMRDockerConnectionForwardsTLSFlags(t *testing.T) { + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + certDir := t.TempDir() + args := []string{ + "--host=tcp://engine:2376", "--tls=true", "--tlsverify=false", + "--tlscacert=" + filepath.Join(certDir, "ca.pem"), + "--tlscert=" + filepath.Join(certDir, "cert.pem"), + "--tlskey=" + filepath.Join(certDir, "key.pem"), + } + require.NoError(t, flags.Parse(args)) + ctx, err := withDMRDockerConnection(t.Context(), nil, flags) + require.NoError(t, err) + cmd := dmrmodels.DockerCommand(ctx, "model", "inspect", "ai/test") + for _, arg := range args { + assert.Contains(t, cmd.Args, arg) + } +} + +func TestDMRDockerConnectionErrorNamesEngine(t *testing.T) { + if _, err := os.Stat("/.dockerenv"); err == nil { + t.Skip("Desktop engine routing is host-only") + } + t.Setenv("MODEL_RUNNER_HOST", "") + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim") + } + dir := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte("#!/bin/sh\necho 'missing engine' >&2\nexit 1\n"), 0o755)) + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + require.NoError(t, flags.Parse([]string{"--host=unix:///missing-dmr-engine.sock"})) + ctx, err := withDMRDockerConnection(t.Context(), nil, flags) + require.NoError(t, err) + baseURL, client := dmrmodels.ResolveBaseURL(ctx, nil, "http://model-runner.docker.internal/engines/v1/") + require.NotNil(t, client) + defer client.CloseIdleConnections() + // The transport retains its selection even when used with a different context. + _, err = dmrmodels.ListModelsAt(context.WithoutCancel(t.Context()), client, baseURL) + require.ErrorContains(t, err, "unix:///missing-dmr-engine.sock") +} + +func TestDMRSelectedDockerStatusFailureDoesNotFallBack(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim") + } + dir := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte("#!/bin/sh\necho 'selected engine unavailable' >&2\nexit 1\n"), 0o755)) + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + t.Setenv("MODEL_RUNNER_HOST", "") + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + require.NoError(t, flags.Parse([]string{"--context=unavailable"})) + ctx, err := withDMRDockerConnection(t.Context(), nil, flags) + require.NoError(t, err) + + _, err = dmrmodels.ListModels(ctx) + require.ErrorContains(t, err, "selected engine unavailable") + require.ErrorContains(t, err, "--context=unavailable") + _, err = dmr.NewClient(ctx, &latest.ModelConfig{Provider: "dmr", Model: "ai/test"}) + require.ErrorContains(t, err, "selected engine unavailable") +} + +// TestDMRDockerHelper emulates Docker's stdio tunnel using the real context resolver. +func TestDMRDockerHelper(t *testing.T) { + if os.Getenv("DMR_DOCKER_HELPER") != "1" { + return + } + var args []string + for i, arg := range os.Args { + if arg == "--" { + args = os.Args[i+1:] + break + } + } + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + flags.SetInterspersed(false) + require.NoError(t, flags.Parse(args)) + opts.SetDefaultOptions(flags) + dockerconfig.SetDir(opts.ConfigDir) + cfg, err := dockerconfig.Load(dockerconfig.Dir()) + require.NoError(t, err) + client, err := command.NewAPIClientFromFlags(opts, cfg) + require.NoError(t, err) + commandArgs := flags.Args() + if len(commandArgs) > 1 && commandArgs[0] == "system" && commandArgs[1] == "dial-stdio" { + conn, err := client.Dialer()(t.Context()) + require.NoError(t, err) + go func() { _, _ = io.Copy(conn, os.Stdin); _ = conn.Close() }() + _, _ = io.Copy(os.Stdout, conn) + os.Exit(0) + } + require.NoError(t, os.WriteFile(os.Getenv("DMR_ARGS_FILE"), []byte(strings.Join(args, "\n")+"\n"), 0o600)) + _, _ = os.Stdout.WriteString(`{"endpoint":"http://model-runner.docker.internal/engines/v1/"}`) + os.Exit(0) +} + +func TestDMRDockerTunnelCancellation(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim") + } + if _, err := os.Stat("/.dockerenv"); err == nil { + t.Skip("Desktop engine routing is host-only") + } + t.Setenv("MODEL_RUNNER_HOST", "") + dir := t.TempDir() + // Read forever without answering HTTP, as a stuck engine/TLS handshake would. + require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte("#!/bin/sh\nwhile read line; do :; done\n"), 0o755)) + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + ctx, err := withDMRDockerConnection(t.Context(), nil, nil) + require.NoError(t, err) + baseURL, client := dmrmodels.ResolveBaseURL(ctx, nil, "http://model-runner.docker.internal/engines/v1/") + require.NotNil(t, client) + defer client.CloseIdleConnections() + ctx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + _, err = dmrmodels.ListModelsAt(ctx, client, baseURL) + require.ErrorIs(t, err, context.DeadlineExceeded) +} + +func TestDMRDockerConnectionAnchorsPaths(t *testing.T) { + t.Chdir(t.TempDir()) + original, err := os.Getwd() + require.NoError(t, err) + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + require.NoError(t, flags.Parse([]string{ + "--context=desktop-test", "--config=./docker-config", "--tlsverify=true", + "--tlscacert=./tls/ca.pem", "--tlscert=./tls/cert.pem", "--tlskey=./tls/key.pem", + })) + ctx, err := withDMRDockerConnection(t.Context(), nil, flags) + require.NoError(t, err) + t.Chdir(t.TempDir()) + args := dmrmodels.DockerCommand(ctx, "model", "status", "--json").Args + for flag, path := range map[string]string{ + "config": "docker-config", "tlscacert": "tls/ca.pem", "tlscert": "tls/cert.pem", "tlskey": "tls/key.pem", + } { + assert.Contains(t, args, "--"+flag+"="+filepath.Join(original, path)) + } + assert.Contains(t, args, "--context=desktop-test") +} + +func TestDMRDockerConnectionTLSDefaults(t *testing.T) { + t.Chdir(t.TempDir()) + original, err := os.Getwd() + require.NoError(t, err) + for _, certExists := range []bool{false, true} { + t.Run(strconv.FormatBool(certExists), func(t *testing.T) { + opts := cliflags.NewClientOptions() + flags := pflag.NewFlagSet("docker", pflag.ContinueOnError) + opts.InstallFlags(flags) + // Emulate defaults from a relative DOCKER_CERT_PATH without marking flags changed. + opts.TLSOptions.CAFile = "ca.pem" + opts.TLSOptions.CertFile = "cert.pem" + opts.TLSOptions.KeyFile = "key.pem" + if certExists { + require.NoError(t, os.WriteFile("cert.pem", nil, 0o600)) + require.NoError(t, os.WriteFile("key.pem", nil, 0o600)) + } + require.NoError(t, flags.Parse([]string{"--host=tcp://engine:2376", "--tlsverify=false"})) + ctx, err := withDMRDockerConnection(t.Context(), nil, flags) + require.NoError(t, err) + args := dmrmodels.DockerCommand(ctx, "model", "status", "--json").Args + assert.Contains(t, args, "--tlscacert="+filepath.Join(original, "ca.pem")) + for flag, file := range map[string]string{"tlscert": "cert.pem", "tlskey": "key.pem"} { + want := "" + if certExists { + want = filepath.Join(original, file) + } + assert.Contains(t, args, "--"+flag+"="+want) + } + }) + } +} diff --git a/cmd/root/flags.go b/cmd/root/flags.go index 7cbb71427..08712e441 100644 --- a/cmd/root/flags.go +++ b/cmd/root/flags.go @@ -115,8 +115,6 @@ func addGatewayFlags(cmd *cobra.Command, runConfig *config.RuntimeConfig, loadUs persistentPreRunE := cmd.PersistentPreRunE cmd.PersistentPreRunE = func(_ *cobra.Command, args []string) error { - ctx := cmd.Context() - // Run any inherited PersistentPreRunE first so directory // overrides (--config-dir, --cache-dir, --data-dir) and other // global setup land before we materialise the env provider — @@ -126,6 +124,7 @@ func addGatewayFlags(cmd *cobra.Command, runConfig *config.RuntimeConfig, loadUs if err := runParentPreRun(cmd, persistentPreRunE, args); err != nil { return err } + ctx := cmd.Context() userCfg, err := loadUserConfig() if err != nil { diff --git a/cmd/root/root.go b/cmd/root/root.go index 3f8494b27..b956011d9 100644 --- a/cmd/root/root.go +++ b/cmd/root/root.go @@ -223,6 +223,13 @@ func Execute(ctx context.Context, stdin io.Reader, stdout, stderr io.Writer, arg rootCmd.SetArgs(args) runningStandalone := plugin.RunningStandalone() + if runningStandalone { + var err error + ctx, err = withDMRDockerConnection(ctx, nil, nil) + if err != nil { + return err + } + } visitAll(rootCmd, func(cmd *cobra.Command) { cmd.SetContext(ctx) @@ -240,7 +247,7 @@ func Execute(ctx context.Context, stdin io.Reader, stdout, stderr io.Writer, arg return rootCmd.Execute() } - plugin.Run(func(command.Cli) *cobra.Command { + plugin.Run(func(dockerCLI command.Cli) *cobra.Command { // Force to the name of the docker command rootCmd.Use = "agent" @@ -252,6 +259,11 @@ func Execute(ctx context.Context, stdin io.Reader, stdout, stderr io.Writer, arg if err := plugin.PersistentPreRunE(cmd, args); err != nil { return err } + dmrCtx, err := withDMRDockerConnection(cmd.Context(), dockerCLI, cmd.Root().Flags()) + if err != nil { + return err + } + visitAll(rootCmd, func(c *cobra.Command) { c.SetContext(dmrCtx) }) if originalPreRun != nil { return originalPreRun(cmd, args) } diff --git a/docs/providers/dmr/index.md b/docs/providers/dmr/index.md index ece1415c8..40b16848c 100644 --- a/docs/providers/dmr/index.md +++ b/docs/providers/dmr/index.md @@ -10,20 +10,40 @@ _Run AI models locally with Docker — no API keys, no costs, full data privacy. ## Overview -Docker Model Runner (DMR) lets you run open-source AI models directly on your machine. Models run in Docker, so there's no API key needed and no data leaves your computer. +Docker Model Runner (DMR) lets you run open-source AI models directly on your machine. Models run in Docker, so there's no API key needed; with a local runner, prompts stay on your computer. Docker Agent automatically discovers models you have already pulled from DMR. When no model is explicitly configured, auto-selection prefers a locally-installed model (choosing the model specified via the `model:` key in the agent YAML if it is already pulled locally, or otherwise the first available non-embedding model) rather than always defaulting to `ai/qwen3:latest` and triggering a pull prompt. > [!TIP] > **No API key needed** > -> DMR runs models locally — your data never leaves your machine. Great for development, sensitive data, or offline use. +> With a local DMR endpoint, your data stays on your machine. Great for development, sensitive data, or offline use. ## Prerequisites - [Docker Desktop](https://www.docker.com/products/docker-desktop/) with the Model Runner feature enabled - Verify with: `docker model status --json` +## Connection selection + +For Docker Desktop, `docker agent doctor` and DMR-backed agents use the Docker +connection selected by the CLI, including `--context`, `--host`, +`DOCKER_CONTEXT`, and `DOCKER_HOST`. Requests go through that engine's socket; +enabling the +legacy `/var/run/docker.sock` symlink is not required. + +```console +docker --context desktop-linux agent doctor +``` + +The standalone `docker-agent` binary uses Docker's environment variables and +current context. Status, model inspection, and pulls use the same selection. +A failed selected Desktop connection is not replaced with a local runner. + +An explicit model `base_url` or `MODEL_RUNNER_HOST` bypasses local discovery; +a models gateway also bypasses it for inference. If you select a remote engine, +URL, or gateway, prompts are sent there rather than staying on your machine. + ## Configuration ### Inline @@ -269,6 +289,19 @@ models: ## Troubleshooting -- **Plugin not found:** Ensure Docker Model Runner is enabled in Docker Desktop. Docker Agent will fall back to the default URL. +- **Plugin not found:** Ensure Docker Model Runner is enabled in Docker Desktop and `docker model status --json` works with the same Docker context. CLI discovery failures are reported rather than falling back to a different local runner. - **Endpoint empty:** Verify the Model Runner is running with `docker model status --json`. - **Performance:** Use `runtime_flags` to tune GPU layers (`--ngl`) and thread count (`--threads`). + +For interrupted or corrupt downloads, retry using the same Docker connection +and config directory as the agent, for example: + +```console +docker --context desktop-linux --config /path/to/docker-config model pull ai/qwen3 +``` + +Docker Agent does not offer to delete local partial-download files when using a +CLI-selected connection: the runner's content store may be remote or belong to a +different Docker configuration. If a pull repeatedly fails with HTTP 416, inspect +the selected runner's logs and content store rather than deleting files from an +unrelated local `~/.docker/models` directory. diff --git a/pkg/model/provider/dmr/client.go b/pkg/model/provider/dmr/client.go index f0431b233..de003fe01 100644 --- a/pkg/model/provider/dmr/client.go +++ b/pkg/model/provider/dmr/client.go @@ -103,6 +103,8 @@ func NewClient(ctx context.Context, cfg *latest.ModelConfig, opts ...options.Opt case dmrmodels.IsNotInstalledError(err): slog.DebugContext(ctx, "docker model status query failed", "error", err) return nil, ErrNotInstalled + case dmrmodels.HasDockerConnection(ctx): + return nil, err default: // The `docker model` CLI is unusable (broken plugin, docker not on // PATH, ...) but the DMR endpoint may still be up: check model diff --git a/pkg/model/provider/dmr/dmrmodels/docker.go b/pkg/model/provider/dmr/dmrmodels/docker.go new file mode 100644 index 000000000..235c94876 --- /dev/null +++ b/pkg/model/provider/dmr/dmrmodels/docker.go @@ -0,0 +1,35 @@ +package dmrmodels + +import ( + "context" + "net" + "os/exec" + "slices" +) + +type dockerConnectionKey struct{} + +type dockerConnection struct { + args []string + dial func(context.Context) (net.Conn, error) +} + +// ContextWithDockerConnection supplies the Docker CLI's connection settings for +// both Model Runner subprocesses and HTTP requests through the engine. +func ContextWithDockerConnection(ctx context.Context, args []string, dial func(context.Context) (net.Conn, error)) context.Context { + return context.WithValue(ctx, dockerConnectionKey{}, &dockerConnection{args: slices.Clone(args), dial: dial}) +} + +// HasDockerConnection reports whether the caller supplied a Docker connection. +func HasDockerConnection(ctx context.Context) bool { + _, ok := ctx.Value(dockerConnectionKey{}).(*dockerConnection) + return ok +} + +// DockerCommand creates a Docker subprocess using the selected connection. +func DockerCommand(ctx context.Context, args ...string) *exec.Cmd { + if conn, ok := ctx.Value(dockerConnectionKey{}).(*dockerConnection); ok { + args = slices.Concat(conn.args, args) + } + return exec.CommandContext(ctx, "docker", args...) +} diff --git a/pkg/model/provider/dmr/dmrmodels/list.go b/pkg/model/provider/dmr/dmrmodels/list.go index 3d45f27b3..8b17eeb5e 100644 --- a/pkg/model/provider/dmr/dmrmodels/list.go +++ b/pkg/model/provider/dmr/dmrmodels/list.go @@ -64,6 +64,9 @@ func ListModelsWithMetadata(ctx context.Context) ([]Model, error) { if IsNotInstalledError(err) { return nil, ErrNotInstalled } + if HasDockerConnection(ctx) { + return nil, err + } // Otherwise the docker CLI plugin may simply be unavailable while // the engine still serves DMR on a default endpoint, so fall // through and let ResolveBaseURL probe the defaults. @@ -75,6 +78,8 @@ func ListModelsWithMetadata(ctx context.Context) ([]Model, error) { baseURL, httpClient := ResolveBaseURL(ctx, &latest.ModelConfig{}, endpoint) if httpClient == nil { httpClient = &http.Client{} //rubocop:disable Lint/HTTPClientTransport // DMR local service; default transport is appropriate + } else { + defer httpClient.CloseIdleConnections() } return ListModelsWithMetadataAt(ctx, httpClient, baseURL) diff --git a/pkg/model/provider/dmr/dmrmodels/resolve.go b/pkg/model/provider/dmr/dmrmodels/resolve.go index d7d54b0b3..922bc1736 100644 --- a/pkg/model/provider/dmr/dmrmodels/resolve.go +++ b/pkg/model/provider/dmr/dmrmodels/resolve.go @@ -17,7 +17,6 @@ import ( "net/http" "net/url" "os" - "os/exec" "strings" "time" @@ -131,7 +130,7 @@ func getDMRFallbackURLs(containerized bool) []string { // High‑level rules: // - If the user explicitly configured a BaseURL or MODEL_RUNNER_HOST, use that (no fallbacks). // - For Desktop endpoints (model-runner.docker.internal) on the host, route -// through the Docker Engine experimental endpoints prefix over the Unix socket. +// through the selected Docker Engine's experimental endpoints prefix. // - For standalone / offload endpoints like http://172.17.0.1:12435/engines/v1/, // use localhost:/engines/v1/ on the host, and the gateway IP:port inside containers. // - Keep a small compatibility workaround for the legacy http://:0/engines/v1/ endpoint. @@ -151,7 +150,12 @@ func ResolveBaseURL(ctx context.Context, cfg *latest.ModelConfig, endpoint strin } // Resolve primary URL based on endpoint - baseURL, httpClient := resolvePrimaryDMRURL(endpoint) + baseURL, httpClient := resolvePrimaryDMRURL(ctx, endpoint) + + // Never substitute a local runner for an engine selected by the CLI. + if httpClient != nil && HasDockerConnection(ctx) { + return baseURL, httpClient + } // Test connectivity and try fallbacks if needed testClient := cmp.Or(httpClient, &http.Client{}) //rubocop:disable Lint/HTTPClientTransport // DMR connectivity probe; default transport is appropriate @@ -183,7 +187,7 @@ func ResolveBaseURL(ctx context.Context, cfg *latest.ModelConfig, endpoint strin // resolvePrimaryDMRURL resolves the primary DMR URL based on the endpoint string. // This handles the various endpoint formats and platform-specific routing without // connectivity testing or fallbacks. -func resolvePrimaryDMRURL(endpoint string) (string, *http.Client) { +func resolvePrimaryDMRURL(ctx context.Context, endpoint string) (string, *http.Client) { ep := strings.TrimSpace(endpoint) // Legacy bug workaround: old DMR versions <= 0.1.44 could report http://:0/engines/v1/. @@ -197,21 +201,26 @@ func resolvePrimaryDMRURL(endpoint string) (string, *http.Client) { u, err := url.Parse(ep) if err != nil { - slog.Debug("failed to parse DMR endpoint, falling back to defaults", "endpoint", ep, "error", err) + slog.DebugContext(ctx, "failed to parse DMR endpoint, falling back to defaults", "endpoint", ep, "error", err) return defaultForEnvironment(), nil } host := u.Hostname() port := u.Port() - // Desktop endpoint on the host — route through Docker Engine's Unix socket. + // Desktop endpoint on the host — route through the selected Docker Engine. if host == "model-runner.docker.internal" && !inContainer() { expPrefix := strings.TrimPrefix(dmrExperimentalEndpointsPrefix, "/") baseURL := fmt.Sprintf("http://_/%s%s/v1", expPrefix, dmrInferencePrefix) + connection, _ := ctx.Value(dockerConnectionKey{}).(*dockerConnection) httpClient := &http.Client{ Transport: &http.Transport{ + IdleConnTimeout: 30 * time.Second, DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) { + if connection != nil { + return connection.dial(ctx) + } var d net.Dialer return d.DialContext(ctx, "unix", "/var/run/docker.sock") }, @@ -240,12 +249,12 @@ func resolvePrimaryDMRURL(endpoint string) (string, *http.Client) { // DockerModelEndpointAndEngine shells out to `docker model status --json` // and returns the resolved endpoint URL and the active inference engine name. func DockerModelEndpointAndEngine(ctx context.Context) (endpoint, engine string, err error) { - cmd := exec.CommandContext(ctx, "docker", "model", "status", "--json") + cmd := DockerCommand(ctx, "model", "status", "--json") var stdout, stderr bytes.Buffer cmd.Stdout = &stdout cmd.Stderr = &stderr if err := cmd.Run(); err != nil { - return "", "", errors.New(strings.TrimSpace(stderr.String())) + return "", "", fmt.Errorf("%s: %s: %w", strings.Join(cmd.Args, " "), strings.TrimSpace(stderr.String()), err) } var st struct { diff --git a/pkg/model/provider/dmr/dmrmodels/resolve_test.go b/pkg/model/provider/dmr/dmrmodels/resolve_test.go index ec32dded9..18a85528f 100644 --- a/pkg/model/provider/dmr/dmrmodels/resolve_test.go +++ b/pkg/model/provider/dmr/dmrmodels/resolve_test.go @@ -1,6 +1,9 @@ package dmrmodels import ( + "context" + "errors" + "net" "net/http" "net/http/httptest" "testing" @@ -77,3 +80,43 @@ func TestDMRConnectivity(t *testing.T) { assert.False(t, result) }) } + +func TestResolvedDockerTransportRetainsConnection(t *testing.T) { + if inContainer() { + t.Skip("Desktop engine routing is host-only") + } + t.Setenv("MODEL_RUNNER_HOST", "") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/exp/vDD4.40/engines/v1/models", r.URL.Path) + _, _ = w.Write([]byte(`{"data":[{"id":"ai/test"}]}`)) + })) + defer server.Close() + ctx := ContextWithDockerConnection(t.Context(), nil, func(ctx context.Context) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, "tcp", server.Listener.Addr().String()) + }) + baseURL, client := ResolveBaseURL(ctx, nil, defaultContainerURL()) + require.NotNil(t, client) + defer client.CloseIdleConnections() + models, err := ListModelsAt(t.Context(), client, baseURL) + require.NoError(t, err) + assert.Equal(t, []string{"ai/test"}, models) +} + +func TestResolveSelectedDockerDoesNotProbeFallbacks(t *testing.T) { + if inContainer() { + t.Skip("Desktop engine routing is host-only") + } + t.Setenv("MODEL_RUNNER_HOST", "") + calls := 0 + ctx := ContextWithDockerConnection(t.Context(), nil, func(context.Context) (net.Conn, error) { + calls++ + return nil, errors.New("selected engine unavailable") + }) + baseURL, client := ResolveBaseURL(ctx, nil, defaultContainerURL()) + require.NotNil(t, client) + defer client.CloseIdleConnections() + assert.Zero(t, calls, "selection must not probe or switch to a local runner") + _, err := ListModelsAt(t.Context(), client, baseURL) + require.ErrorContains(t, err, "selected engine unavailable") + assert.Equal(t, 1, calls) +} diff --git a/pkg/model/provider/dmr/pull.go b/pkg/model/provider/dmr/pull.go index 18b2bde24..41c5e49b2 100644 --- a/pkg/model/provider/dmr/pull.go +++ b/pkg/model/provider/dmr/pull.go @@ -8,7 +8,6 @@ import ( "io" "log/slog" "os" - "os/exec" "path/filepath" "regexp" "strings" @@ -16,6 +15,7 @@ import ( "golang.org/x/term" "github.com/docker/docker-agent/pkg/input" + "github.com/docker/docker-agent/pkg/model/provider/dmr/dmrmodels" ) func pullDockerModelIfNeeded(ctx context.Context, model string) error { @@ -67,6 +67,11 @@ func pullWithRecovery(ctx context.Context, model string, stdout, stderr io.Write if !ok { return err } + // An injected connection may use a remote runner or another config directory; + // its content store cannot safely be inferred from the local environment. + if dmrmodels.HasDockerConnection(ctx) { + return pfe + } if path, size, ok := corruptPartial(pfe.Detail); ok { pfe.CorruptPartial = path if confirmRemoveCorruptPartial(ctx, stdout, path, size) { @@ -89,7 +94,7 @@ func runModelPull(ctx context.Context, model string, stdout, stderr io.Writer) e slog.InfoContext(ctx, "Pulling DMR model", "model", model) fmt.Fprintf(stdout, "Pulling model %s...\n", model) - cmd := exec.CommandContext(ctx, "docker", "model", "pull", model) + cmd := dmrmodels.DockerCommand(ctx, "model", "pull", model) cmd.Stdout = stdout // Tee stderr so the live pull output still reaches the terminal while we // also capture it, otherwise the real cause (e.g. a registry error) is lost @@ -130,7 +135,7 @@ func confirmModelPull(ctx context.Context, model string, out io.Writer) error { } func modelExists(ctx context.Context, model string) bool { - cmd := exec.CommandContext(ctx, "docker", "model", "inspect", model) + cmd := dmrmodels.DockerCommand(ctx, "model", "inspect", model) var stderr bytes.Buffer cmd.Stdout = io.Discard cmd.Stderr = &stderr diff --git a/pkg/model/provider/dmr/pull_test.go b/pkg/model/provider/dmr/pull_test.go index ce2255830..a664b7548 100644 --- a/pkg/model/provider/dmr/pull_test.go +++ b/pkg/model/provider/dmr/pull_test.go @@ -3,13 +3,17 @@ package dmr import ( "errors" "fmt" + "io" "os" "path/filepath" + "runtime" "strings" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/model/provider/dmr/dmrmodels" ) const testQwenBlobURL = "https://production.cloudfront.docker.com/registry-v2/docker/registry/v2/blobs/sha256/b5/b505f0cf69207567fdc6acec5a6d36303673a7da8cddf030f041677c85681729/data?Expires=1&Signature=x" @@ -175,3 +179,56 @@ func TestPullFailedError(t *testing.T) { assert.NotContains(t, summary, "\n") }) } + +func TestPullUsesSelectedDockerConnection(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim") + } + dir := t.TempDir() + argsFile := filepath.Join(dir, "args") + t.Setenv("DMR_ARGS_FILE", argsFile) + script := `#!/bin/sh +printf '%s\n' "$*" >> "$DMR_ARGS_FILE" +[ "$1" = "--context=desktop-test" ] || exit 2 +[ "$3" = "inspect" ] && exit 1 +exit 0 +` + require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte(script), 0o755)) + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + ctx := dmrmodels.ContextWithDockerConnection(t.Context(), []string{"--context=desktop-test"}, nil) + require.NoError(t, PullTo(ctx, "ai/test", io.Discard, io.Discard)) + got, err := os.ReadFile(argsFile) + require.NoError(t, err) + assert.Equal(t, "--context=desktop-test model inspect ai/test\n--context=desktop-test model pull ai/test\n", string(got)) +} + +func TestPullSelectedConnectionDoesNotRecoverLocalPartial(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses a shell Docker shim") + } + dir := t.TempDir() + script := "#!/bin/sh\nprintf '%s\n' '" + testQwenBlobURL + ": 416 Requested Range Not Satisfiable' >&2\nexit 1\n" + require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte(script), 0o755)) + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + localConfig := t.TempDir() + t.Setenv("DOCKER_CONFIG", localConfig) + blobDir := filepath.Join(localConfig, "models", "blobs", "sha256") + require.NoError(t, os.MkdirAll(blobDir, 0o755)) + partial := filepath.Join(blobDir, "b505f0cf69207567fdc6acec5a6d36303673a7da8cddf030f041677c85681729.incomplete") + require.NoError(t, os.WriteFile(partial, []byte("unrelated download"), 0o600)) + + for _, args := range [][]string{ + {"--context=remote-desktop"}, + {"--context=desktop-linux", "--config=" + t.TempDir()}, + } { + ctx := dmrmodels.ContextWithDockerConnection(t.Context(), args, nil) + err := PullTo(ctx, "ai/test", io.Discard, io.Discard) + var failure *PullFailedError + require.ErrorAs(t, err, &failure) + assert.Empty(t, failure.CorruptPartial) + assert.NotContains(t, err.Error(), partial) + data, err := os.ReadFile(partial) + require.NoError(t, err) + assert.Equal(t, "unrelated download", string(data)) + } +} diff --git a/pkg/server/server.go b/pkg/server/server.go index 59388863b..562a80db6 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -147,6 +147,7 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { // manager) start fresh trace ids per request rather than // chaining onto the calling client's trace. srv := http.Server{ + BaseContext: func(net.Listener) context.Context { return ctx }, Handler: otelhttp.NewHandler(s.e, "agent-api"), ReadHeaderTimeout: 10 * time.Second, } diff --git a/pkg/server/server_test.go b/pkg/server/server_test.go index 80be4ff22..6149e764a 100644 --- a/pkg/server/server_test.go +++ b/pkg/server/server_test.go @@ -16,6 +16,7 @@ import ( "testing" "time" + "github.com/labstack/echo/v4" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -618,3 +619,26 @@ func (s mockStore) GetSessions(context.Context) ([]*session.Session, error) { func (s mockStore) GetSessionSummaries(context.Context) ([]session.Summary, error) { return nil, nil } + +func TestServerPreservesServingContext(t *testing.T) { + t.Parallel() + type contextKey struct{} + ctx := context.WithValue(t.Context(), contextKey{}, "selected Docker connection") + srv := NewWithManager(nil, "") + srv.e.GET("/context", func(c echo.Context) error { + return c.String(http.StatusOK, c.Request().Context().Value(contextKey{}).(string)) + }) + var lc net.ListenConfig + ln, err := lc.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + go func() { _ = srv.Serve(ctx, ln) }() + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://"+ln.Addr().String()+"/context", http.NoBody) + require.NoError(t, err) + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, "selected Docker connection", string(body)) +}