Skip to content
Draft
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
222 changes: 222 additions & 0 deletions pkg/distribution/oci/remote/docker_token_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,222 @@
package remote

import (
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"fmt"
"io"
"math/big"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync/atomic"
"testing"
"time"

"github.com/docker/model-runner/pkg/distribution/oci/reference"
)

func TestTrustedDockerTokenEndpoint(t *testing.T) {
for _, registry := range []string{"docker.io", "index.docker.io", "registry-1.docker.io"} {
for _, endpoint := range []string{"https://auth.docker.io/token", "https://AUTH.DOCKER.IO:443/token?service=registry.docker.io"} {
u, err := url.Parse(endpoint)
if err != nil {
t.Fatal(err)
}
if !isTrustedDockerTokenEndpoint(u, registry) {
t.Errorf("trusted Docker endpoint rejected: registry=%s endpoint=%s", registry, endpoint)
}
}
}
for _, endpoint := range []string{
"http://auth.docker.io/token",
"https://auth.docker.io:444/token",
"https://auth.docker.io.evil.example/token",
"https://evil-auth.docker.io/token",
"https://auth.docker.io./token",
"https://user@auth.docker.io/token",
"https://auth.docker.io/admin",
"https://auth.docker.io/%74oken",
"https://auth.docker.io/token#fragment",
"https://127.0.0.1/token",
} {
u, err := url.Parse(endpoint)
if err != nil {
t.Fatal(err)
}
if isTrustedDockerTokenEndpoint(u, "registry-1.docker.io") {
t.Errorf("untrusted endpoint accepted: %s", endpoint)
}
}
u, _ := url.Parse("https://auth.docker.io/token")
for _, registry := range []string{"", "evil.example", "registry-1.docker.io.evil.example", "registry-1.docker.io:444"} {
if isTrustedDockerTokenEndpoint(u, registry) {
t.Errorf("Docker exception applied to unrelated registry: %s", registry)
}
}
}

// dockerTokenProxy serves a TLS-verified Docker token endpoint behind a local
// CONNECT proxy. Its custom dialer models Desktop's supplied proxy transport.
func dockerTokenProxy(t *testing.T, handler http.Handler) *http.Transport {
t.Helper()
publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
certificate := &x509.Certificate{
SerialNumber: big.NewInt(1),
DNSNames: []string{"auth.docker.io", "registry-1.docker.io", "index.docker.io"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
}
der, err := x509.CreateCertificate(rand.Reader, certificate, certificate, publicKey, privateKey)
if err != nil {
t.Fatal(err)
}
parsed, err := x509.ParseCertificate(der)
if err != nil {
t.Fatal(err)
}
roots := x509.NewCertPool()
roots.AddCert(parsed)
endpoint := httptest.NewUnstartedServer(handler)
endpoint.TLS = &tls.Config{Certificates: []tls.Certificate{{Certificate: [][]byte{der}, PrivateKey: privateKey}}, MinVersion: tls.VersionTLS12}
endpoint.StartTLS()
t.Cleanup(endpoint.Close)
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodConnect {
http.Error(w, "CONNECT required", http.StatusBadRequest)
return
}
upstream, err := (&net.Dialer{}).DialContext(r.Context(), "tcp", endpoint.Listener.Addr().String())
if err != nil {
http.Error(w, err.Error(), http.StatusBadGateway)
return
}
defer upstream.Close()
client, buffered, err := w.(http.Hijacker).Hijack()
if err != nil {
return
}
defer client.Close()
fmt.Fprint(buffered, "HTTP/1.1 200 Connection Established\r\n\r\n")
buffered.Flush()
done := make(chan struct{})
go func() {
io.Copy(client, upstream)
client.Close()
close(done)
}()
io.Copy(upstream, buffered)
upstream.Close()
<-done
}))
t.Cleanup(proxy.Close)
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.Proxy = http.ProxyURL(&url.URL{Scheme: "http", Host: "desktop-proxy:80"})
transport.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) {
if address != "desktop-proxy:80" {
return nil, fmt.Errorf("supplied dialer expected proxy address, got %s", address)
}
return (&net.Dialer{}).DialContext(ctx, network, proxy.Listener.Addr().String())
}
transport.TLSClientConfig = &tls.Config{RootCAs: roots, MinVersion: tls.VersionTLS12}
t.Cleanup(transport.CloseIdleConnections)
return transport
}

func TestDockerTokenWithoutLocalDNS(t *testing.T) {
var lookups atomic.Int32
originalResolver := net.DefaultResolver
net.DefaultResolver = &net.Resolver{PreferGo: true, Dial: func(context.Context, string, string) (net.Conn, error) {
lookups.Add(1)
return nil, fmt.Errorf("public DNS unavailable")
}}
t.Cleanup(func() { net.DefaultResolver = originalResolver })
var hits atomic.Int32
transport := dockerTokenProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hits.Add(1)
if r.Host != "auth.docker.io" || r.URL.Path != "/token" {
http.Error(w, "unexpected token endpoint", http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "application/json")
fmt.Fprintln(w, `{"token":"via-proxy"}`)
}))
client := newGuardedAuthClient(transport, "registry-1.docker.io")
guard := client.Transport.(*guardedAuthTransport)
t.Cleanup(guard.proxied.CloseIdleConnections)
t.Cleanup(guard.direct.CloseIdleConnections)
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "https://auth.docker.io/token", http.NoBody)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatalf("guarded Docker token request failed: %v", err)
}
resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("token request status: %d", resp.StatusCode)
}
ref, err := reference.ParseReference("ai/gemma3:4b")
if err != nil {
t.Fatal(err)
}
token, err := Exchange(t.Context(), ref.Context().Registry, nil, transport, nil, &PingResponse{
WWWAuthenticate: WWWAuthenticate{Realm: "https://auth.docker.io/token"},
})
if err != nil || token.Token != "via-proxy" {
t.Fatalf("Exchange failed: token=%v error=%v", token, err)
}
if hits.Load() != 2 || lookups.Load() != 0 {
t.Fatalf("expected two proxied requests without DNS, got hits=%d DNS=%d", hits.Load(), lookups.Load())
}
// Unrelated registries still fail closed, even for the Docker token URL.
otherClient := newGuardedAuthClient(transport, "evil.example")
_, err = otherClient.Do(req)
if err == nil || !strings.Contains(err.Error(), "resolving realm hostname") {
t.Fatalf("untrusted registry should require local DNS: %v", err)
}
// Without a proxy, the validating dialer still requires local DNS.
directTransport := transport.Clone()
directTransport.Proxy = nil
_, err = newGuardedAuthClient(directTransport, "registry-1.docker.io").Do(req)
if err == nil || !strings.Contains(err.Error(), "resolving realm hostname") {
t.Fatalf("direct connection should require validated DNS: %v", err)
}
if hits.Load() != 2 {
t.Fatalf("blocked requests reached proxy: hits=%d", hits.Load())
}
}

func TestDockerTokenRedirectToPrivateRealmBlocked(t *testing.T) {
var hits atomic.Int32
transport := dockerTokenProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hits.Add(1)
http.Redirect(w, r, "https://169.254.169.254/token", http.StatusFound)
}))
client := newGuardedAuthClient(transport, "registry-1.docker.io")
guard := client.Transport.(*guardedAuthTransport)
t.Cleanup(guard.proxied.CloseIdleConnections)
t.Cleanup(guard.direct.CloseIdleConnections)
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "https://auth.docker.io/token", http.NoBody)
if err != nil {
t.Fatal(err)
}
_, err = client.Do(req)
if err == nil || !strings.Contains(err.Error(), "disallowed IP address") {
t.Fatalf("redirect to private realm should be rejected: %v", err)
}
if hits.Load() != 1 {
t.Fatalf("redirect reached proxy endpoint: hits=%d", hits.Load())
}
}
4 changes: 2 additions & 2 deletions pkg/distribution/oci/remote/guard_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ func TestNewGuardedAuthClientBlocksLoopback(t *testing.T) {
}))
defer internalService.Close()

client := newGuardedAuthClient(nil)
client := newGuardedAuthClient(nil, "")
resp, err := client.Get(internalService.URL) //nolint:noctx
if err == nil {
resp.Body.Close()
Expand Down Expand Up @@ -50,7 +50,7 @@ func TestGuardedAuthClientHonorsProxyOnPrivateAddress(t *testing.T) {
base := http.DefaultTransport.(*http.Transport).Clone()
base.Proxy = http.ProxyURL(proxyURL)

client := newGuardedAuthClient(base)
client := newGuardedAuthClient(base, "")

// A public realm must be reachable through the loopback proxy. 203.0.113.0/24
// is TEST-NET-3: never routable, so a hit proves the request went via the proxy.
Expand Down
4 changes: 2 additions & 2 deletions pkg/distribution/oci/remote/remote.go
Original file line number Diff line number Diff line change
Expand Up @@ -425,7 +425,7 @@ type resolverComponents struct {
func createResolver(o *options, ref reference.Reference) resolverComponents {
authorizer := docker.NewDockerAuthorizer(
docker.WithAuthCreds(credentialsFunc(o, ref)),
docker.WithAuthClient(newGuardedAuthClient(o.transport)))
docker.WithAuthClient(newGuardedAuthClient(o.transport, ref.Context().Registry.RegistryStr())))

// Wrap transport with Range header support for resumable downloads
// and User-Agent header for registry compatibility (required by HuggingFace)
Expand Down Expand Up @@ -530,7 +530,7 @@ func createResolverWithPushScope(o *options, ref reference.Reference) (resolverC
}
return cfg.Username, cfg.Password, nil
}),
docker.WithAuthClient(newGuardedAuthClient(o.transport)))
docker.WithAuthClient(newGuardedAuthClient(o.transport, ref.Context().Registry.RegistryStr())))

resolver := docker.NewResolver(docker.ResolverOptions{
Hosts: docker.ConfigureDefaultRegistries(
Expand Down
56 changes: 37 additions & 19 deletions pkg/distribution/oci/remote/transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -90,13 +90,29 @@ func isDisallowedIP(ip net.IP) bool {
return false
}

// validateTokenEndpointURL validates the host of a token-endpoint URL against
// the internal-hostname blocklist and the private/loopback/link-local ranges.
// The local DNS resolution this performs is deliberate even when a proxy will
// resolve the name itself: checking the resolved IPs is the validation, and a
// name that cannot be resolved locally is rejected (fail closed) rather than
// forwarded unchecked.
func validateTokenEndpointURL(u *url.URL) error {
// isTrustedDockerTokenEndpoint recognizes Docker Hub's token endpoint only when
// authenticating to Docker Hub. Exact HTTPS authority and path matching avoids
// extending this exception to arbitrary hosts, ports, or authentication services.
func isTrustedDockerTokenEndpoint(u *url.URL, registryHost string) bool {
switch strings.ToLower(registryHost) {
case "docker.io", "index.docker.io", "registry-1.docker.io":
default:
return false
}
return u.Scheme == "https" && u.User == nil && u.Fragment == "" &&
(strings.EqualFold(u.Host, "auth.docker.io") || strings.EqualFold(u.Host, "auth.docker.io:443")) &&
u.EscapedPath() == "/token"
}

// validateTokenEndpointURL validates untrusted token endpoints against the
// internal-hostname and private/loopback/link-local blocklists. Docker Hub's
// trusted HTTPS token endpoint does not require local DNS resolution: a proxy
// may resolve it on networks where public DNS is unavailable to the client.
// Direct connections still validate and pin the resolved IP in their dialer.
func validateTokenEndpointURL(u *url.URL, registryHost string) error {
if isTrustedDockerTokenEndpoint(u, registryHost) {
return nil
}
port := u.Port()
if port == "" {
if u.Scheme == "https" {
Expand Down Expand Up @@ -153,15 +169,15 @@ func resolveAndValidateHost(hostname, port string) (dialAddr string, err error)
// both by containerd's authorizer (via docker.WithAuthClient) and by the
// hand-rolled Exchange(). The realm URL in a registry's WWW-Authenticate
// challenge is attacker-controlled, so every request this client makes is
// validated against the internal-hostname blocklist and the private/loopback/
// link-local IP ranges before a connection is established.
// validated before a connection is established, with a narrow exception for
// Docker Hub's trusted token endpoint when using a proxy.
//
// How the connection is guarded depends on whether a proxy applies to the
// request (see guardedAuthTransport). A dial-time-only guard would break every
// proxied deployment: with a proxy configured, the dialer sees the proxy's
// address — commonly a private or loopback IP — rather than the realm's, and
// would reject the proxy itself.
func newGuardedAuthClient(base http.RoundTripper) *http.Client {
func newGuardedAuthClient(base http.RoundTripper, registryHost string) *http.Client {
var proxied *http.Transport
if t, ok := base.(*http.Transport); ok {
proxied = t.Clone()
Expand All @@ -185,7 +201,7 @@ func newGuardedAuthClient(base http.RoundTripper) *http.Client {
return (&net.Dialer{}).DialContext(ctx, network, dialAddr)
}

return &http.Client{Transport: &guardedAuthTransport{proxied: proxied, direct: direct}}
return &http.Client{Transport: &guardedAuthTransport{proxied: proxied, direct: direct, registryHost: registryHost}}
}

// guardedAuthTransport validates every token-endpoint request against the SSRF
Expand All @@ -195,12 +211,14 @@ func newGuardedAuthClient(base http.RoundTripper) *http.Client {
// IP just before connecting and dials that exact address, so DNS rebinding
// cannot slip an internal address past the check.
// - Proxied connections go through a transport with the proxy configuration
// intact and a stock dialer: the proxy is the one connecting to the realm,
// intact and the supplied dialer: the proxy is the one connecting to the realm,
// so pinning the dial address is neither possible nor meaningful. The
// realm host is validated here at the request level instead.
// realm host is validated here at the request level instead, except for
// Docker Hub's trusted HTTPS token endpoint. Redirects are checked as well.
type guardedAuthTransport struct {
proxied *http.Transport // proxy settings intact, stock dialer
direct *http.Transport // no proxy, validating dialer pinned to the resolved IP
proxied *http.Transport // proxy settings and supplied dialer intact
direct *http.Transport // no proxy, validating dialer pinned to the resolved IP
registryHost string
}

func (g *guardedAuthTransport) RoundTrip(req *http.Request) (*http.Response, error) {
Expand All @@ -210,7 +228,7 @@ func (g *guardedAuthTransport) RoundTrip(req *http.Request) (*http.Response, err
return nil, fmt.Errorf("determining proxy for token endpoint: %w", err)
}
if proxyURL != nil {
if err := validateTokenEndpointURL(req.URL); err != nil {
if err := validateTokenEndpointURL(req.URL, g.registryHost); err != nil {
return nil, fmt.Errorf("realm URL rejected: %w", err)
}
return g.proxied.RoundTrip(req)
Expand Down Expand Up @@ -288,8 +306,8 @@ func parseWWWAuthenticate(header string) WWWAuthenticate {
// the registry's WWW-Authenticate challenge and is therefore untrusted; the
// guarded client rejects realms on internal hostnames or private/loopback
// addresses and honors any configured proxy.
func Exchange(ctx context.Context, _ reference.Registry, auth authn.Authenticator, transport http.RoundTripper, scopes []string, pr *PingResponse) (*Token, error) {
client := newGuardedAuthClient(transport)
func Exchange(ctx context.Context, reg reference.Registry, auth authn.Authenticator, transport http.RoundTripper, scopes []string, pr *PingResponse) (*Token, error) {
client := newGuardedAuthClient(transport, reg.RegistryStr())

// Build token request URL
tokenURL, err := url.Parse(pr.WWWAuthenticate.Realm)
Expand All @@ -300,7 +318,7 @@ func Exchange(ctx context.Context, _ reference.Registry, auth authn.Authenticato
// Validate the realm before any request is made so a blocked realm fails
// fast with a clear error. The guarded client re-validates at connection
// time (or per request when proxied), closing the TOCTOU window.
if err := validateTokenEndpointURL(tokenURL); err != nil {
if err := validateTokenEndpointURL(tokenURL, reg.RegistryStr()); err != nil {
return nil, fmt.Errorf("realm URL rejected: %w", err)
}

Expand Down
Loading