146 lines
4.7 KiB
Go
146 lines
4.7 KiB
Go
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
|
|
}
|