Preserve transport error identities
This commit is contained in:
@@ -25,6 +25,18 @@ var (
|
||||
ErrMalformedResponse = errors.New("malformed llm response")
|
||||
)
|
||||
|
||||
type requestFailedError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *requestFailedError) Error() string {
|
||||
return ErrRequestFailed.Error()
|
||||
}
|
||||
|
||||
func (e *requestFailedError) Unwrap() []error {
|
||||
return []error{ErrRequestFailed, e.cause}
|
||||
}
|
||||
|
||||
type OpenAICompatibleConfig struct {
|
||||
BaseURL string
|
||||
Model string
|
||||
@@ -130,7 +142,7 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
|
||||
httpResp, err := httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err)
|
||||
return nil, &requestFailedError{cause: err}
|
||||
}
|
||||
defer httpResp.Body.Close()
|
||||
|
||||
|
||||
@@ -4,9 +4,11 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -20,10 +22,14 @@ var errTransportStopped = errors.New("transport stopped after request inspection
|
||||
type deadlineCapturingTransport struct {
|
||||
deadline time.Time
|
||||
hasDeadline bool
|
||||
err error
|
||||
}
|
||||
|
||||
func (t *deadlineCapturingTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
t.deadline, t.hasDeadline = req.Context().Deadline()
|
||||
if t.err != nil {
|
||||
return nil, t.err
|
||||
}
|
||||
return nil, errTransportStopped
|
||||
}
|
||||
|
||||
@@ -33,6 +39,19 @@ func (contextErrorTransport) RoundTrip(req *http.Request) (*http.Response, error
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
type waitingContextTransport struct{}
|
||||
|
||||
func (waitingContextTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
<-req.Context().Done()
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
|
||||
func assertDeadlineNear(t *testing.T, deadline, before, after time.Time, duration time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
@@ -762,6 +781,9 @@ func TestOpenAICompatibleClientOmittedTimeoutUsesClientTimeout(t *testing.T) {
|
||||
if !errors.Is(err, ErrRequestFailed) {
|
||||
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, errTransportStopped) {
|
||||
t.Fatalf("expected transport cause, got %v", err)
|
||||
}
|
||||
if !transport.hasDeadline {
|
||||
t.Fatal("expected client timeout to set a transport deadline")
|
||||
}
|
||||
@@ -1027,7 +1049,7 @@ func TestOpenAICompatibleClientMalformedResponseMissingChoices(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T) {
|
||||
transport := &deadlineCapturingTransport{}
|
||||
transport := &deadlineCapturingTransport{err: context.DeadlineExceeded}
|
||||
generationTimeout := 2 * time.Second
|
||||
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
@@ -1056,6 +1078,9 @@ func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T
|
||||
if !errors.Is(err, ErrRequestFailed) {
|
||||
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("expected generation deadline identity, got %v", err)
|
||||
}
|
||||
if !transport.hasDeadline {
|
||||
t.Fatal("expected generation timeout to set a transport deadline")
|
||||
}
|
||||
@@ -1063,7 +1088,7 @@ func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientCallerDeadlineTakesPrecedence(t *testing.T) {
|
||||
transport := &deadlineCapturingTransport{}
|
||||
transport := &deadlineCapturingTransport{err: context.DeadlineExceeded}
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "http://example.com/v1",
|
||||
Model: "m",
|
||||
@@ -1092,6 +1117,9 @@ func TestOpenAICompatibleClientCallerDeadlineTakesPrecedence(t *testing.T) {
|
||||
if !errors.Is(err, ErrRequestFailed) {
|
||||
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("expected caller deadline identity, got %v", err)
|
||||
}
|
||||
if !transport.hasDeadline {
|
||||
t.Fatal("expected caller context to set a transport deadline")
|
||||
}
|
||||
@@ -1124,8 +1152,87 @@ func TestOpenAICompatibleClientCancellationReturnsRequestFailure(t *testing.T) {
|
||||
if !errors.Is(err, ErrRequestFailed) {
|
||||
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), context.Canceled.Error()) {
|
||||
t.Fatalf("expected cancellation detail, got %v", err)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected cancellation identity, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientExpiredCallerDeadlineReturnsRequestFailure(t *testing.T) {
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "http://example.com/v1",
|
||||
Model: "m",
|
||||
HTTPClient: &http.Client{
|
||||
Transport: contextErrorTransport{},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithDeadline(context.Background(), time.Unix(1, 0))
|
||||
defer cancel()
|
||||
_, err = client.Generate(ctx, domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||
})
|
||||
if !errors.Is(err, ErrRequestFailed) || !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("expected request failure and caller deadline identities, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientWholeRequestTimeoutReturnsRequestFailure(t *testing.T) {
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "http://example.com/v1",
|
||||
Model: "m",
|
||||
HTTPClient: &http.Client{
|
||||
Timeout: time.Millisecond,
|
||||
Transport: waitingContextTransport{},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||
})
|
||||
if !errors.Is(err, ErrRequestFailed) || !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("expected request failure and client timeout identities, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientTransportFailurePreservesCauseWithoutSensitiveText(t *testing.T) {
|
||||
transportCause := errors.New("transport diagnostic")
|
||||
const (
|
||||
endpoint = "http://sensitive-endpoint.example/private"
|
||||
apiKey = "sensitive-api-key"
|
||||
content = "sensitive prompt content"
|
||||
)
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: endpoint,
|
||||
Model: "m",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return nil, transportCause
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: content}}},
|
||||
Target: domain.ExecutionTarget{APIKey: apiKey},
|
||||
})
|
||||
if !errors.Is(err, ErrRequestFailed) || !errors.Is(err, transportCause) {
|
||||
t.Fatalf("expected request failure and transport cause identities, got %v", err)
|
||||
}
|
||||
var requestErr *url.Error
|
||||
if !errors.As(err, &requestErr) {
|
||||
t.Fatalf("expected underlying http.Client.Do URL error, got %T: %v", err, err)
|
||||
}
|
||||
for _, sensitive := range []string{endpoint, "sensitive-endpoint.example", apiKey, content, transportCause.Error()} {
|
||||
if strings.Contains(err.Error(), sensitive) {
|
||||
t.Fatalf("transport error exposed %q: %v", sensitive, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1170,23 +1277,38 @@ func TestOpenAICompatibleClientRejectsInvalidExecutionSettings(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) {
|
||||
var selectedURL string
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "",
|
||||
Model: "m",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
selectedURL = req.URL.String()
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
`{"choices":[{"message":{"content":"request endpoint"}}]}`,
|
||||
)),
|
||||
Request: req,
|
||||
}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected empty configured base URL to be allowed, got %v", err)
|
||||
}
|
||||
|
||||
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||
response, err := client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||
Target: domain.ExecutionTarget{Endpoint: "http://localhost:9999/v1"},
|
||||
Target: domain.ExecutionTarget{Endpoint: "http://request-endpoint.example/v1"},
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected request failure due to unreachable endpoint")
|
||||
if err != nil {
|
||||
t.Fatalf("generate with request endpoint: %v", err)
|
||||
}
|
||||
if !errors.Is(err, ErrRequestFailed) {
|
||||
t.Fatalf("expected ErrRequestFailed with request endpoint override, got %v", err)
|
||||
if response.Content != "request endpoint" {
|
||||
t.Fatalf("response content = %q, want request endpoint", response.Content)
|
||||
}
|
||||
if selectedURL != "http://request-endpoint.example/v1/chat/completions" {
|
||||
t.Fatalf("selected URL = %q", selectedURL)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user