diff --git a/internal/remoteconfig/client.go b/internal/remoteconfig/client.go index 83c4f764b..86a2f0ac6 100644 --- a/internal/remoteconfig/client.go +++ b/internal/remoteconfig/client.go @@ -68,6 +68,21 @@ type Response struct { const maxHTTPErrorBodyBytes = 4096 +func decodeLimitedJSONResponse(resp *http.Response, maxBytes int64, destination any) error { + if err := securityutil.LimitResponseBody(resp, maxBytes); err != nil { + return err + } + + encoded, err := io.ReadAll(resp.Body) + if err != nil { + return err + } + if err := json.Unmarshal(encoded, destination); err != nil { + return err + } + return nil +} + // New creates a new remote config client. func New(cfg Config) *Client { cfg, cfgErr := normalizeConfig(cfg) @@ -168,7 +183,7 @@ func (c *Client) Fetch(ctx context.Context) (map[string]interface{}, *bool, erro } var configResp Response - if err := json.NewDecoder(io.LimitReader(resp.Body, maxConfigResponseBodyBytes)).Decode(&configResp); err != nil { + if err := decodeLimitedJSONResponse(resp, maxConfigResponseBodyBytes, &configResp); err != nil { logger.Warn(). Err(err). Str("action", "decode_response_failed"). @@ -346,7 +361,7 @@ func (c *Client) resolveAgentID(ctx context.Context) (string, error) { ID string `json:"id"` } `json:"agent"` } - if err := json.NewDecoder(io.LimitReader(resp.Body, maxAgentLookupResponseBodyBytes)).Decode(&payload); err != nil { + if err := decodeLimitedJSONResponse(resp, maxAgentLookupResponseBodyBytes, &payload); err != nil { logger.Warn(). Err(err). Str("action", "agent_lookup_decode_failed"). diff --git a/internal/remoteconfig/client_response_limit_test.go b/internal/remoteconfig/client_response_limit_test.go new file mode 100644 index 000000000..21f2d623a --- /dev/null +++ b/internal/remoteconfig/client_response_limit_test.go @@ -0,0 +1,79 @@ +package remoteconfig + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestClientFetchRejectsUndeclaredOversizedResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/agents/agent/agent-1/config" { + http.NotFound(w, r) + return + } + writeStreamedJSONWithPadding(w, + `{"success":true,"agentId":"agent-1","config":{"settings":{"interval":"1m"}}}`, + maxConfigResponseBodyBytes, + ) + })) + defer server.Close() + + client := New(Config{ + PulseURL: server.URL, + APIToken: "token", + AgentID: "agent-1", + }) + defer client.Close() + + _, _, err := client.Fetch(context.Background()) + if err == nil || !strings.Contains(err.Error(), fmt.Sprintf("response body exceeds %d bytes", maxConfigResponseBodyBytes)) { + t.Fatalf("Fetch() error = %v, want response size limit error", err) + } +} + +func TestClientResolveAgentIDRejectsUndeclaredOversizedResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != agentLookupPath { + http.NotFound(w, r) + return + } + writeStreamedJSONWithPadding(w, + `{"success":true,"agent":{"id":"agent-1"}}`, + maxAgentLookupResponseBodyBytes, + ) + })) + defer server.Close() + + client := New(Config{ + PulseURL: server.URL, + APIToken: "token", + Hostname: "node-1", + }) + defer client.Close() + + _, err := client.resolveAgentID(context.Background()) + if err == nil || !strings.Contains(err.Error(), fmt.Sprintf("response body exceeds %d bytes", maxAgentLookupResponseBodyBytes)) { + t.Fatalf("resolveAgentID() error = %v, want response size limit error", err) + } +} + +// writeStreamedJSONWithPadding flushes the valid JSON prefix before writing the +// padding so net/http cannot provide a Content-Length. This exercises the +// streaming-response boundary rather than the declared-size fast path. +func writeStreamedJSONWithPadding(w http.ResponseWriter, payload string, paddingBytes int64) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + if _, err := w.Write([]byte(payload)); err != nil { + return + } + if flusher, ok := w.(http.Flusher); ok { + flusher.Flush() + } + // The client is expected to close the response once the cap is crossed, so + // a write error here is a successful rejection rather than a server failure. + _, _ = w.Write([]byte(strings.Repeat(" ", int(paddingBytes)))) +}