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 body string want providerErrorDetails }{ { name: "all supported string fields", body: `{"error":{"message":"diagnostic","type":"invalid_request_error","code":"unsupported_parameter"}}`, want: providerErrorDetails{providerMessage: "diagnostic", providerType: "invalid_request_error", providerCode: "unsupported_parameter"}, }, { name: "integer code", body: `{"error":{"code":17}}`, want: providerErrorDetails{providerCode: "17"}, }, { name: "fractional code", body: `{"error":{"code":1.25}}`, want: providerErrorDetails{providerCode: "1.25"}, }, { name: "exponent code", body: `{"error":{"code":6.02e+23}}`, want: providerErrorDetails{providerCode: "6.02e+23"}, }, { name: "invalid fields do not discard valid fields", body: `{"error":{"message":null,"type":"invalid_request_error","code":false}}`, want: providerErrorDetails{providerType: "invalid_request_error"}, }, { name: "unknown fields are ignored", body: `{"trace":"do not retain","error":{"param":"temperature","metadata":{"secret":"x"}}}`, want: providerErrorDetails{}, }, {name: "missing error", body: `{}`, want: providerErrorDetails{}}, {name: "null error", body: `{"error":null}`, want: providerErrorDetails{}}, {name: "scalar error", body: `{"error":"nope"}`, want: providerErrorDetails{}}, {name: "empty error", body: `{"error":{}}`, want: providerErrorDetails{}}, {name: "malformed", body: `{"error":`, want: providerErrorDetails{}}, {name: "truncated", body: `{"error":{"message":"x"`, want: providerErrorDetails{}}, {name: "trailing garbage", body: `{"error":{"message":"x"}} garbage`, want: providerErrorDetails{}}, {name: "second document", body: `{"error":{"message":"x"}} {}`, want: providerErrorDetails{}}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { if got := parseProviderErrorEnvelope([]byte(tc.body)); !reflect.DeepEqual(got, tc.want) { t.Fatalf("parseProviderErrorEnvelope() = %#v, want %#v", got, tc.want) } }) } } func TestProviderErrorTextNormalizationAndLimits(t *testing.T) { validIdentifier := strings.Repeat("界", maxProviderErrorIdentifierRunes) validMessage := strings.Repeat("界", maxProviderErrorMessageRunes) tests := []struct { name string got string want string }{ {name: "multibyte text", got: "Grüße 世界", want: "Grüße 世界"}, {name: "invalid UTF-8", got: string([]byte{'a', 0xff, 'b'}), want: "a�b"}, {name: "whitespace control and format runs", got: " \n\talpha\x00\u200b\u200bbeta \r ", want: "alpha beta"}, {name: "blank normalization", got: "\t\u200b\n", want: ""}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { if got := normalizeProviderErrorText(tc.got); got != tc.want { t.Fatalf("normalizeProviderErrorText() = %q, want %q", got, tc.want) } }) } if got := normalizeProviderErrorIdentifier(validIdentifier); got != validIdentifier { t.Fatalf("exact identifier boundary = %q, want retained value", got) } if got := normalizeProviderErrorIdentifier(validIdentifier + "界"); got != "" { t.Fatalf("overlong identifier = %q, want empty", got) } if got := normalizeProviderErrorMessage(validMessage); got != validMessage { t.Fatalf("exact message boundary = %q, want retained value", got) } wantTruncatedMessage := strings.Repeat("界", maxProviderErrorMessageRunes-1) + "…" if got := normalizeProviderErrorMessage(validMessage + "界"); got != wantTruncatedMessage { t.Fatalf("overlong message length = %d, want %d", utf8.RuneCountInString(got), maxProviderErrorMessageRunes) } } func TestProviderHTTPErrorIdentityAndFormatting(t *testing.T) { const marker = "provider-secret-marker" err := newProviderHTTPError(429, providerErrorDetails{ providerCode: marker + "-code", providerType: marker + "-type", providerMessage: marker + "-message", }) if err.StatusCode() != 429 || err.ProviderCode() != marker+"-code" || err.ProviderType() != marker+"-type" || err.ProviderMessage() != marker+"-message" { t.Fatalf("accessors returned unexpected values: %#v", err) } if !errors.Is(err, ErrUnexpectedStatus) { t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err) } for _, rendered := range []string{fmt.Sprintf("%v", err), fmt.Sprintf("%+v", err), fmt.Sprintf("%#v", err)} { if rendered != "llm returned non-success status: status=429" { t.Fatalf("formatted error = %q", rendered) } if strings.Contains(rendered, marker) { t.Fatalf("formatted error exposed provider marker: %q", rendered) } } var nilError *ProviderHTTPError if nilError.StatusCode() != 0 || nilError.ProviderCode() != "" || nilError.ProviderType() != "" || nilError.ProviderMessage() != "" { t.Fatal("nil accessors returned provider values") } if nilError.Error() != "llm returned non-success status" || nilError.GoString() != "llm returned non-success status" || !errors.Is(nilError, ErrUnexpectedStatus) { t.Fatalf("nil error behavior is not safe: %v", nilError) } zero := &ProviderHTTPError{} if zero.Error() != "llm returned non-success status" || zero.GoString() != "llm returned non-success status" || !errors.Is(zero, ErrUnexpectedStatus) { 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) } } }) } }