diff --git a/docs/roadmap/implementation.md b/docs/roadmap/implementation.md index 5ecfb68..0f82041 100644 --- a/docs/roadmap/implementation.md +++ b/docs/roadmap/implementation.md @@ -218,6 +218,8 @@ go test ./... Stage 2 is complete when every read path is deterministically bounded and the HTTP client's current branch is still untouched. +**Status:** Complete. + ## Stage 3: Integrate Structured Status Errors into the Built-In Client ### Objective diff --git a/internal/llm/provider_http_error.go b/internal/llm/provider_http_error.go index 2bcbd47..a12892d 100644 --- a/internal/llm/provider_http_error.go +++ b/internal/llm/provider_http_error.go @@ -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() diff --git a/internal/llm/provider_http_error_test.go b/internal/llm/provider_http_error_test.go index 7f8088a..cf14b22 100644 --- a/internal/llm/provider_http_error_test.go +++ b/internal/llm/provider_http_error_test.go @@ -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) + } + } + }) + } +}