Bound provider error response body reads
This commit is contained in:
@@ -218,6 +218,8 @@ go test ./...
|
|||||||
Stage 2 is complete when every read path is deterministically bounded and the
|
Stage 2 is complete when every read path is deterministically bounded and the
|
||||||
HTTP client's current branch is still untouched.
|
HTTP client's current branch is still untouched.
|
||||||
|
|
||||||
|
**Status:** Complete.
|
||||||
|
|
||||||
## Stage 3: Integrate Structured Status Errors into the Built-In Client
|
## Stage 3: Integrate Structured Status Errors into the Built-In Client
|
||||||
|
|
||||||
### Objective
|
### Objective
|
||||||
|
|||||||
@@ -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 {
|
func parseProviderErrorEnvelope(body []byte) providerErrorDetails {
|
||||||
decoder := json.NewDecoder(strings.NewReader(string(body)))
|
decoder := json.NewDecoder(strings.NewReader(string(body)))
|
||||||
decoder.UseNumber()
|
decoder.UseNumber()
|
||||||
|
|||||||
@@ -3,12 +3,39 @@ package llm
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"reflect"
|
"reflect"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"unicode/utf8"
|
"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) {
|
func TestProviderHTTPErrorEnvelopeParsing(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -136,3 +163,93 @@ func TestProviderHTTPErrorIdentityAndFormatting(t *testing.T) {
|
|||||||
t.Fatalf("zero error behavior is not safe: %v", zero)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user