package llm import ( "context" "errors" "io" "net/http" "strings" "testing" ) func TestOpenAICompatibleClientStructuredNonSuccessResponse(t *testing.T) { body := `{"error":{"message":" provider\nmessage\u200b","type":"invalid\ttype","code":1.5e+4}}` responseBody := &countingReadCloser{reader: strings.NewReader(body)} client := newNonSuccessResponseClient(t, http.StatusBadRequest, int64(len(body)), responseBody) response, err := client.Generate(context.Background(), ordinaryGenerateRequest()) if response != nil { t.Fatalf("response = %#v, want nil", response) } if !errors.Is(err, ErrUnexpectedStatus) { t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err) } var providerHTTPError *ProviderHTTPError if !errors.As(err, &providerHTTPError) { t.Fatalf("error = %T, want *ProviderHTTPError", err) } if providerHTTPError.StatusCode() != http.StatusBadRequest || providerHTTPError.ProviderCode() != "1.5e+4" || providerHTTPError.ProviderType() != "invalid type" || providerHTTPError.ProviderMessage() != "provider message" { t.Fatalf("provider error = %#v", providerHTTPError) } if !responseBody.closed { t.Fatal("non-success response body was not closed") } } func TestOpenAICompatibleClientNonSuccessBodyOwnership(t *testing.T) { const marker = "provider-body-marker" normalBody := `{"error":{"message":"` + marker + `"}}` overLimitBody := normalBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(normalBody)) tests := []struct { name string contentLength int64 reader io.Reader wantRead int64 wantMessage string }{ { name: "normal", contentLength: int64(len(normalBody)), reader: strings.NewReader(normalBody), wantRead: int64(len(normalBody)), wantMessage: marker, }, { name: "declared oversize", contentLength: maxProviderErrorResponseBytes + 1, reader: strings.NewReader(normalBody), wantRead: 0, }, { name: "streamed oversize", contentLength: -1, reader: &guardedReader{ reader: strings.NewReader(overLimitBody), remaining: maxProviderErrorResponseBytes + 1, }, wantRead: maxProviderErrorResponseBytes + 1, }, { name: "underreported oversize", contentLength: maxProviderErrorResponseBytes, reader: &guardedReader{ reader: strings.NewReader(overLimitBody), remaining: maxProviderErrorResponseBytes + 1, }, wantRead: maxProviderErrorResponseBytes + 1, }, { name: "malformed", contentLength: 1, reader: strings.NewReader("{"), wantRead: 1, }, { name: "read failure", contentLength: -1, reader: failingReader{err: errors.New("response read failed")}, wantRead: 0, }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { body := &countingReadCloser{reader: tc.reader} client := newNonSuccessResponseClient(t, http.StatusBadGateway, tc.contentLength, body) response, err := client.Generate(context.Background(), ordinaryGenerateRequest()) if response != nil { t.Fatalf("response = %#v, want nil", response) } var providerHTTPError *ProviderHTTPError if !errors.As(err, &providerHTTPError) { t.Fatalf("error = %T, want *ProviderHTTPError", err) } if !body.closed { t.Fatal("response body was not closed") } if body.bytesRead != tc.wantRead { t.Fatalf("body bytes read = %d, want %d", body.bytesRead, tc.wantRead) } if body.bytesRead > maxProviderErrorResponseBytes+1 { t.Fatalf("body bytes read = %d, exceeds overflow probe", body.bytesRead) } if guarded, ok := tc.reader.(*guardedReader); ok && guarded.violated { t.Fatal("body reader was asked to read beyond the overflow probe") } if got := providerHTTPError.ProviderMessage(); got != tc.wantMessage { t.Fatalf("provider message = %q, want %q", got, tc.wantMessage) } if tc.wantMessage == "" && (providerHTTPError.ProviderCode() != "" || providerHTTPError.ProviderType() != "" || strings.Contains(providerHTTPError.Error(), marker)) { t.Fatalf("discarded details were retained: %#v", providerHTTPError) } }) } } func newNonSuccessResponseClient(t *testing.T, statusCode int, contentLength int64, body io.ReadCloser) *OpenAICompatibleClient { t.Helper() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "https://provider.example/v1", Model: "m", HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return &http.Response{ StatusCode: statusCode, ContentLength: contentLength, Body: body, }, nil })}, }) if err != nil { t.Fatalf("construct client: %v", err) } return client }