From 2d6bc194426d45e55eb299736f3a76172dae3974 Mon Sep 17 00:00:00 2001 From: dawn Date: Sat, 22 Aug 2026 00:42:29 +0900 Subject: [PATCH] xrpc: retain unexpected error response bodies Signed-off-by: dawn --- xrpc/xrpc.go | 44 +++++++++++++++++++++++------ xrpc/xrpc_test.go | 71 +++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 106 insertions(+), 9 deletions(-) diff --git a/xrpc/xrpc.go b/xrpc/xrpc.go index 30c4b86e..2a661fb9 100644 --- a/xrpc/xrpc.go +++ b/xrpc/xrpc.go @@ -40,6 +40,8 @@ var ( Procedure = http.MethodPost ) +const maxErrorResponseBody = 4 << 10 + type AuthInfo struct { AccessJwt string `json:"accessJwt"` RefreshJwt string `json:"refreshJwt"` @@ -57,21 +59,33 @@ func (xe *XRPCError) Error() string { } type Error struct { - StatusCode int - Wrapped error - Ratelimit *RatelimitInfo + StatusCode int + Wrapped error + Ratelimit *RatelimitInfo + ResponseBody string + ResponseBodyTruncated bool } func (e *Error) Error() string { // Preserving "XRPC ERROR %d" prefix for compatibility - previously matching this string was the only way // to obtain the status code. - if e.Wrapped == nil { + detail := "" + if e.ResponseBody != "" { + suffix := "" + if e.ResponseBodyTruncated { + suffix = " (truncated)" + } + detail = fmt.Sprintf("unexpected response body %q%s", e.ResponseBody, suffix) + } else if e.Wrapped != nil { + detail = e.Wrapped.Error() + } + if detail == "" { return fmt.Sprintf("XRPC ERROR %d", e.StatusCode) } if e.StatusCode == http.StatusTooManyRequests && e.Ratelimit != nil { - return fmt.Sprintf("XRPC ERROR %d: %s (throttled until %s)", e.StatusCode, e.Wrapped, e.Ratelimit.Reset.Local()) + return fmt.Sprintf("XRPC ERROR %d: %s (throttled until %s)", e.StatusCode, detail, e.Ratelimit.Reset.Local()) } - return fmt.Sprintf("XRPC ERROR %d: %s", e.StatusCode, e.Wrapped) + return fmt.Sprintf("XRPC ERROR %d: %s", e.StatusCode, detail) } func (e *Error) Unwrap() error { @@ -85,7 +99,7 @@ func (e *Error) IsThrottled() bool { return e.StatusCode == http.StatusTooManyRequests } -func errorFromHTTPResponse(resp *http.Response, err error) error { +func errorFromHTTPResponse(resp *http.Response, err error) *Error { r := &Error{ StatusCode: resp.StatusCode, Wrapped: err, @@ -197,9 +211,21 @@ func (c *Client) Do(ctx context.Context, kind string, inpenc string, method stri defer resp.Body.Close() if resp.StatusCode != 200 { + body, err := io.ReadAll(io.LimitReader(resp.Body, maxErrorResponseBody+1)) + if err != nil { + return errorFromHTTPResponse(resp, fmt.Errorf("failed to read xrpc error response: %w", err)) + } + truncated := len(body) > maxErrorResponseBody + if truncated { + body = body[:maxErrorResponseBody] + } + var xe XRPCError - if err := json.NewDecoder(resp.Body).Decode(&xe); err != nil { - return errorFromHTTPResponse(resp, fmt.Errorf("failed to decode xrpc error message: %w", err)) + if err := json.NewDecoder(bytes.NewReader(body)).Decode(&xe); err != nil { + xrpcErr := errorFromHTTPResponse(resp, fmt.Errorf("failed to decode xrpc error message: %w", err)) + xrpcErr.ResponseBody = strings.TrimSpace(strings.ToValidUTF8(string(body), "�")) + xrpcErr.ResponseBodyTruncated = truncated + return xrpcErr } return errorFromHTTPResponse(resp, &xe) } diff --git a/xrpc/xrpc_test.go b/xrpc/xrpc_test.go index 97dc1b2c..39d45b45 100644 --- a/xrpc/xrpc_test.go +++ b/xrpc/xrpc_test.go @@ -1,6 +1,11 @@ package xrpc import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "strings" "testing" ) @@ -57,3 +62,69 @@ func TestMakeParams(t *testing.T) { }) } } + +func TestErrorResponseBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotFound) + _, _ = w.Write([]byte("404 page not found\n")) + })) + defer server.Close() + + client := &Client{Host: server.URL, Client: server.Client()} + err := client.Do(context.Background(), Query, "", "test.method", nil, nil, nil) + if err == nil { + t.Fatal("expected request to fail") + } + + var xrpcErr *Error + if !errors.As(err, &xrpcErr) { + t.Fatalf("error type = %T, want *Error", err) + } + if xrpcErr.ResponseBody != "404 page not found" { + t.Fatalf("response body = %q, want %q", xrpcErr.ResponseBody, "404 page not found") + } + if got, want := err.Error(), `XRPC ERROR 404: unexpected response body "404 page not found"`; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestErrorResponseBodyIsBounded(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(strings.Repeat("x", maxErrorResponseBody+1))) + })) + defer server.Close() + + client := &Client{Host: server.URL, Client: server.Client()} + err := client.Do(context.Background(), Query, "", "test.method", nil, nil, nil) + + var xrpcErr *Error + if !errors.As(err, &xrpcErr) { + t.Fatalf("error type = %T, want *Error", err) + } + if len(xrpcErr.ResponseBody) != maxErrorResponseBody { + t.Fatalf("response body length = %d, want %d", len(xrpcErr.ResponseBody), maxErrorResponseBody) + } + if !xrpcErr.ResponseBodyTruncated { + t.Fatal("expected response body truncation") + } +} + +func TestStructuredXRPCErrorRemainsTyped(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(`{"error":"Forbidden","message":"key cannot push"}`)) + })) + defer server.Close() + + client := &Client{Host: server.URL, Client: server.Client()} + err := client.Do(context.Background(), Query, "", "test.method", nil, nil, nil) + + var named *XRPCError + if !errors.As(err, &named) { + t.Fatalf("error type = %T, want wrapped *XRPCError", err) + } + if named.ErrStr != "Forbidden" || named.Message != "key cannot push" { + t.Fatalf("XRPC error = %#v", named) + } +} -- 2.51.2