diff --git a/pkg/model/provider/dmr/dmrmodels/docker.go b/pkg/model/provider/dmr/dmrmodels/docker.go index 235c94876..232bceb76 100644 --- a/pkg/model/provider/dmr/dmrmodels/docker.go +++ b/pkg/model/provider/dmr/dmrmodels/docker.go @@ -3,8 +3,10 @@ package dmrmodels import ( "context" "net" + "net/http" "os/exec" "slices" + "strings" ) type dockerConnectionKey struct{} @@ -33,3 +35,17 @@ func DockerCommand(ctx context.Context, args ...string) *exec.Cmd { } return exec.CommandContext(ctx, "docker", args...) } + +type dockerTransport struct { + *http.Transport + + args []string +} + +func (t *dockerTransport) RoundTrip(req *http.Request) (*http.Response, error) { + resp, err := t.Transport.RoundTrip(req) + if err != nil { + return nil, &net.OpError{Op: "docker connection", Net: strings.Join(t.args, " "), Err: err} + } + return resp, nil +} diff --git a/pkg/model/provider/dmr/dmrmodels/resolve.go b/pkg/model/provider/dmr/dmrmodels/resolve.go index 922bc1736..019582fc9 100644 --- a/pkg/model/provider/dmr/dmrmodels/resolve.go +++ b/pkg/model/provider/dmr/dmrmodels/resolve.go @@ -214,18 +214,20 @@ func resolvePrimaryDMRURL(ctx context.Context, endpoint string) (string, *http.C 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") - }, + 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") }, } + httpClient := &http.Client{Transport: transport} + if connection != nil { + httpClient.Transport = &dockerTransport{Transport: transport, args: connection.args} + } return baseURL, httpClient } diff --git a/pkg/model/provider/dmr/dmrmodels/resolve_test.go b/pkg/model/provider/dmr/dmrmodels/resolve_test.go index 18a85528f..677df3136 100644 --- a/pkg/model/provider/dmr/dmrmodels/resolve_test.go +++ b/pkg/model/provider/dmr/dmrmodels/resolve_test.go @@ -3,9 +3,13 @@ package dmrmodels import ( "context" "errors" + "io" "net" "net/http" "net/http/httptest" + "net/url" + "os" + "syscall" "testing" "github.com/stretchr/testify/assert" @@ -102,6 +106,59 @@ func TestResolvedDockerTransportRetainsConnection(t *testing.T) { assert.Equal(t, []string{"ai/test"}, models) } +func TestSelectedDockerTransportErrorsNameEngine(t *testing.T) { + if inContainer() { + t.Skip("Desktop engine routing is host-only") + } + t.Setenv("MODEL_RUNNER_HOST", "") + for _, tt := range []struct { + name string + err error + dial bool + }{ + {name: "dial", err: net.ErrClosed, dial: true}, + {name: "closed pipe", err: os.ErrClosed}, + {name: "broken pipe", err: syscall.EPIPE}, + {name: "EOF", err: io.EOF, dial: true}, + {name: "canceled", err: context.Canceled}, + {name: "deadline", err: context.DeadlineExceeded, dial: true}, + } { + t.Run(tt.name, func(t *testing.T) { + ctx := ContextWithDockerConnection(t.Context(), []string{"--host=unix:///missing-dmr-engine.sock"}, func(context.Context) (net.Conn, error) { + if tt.dial { + return nil, tt.err + } + conn, peer := net.Pipe() + t.Cleanup(func() { _ = peer.Close() }) + return &failingDockerConn{Conn: conn, err: tt.err}, nil + }) + baseURL, client := ResolveBaseURL(ctx, nil, defaultContainerURL()) + require.NotNil(t, client) + defer client.CloseIdleConnections() + _, err := ListModelsAt(t.Context(), client, baseURL) + require.ErrorIs(t, err, tt.err) + require.ErrorContains(t, err, "unix:///missing-dmr-engine.sock") + var requestErr *url.Error + require.ErrorAs(t, err, &requestErr) + assert.Equal(t, errors.Is(tt.err, context.DeadlineExceeded), requestErr.Timeout()) + }) + } +} + +type failingDockerConn struct { + net.Conn + + err error +} + +func (c *failingDockerConn) Read([]byte) (int, error) { + return 0, c.err +} + +func (c *failingDockerConn) Write([]byte) (int, error) { + return 0, c.err +} + func TestResolveSelectedDockerDoesNotProbeFallbacks(t *testing.T) { if inContainer() { t.Skip("Desktop engine routing is host-only")