Bound provider error response body reads

This commit is contained in:
2026-08-23 18:56:15 +00:00
parent af0bd3f31a
commit fa2384e696
3 changed files with 135 additions and 0 deletions

View File

@@ -81,6 +81,22 @@ func newProviderHTTPError(statusCode int, details providerErrorDetails) *Provide
}
}
func providerHTTPErrorFromBody(statusCode int, contentLength int64, body io.Reader) *ProviderHTTPError {
if contentLength > maxProviderErrorResponseBytes {
return newProviderHTTPError(statusCode, providerErrorDetails{})
}
limited := &io.LimitedReader{
R: body,
N: maxProviderErrorResponseBytes + 1,
}
contents, err := io.ReadAll(limited)
if err != nil || limited.N == 0 {
return newProviderHTTPError(statusCode, providerErrorDetails{})
}
return newProviderHTTPError(statusCode, parseProviderErrorEnvelope(contents))
}
func parseProviderErrorEnvelope(body []byte) providerErrorDetails {
decoder := json.NewDecoder(strings.NewReader(string(body)))
decoder.UseNumber()

View File

@@ -3,12 +3,39 @@ package llm
import (
"errors"
"fmt"
"io"
"reflect"
"strings"
"testing"
"unicode/utf8"
)
type guardedReader struct {
reader io.Reader
remaining int64
bytes int64
violated bool
}
func (r *guardedReader) Read(buffer []byte) (int, error) {
if int64(len(buffer)) > r.remaining {
r.violated = true
return 0, errors.New("reader was read past its allowed boundary")
}
n, err := r.reader.Read(buffer)
r.bytes += int64(n)
r.remaining -= int64(n)
return n, err
}
type failingReader struct {
err error
}
func (r failingReader) Read([]byte) (int, error) {
return 0, r.err
}
func TestProviderHTTPErrorEnvelopeParsing(t *testing.T) {
tests := []struct {
name string
@@ -136,3 +163,93 @@ func TestProviderHTTPErrorIdentityAndFormatting(t *testing.T) {
t.Fatalf("zero error behavior is not safe: %v", zero)
}
}
func TestProviderHTTPErrorBodyBounds(t *testing.T) {
const (
statusCode = 502
marker = "provider-body-marker"
)
ordinaryBody := `{"error":{"message":"` + marker + `"}}`
exactLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)-len(ordinaryBody))
overLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(ordinaryBody))
tests := []struct {
name string
contentLength int64
reader io.Reader
wantRead int64
wantMessage string
}{
{
name: "recognized envelope",
contentLength: int64(len(ordinaryBody)),
reader: strings.NewReader(ordinaryBody),
wantRead: int64(len(ordinaryBody)),
wantMessage: marker,
},
{
name: "exact limit",
contentLength: maxProviderErrorResponseBytes,
reader: strings.NewReader(exactLimitBody),
wantRead: maxProviderErrorResponseBytes,
wantMessage: marker,
},
{
name: "declared oversize does not read",
contentLength: maxProviderErrorResponseBytes + 1,
reader: strings.NewReader(ordinaryBody),
wantRead: 0,
},
{
name: "unknown length oversize",
contentLength: -1,
reader: strings.NewReader(overLimitBody),
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "underreported oversize",
contentLength: maxProviderErrorResponseBytes,
reader: strings.NewReader(overLimitBody),
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "read failure",
contentLength: -1,
reader: failingReader{err: errors.New("read failure")},
wantRead: 0,
},
{
name: "empty body",
contentLength: 0,
reader: strings.NewReader(""),
wantRead: 0,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
reader := &guardedReader{
reader: tc.reader,
remaining: maxProviderErrorResponseBytes + 1,
}
err := providerHTTPErrorFromBody(statusCode, tc.contentLength, reader)
if err == nil || err.StatusCode() != statusCode {
t.Fatalf("error status = %v, want %d", err, statusCode)
}
if reader.bytes != tc.wantRead {
t.Fatalf("body bytes read = %d, want %d", reader.bytes, tc.wantRead)
}
if reader.violated {
t.Fatal("body reader was asked to read beyond the overflow probe")
}
if got := err.ProviderMessage(); got != tc.wantMessage {
t.Fatalf("provider message = %q, want %q", got, tc.wantMessage)
}
if tc.wantMessage == "" {
if err.ProviderCode() != "" || err.ProviderType() != "" || strings.Contains(err.Error(), marker) {
t.Fatalf("discarded details were retained: %#v", err)
}
}
})
}
}