139 lines
5.3 KiB
Go
139 lines
5.3 KiB
Go
package llm
|
||
|
||
import (
|
||
"errors"
|
||
"fmt"
|
||
"reflect"
|
||
"strings"
|
||
"testing"
|
||
"unicode/utf8"
|
||
)
|
||
|
||
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)
|
||
}
|
||
}
|