From 6383a0fa9cbf97494b847226e189f6e36b401a08 Mon Sep 17 00:00:00 2001 From: Patrick Devine Date: Fri, 18 Sep 2026 17:15:35 -0700 Subject: [PATCH] server: allow registry cross-host redirects among allowlisted hosts (#18533) --- server/images.go | 26 ++++++++++--- server/images_test.go | 85 +++++++++++++++++++++++++++++++++++++++++-- 2 files changed, 102 insertions(+), 9 deletions(-) diff --git a/server/images.go b/server/images.go index f2e985e49..1c19d8154 100644 --- a/server/images.go +++ b/server/images.go @@ -1376,6 +1376,18 @@ var testMakeRequestDialContext func(ctx context.Context, network, addr string) ( var errBlockedRedirect = errors.New("blocked redirect to a different host") +// isAllowedHost reports whether host may receive cross-host redirects. +var allowedRedirectHosts = []string{"ollama.com", "ollama.ai", "hf.co", "huggingface.co"} + +func isAllowedHost(host string) bool { + for _, h := range allowedRedirectHosts { + if host == h || strings.HasSuffix(host, "."+h) { + return true + } + } + return false +} + func makeRequest(ctx context.Context, method string, requestURL *url.URL, headers http.Header, body io.Reader, regOpts *registryOptions) (*http.Response, error) { if requestURL.Scheme != "http" && regOpts != nil && regOpts.Insecure { requestURL.Scheme = "http" @@ -1416,16 +1428,20 @@ func makeRequest(ctx context.Context, method string, requestURL *url.URL, header if checkRedirect == nil { insecure := regOpts != nil && regOpts.Insecure // Default redirect policy: same-host only, so a registry can't steer - // manifest or blob requests at internal addresses. --insecure opts out - // for trusted LAN/local registries. + // manifest or blob requests at internal addresses. CDN-backed + // registries redirect among their own hosts, allowed via + // isAllowedHost. --insecure opts out for trusted LAN/local registries. checkRedirect = func(req *http.Request, via []*http.Request) error { if len(via) > 10 { return errMaxRedirectsExceeded } - if !insecure && req.URL.Host != via[0].URL.Host { - return errBlockedRedirect + if insecure || req.URL.Host == via[0].URL.Host { + return nil } - return nil + if isAllowedHost(via[0].URL.Hostname()) && isAllowedHost(req.URL.Hostname()) { + return nil + } + return errBlockedRedirect } } diff --git a/server/images_test.go b/server/images_test.go index 7852f1cac..a5502e26e 100644 --- a/server/images_test.go +++ b/server/images_test.go @@ -1,9 +1,11 @@ package server import ( + "context" "crypto/sha256" "errors" "fmt" + "net" "net/http" "net/http/httptest" "net/url" @@ -916,10 +918,9 @@ func TestPullModelDuplicateDigestVerifiesBlob(t *testing.T) { } } -// TestPullManifestRejectsCrossHostRedirect: a manifest GET that the registry -// redirects to a different host must be refused by default, so a malicious -// registry can't turn a pull into a request to an internal address. -// --insecure opts out for trusted registries. +// TestPullManifestRejectsCrossHostRedirect: a registry can't redirect a +// pull at an internal address; cross-host redirects to public addresses +// (hf.co's CDN) are fine. --insecure opts out. func TestPullManifestRejectsCrossHostRedirect(t *testing.T) { t.Setenv("OLLAMA_MODELS", t.TempDir()) @@ -966,3 +967,79 @@ func TestPullManifestRejectsCrossHostRedirect(t *testing.T) { t.Fatal("redirect target was not reached with Insecure set") } } + +// TestPullManifestRedirectPolicy: cross-host redirects are blocked by +// default except between allowlisted hosts. +func TestPullManifestRedirectPolicy(t *testing.T) { + for _, tc := range []struct { + name string + origin string // registry host receiving the initial request + target string // redirect target + allowed bool + }{ + {name: "hf to cdn sibling", origin: "hf.co", target: "us.aws.cdn.hf.co", allowed: true}, + {name: "hf to huggingface", origin: "hf.co", target: "huggingface.co", allowed: true}, + {name: "ollama registry to cdn", origin: "registry.ollama.ai", target: "cdn.ollama.com", allowed: true}, + {name: "public third party", origin: "hf.co", target: "93.184.216.34", allowed: false}, + {name: "other registry cross-host", origin: "registry.example.com", target: "cdn.example.com", allowed: false}, + } { + t.Run(tc.name, func(t *testing.T) { + var hit bool + cdn := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hit = true + w.Write([]byte("ok")) + })) + defer cdn.Close() + _, cdnPort, err := net.SplitHostPort(strings.TrimPrefix(cdn.URL, "http://")) + if err != nil { + t.Fatal(err) + } + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "http://"+net.JoinHostPort(tc.target, cdnPort)+r.URL.Path, http.StatusFound) + })) + defer ts.Close() + + // Steer all dials at the local servers so tests stay offline. + prev := testMakeRequestDialContext + testMakeRequestDialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + host, _, err := net.SplitHostPort(addr) + if err != nil { + return nil, err + } + if host == tc.target { + addr = net.JoinHostPort("127.0.0.1", cdnPort) + } else { + _, port, _ := net.SplitHostPort(strings.TrimPrefix(ts.URL, "http://")) + addr = net.JoinHostPort("127.0.0.1", port) + } + return new(net.Dialer).DialContext(ctx, network, addr) + } + defer func() { testMakeRequestDialContext = prev }() + + requestURL, err := url.Parse(ts.URL + "/v2/unsloth/model/manifests/latest") + if err != nil { + t.Fatal(err) + } + requestURL.Host = net.JoinHostPort(tc.origin, requestURL.Port()) + + resp, err := makeRequest(t.Context(), http.MethodGet, requestURL, nil, nil, ®istryOptions{}) + if tc.allowed { + if err != nil { + t.Fatalf("makeRequest = %v, want %s -> %s followed", err, tc.origin, tc.target) + } + resp.Body.Close() + if !hit { + t.Fatal("redirect target not reached") + } + return + } + if !errors.Is(err, errBlockedRedirect) { + t.Fatalf("makeRequest = %v, want errBlockedRedirect", err) + } + if hit { + t.Fatal("blocked redirect target received a request") + } + }) + } +}