Preserve transport error identities

This commit is contained in:
2026-08-11 23:28:00 +00:00
parent 350b0e76d9
commit c281f721bc
5 changed files with 171 additions and 20 deletions

View File

@@ -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()

View File

@@ -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)
}
}