Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
75 changes: 70 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,21 @@ 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()
transport.Proxy = nil
Comment thread
monkeyscan[bot] marked this conversation as resolved.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

移除环境代理后,依赖正向代理出网的环境将无法拉取 HTTP(S) 脚本源且无明确诊断或恢复开关

改动将 transport.Proxy 置为 nil 并删除 scriptSourceTransport 封装。对可直连公网的环境行为不变;但在必须经正向代理出网的环境(进程设置了 HTTP_PROXY/HTTPS_PROXY,且防火墙阻断直连)中,HTTP(S) 脚本源拉取会从原来的“经代理成功”变成“直连超时/连接被拒”。基线代码的注释与专测(TestSafeScriptDialContextAllowsConfiguredProxyEndpoint)表明这一场景此前是被有意支持的。现在既没有配置开关/白名单可恢复,用户侧也只看到通用的 dial 错误(如 i/o timeout),无法得知是脚本源有意禁用了代理,排障成本高。

Problem code:

Changed code at pkg/compose/script_source.go:52

Recommendation:
建议在初始化或 Resolve/readHTTP 出错时检测进程环境中的 HTTP_PROXY/HTTPS_PROXY,若存在则附加“脚本源出于 SSRF 防护有意不使用环境代理”的明确提示;如确有受信代理需求,评估增加默认关闭的显式受信代理配置项,并在文档中标注对应 SSRF 风险。

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 +67,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 +240,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 +259,62 @@ 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) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, fmt.Errorf("parse script endpoint: %w", err)
}
addresses, err := net.DefaultResolver.LookupIPAddr(ctx, host)
if err != nil {
return nil, fmt.Errorf("resolve script endpoint: %w", err)
}
dialer := net.Dialer{}
var dialErrs error
for _, candidate := range addresses {
if !isPublicScriptAddress(candidate.IP) {
continue
}
conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(candidate.IP.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() {
return false
}
// Go's IsPrivate intentionally excludes the carrier-grade NAT range.
cgnat := &net.IPNet{IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}
return !cgnat.Contains(ip)
}
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
54 changes: 44 additions & 10 deletions pkg/compose/script_source_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"compress/gzip"
"context"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/url"
Expand All @@ -17,6 +18,28 @@ import (
"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",
} {
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 TestDefaultScriptSourceResolverReadsFilesAndFileURLs(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "scheduler script.js")
Expand All @@ -27,7 +50,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 +88,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 +127,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 +144,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 +157,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 +170,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 +181,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 +204,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 +224,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