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
2 changes: 1 addition & 1 deletion cmd/octobus/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -125,7 +125,7 @@ func serve(opts serveOptions) error {
if err := startupInventory(ctx, logger, st); err != nil {
return err
}
adminServer := &admin.Server{Store: st, Importer: &packageimport.Importer{DataDir: dataDir, Store: st}, Supervisor: sup, Gateway: gateway, AccessLogPath: filepath.Join(dataDir, accesslog.FileName), Logger: logger}
adminServer := &admin.Server{Store: st, Importer: &packageimport.Importer{DataDir: dataDir, Store: st, RemoteTargetValidator: packageimport.DefaultRemoteTargetValidator}, Supervisor: sup, Gateway: gateway, AccessLogPath: filepath.Join(dataDir, accesslog.FileName), Logger: logger}
grpcServer := protocol.GRPCServer(gateway)
publicServer := admin.NewHTTPServer(opts.addr, h2c.NewHandler(server.CombinedHandler(adminServer.Handler(), grpcServer, gateway), &http2.Server{}))
publicListener, err := net.Listen("tcp", opts.addr)
Expand Down
8 changes: 7 additions & 1 deletion internal/packageimport/git_source.go
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,11 @@ func (i *Importer) prepareGitSource(ctx context.Context, rawSource, staging stri
if err := os.MkdirAll(repoDir, 0o755); err != nil {
return preparedSource{}, err
}
if i.RemoteTargetValidator != nil {
Comment thread
monkeyscan[bot] marked this conversation as resolved.
if err := i.RemoteTargetValidator(ctx, src.CredentialURL); err != nil {
return preparedSource{}, fmt.Errorf("validate Git remote: %w", err)
}
}
if err := runner.run(ctx, repoDir, "init", "--bare", "."); err != nil {
return preparedSource{}, err
}
Expand Down Expand Up @@ -318,7 +323,8 @@ func (r *gitRunner) run(ctx context.Context, dir string, args ...string) error {
}

func (r *gitRunner) output(ctx context.Context, dir string, args ...string) (string, error) {
cmd := exec.CommandContext(ctx, "git", args...)
gitArgs := append([]string{"-c", "http.followRedirects=false"}, args...)
cmd := exec.CommandContext(ctx, "git", gitArgs...)
cmd.Dir = dir
cmd.Env = r.env
var out strings.Builder
Expand Down
40 changes: 8 additions & 32 deletions internal/packageimport/importer.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@ import (
"fmt"
"io"
"io/fs"
"net/http"
"net/url"
"os"
"os/exec"
Expand All @@ -27,6 +26,11 @@ import (
type Importer struct {
DataDir string
Store *store.Store

// RemoteTargetValidator is configured by the daemon to enforce the
// network policy for server-side remote imports. Tests and local-only
// importers may leave it unset.
RemoteTargetValidator func(context.Context, string) error
}

type Options struct {
Expand Down Expand Up @@ -608,7 +612,7 @@ func (i *Importer) prepareSource(ctx context.Context, opts Options, staging stri
prepared.ServiceRoot = serviceRoot
return prepared, nil
case sourceRemoteArchive:
return prepareRemoteArchiveSource(ctx, source, serviceRoot, staging)
return i.prepareRemoteArchiveSource(ctx, source, serviceRoot, staging)
case sourceHTTPSGit:
return i.prepareGitSource(ctx, opts.Source, staging)
case sourceUnsupportedGit:
Expand Down Expand Up @@ -771,13 +775,13 @@ func hashFile(path string) (string, error) {
return domain.HashBytes(b), nil
}

func prepareRemoteArchiveSource(ctx context.Context, source, serviceRoot, staging string) (preparedSource, error) {
func (i *Importer) prepareRemoteArchiveSource(ctx context.Context, source, serviceRoot, staging string) (preparedSource, error) {
artifactName, err := remoteArchiveArtifactName(source)
if err != nil {
return preparedSource{}, err
}
artifactPath := filepath.Join(staging, artifactName)
if err := downloadRemoteArchive(ctx, source, artifactPath); err != nil {
if err := downloadRemoteArchive(ctx, source, artifactPath, i.RemoteTargetValidator); err != nil {
return preparedSource{}, err
}
packageDir := filepath.Join(staging, "package")
Expand Down Expand Up @@ -819,34 +823,6 @@ func remoteArchiveArtifactName(source string) (string, error) {
return "", fmt.Errorf("unsupported remote package source %q: must end with .tgz, .tar.gz, or .zip", redactedRemoteArchiveSource(source))
}

func downloadRemoteArchive(ctx context.Context, source, artifactPath string) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, source, nil)
if err != nil {
return fmt.Errorf("download remote package %q: %w", redactedRemoteArchiveSource(source), err)
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("download remote package %q: %w", redactedRemoteArchiveSource(source), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("download remote package %q: HTTP %d", redactedRemoteArchiveSource(source), resp.StatusCode)
}
if err := os.MkdirAll(filepath.Dir(artifactPath), 0o755); err != nil {
return err
}
out, err := os.OpenFile(artifactPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o644)
if err != nil {
return err
}
_, copyErr := io.Copy(out, resp.Body)
closeErr := out.Close()
if copyErr != nil {
return copyErr
}
return closeErr
}

func redactedRemoteArchiveSource(source string) string {
u, err := url.Parse(source)
if err != nil {
Expand Down
179 changes: 179 additions & 0 deletions internal/packageimport/remote_security.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,179 @@
package packageimport

import (
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"time"
)

const remoteImportTimeout = 10 * time.Minute

var forbiddenRemoteNetworks = mustParseRemoteNetworks(
"0.0.0.0/8",
"10.0.0.0/8",
"100.64.0.0/10",
"127.0.0.0/8",
"169.254.0.0/16",
"172.16.0.0/12",
"192.0.0.0/24",
"192.0.2.0/24",
"192.168.0.0/16",
"198.18.0.0/15",
"198.51.100.0/24",
"203.0.113.0/24",
"224.0.0.0/4",
"240.0.0.0/4",
"::/128",
"::1/128",
"fc00::/7",
"fe80::/10",
"ff00::/8",
"2001:db8::/32",
)

func mustParseRemoteNetworks(cidrs ...string) []*net.IPNet {
networks := make([]*net.IPNet, 0, len(cidrs))
for _, cidr := range cidrs {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
panic(err)
}
networks = append(networks, network)
}
return networks
}

func DefaultRemoteTargetValidator(ctx context.Context, raw string) error {
u, err := url.Parse(raw)
if err != nil {
return fmt.Errorf("invalid remote URL: %w", err)
}
if u.Scheme != "http" && u.Scheme != "https" {
return errors.New("only HTTP(S) remote targets are allowed")
}
if u.Hostname() == "" {
return errors.New("remote URL host is required")
}
return validateRemoteHost(ctx, u.Hostname())
}

func validateRemoteHost(ctx context.Context, hostname string) error {
if ip := net.ParseIP(hostname); ip != nil {
if isForbiddenRemoteIP(ip) {
return fmt.Errorf("remote target %q resolves to a private or special address", hostname)
}
return nil
}

ips, err := net.DefaultResolver.LookupIPAddr(ctx, hostname)
if err != nil {
return fmt.Errorf("resolve remote target %q: %w", hostname, err)
}
if len(ips) == 0 {
return fmt.Errorf("remote target %q has no address", hostname)
}
for _, item := range ips {
if isForbiddenRemoteIP(item.IP) {
return fmt.Errorf("remote target %q resolves to a private or special address", hostname)
}
}
return nil
}

func isForbiddenRemoteIP(ip net.IP) bool {
if ip == nil || !ip.IsGlobalUnicast() {
return true
}
for _, network := range forbiddenRemoteNetworks {
if network.Contains(ip) {
return true
}
}
return false
}

func newRemoteHTTPClient(validate func(context.Context, string) error) *http.Client {
transport := http.DefaultTransport.(*http.Transport).Clone()
client := &http.Client{
Transport: transport,
Timeout: remoteImportTimeout,
}
if validate == nil {
return client
}
transport.DialContext = safeRemoteDialContext(validate)
client.CheckRedirect = func(req *http.Request, _ []*http.Request) error {
return validate(req.Context(), req.URL.String())
}
return client
}

func safeRemoteDialContext(validate func(context.Context, string) error) func(context.Context, string, string) (net.Conn, error) {
dialer := &net.Dialer{}
return func(ctx context.Context, network, address string) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, err
}
if err := validate(ctx, "https://"+net.JoinHostPort(host, port)); err != nil {
return nil, err
}

ips, err := net.DefaultResolver.LookupIPAddr(ctx, host)
if err != nil {
return nil, err
}
for _, item := range ips {
if isForbiddenRemoteIP(item.IP) {
continue
}
conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(item.IP.String(), port))
if err == nil {
return conn, nil
}
}
return nil, fmt.Errorf("unable to connect to allowed address for %s", host)
}
}

func downloadRemoteArchive(ctx context.Context, source, artifactPath string, validate func(context.Context, string) error) error {
if validate != nil {
if err := validate(ctx, source); err != nil {
return fmt.Errorf("validate remote package target: %w", err)
}
}
ctx, cancel := context.WithTimeout(ctx, remoteImportTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, source, nil)
if err != nil {
return fmt.Errorf("download remote package %q: %w", redactedRemoteArchiveSource(source), err)
}
resp, err := newRemoteHTTPClient(validate).Do(req)
if err != nil {
return fmt.Errorf("download remote package %q: %w", redactedRemoteArchiveSource(source), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("download remote package %q: HTTP %d", redactedRemoteArchiveSource(source), resp.StatusCode)
}
if err := os.MkdirAll(filepath.Dir(artifactPath), 0o755); err != nil {
return err
}
out, err := os.OpenFile(artifactPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o644)
if err != nil {
return err
}
_, copyErr := io.Copy(out, resp.Body)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

远程归档下载无响应体大小上限,存在磁盘耗尽风险

downloadRemoteArchive 使用 io.Copy(out, resp.Body) 将响应体无任何大小上限地写入磁盘(remote_security.go 第 173 行)。服务端导入的 source 是用户可控的,而校验器只限制目标地址,不限制响应体大小。任意公网 URL 可以在 remoteImportTimeout(10 分钟)内持续返回数据,导致磁盘被写满;这是把远程导入作为服务端功能的可用性/稳定性风险。该函数从 importer.go 迁移到新模块时原样保留了无上限拷贝,本次 SSRF 加固并未覆盖这一点。

Problem code:

Changed code at internal/packageimport/remote_security.go:173

Recommendation:
对下载响应体设置大小上限:使用 io.LimitReader 包裹 resp.Body(例如基于 Content-Length 或固定上限,超出即中止并清理已写文件),并在后续解压/展开阶段同样限制总大小与条目数量,防止磁盘耗尽。

Suggested diff:

if _, copyErr := io.Copy(out, io.LimitReader(resp.Body, maxRemoteArchiveBytes)); copyErr != nil {
    out.Close()
    os.Remove(artifactPath)
    return fmt.Errorf("download remote package %q: %w", redactedRemoteArchiveSource(source), copyErr)
}

closeErr := out.Close()
if copyErr != nil {
return copyErr
}
return closeErr
}
59 changes: 59 additions & 0 deletions internal/packageimport/remote_security_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
package packageimport

import (
"context"
"net/http"
"net/url"
"strings"
"testing"
)

func TestDefaultRemoteTargetValidatorRejectsPrivateAndSpecialAddresses(t *testing.T) {
for _, raw := range []string{
"http://127.0.0.1/package.tgz",
"https://localhost/package.tgz",
"http://169.254.169.254/package.tgz",
"http://192.0.2.8/package.tgz",
"http://[::1]/package.tgz",
} {
err := DefaultRemoteTargetValidator(context.Background(), raw)
if err == nil {
t.Fatalf("validator accepted %s", raw)
}
if !strings.Contains(err.Error(), "private or special") {
t.Fatalf("unexpected error for %s: %v", raw, err)
}
}
}

func TestDefaultRemoteTargetValidatorRejectsUnsupportedURLs(t *testing.T) {
if err := DefaultRemoteTargetValidator(context.Background(), "file:///tmp/package.tgz"); err == nil {
t.Fatal("validator accepted a non-HTTP URL")
}
}

func TestRemoteHTTPClientRevalidatesRedirects(t *testing.T) {
validator := func(_ context.Context, raw string) error {
if strings.HasSuffix(raw, "/internal") {
return context.Canceled
}
return nil
}
client := newRemoteHTTPClient(validator)
redirect, err := url.Parse("https://public.example/internal")
if err != nil {
t.Fatal(err)
}
err = client.CheckRedirect(&http.Request{URL: redirect, Method: http.MethodGet}, nil)
if err == nil || !strings.Contains(err.Error(), context.Canceled.Error()) {
t.Fatalf("redirect validation error = %v", err)
}
}

func TestPrepareGitSourceRejectsPrivateRemoteBeforeGitFetch(t *testing.T) {
imp := &Importer{RemoteTargetValidator: DefaultRemoteTargetValidator}
_, err := imp.prepareGitSource(context.Background(), "https://127.0.0.1/repo.git", t.TempDir())
if err == nil || !strings.Contains(err.Error(), "private or special") {
t.Fatalf("private Git remote error = %v", err)
}
}
Loading