server: allow registry cross-host redirects among allowlisted hosts (#18533)

This commit is contained in:
Patrick Devine
2026-09-18 17:15:35 -07:00
committed by GitHub
parent d0c8cdb795
commit 6383a0fa9c
2 changed files with 102 additions and 9 deletions
+21 -5
View File
@@ -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
}
}
+81 -4
View File
@@ -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, &registryOptions{})
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")
}
})
}
}