From 6dc03898e77ae2f198bfb5fdfa895e1dc57a48a3 Mon Sep 17 00:00:00 2001 From: "Wen.Vale" <224270013+SEVEN-us@users.noreply.github.com> Date: Wed, 30 Sep 2026 15:45:09 +0800 Subject: [PATCH 1/2] fix(multimodal): bound remote image downloads by the file size limit --- internal/utils/httputil.go | 10 ++- internal/utils/httputil_test.go | 111 ++++++++++++++++++++++++++++++++ 2 files changed, 120 insertions(+), 1 deletion(-) create mode 100644 internal/utils/httputil_test.go diff --git a/internal/utils/httputil.go b/internal/utils/httputil.go index 9fd05cca4e0..8a8da886b93 100644 --- a/internal/utils/httputil.go +++ b/internal/utils/httputil.go @@ -15,6 +15,7 @@ var defaultHTTPClient = NewSSRFSafeHTTPClient(SSRFSafeHTTPClientConfig{ // DownloadBytes fetches the content at the given HTTP(S) URL and returns the // raw bytes. It reuses a package-level http.Client with a 60-second timeout. +// Responses exceeding the configured MAX_FILE_SIZE_MB limit are rejected. func DownloadBytes(url string) ([]byte, error) { if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") { return nil, fmt.Errorf("unsupported URL scheme: %s", url) @@ -30,9 +31,16 @@ func DownloadBytes(url string) ([]byte, error) { if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("HTTP %d for %s", resp.StatusCode, url) } - data, err := io.ReadAll(resp.Body) + maxBytes := GetMaxFileSize() + if resp.ContentLength > maxBytes { + return nil, fmt.Errorf("download size exceeds limit of %d bytes (MAX_FILE_SIZE_MB)", maxBytes) + } + data, err := io.ReadAll(io.LimitReader(resp.Body, maxBytes+1)) if err != nil { return nil, fmt.Errorf("read body: %w", err) } + if int64(len(data)) > maxBytes { + return nil, fmt.Errorf("download size exceeds limit of %d bytes (MAX_FILE_SIZE_MB)", maxBytes) + } return data, nil } diff --git a/internal/utils/httputil_test.go b/internal/utils/httputil_test.go new file mode 100644 index 00000000000..471a32ab2e4 --- /dev/null +++ b/internal/utils/httputil_test.go @@ -0,0 +1,111 @@ +package utils + +import ( + "errors" + "io" + "net/http" + "strings" + "testing" +) + +type downloadTestTransport func(*http.Request) (*http.Response, error) + +func (f downloadTestTransport) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +type downloadTestBody struct { + remaining int64 + read int64 + closed bool + err error +} + +func (b *downloadTestBody) Read(p []byte) (int, error) { + if b.err != nil { + return 0, b.err + } + if b.remaining == 0 { + return 0, io.EOF + } + n := int64(len(p)) + if n > b.remaining { + n = b.remaining + } + clear(p[:int(n)]) + b.remaining -= n + b.read += n + return int(n), nil +} + +func (b *downloadTestBody) Close() error { + b.closed = true + return nil +} + +func TestDownloadBytesSizeLimit(t *testing.T) { + const limit = int64(1024 * 1024) + t.Setenv("MAX_FILE_SIZE_MB", "1") + t.Setenv("SSRF_WHITELIST", "download.example") + t.Setenv("SSRF_WHITELIST_EXTRA", "") + ResetSSRFWhitelistForTest() + t.Cleanup(ResetSSRFWhitelistForTest) + originalClient := defaultHTTPClient + t.Cleanup(func() { defaultHTTPClient = originalClient }) + + tests := []struct { + name string + contentLength int64 + size int64 + status int + readErr error + wantErr string + wantRead int64 + }{ + {name: "small body", contentLength: 4, size: 4, wantRead: 4}, + {name: "empty body"}, + {name: "exact limit", contentLength: limit, size: limit, wantRead: limit}, + {name: "exact limit without length", contentLength: -1, size: limit, wantRead: limit}, + {name: "oversized content length", contentLength: limit + 1, size: limit + 1, wantErr: "exceeds"}, + {name: "oversized stream without length", contentLength: -1, size: 2 * limit, wantErr: "exceeds", wantRead: limit + 1}, + {name: "understated content length", contentLength: 1, size: 2 * limit, wantErr: "exceeds", wantRead: limit + 1}, + {name: "HTTP failure", status: http.StatusNotFound, contentLength: 4, size: 4, wantErr: "HTTP 404"}, + {name: "read failure", contentLength: -1, readErr: errors.New("broken stream"), wantErr: "read body: broken stream"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := &downloadTestBody{remaining: tt.size, err: tt.readErr} + status := tt.status + if status == 0 { + status = http.StatusOK + } + defaultHTTPClient = &http.Client{Transport: downloadTestTransport(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: status, + ContentLength: tt.contentLength, + Body: body, + Header: make(http.Header), + Request: req, + }, nil + })} + + data, err := DownloadBytes("https://download.example/image.png") + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("error = %v, want %q", err, tt.wantErr) + } + if data != nil { + t.Fatal("failed download returned partial data") + } + } else if err != nil || int64(len(data)) != tt.size { + t.Fatalf("download length = %d, error = %v, want %d bytes", len(data), err, tt.size) + } + if body.read != tt.wantRead { + t.Errorf("read %d bytes, want %d", body.read, tt.wantRead) + } + if !body.closed { + t.Error("response body was not closed") + } + }) + } +} From d8101a6867e98b27026133c925e8fd6444717264 Mon Sep 17 00:00:00 2001 From: wizardchen Date: Thu, 8 Oct 2026 00:18:23 +0800 Subject: [PATCH 2/2] style(utils): wrap over-long lines in the DownloadBytes test --- internal/utils/httputil_test.go | 20 +++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/internal/utils/httputil_test.go b/internal/utils/httputil_test.go index 471a32ab2e4..8ff4905f04e 100644 --- a/internal/utils/httputil_test.go +++ b/internal/utils/httputil_test.go @@ -67,10 +67,19 @@ func TestDownloadBytesSizeLimit(t *testing.T) { {name: "exact limit", contentLength: limit, size: limit, wantRead: limit}, {name: "exact limit without length", contentLength: -1, size: limit, wantRead: limit}, {name: "oversized content length", contentLength: limit + 1, size: limit + 1, wantErr: "exceeds"}, - {name: "oversized stream without length", contentLength: -1, size: 2 * limit, wantErr: "exceeds", wantRead: limit + 1}, - {name: "understated content length", contentLength: 1, size: 2 * limit, wantErr: "exceeds", wantRead: limit + 1}, + { + name: "oversized stream without length", contentLength: -1, size: 2 * limit, + wantErr: "exceeds", wantRead: limit + 1, + }, + { + name: "understated content length", contentLength: 1, size: 2 * limit, + wantErr: "exceeds", wantRead: limit + 1, + }, {name: "HTTP failure", status: http.StatusNotFound, contentLength: 4, size: 4, wantErr: "HTTP 404"}, - {name: "read failure", contentLength: -1, readErr: errors.New("broken stream"), wantErr: "read body: broken stream"}, + { + name: "read failure", contentLength: -1, readErr: errors.New("broken stream"), + wantErr: "read body: broken stream", + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -79,7 +88,7 @@ func TestDownloadBytesSizeLimit(t *testing.T) { if status == 0 { status = http.StatusOK } - defaultHTTPClient = &http.Client{Transport: downloadTestTransport(func(req *http.Request) (*http.Response, error) { + transport := downloadTestTransport(func(req *http.Request) (*http.Response, error) { return &http.Response{ StatusCode: status, ContentLength: tt.contentLength, @@ -87,7 +96,8 @@ func TestDownloadBytesSizeLimit(t *testing.T) { Header: make(http.Header), Request: req, }, nil - })} + }) + defaultHTTPClient = &http.Client{Transport: transport} data, err := DownloadBytes("https://download.example/image.png") if tt.wantErr != "" {