Skip to content
Open
Show file tree
Hide file tree
Changes from 2 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
126 changes: 121 additions & 5 deletions pkg/compose/script_source.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
Expand Down Expand Up @@ -35,16 +36,23 @@ func (f ScriptSourceResolverFunc) Resolve(ctx context.Context, source sources.So
}

type defaultScriptSourceResolver struct {
client *http.Client
env map[string]string
client *http.Client
env map[string]string
validateNetworkTarget func(context.Context, *url.URL) error
}

// NewDefaultScriptSourceResolver returns the bounded file and HTTP(S) resolver
// used by CLI compose loading.
func NewDefaultScriptSourceResolver(env map[string]string) ScriptSourceResolver {
resolver := &defaultScriptSourceResolver{env: env}
resolver := &defaultScriptSourceResolver{env: env, validateNetworkTarget: validateScriptNetworkTarget}
transport := http.DefaultTransport.(*http.Transport).Clone()
// Keep the standard proxy behavior for environments that require a forward
Comment thread
monkeyscan[bot] marked this conversation as resolved.
Outdated
Comment thread
monkeyscan[bot] marked this conversation as resolved.
Outdated
// proxy. The custom dialer still validates every resolved destination before
// connecting, so proxy configuration cannot bypass the private-address check.
transport.DialContext = safeScriptDialContext
resolver.client = &http.Client{
Timeout: defaultScriptSourceTimeout,
Timeout: defaultScriptSourceTimeout,
Transport: transport,
CheckRedirect: func(req *http.Request, via []*http.Request) error {
if len(via) > maxScriptSourceRedirects {
return fmt.Errorf("too many redirects (maximum %d)", maxScriptSourceRedirects)
Expand All @@ -61,7 +69,7 @@ func NewDefaultScriptSourceResolver(env map[string]string) ScriptSourceResolver
if len(via) > 0 && via[len(via)-1].URL.Scheme == "https" && req.URL.Scheme == "http" {
return errors.New("HTTPS redirect downgrade to HTTP is not allowed")
}
return nil
return resolver.validateNetworkTarget(req.Context(), req.URL)
},
}
return resolver
Expand Down Expand Up @@ -234,6 +242,9 @@ func readScriptFile(path string) ([]byte, error) {
}

func (r *defaultScriptSourceResolver) readHTTP(ctx context.Context, location *url.URL, source sources.Source) ([]byte, error) {
if err := r.validateNetworkTarget(ctx, location); err != nil {
return nil, fmt.Errorf("fetch script from %s: %w", redactedScriptURL(location), err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, location.String(), nil)
if err != nil {
return nil, fmt.Errorf("create script request for %s", redactedScriptURL(location))
Expand All @@ -250,6 +261,111 @@ func (r *defaultScriptSourceResolver) readHTTP(ctx context.Context, location *ur
return readLimitedScript(resp.Body)
}

func validateScriptNetworkTarget(ctx context.Context, target *url.URL) error {
host := strings.TrimSpace(target.Hostname())
if host == "" {
return errors.New("script URL requires a valid host")
}
addresses, err := net.DefaultResolver.LookupIPAddr(ctx, host)
if err != nil {
return fmt.Errorf("resolve script URL host: %w", err)
}
if len(addresses) == 0 {
return errors.New("script URL host resolved to no addresses")
}
for _, address := range addresses {
if !isPublicScriptAddress(address.IP) {
return fmt.Errorf("script URL host resolves to prohibited address %s", address.IP)
}
}
return nil
}

func safeScriptDialContext(ctx context.Context, network, address string) (net.Conn, error) {
dialer := net.Dialer{}
return safeScriptDialContextWithResolver(ctx, network, address,
func(ctx context.Context, host string) ([]net.IP, error) {
addresses, err := net.DefaultResolver.LookupIPAddr(ctx, host)
if err != nil {
return nil, err
}
ips := make([]net.IP, 0, len(addresses))
for _, address := range addresses {
ips = append(ips, address.IP)
}
return ips, nil
}, dialer.DialContext)
}

func safeScriptDialContextWithResolver(
ctx context.Context,
network, address string,
lookup func(context.Context, string) ([]net.IP, error),
dial func(context.Context, string, string) (net.Conn, error),
) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, fmt.Errorf("parse script endpoint: %w", err)
}
addresses, err := lookup(ctx, host)
if err != nil {
return nil, fmt.Errorf("resolve script endpoint: %w", err)
}
var dialErrs error
for _, candidate := range addresses {
if !isPublicScriptAddress(candidate) {
continue
}
conn, dialErr := dial(ctx, network, net.JoinHostPort(candidate.String(), port))
if dialErr == nil {
return conn, nil
}
dialErrs = errors.Join(dialErrs, dialErr)
}
if dialErrs != nil {
return nil, dialErrs
}
return nil, errors.New("script endpoint has no permitted public address")
}
Comment thread
monkeyscan[bot] marked this conversation as resolved.

func isPublicScriptAddress(ip net.IP) bool {
if ip == nil || !ip.IsGlobalUnicast() || ip.IsPrivate() || ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
return false
}
for _, network := range prohibitedScriptNetworks {
if network.Contains(ip) {
return false
}
}
return true
}

var prohibitedScriptNetworks = mustParseScriptNetworks([]string{
"0.0.0.0/8",
"100.64.0.0/10", // RFC 6598 carrier-grade NAT.
"192.0.0.0/24",
"192.0.2.0/24", // RFC 5737 documentation.
"198.18.0.0/15", // RFC 2544 benchmarking.
"198.51.100.0/24", // RFC 5737 documentation.
"203.0.113.0/24", // RFC 5737 documentation.
"224.0.0.0/4",
"240.0.0.0/4",
"fec0::/10", // deprecated IPv6 site-local.
"2001:db8::/32", // IPv6 documentation.
})

func mustParseScriptNetworks(cidrs []string) []*net.IPNet {
networks := make([]*net.IPNet, 0, len(cidrs))
for _, cidr := range cidrs {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
panic(fmt.Sprintf("invalid script network %q: %v", cidr, err))
}
networks = append(networks, network)
}
return networks
}
Comment thread
monkeyscan[bot] marked this conversation as resolved.

func readLimitedScript(reader io.Reader) ([]byte, error) {
data, err := io.ReadAll(io.LimitReader(reader, maxScriptSourceBytes+1))
if err != nil {
Expand Down
101 changes: 91 additions & 10 deletions pkg/compose/script_source_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,20 +3,90 @@ package compose
import (
"compress/gzip"
"context"
"errors"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/url"
"os"
"os/exec"
"path/filepath"
"reflect"
"strings"
"testing"
"time"

"github.com/chaitin/agent-compose/pkg/sources"
)

func newTestScriptSourceResolver() *defaultScriptSourceResolver {
resolver := NewDefaultScriptSourceResolver(nil).(*defaultScriptSourceResolver)
resolver.validateNetworkTarget = func(context.Context, *url.URL) error { return nil }
transport := resolver.client.Transport.(*http.Transport).Clone()
transport.DialContext = (&net.Dialer{}).DialContext
resolver.client.Transport = transport
return resolver
}

func TestDefaultScriptSourceResolverRejectsPrivateHTTPTargets(t *testing.T) {
resolver := NewDefaultScriptSourceResolver(nil)
for _, location := range []string{
"http://127.0.0.1/script.js",
"http://169.254.169.254/latest/meta-data/",
"http://10.0.0.1/script.js",
"http://100.64.0.1/script.js",
"http://[::1]/script.js",
"http://0.0.0.1/script.js",
"http://198.18.0.1/script.js",
"http://240.0.0.1/script.js",
"http://[fec0::1]/script.js",
} {
if _, err := resolver.Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: location}); err == nil || !strings.Contains(err.Error(), "prohibited address") {
t.Fatalf("Resolve(%q) error = %v, want prohibited address error", location, err)
}
}
}

func TestIsPublicScriptAddressRejectsSpecialUseNetworks(t *testing.T) {
for _, test := range []struct {
address string
public bool
}{
{address: "8.8.8.8", public: true},
{address: "100.64.0.1"},
{address: "198.18.0.1"},
{address: "240.0.0.1"},
{address: "0.0.0.1"},
{address: "fec0::1"},
{address: "::1"},
} {
if got := isPublicScriptAddress(net.ParseIP(test.address)); got != test.public {
t.Errorf("isPublicScriptAddress(%q) = %t, want %t", test.address, got, test.public)
}
}
}

func TestSafeScriptDialContextSkipsPrivateAddresses(t *testing.T) {
var dialed []string
_, err := safeScriptDialContextWithResolver(
context.Background(), "tcp", "mixed.example:443",
func(context.Context, string) ([]net.IP, error) {
return []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}, nil
},
func(_ context.Context, _, address string) (net.Conn, error) {
dialed = append(dialed, address)
return nil, errors.New("dial intentionally stopped")
},
)
if err == nil || !strings.Contains(err.Error(), "dial intentionally stopped") {
t.Fatalf("dial error = %v, want intentional dial error", err)
}
if !reflect.DeepEqual(dialed, []string{"8.8.8.8:443"}) {
t.Fatalf("dialed addresses = %v, want only public address", dialed)
}
}

func TestDefaultScriptSourceResolverReadsFilesAndFileURLs(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "scheduler script.js")
Expand All @@ -27,7 +97,7 @@ func TestDefaultScriptSourceResolverReadsFilesAndFileURLs(t *testing.T) {
if err := os.Symlink(path, link); err != nil {
t.Fatal(err)
}
resolver := NewDefaultScriptSourceResolver(nil)
resolver := newTestScriptSourceResolver()
for _, location := range []string{path, (&url.URL{Scheme: "file", Path: link}).String()} {
data, err := resolver.Resolve(context.Background(), sources.Source{Provider: sources.ProviderFile, Path: location})
if err != nil || !strings.Contains(string(data), "scheduler.interval") {
Expand Down Expand Up @@ -65,7 +135,7 @@ func TestDefaultScriptSourceResolverReadsGitFile(t *testing.T) {
t.Fatalf("git %s failed: %v\n%s", strings.Join(args, " "), err, output)
}
}
data, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{
data, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{
Provider: sources.ProviderGit,
URL: repository,
Ref: "main",
Expand Down Expand Up @@ -104,7 +174,7 @@ func TestDefaultScriptSourceResolverRejectsEscapingGitSymlink(t *testing.T) {
}
}

data, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{
data, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{
Provider: sources.ProviderGit,
URL: repository,
Ref: "main",
Expand All @@ -121,7 +191,7 @@ func TestDefaultScriptSourceResolverHTTPFailures(t *testing.T) {
w.WriteHeader(http.StatusBadGateway)
}))
defer server.Close()
_, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL + "/scheduler.js?token=super-secret"})
_, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL + "/scheduler.js?token=super-secret"})
if err == nil || !strings.Contains(err.Error(), "status 502") || strings.Contains(err.Error(), "super-secret") {
t.Fatalf("Resolve error = %v", err)
}
Expand All @@ -134,7 +204,7 @@ func TestDefaultScriptSourceResolverHTTPFailures(t *testing.T) {
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
defer cancel()
_, err := NewDefaultScriptSourceResolver(nil).Resolve(ctx, sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
_, err := newTestScriptSourceResolver().Resolve(ctx, sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
if err == nil || !strings.Contains(err.Error(), "deadline exceeded") {
t.Fatalf("Resolve timeout error = %v", err)
}
Expand All @@ -147,7 +217,7 @@ func TestDefaultScriptSourceResolverHTTPFailures(t *testing.T) {
http.Redirect(w, r, fmt.Sprintf("/%d", n+1), http.StatusFound)
}))
defer server.Close()
_, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL + "/0"})
_, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL + "/0"})
if err == nil || !strings.Contains(err.Error(), "too many redirects") {
t.Fatalf("Resolve redirects error = %v", err)
}
Expand All @@ -158,7 +228,7 @@ func TestDefaultScriptSourceResolverHTTPFailures(t *testing.T) {
http.Redirect(w, r, "file:///tmp/scheduler.js", http.StatusFound)
}))
defer server.Close()
_, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
_, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
if err == nil || !strings.Contains(err.Error(), "not supported") {
t.Fatalf("Resolve redirect error = %v", err)
}
Expand All @@ -181,7 +251,10 @@ agents:
provider: http
url: %s
`, location))
normalized, err := Normalize(spec, NormalizeOptions{ResolveScriptURLs: true})
normalized, err := Normalize(spec, NormalizeOptions{
ResolveScriptURLs: true,
ScriptSourceResolver: newTestScriptSourceResolver(),
})
if err != nil {
t.Fatalf("Normalize returned error: %v", err)
}
Expand All @@ -198,14 +271,22 @@ func TestDefaultScriptSourceResolverLimitsDecodedHTTPContent(t *testing.T) {
_ = writer.Close()
}))
defer server.Close()
_, err := NewDefaultScriptSourceResolver(nil).Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
_, err := newTestScriptSourceResolver().Resolve(context.Background(), sources.Source{Provider: sources.ProviderHTTP, URL: server.URL})
if err == nil || !strings.Contains(err.Error(), "exceeds") {
t.Fatalf("Resolve oversized content error = %v", err)
}
}

func TestDefaultScriptSourceResolverRejectsHTTPSDowngrade(t *testing.T) {
func TestDefaultScriptSourceResolverRejectsPrivateRedirect(t *testing.T) {
resolver := NewDefaultScriptSourceResolver(nil).(*defaultScriptSourceResolver)
req := httptest.NewRequest(http.MethodGet, "http://127.0.0.1/script.js", nil)
if err := resolver.client.CheckRedirect(req, nil); err == nil || !strings.Contains(err.Error(), "prohibited address") {
t.Fatalf("private redirect error = %v, want prohibited address error", err)
}
}

func TestDefaultScriptSourceResolverRejectsHTTPSDowngrade(t *testing.T) {
resolver := newTestScriptSourceResolver()
httpsRequest := httptest.NewRequest(http.MethodGet, "https://example.test/source", nil)
httpRequest := httptest.NewRequest(http.MethodGet, "http://example.test/target", nil)
err := resolver.client.CheckRedirect(httpRequest, []*http.Request{httpsRequest})
Expand Down
Loading