diff --git a/internal/securityutil/responsebody.go b/internal/securityutil/responsebody.go new file mode 100644 index 000000000..d837fc964 --- /dev/null +++ b/internal/securityutil/responsebody.go @@ -0,0 +1,68 @@ +package securityutil + +import ( + "fmt" + "io" + "net/http" +) + +// LimitResponseBody bounds the bytes a caller can read from an HTTP response. +// It closes responses whose declared size already exceeds the limit. Responses +// without a trustworthy Content-Length remain bounded while they are read. +func LimitResponseBody(resp *http.Response, limit int64) error { + if resp == nil || resp.Body == nil { + return fmt.Errorf("response body is required") + } + if limit < 0 { + return fmt.Errorf("response body limit must not be negative") + } + if resp.ContentLength > limit { + _ = resp.Body.Close() + return fmt.Errorf("response body exceeds %d bytes", limit) + } + + resp.Body = &limitedResponseBody{ + body: resp.Body, + remaining: limit, + limit: limit, + } + return nil +} + +type limitedResponseBody struct { + body io.ReadCloser + remaining int64 + limit int64 + exceeded bool +} + +func (r *limitedResponseBody) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if r.exceeded { + return 0, fmt.Errorf("response body exceeds %d bytes", r.limit) + } + if r.remaining > 0 { + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + n, err := r.body.Read(p) + r.remaining -= int64(n) + return n, err + } + + // Probe for one additional byte. This distinguishes a body exactly at the + // limit from an oversized body without exposing bytes beyond the boundary. + var probe [1]byte + n, err := r.body.Read(probe[:]) + if n > 0 { + r.exceeded = true + return 0, fmt.Errorf("response body exceeds %d bytes", r.limit) + } + return 0, err +} + +func (r *limitedResponseBody) Close() error { + return r.body.Close() +} diff --git a/internal/securityutil/responsebody_test.go b/internal/securityutil/responsebody_test.go new file mode 100644 index 000000000..75da612a7 --- /dev/null +++ b/internal/securityutil/responsebody_test.go @@ -0,0 +1,96 @@ +package securityutil + +import ( + "io" + "net/http" + "strings" + "testing" +) + +type trackingReadCloser struct { + io.Reader + closed bool +} + +func (r *trackingReadCloser) Close() error { + r.closed = true + return nil +} + +func TestLimitResponseBody(t *testing.T) { + t.Run("rejects declared oversize and closes body", func(t *testing.T) { + body := &trackingReadCloser{Reader: strings.NewReader("oversized")} + resp := &http.Response{Body: body, ContentLength: 9} + + err := LimitResponseBody(resp, 8) + if err == nil || !strings.Contains(err.Error(), "response body exceeds 8 bytes") { + t.Fatalf("LimitResponseBody() error = %v", err) + } + if !body.closed { + t.Fatal("oversized response body was not closed") + } + }) + + t.Run("allows body exactly at limit", func(t *testing.T) { + resp := &http.Response{ + Body: io.NopCloser(strings.NewReader("12345678")), + ContentLength: -1, + } + if err := LimitResponseBody(resp, 8); err != nil { + t.Fatalf("LimitResponseBody() error = %v", err) + } + + got, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if string(got) != "12345678" { + t.Fatalf("ReadAll() = %q", got) + } + }) + + t.Run("rejects streamed oversize at boundary", func(t *testing.T) { + resp := &http.Response{ + Body: io.NopCloser(strings.NewReader("123456789")), + ContentLength: -1, + } + if err := LimitResponseBody(resp, 8); err != nil { + t.Fatalf("LimitResponseBody() error = %v", err) + } + + got, err := io.ReadAll(resp.Body) + if err == nil || !strings.Contains(err.Error(), "response body exceeds 8 bytes") { + t.Fatalf("ReadAll() error = %v", err) + } + if string(got) != "12345678" { + t.Fatalf("ReadAll() returned bytes beyond limit: %q", got) + } + }) + + t.Run("preserves close", func(t *testing.T) { + body := &trackingReadCloser{Reader: strings.NewReader("ok")} + resp := &http.Response{Body: body, ContentLength: 2} + if err := LimitResponseBody(resp, 8); err != nil { + t.Fatalf("LimitResponseBody() error = %v", err) + } + if err := resp.Body.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + if !body.closed { + t.Fatal("underlying response body was not closed") + } + }) + + t.Run("validates arguments", func(t *testing.T) { + if err := LimitResponseBody(nil, 8); err == nil { + t.Fatal("expected nil response error") + } + resp := &http.Response{Body: io.NopCloser(strings.NewReader(""))} + if err := LimitResponseBody(resp, -1); err == nil { + t.Fatal("expected negative limit error") + } + if err := LimitResponseBody(&http.Response{}, 8); err == nil { + t.Fatal("expected missing body error") + } + }) +} diff --git a/pkg/pbs/client.go b/pkg/pbs/client.go index 94e4d2197..6e1844465 100644 --- a/pkg/pbs/client.go +++ b/pkg/pbs/client.go @@ -272,6 +272,9 @@ func (c *Client) handleAuthResponse(resp *http.Response) error { } return &authHTTPError{status: resp.StatusCode, body: string(body)} } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return err + } var result struct { Data struct { @@ -423,6 +426,9 @@ func (c *Client) request(ctx context.Context, method, path string, data url.Valu return nil, &apiHTTPError{status: resp.StatusCode, body: string(body)} } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return nil, err + } return resp, nil } diff --git a/pkg/pbs/client_security_test.go b/pkg/pbs/client_security_test.go index 3b684b5d9..03e2fd66f 100644 --- a/pkg/pbs/client_security_test.go +++ b/pkg/pbs/client_security_test.go @@ -34,3 +34,32 @@ func TestGetVersionRejectsOversizedErrorBody(t *testing.T) { t.Fatalf("expected size-limit error, got: %v", err) } } + +func TestGetVersionRejectsOversizedSuccessBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api2/json/version" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + _, _ = w.Write([]byte(`{"data":{"version":"3.4"},"padding":"`)) + _, _ = w.Write([]byte(strings.Repeat("x", int(maxResponseBodyBytes)))) + _, _ = w.Write([]byte(`"}`)) + })) + defer server.Close() + + client, err := NewClient(ClientConfig{ + Host: server.URL, + TokenName: "root@pbs!pulse-token", + TokenValue: "secret", + }) + if err != nil { + t.Fatalf("NewClient() error = %v", err) + } + + _, err = client.GetVersion(context.Background()) + if err == nil { + t.Fatal("expected oversized body error, got nil") + } + if !strings.Contains(err.Error(), "response body exceeds") { + t.Fatalf("expected size-limit error, got: %v", err) + } +} diff --git a/pkg/pmg/client.go b/pkg/pmg/client.go index 37ab0f19b..1d69902ac 100644 --- a/pkg/pmg/client.go +++ b/pkg/pmg/client.go @@ -331,6 +331,9 @@ func (c *Client) handleAuthResponse(resp *http.Response) error { } return &authHTTPError{status: resp.StatusCode, body: string(body)} } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return err + } var result struct { Data struct { @@ -459,6 +462,9 @@ func (c *Client) request(ctx context.Context, method, path string, params url.Va return nil, apiErr } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return nil, err + } return resp, nil } diff --git a/pkg/pmg/client_security_test.go b/pkg/pmg/client_security_test.go index 59ccadee8..ab66aaec2 100644 --- a/pkg/pmg/client_security_test.go +++ b/pkg/pmg/client_security_test.go @@ -34,3 +34,32 @@ func TestGetVersionRejectsOversizedErrorBody(t *testing.T) { t.Fatalf("expected size-limit error, got: %v", err) } } + +func TestGetVersionRejectsOversizedSuccessBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api2/json/version" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + _, _ = w.Write([]byte(`{"data":{"version":"8.2"},"padding":"`)) + _, _ = w.Write([]byte(strings.Repeat("x", int(maxResponseBodyBytes)))) + _, _ = w.Write([]byte(`"}`)) + })) + defer server.Close() + + client, err := NewClient(ClientConfig{ + Host: server.URL, + TokenName: "root@pmg!pulse-token", + TokenValue: "secret", + }) + if err != nil { + t.Fatalf("NewClient() error = %v", err) + } + + _, err = client.GetVersion(context.Background()) + if err == nil { + t.Fatal("expected oversized body error, got nil") + } + if !strings.Contains(err.Error(), "response body exceeds") { + t.Fatalf("expected size-limit error, got: %v", err) + } +} diff --git a/pkg/proxmox/client.go b/pkg/proxmox/client.go index aa0443d7b..34367a0ba 100644 --- a/pkg/proxmox/client.go +++ b/pkg/proxmox/client.go @@ -466,6 +466,9 @@ func (c *Client) handleAuthResponse(resp *http.Response) error { } return &authHTTPError{status: resp.StatusCode, body: string(body)} } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return err + } var result struct { Data struct { @@ -629,6 +632,9 @@ func (c *Client) requestWithRetry(ctx context.Context, method, path string, data return nil, apiErr } + if err := securityutil.LimitResponseBody(resp, maxResponseBodyBytes); err != nil { + return nil, err + } return resp, nil } diff --git a/pkg/proxmox/client_security_test.go b/pkg/proxmox/client_security_test.go index 946c82b79..2f07eddec 100644 --- a/pkg/proxmox/client_security_test.go +++ b/pkg/proxmox/client_security_test.go @@ -34,3 +34,32 @@ func TestGetNodesRejectsOversizedErrorBody(t *testing.T) { t.Fatalf("expected size-limit error, got: %v", err) } } + +func TestGetNodesRejectsOversizedSuccessBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api2/json/nodes" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + _, _ = w.Write([]byte(`{"data":[],"padding":"`)) + _, _ = w.Write([]byte(strings.Repeat("x", int(maxResponseBodyBytes)))) + _, _ = w.Write([]byte(`"}`)) + })) + defer server.Close() + + client, err := NewClient(ClientConfig{ + Host: server.URL, + TokenName: "root@pam!pulse-token", + TokenValue: "secret", + }) + if err != nil { + t.Fatalf("NewClient() error = %v", err) + } + + _, err = client.GetNodes(context.Background()) + if err == nil { + t.Fatal("expected oversized body error, got nil") + } + if !strings.Contains(err.Error(), "response body exceeds") { + t.Fatalf("expected size-limit error, got: %v", err) + } +}