Files
promptkit/internal/llm/openai_compatible_client_test.go

1698 lines
54 KiB
Go

package llm
import (
"context"
"encoding/json"
"errors"
"io"
"math"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"strings"
"sync"
"testing"
"time"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
var (
errTransportStopped = errors.New("transport stopped after request inspection")
errResponseReadPastLimit = errors.New("response reader was read past the allowed boundary")
)
const successfulProviderResponse = `{"choices":[{"message":{"content":"ok"}}]}`
type recordedProviderRequest struct {
method string
url string
header http.Header
fields map[string]json.RawMessage
}
func (r recordedProviderRequest) decodeField(t *testing.T, name string, destination any) bool {
t.Helper()
raw, exists := r.fields[name]
if !exists {
return false
}
if err := json.Unmarshal(raw, destination); err != nil {
t.Fatalf("decode request field %q: %v", name, err)
}
return true
}
type recordingProvider struct {
t *testing.T
server *httptest.Server
mu sync.Mutex
statusCode int
responseBody string
requests []recordedProviderRequest
}
func newRecordingProvider(t *testing.T) *recordingProvider {
t.Helper()
provider := &recordingProvider{
t: t,
statusCode: http.StatusOK,
responseBody: successfulProviderResponse,
}
provider.server = httptest.NewServer(http.HandlerFunc(provider.handle))
t.Cleanup(provider.server.Close)
return provider
}
func (p *recordingProvider) handle(w http.ResponseWriter, request *http.Request) {
defer request.Body.Close()
fields := make(map[string]json.RawMessage)
if err := json.NewDecoder(request.Body).Decode(&fields); err != nil {
p.t.Errorf("decode provider request: %v", err)
}
p.mu.Lock()
p.requests = append(p.requests, recordedProviderRequest{
method: request.Method,
url: request.URL.String(),
header: request.Header.Clone(),
fields: fields,
})
statusCode := p.statusCode
responseBody := p.responseBody
p.mu.Unlock()
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(statusCode)
if _, err := io.WriteString(w, responseBody); err != nil {
p.t.Errorf("write provider response: %v", err)
}
}
func (p *recordingProvider) endpoint(path string) string {
return p.server.URL + path
}
func (p *recordingProvider) respond(statusCode int, body string) {
p.mu.Lock()
defer p.mu.Unlock()
p.statusCode = statusCode
p.responseBody = body
}
func (p *recordingProvider) requestCount() int {
p.mu.Lock()
defer p.mu.Unlock()
return len(p.requests)
}
func (p *recordingProvider) lastRequest(t *testing.T) recordedProviderRequest {
t.Helper()
p.mu.Lock()
defer p.mu.Unlock()
if len(p.requests) == 0 {
t.Fatal("provider received no request")
}
return p.requests[len(p.requests)-1]
}
func newProviderClient(t *testing.T, provider *recordingProvider, config OpenAICompatibleConfig) *OpenAICompatibleClient {
t.Helper()
if config.BaseURL == "" {
config.BaseURL = provider.endpoint("/v1")
}
client, err := NewOpenAICompatibleClient(config)
if err != nil {
t.Fatalf("construct client: %v", err)
}
return client
}
func ordinaryGenerateRequest() domain.GenerateRequest {
return domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Model: "model"},
}
}
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
}
type contextErrorTransport struct{}
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()
}
type countingReadCloser struct {
reader io.Reader
bytesRead int64
closed bool
}
func (r *countingReadCloser) Read(p []byte) (int, error) {
n, err := r.reader.Read(p)
r.bytesRead += int64(n)
return n, err
}
func (r *countingReadCloser) Close() error {
r.closed = true
return nil
}
type repeatingByteReader byte
func (r repeatingByteReader) Read(p []byte) (int, error) {
for i := range p {
p[i] = byte(r)
}
return len(p), nil
}
type guardedRepeatingReader struct {
value byte
remaining int64
readPastLimit bool
}
func (r *guardedRepeatingReader) Read(p []byte) (int, error) {
if r.remaining == 0 {
r.readPastLimit = true
return 0, errResponseReadPastLimit
}
if int64(len(p)) > r.remaining {
p = p[:r.remaining]
}
for i := range p {
p[i] = r.value
}
r.remaining -= int64(len(p))
return len(p), nil
}
const (
successResponsePrefix = `{"choices":[{"message":{"content":"`
successResponseSuffix = `"}}]}`
responseContentMarker = "provider-secret-fragment"
)
func sizedSuccessResponseBody(size int64) *countingReadCloser {
contentBytes := size - int64(len(successResponsePrefix)+len(successResponseSuffix))
if contentBytes < int64(len(responseContentMarker)) {
panic("successful response size is too small")
}
return &countingReadCloser{reader: io.MultiReader(
strings.NewReader(successResponsePrefix),
strings.NewReader(responseContentMarker),
io.LimitReader(repeatingByteReader('x'), contentBytes-int64(len(responseContentMarker))),
strings.NewReader(successResponseSuffix),
)}
}
func assertDeadlineNear(t *testing.T, deadline, before, after time.Time, duration time.Duration) {
t.Helper()
const tolerance = 100 * time.Millisecond
earliest := before.Add(duration - tolerance)
latest := after.Add(duration + tolerance)
if deadline.Before(earliest) || deadline.After(latest) {
t.Fatalf("expected deadline between %v and %v, got %v", earliest, latest, deadline)
}
}
func TestNewOpenAICompatibleClientRejectsInvalidBaseURL(t *testing.T) {
for _, endpoint := range []string{
"://invalid",
"/v1",
"https:///v1",
"ftp://provider.example/v1",
"https://user@provider.example/v1",
"https://provider.example/v1?mode=chat",
"https://provider.example/v1#chat",
} {
t.Run(endpoint, func(t *testing.T) {
_, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: endpoint})
if !errors.Is(err, ErrInvalidConfig) {
t.Fatalf("expected ErrInvalidConfig, got %v", err)
}
})
}
}
func TestNewOpenAICompatibleClientDoesNotMutateSuppliedZeroTimeoutClient(t *testing.T) {
transport := http.DefaultTransport
supplied := &http.Client{Transport: transport}
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
HTTPClient: supplied,
})
if err != nil {
t.Fatalf("unexpected constructor error: %v", err)
}
if supplied.Timeout != 0 {
t.Fatalf("expected supplied client timeout to remain zero, got %v", supplied.Timeout)
}
if client.httpClient == supplied {
t.Fatal("expected constructed client to use a cloned HTTP client")
}
if client.httpClient.Timeout <= 0 {
t.Fatalf("expected constructed client to use a positive default timeout, got %v", client.httpClient.Timeout)
}
if client.httpClient.Transport != transport {
t.Fatal("expected cloned client to preserve the supplied transport")
}
}
func TestNewOpenAICompatibleClientDoesNotMutateSuppliedNonzeroTimeoutClient(t *testing.T) {
transport := http.DefaultTransport
suppliedTimeout := 37 * time.Second
supplied := &http.Client{
Timeout: suppliedTimeout,
Transport: transport,
}
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
Timeout: 2 * time.Second,
HTTPClient: supplied,
})
if err != nil {
t.Fatalf("unexpected constructor error: %v", err)
}
if supplied.Timeout != suppliedTimeout {
t.Fatalf("expected supplied client timeout to remain %v, got %v", suppliedTimeout, supplied.Timeout)
}
if client.httpClient == supplied {
t.Fatal("expected constructed client to use a cloned HTTP client")
}
if client.httpClient.Timeout != suppliedTimeout {
t.Fatalf("expected cloned client timeout %v, got %v", suppliedTimeout, client.httpClient.Timeout)
}
if client.httpClient.Transport != transport {
t.Fatal("expected cloned client to preserve the supplied transport")
}
}
func TestNewOpenAICompatibleClientTreatsSuppliedNegativeTimeoutAsUnset(t *testing.T) {
transport := http.DefaultTransport
supplied := &http.Client{
Timeout: -time.Second,
Transport: transport,
}
configuredTimeout := 23 * time.Second
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
Timeout: configuredTimeout,
HTTPClient: supplied,
})
if err != nil {
t.Fatalf("unexpected constructor error: %v", err)
}
if supplied.Timeout != -time.Second {
t.Fatalf("expected supplied client timeout to remain negative, got %v", supplied.Timeout)
}
if client.httpClient == supplied {
t.Fatal("expected constructed client to use a cloned HTTP client")
}
if client.httpClient.Timeout != configuredTimeout {
t.Fatalf("expected cloned client timeout %v, got %v", configuredTimeout, client.httpClient.Timeout)
}
if client.httpClient.Transport != transport {
t.Fatal("expected cloned client to preserve the supplied transport")
}
}
func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
tests := []struct {
name string
run func(*testing.T)
}{
{name: "complete request and response mapping", run: checkCompleteRequestAndResponseMapping},
{name: "cache-controlled message", run: checkCacheControlledMessageMapping},
{name: "empty cache-control TTL", run: checkEmptyCacheControlTTLOmission},
{name: "session ID", run: checkSessionIDMapping},
{name: "optional response fields", run: checkOptionalResponseFieldOmission},
{name: "reasoning and extra parameters", run: checkReasoningAndExtraParameterMapping},
{name: "optional request field presence", run: checkRequestFieldPresence},
{name: "configured model fallback", run: checkConfiguredModelFallback},
}
for _, tc := range tests {
t.Run(tc.name, tc.run)
}
}
func checkCompleteRequestAndResponseMapping(t *testing.T) {
provider := newRecordingProvider(t)
provider.respond(http.StatusOK, `{
"choices": [{"message": {"role": "assistant", "content": "hello from model"}}],
"usage": {"prompt_tokens": 11, "completion_tokens": 22, "total_tokens": 33}
}`)
client := newProviderClient(t, provider, OpenAICompatibleConfig{
Timeout: 2 * time.Second,
})
t.Setenv("PROMPTKIT_TEST_API_KEY", "secret-key")
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{Role: "system", Content: "You are helpful."},
{Role: "user", Content: "Say hello"},
}},
Target: domain.ExecutionTarget{
Model: "gpt-test",
Temperature: 0.4,
MaxTokens: 123,
TopP: 0.7,
ServiceTier: "priority",
APIKeyEnv: "PROMPTKIT_TEST_API_KEY",
},
StructuredOutput: &domain.StructuredOutputSpec{
Type: domain.StructuredOutputJSONSchema,
JSONSchema: &domain.StructuredOutputJSONSpec{
Name: "weather_schema",
Strict: true,
Schema: map[string]any{
"type": "object",
"properties": map[string]any{
"location": map[string]any{"type": "string"},
},
"required": []any{"location"},
},
},
},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if resp.Content != "hello from model" {
t.Fatalf("unexpected content: %q", resp.Content)
}
if resp.Usage.PromptTokens != 11 || resp.Usage.CompletionTokens != 22 || resp.Usage.TotalTokens != 33 {
t.Fatalf("unexpected usage: %+v", resp.Usage)
}
if resp.Usage.CachedTokens != 0 || resp.Usage.CacheWriteTokens != 0 {
t.Fatalf("expected absent cache usage fields to remain zero, got %+v", resp.Usage)
}
request := provider.lastRequest(t)
if request.method != http.MethodPost {
t.Fatalf("method = %q, want POST", request.method)
}
if request.url != "/v1/chat/completions" {
t.Fatalf("URL = %q, want /v1/chat/completions", request.url)
}
if got := request.header.Get("Content-Type"); got != "application/json" {
t.Fatalf("Content-Type = %q, want application/json", got)
}
if got := request.header.Get("Authorization"); got != "Bearer secret-key" {
t.Fatalf("Authorization = %q, want Bearer secret-key", got)
}
var model string
if !request.decodeField(t, "model", &model) || model != "gpt-test" {
t.Fatalf("model = %q, want gpt-test", model)
}
var temperature float64
if !request.decodeField(t, "temperature", &temperature) || temperature != 0.4 {
t.Fatalf("temperature = %v, want 0.4", temperature)
}
var maxTokens int
if !request.decodeField(t, "max_tokens", &maxTokens) || maxTokens != 123 {
t.Fatalf("max_tokens = %d, want 123", maxTokens)
}
var topP float64
if !request.decodeField(t, "top_p", &topP) || topP != 0.7 {
t.Fatalf("top_p = %v, want 0.7", topP)
}
var serviceTier string
if !request.decodeField(t, "service_tier", &serviceTier) || serviceTier != "priority" {
t.Fatalf("service_tier = %q, want priority", serviceTier)
}
var messages []struct {
Role string `json:"role"`
Content string `json:"content"`
}
if !request.decodeField(t, "messages", &messages) || len(messages) != 2 {
t.Fatalf("messages = %#v, want two entries", messages)
}
if messages[0].Role != "system" || messages[0].Content != "You are helpful." {
t.Fatalf("unexpected first message: %#v", messages[0])
}
if messages[1].Role != "user" || messages[1].Content != "Say hello" {
t.Fatalf("unexpected second message: %#v", messages[1])
}
var responseFormat struct {
Type string `json:"type"`
JSONSchema struct {
Name string `json:"name"`
Strict bool `json:"strict"`
Schema map[string]any `json:"schema"`
} `json:"json_schema"`
}
if !request.decodeField(t, "response_format", &responseFormat) {
t.Fatal("response_format was omitted")
}
if responseFormat.Type != "json_schema" || responseFormat.JSONSchema.Name != "weather_schema" || !responseFormat.JSONSchema.Strict {
t.Fatalf("unexpected response_format: %#v", responseFormat)
}
if responseFormat.JSONSchema.Schema["type"] != "object" {
t.Fatalf("unexpected JSON schema: %#v", responseFormat.JSONSchema.Schema)
}
}
func TestOpenAICompatibleClientAuthentication(t *testing.T) {
tests := []struct {
name string
configureEnv func(*testing.T)
target domain.ExecutionTarget
wantAuth string
wantErr error
wantCallCount int
}{
{
name: "direct key takes precedence over environment",
configureEnv: func(t *testing.T) {
t.Setenv("PROMPTKIT_TEST_API_KEY", "env-key")
},
target: domain.ExecutionTarget{
APIKeyEnv: "PROMPTKIT_TEST_API_KEY",
APIKey: "direct-llm-key",
},
wantAuth: "Bearer direct-llm-key",
wantCallCount: 1,
},
{
name: "no key omits authorization",
wantCallCount: 1,
},
{
name: "missing environment key fails before transport",
configureEnv: func(t *testing.T) {
t.Setenv("PROMPTKIT_MISSING_KEY", "")
},
target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_MISSING_KEY"},
wantErr: ErrInvalidRequest,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
if tc.configureEnv != nil {
tc.configureEnv(t)
}
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "model"})
request := ordinaryGenerateRequest()
request.Target = tc.target
_, err := client.Generate(context.Background(), request)
if tc.wantErr != nil {
if !errors.Is(err, tc.wantErr) {
t.Fatalf("error = %v, want %v", err, tc.wantErr)
}
} else if err != nil {
t.Fatalf("generate: %v", err)
}
if got := provider.requestCount(); got != tc.wantCallCount {
t.Fatalf("provider calls = %d, want %d", got, tc.wantCallCount)
}
if tc.wantCallCount == 1 {
if got := provider.lastRequest(t).header.Get("Authorization"); got != tc.wantAuth {
t.Fatalf("Authorization = %q, want %q", got, tc.wantAuth)
}
}
})
}
}
func checkCacheControlledMessageMapping(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{
Role: "system",
Content: "Stable instructions.",
CacheControl: &domain.CacheControl{
Type: domain.CacheControlEphemeral,
TTL: "1h",
},
},
{Role: "user", Content: "Dynamic request."},
}},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
request := provider.lastRequest(t)
for _, forbidden := range []string{"cache_control", "extra_params"} {
if _, exists := request.fields[forbidden]; exists {
t.Fatalf("expected top-level %s to be omitted", forbidden)
}
}
var messages []struct {
Role string `json:"role"`
Content json.RawMessage `json:"content"`
}
if !request.decodeField(t, "messages", &messages) || len(messages) != 2 {
t.Fatalf("messages = %#v, want two entries", messages)
}
if messages[0].Role != "system" {
t.Fatalf("first message role = %q, want system", messages[0].Role)
}
var contentBlocks []struct {
Type string `json:"type"`
Text string `json:"text"`
CacheControl struct {
Type string `json:"type"`
TTL string `json:"ttl"`
} `json:"cache_control"`
}
if err := json.Unmarshal(messages[0].Content, &contentBlocks); err != nil || len(contentBlocks) != 1 {
t.Fatalf("decode content blocks: %v; blocks = %#v", err, contentBlocks)
}
block := contentBlocks[0]
if block.Type != "text" || block.Text != "Stable instructions." {
t.Fatalf("unexpected text content block: %#v", block)
}
if block.CacheControl.Type != string(domain.CacheControlEphemeral) || block.CacheControl.TTL != "1h" {
t.Fatalf("unexpected cache control: %#v", block.CacheControl)
}
var ordinaryContent string
if err := json.Unmarshal(messages[1].Content, &ordinaryContent); err != nil {
t.Fatalf("decode ordinary message content: %v", err)
}
if messages[1].Role != "user" || ordinaryContent != "Dynamic request." {
t.Fatalf("unexpected ordinary message: role=%q content=%q", messages[1].Role, ordinaryContent)
}
}
func checkEmptyCacheControlTTLOmission(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{
Role: "system",
Content: "Stable instructions.",
CacheControl: &domain.CacheControl{
Type: domain.CacheControlEphemeral,
},
},
}},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
request := provider.lastRequest(t)
var messages []struct {
Content []struct {
CacheControl map[string]json.RawMessage `json:"cache_control"`
} `json:"content"`
}
if !request.decodeField(t, "messages", &messages) || len(messages) != 1 || len(messages[0].Content) != 1 {
t.Fatalf("unexpected messages: %#v", messages)
}
cacheControl := messages[0].Content[0].CacheControl
var controlType string
if err := json.Unmarshal(cacheControl["type"], &controlType); err != nil {
t.Fatalf("decode cache control type: %v", err)
}
if controlType != string(domain.CacheControlEphemeral) {
t.Fatalf("cache_control type = %q", controlType)
}
if _, exists := cacheControl["ttl"]; exists {
t.Fatalf("expected empty ttl to be omitted, got %#v", cacheControl)
}
}
func checkSessionIDMapping(t *testing.T) {
tests := []struct {
name string
sessionID string
want string
wantPresent bool
}{
{name: "trimmed value is present", sessionID: " session-123 ", want: "session-123", wantPresent: true},
{name: "empty value is omitted", sessionID: " "},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
request := ordinaryGenerateRequest()
request.Prompt.SessionID = tc.sessionID
if _, err := client.Generate(context.Background(), request); err != nil {
t.Fatalf("generate: %v", err)
}
recorded := provider.lastRequest(t)
var got string
present := recorded.decodeField(t, "session_id", &got)
if present != tc.wantPresent || got != tc.want {
t.Fatalf("session_id = %q, present = %v; want %q, %v", got, present, tc.want, tc.wantPresent)
}
if got := recorded.header.Get("x-session-id"); got != "" {
t.Fatalf("unexpected x-session-id header: %q", got)
}
})
}
}
func TestOpenAICompatibleClientRejectsTooLongSessionID(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Model: "model",
})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{
SessionID: strings.Repeat("x", domain.SessionIDMaxLength+1),
Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}},
},
})
if err == nil {
t.Fatal("expected invalid request error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
}
func TestOpenAICompatibleClientResponseFraming(t *testing.T) {
tests := []struct {
name string
run func(*testing.T)
}{
{name: "usage mapping", run: checkCacheUsageMapping},
{name: "common response failures", run: checkCommonResponseFailures},
{name: "successful response byte boundary", run: checkSuccessfulResponseByteBoundary},
{name: "continuing oversized response", run: checkContinuingOversizedResponse},
{name: "single response document", run: checkSingleResponseDocument},
}
for _, tc := range tests {
t.Run(tc.name, tc.run)
}
}
func checkCacheUsageMapping(t *testing.T) {
provider := newRecordingProvider(t)
provider.respond(http.StatusOK, `{
"choices": [{"message": {"role": "assistant", "content": "ok"}}],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 20,
"total_tokens": 120,
"prompt_tokens_details": {"cached_tokens": 80},
"cache_write_tokens": 60
}
}`)
client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "model"})
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if resp.Usage.PromptTokens != 100 || resp.Usage.CompletionTokens != 20 || resp.Usage.TotalTokens != 120 {
t.Fatalf("unexpected base usage fields: %+v", resp.Usage)
}
if resp.Usage.CachedTokens != 80 || resp.Usage.CacheWriteTokens != 60 {
t.Fatalf("unexpected cache usage fields: %+v", resp.Usage)
}
}
func checkOptionalResponseFieldOmission(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
_, err := client.Generate(context.Background(), ordinaryGenerateRequest())
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
fields := provider.lastRequest(t).fields
if _, exists := fields["response_format"]; exists {
t.Fatal("expected response_format omitted")
}
if _, exists := fields["service_tier"]; exists {
t.Fatal("expected service_tier omitted")
}
}
func checkReasoningAndExtraParameterMapping(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{
Model: "model",
ReasoningEffort: "high",
ExtraParams: map[string]any{
"string_value": "on",
"number_value": 42,
"boolean_value": true,
"object_value": map[string]any{"nested": "value", "count": 2},
"array_value": []any{"first", 3, false},
},
},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
request := provider.lastRequest(t)
var reasoningEffort string
if !request.decodeField(t, "reasoning_effort", &reasoningEffort) || reasoningEffort != "high" {
t.Fatalf("reasoning_effort = %q, want high", reasoningEffort)
}
var stringValue string
if !request.decodeField(t, "string_value", &stringValue) || stringValue != "on" {
t.Fatalf("string_value = %q, want on", stringValue)
}
var numberValue int
if !request.decodeField(t, "number_value", &numberValue) || numberValue != 42 {
t.Fatalf("number_value = %d, want 42", numberValue)
}
var booleanValue bool
if !request.decodeField(t, "boolean_value", &booleanValue) || !booleanValue {
t.Fatalf("boolean_value = %v, want true", booleanValue)
}
var objectValue struct {
Nested string `json:"nested"`
Count int `json:"count"`
}
if !request.decodeField(t, "object_value", &objectValue) || objectValue.Nested != "value" || objectValue.Count != 2 {
t.Fatalf("unexpected object_value: %#v", objectValue)
}
if _, exists := request.fields["extra_params"]; exists {
t.Fatal("expected extra_params wrapper omitted")
}
var arrayValue []any
if !request.decodeField(t, "array_value", &arrayValue) || len(arrayValue) != 3 || arrayValue[0] != "first" || arrayValue[1] != float64(3) || arrayValue[2] != false {
t.Fatalf("unexpected array_value: %#v", arrayValue)
}
}
func checkRequestFieldPresence(t *testing.T) {
tests := []struct {
name string
presence domain.ExecutionTargetPresence
wantAbsent []string
wantZeros bool
}{
{
name: "unset optional fields are omitted",
wantAbsent: []string{"reasoning_effort", "extra_params"},
},
{
name: "explicit numeric zeros are present",
presence: domain.ExecutionTargetPresence{
Temperature: true,
MaxTokens: true,
TopP: true,
},
wantZeros: true,
},
{
name: "implicit numeric zeros are omitted",
wantAbsent: []string{"temperature", "max_tokens", "top_p"},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
request := ordinaryGenerateRequest()
request.TargetPresence = tc.presence
if _, err := client.Generate(context.Background(), request); err != nil {
t.Fatalf("generate: %v", err)
}
recorded := provider.lastRequest(t)
for _, field := range tc.wantAbsent {
if _, exists := recorded.fields[field]; exists {
t.Fatalf("expected field %q to be omitted", field)
}
}
if tc.wantZeros {
var temperature, topP float64
var maxTokens int
if !recorded.decodeField(t, "temperature", &temperature) || temperature != 0 {
t.Fatalf("temperature = %v, want explicit zero", temperature)
}
if !recorded.decodeField(t, "max_tokens", &maxTokens) || maxTokens != 0 {
t.Fatalf("max_tokens = %d, want explicit zero", maxTokens)
}
if !recorded.decodeField(t, "top_p", &topP) || topP != 0 {
t.Fatalf("top_p = %v, want explicit zero", topP)
}
}
})
}
}
func TestOpenAICompatibleClientTimeoutAndErrorIdentity(t *testing.T) {
tests := []struct {
name string
run func(*testing.T)
}{
{name: "omitted generation timeout uses client timeout", run: checkOmittedGenerationTimeout},
{name: "generation timeout sets earlier deadline", run: checkGenerationTimeoutPrecedence},
{name: "caller deadline takes precedence", run: checkCallerDeadlinePrecedence},
{name: "caller cancellation identity", run: checkCallerCancellationIdentity},
{name: "expired caller deadline identity", run: checkExpiredCallerDeadlineIdentity},
{name: "whole-request timeout identity", run: checkWholeRequestTimeoutIdentity},
{name: "transport cause identity and redaction", run: checkTransportFailureIdentityAndRedaction},
}
for _, tc := range tests {
t.Run(tc.name, tc.run)
}
}
func checkOmittedGenerationTimeout(t *testing.T) {
transport := &deadlineCapturingTransport{}
clientTimeout := 5 * time.Second
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Timeout: clientTimeout,
HTTPClient: &http.Client{
Transport: transport,
},
})
if err != nil {
t.Fatal(err)
}
before := time.Now()
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Model: "model", TimeoutSeconds: 0},
})
after := time.Now()
if err == nil {
t.Fatal("expected transport error")
}
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")
}
assertDeadlineNear(t, transport.deadline, before, after, clientTimeout)
}
func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *testing.T) {
tests := []struct {
name string
extraParams map[string]any
want string
}{
{name: "empty key", extraParams: map[string]any{"": "empty"}, want: "key must not be empty"},
{name: "unserializable value", extraParams: map[string]any{"bad": math.Inf(1)}, want: "JSON-serializable"},
}
for _, key := range []string{
"model",
"session_id",
"messages",
"temperature",
"max_tokens",
"top_p",
"service_tier",
"reasoning_effort",
"response_format",
} {
tests = append(tests, struct {
name string
extraParams map[string]any
want string
}{
name: "reserved key " + key,
extraParams: map[string]any{key: "collision"},
want: "reserved request field",
})
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Model: "model", ExtraParams: tc.extraParams},
})
if err == nil {
t.Fatal("expected invalid request error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("expected error to contain %q, got %v", tc.want, err)
}
if calls := provider.requestCount(); calls != 0 {
t.Fatalf("provider calls = %d, want 0", calls)
}
})
}
}
func checkConfiguredModelFallback(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "default-model"})
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
var model string
if !provider.lastRequest(t).decodeField(t, "model", &model) || model != "default-model" {
t.Fatalf("model = %q, want default-model", model)
}
}
func TestOpenAICompatibleClientEndpointComposition(t *testing.T) {
tests := []struct {
name string
run func(*testing.T)
}{
{name: "request endpoint overrides configured endpoint", run: checkEndpointOverride},
{name: "request endpoint works without configured endpoint", run: checkEmptyConfiguredEndpoint},
{name: "completion URL preserves valid base paths", run: checkCompletionURLComposition},
{name: "invalid selected endpoint fails before transport", run: checkInvalidSelectedEndpointRejection},
{name: "an endpoint is required", run: checkRequiredEndpoint},
}
for _, tc := range tests {
t.Run(tc.name, tc.run)
}
}
func checkEndpointOverride(t *testing.T) {
configuredProvider := newRecordingProvider(t)
configuredProvider.respond(http.StatusOK, `{"choices":[{"message":{"content":"default"}}]}`)
selectedProvider := newRecordingProvider(t)
selectedProvider.respond(http.StatusOK, `{"choices":[{"message":{"content":"override"}}]}`)
client := newProviderClient(t, configuredProvider, OpenAICompatibleConfig{Model: "m"})
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Endpoint: selectedProvider.endpoint("/v1")},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if resp.Content != "override" {
t.Fatalf("expected override response, got %q", resp.Content)
}
if got := configuredProvider.requestCount(); got != 0 {
t.Fatalf("configured provider calls = %d, want 0", got)
}
if got := selectedProvider.requestCount(); got != 1 {
t.Fatalf("selected provider calls = %d, want 1", got)
}
if got := selectedProvider.lastRequest(t).url; got != "/v1/chat/completions" {
t.Fatalf("selected URL = %q, want /v1/chat/completions", got)
}
}
func checkCommonResponseFailures(t *testing.T) {
const sensitiveBody = `provider-secret-fragment request_payload_details`
tests := []struct {
name string
statusCode int
body string
wantErr error
wantText string
redact string
}{
{
name: "non-success status is redacted",
statusCode: http.StatusBadRequest,
body: `{"error":"` + sensitiveBody + `"}`,
wantErr: ErrUnexpectedStatus,
wantText: "status=400",
redact: sensitiveBody,
},
{name: "invalid JSON", statusCode: http.StatusOK, body: `{not valid json`, wantErr: ErrMalformedResponse},
{name: "missing choices", statusCode: http.StatusOK, body: `{"choices": []}`, wantErr: ErrMalformedResponse},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
provider := newRecordingProvider(t)
provider.respond(tc.statusCode, tc.body)
client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "m"})
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
if response != nil {
t.Fatalf("expected no partial response, got %+v", response)
}
if !errors.Is(err, tc.wantErr) {
t.Fatalf("error = %v, want %v", err, tc.wantErr)
}
if tc.wantText != "" && !strings.Contains(err.Error(), tc.wantText) {
t.Fatalf("error %q does not contain %q", err, tc.wantText)
}
if tc.redact != "" && strings.Contains(err.Error(), tc.redact) {
t.Fatalf("error exposed provider response content: %v", err)
}
})
}
}
func checkSuccessfulResponseByteBoundary(t *testing.T) {
tests := []struct {
name string
size int64
contentLength int64
wantErr bool
}{
{name: "just below limit with content length", size: maxOpenAIChatResponseBytes - 1, contentLength: maxOpenAIChatResponseBytes - 1},
{name: "exact limit with content length", size: maxOpenAIChatResponseBytes, contentLength: maxOpenAIChatResponseBytes},
{name: "one byte over with content length", size: maxOpenAIChatResponseBytes + 1, contentLength: maxOpenAIChatResponseBytes + 1, wantErr: true},
{name: "just below limit without content length", size: maxOpenAIChatResponseBytes - 1, contentLength: -1},
{name: "exact limit without content length", size: maxOpenAIChatResponseBytes, contentLength: -1},
{name: "one byte over without content length", size: maxOpenAIChatResponseBytes + 1, contentLength: -1, wantErr: true},
{name: "one byte over with underreported content length", size: maxOpenAIChatResponseBytes + 1, contentLength: maxOpenAIChatResponseBytes - 1, wantErr: true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
body := sizedSuccessResponseBody(tc.size)
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://provider.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: body,
ContentLength: tc.contentLength,
Request: req,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
response, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if tc.wantErr {
if response != nil {
t.Fatalf("expected no partial response, got %+v", response)
}
if !errors.Is(err, ErrMalformedResponse) {
t.Fatalf("expected ErrMalformedResponse, got %v", err)
}
if strings.Contains(err.Error(), responseContentMarker) {
t.Fatalf("error exposed provider content: %v", err)
}
} else {
if err != nil {
t.Fatalf("generate: %v", err)
}
wantContentBytes := tc.size - int64(len(successResponsePrefix)+len(successResponseSuffix))
if int64(len(response.Content)) != wantContentBytes || !strings.HasPrefix(response.Content, responseContentMarker) {
t.Fatalf("unexpected response content length or prefix")
}
}
if body.bytesRead > maxOpenAIChatResponseBytes+1 {
t.Fatalf("read %d bytes, limit is %d", body.bytesRead, maxOpenAIChatResponseBytes+1)
}
if tc.wantErr && tc.contentLength > maxOpenAIChatResponseBytes && body.bytesRead != 0 {
t.Fatalf("read %d bytes despite oversized Content-Length", body.bytesRead)
}
if tc.wantErr && tc.contentLength <= maxOpenAIChatResponseBytes && body.bytesRead != maxOpenAIChatResponseBytes+1 {
t.Fatalf("read %d bytes, want one byte beyond the limit", body.bytesRead)
}
if !tc.wantErr && body.bytesRead != tc.size {
t.Fatalf("read %d bytes, want complete %d-byte response", body.bytesRead, tc.size)
}
if !body.closed {
t.Fatal("response body was not closed")
}
})
}
}
func checkContinuingOversizedResponse(t *testing.T) {
prefix := successResponsePrefix + responseContentMarker
continuation := &guardedRepeatingReader{
value: 'x',
remaining: maxOpenAIChatResponseBytes + 1 - int64(len(prefix)),
}
body := &countingReadCloser{reader: io.MultiReader(
strings.NewReader(prefix),
continuation,
)}
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://provider.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: body,
ContentLength: -1,
Request: req,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
type outcome struct {
response *domain.GenerateResponse
err error
}
done := make(chan outcome, 1)
go func() {
response, generateErr := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
done <- outcome{response: response, err: generateErr}
}()
select {
case result := <-done:
if result.response != nil {
t.Fatalf("expected no partial response, got %+v", result.response)
}
if !errors.Is(result.err, ErrMalformedResponse) {
t.Fatalf("expected ErrMalformedResponse, got %v", result.err)
}
if strings.Contains(result.err.Error(), responseContentMarker) {
t.Fatalf("error exposed provider content: %v", result.err)
}
case <-time.After(5 * time.Second):
t.Fatal("timed out rejecting continuing oversized response")
}
if body.bytesRead != maxOpenAIChatResponseBytes+1 {
t.Fatalf("read %d bytes, want %d", body.bytesRead, maxOpenAIChatResponseBytes+1)
}
if continuation.readPastLimit {
t.Fatal("response reader was read past one byte beyond the limit")
}
if !body.closed {
t.Fatal("response body was not closed")
}
}
func checkSingleResponseDocument(t *testing.T) {
validResponse := `{"choices":[{"message":{"content":"ok"}}]}`
tests := []struct {
name string
body string
wantErr bool
}{
{name: "one document", body: validResponse},
{name: "trailing whitespace", body: validResponse + " \n\t\r "},
{name: "trailing garbage", body: validResponse + " " + responseContentMarker, wantErr: true},
{name: "second JSON value", body: validResponse + ` {"detail":"` + responseContentMarker + `"}`, wantErr: true},
{name: "truncated document", body: validResponse[:len(validResponse)-2], wantErr: true},
{name: "missing choices", body: `{"choices":[]}`, wantErr: true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
body := &countingReadCloser{reader: strings.NewReader(tc.body)}
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://provider.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: body,
ContentLength: int64(len(tc.body)),
Request: req,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
response, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if tc.wantErr {
if response != nil {
t.Fatalf("expected no partial response, got %+v", response)
}
if !errors.Is(err, ErrMalformedResponse) {
t.Fatalf("expected ErrMalformedResponse, got %v", err)
}
if strings.Contains(err.Error(), responseContentMarker) {
t.Fatalf("error exposed provider content: %v", err)
}
} else if err != nil || response == nil || response.Content != "ok" {
t.Fatalf("response = %+v, error = %v", response, err)
}
if !body.closed {
t.Fatal("response body was not closed")
}
})
}
}
func checkGenerationTimeoutPrecedence(t *testing.T) {
transport := &deadlineCapturingTransport{err: context.DeadlineExceeded}
generationTimeout := 2 * time.Second
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Model: "m",
HTTPClient: &http.Client{
Timeout: 10 * time.Second,
Transport: transport,
},
})
if err != nil {
t.Fatal(err)
}
before := time.Now()
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{
TimeoutSeconds: int(generationTimeout / time.Second),
},
})
after := time.Now()
if err == nil {
t.Fatal("expected transport error")
}
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")
}
assertDeadlineNear(t, transport.deadline, before, after, generationTimeout)
}
func checkCallerDeadlinePrecedence(t *testing.T) {
transport := &deadlineCapturingTransport{err: context.DeadlineExceeded}
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Model: "m",
HTTPClient: &http.Client{
Timeout: 10 * time.Second,
Transport: transport,
},
})
if err != nil {
t.Fatal(err)
}
callerDeadline := time.Now().Add(time.Second)
ctx, cancel := context.WithDeadline(context.Background(), callerDeadline)
defer cancel()
_, err = client.Generate(ctx, domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{
TimeoutSeconds: 2,
},
})
if err == nil {
t.Fatal("expected transport error")
}
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")
}
if !transport.deadline.Equal(callerDeadline) {
t.Fatalf("expected caller deadline %v, got %v", callerDeadline, transport.deadline)
}
}
func checkCallerCancellationIdentity(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.WithCancel(context.Background())
cancel()
_, err = client.Generate(ctx, domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if err == nil {
t.Fatal("expected cancellation error")
}
if !errors.Is(err, ErrRequestFailed) {
t.Fatalf("expected ErrRequestFailed, got %v", err)
}
if !errors.Is(err, context.Canceled) {
t.Fatalf("expected cancellation identity, got %v", err)
}
}
func checkExpiredCallerDeadlineIdentity(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 checkWholeRequestTimeoutIdentity(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 checkTransportFailureIdentityAndRedaction(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)
}
}
}
func TestOpenAICompatibleClientRejectsInvalidExecutionSettings(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Model: "m",
})
if err != nil {
t.Fatal(err)
}
type testCase struct {
name string
target domain.ExecutionTarget
}
tests := []testCase{
{name: "non-finite temperature", target: domain.ExecutionTarget{Temperature: math.NaN()}},
{name: "negative max tokens", target: domain.ExecutionTarget{MaxTokens: -1}},
{name: "non-finite top p", target: domain.ExecutionTarget{TopP: math.Inf(1)}},
{name: "negative timeout", target: domain.ExecutionTarget{TimeoutSeconds: -1}},
}
if strconv.IntSize == 64 {
durationLimit := int64(math.MaxInt64 / int64(time.Second))
tests = append(tests, testCase{name: "unrepresentable timeout", target: domain.ExecutionTarget{TimeoutSeconds: int(durationLimit) + 1}})
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: tt.target,
})
if err == nil {
t.Fatal("expected invalid request error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
})
}
}
func checkEmptyConfiguredEndpoint(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)
}
response, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Endpoint: "http://request-endpoint.example/v1"},
})
if err != nil {
t.Fatalf("generate with request endpoint: %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)
}
}
func checkCompletionURLComposition(t *testing.T) {
tests := []struct {
name string
baseURL string
wantURL string
}{
{name: "HTTP host", baseURL: "http://provider.example", wantURL: "http://provider.example/chat/completions"},
{name: "HTTPS host", baseURL: "https://provider.example", wantURL: "https://provider.example/chat/completions"},
{name: "nested path", baseURL: "https://provider.example/api/openai/v1", wantURL: "https://provider.example/api/openai/v1/chat/completions"},
{name: "trailing slash", baseURL: "https://provider.example/v1/", wantURL: "https://provider.example/v1/chat/completions"},
{name: "repeated trailing slashes", baseURL: " https://provider.example/api/v1/// ", wantURL: "https://provider.example/api/v1/chat/completions"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
var selectedURL string
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: tc.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":"ok"}}]}`)),
Request: req,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if err != nil {
t.Fatalf("generate: %v", err)
}
if selectedURL != tc.wantURL {
t.Fatalf("selected URL = %q, want %q", selectedURL, tc.wantURL)
}
})
}
}
func checkInvalidSelectedEndpointRejection(t *testing.T) {
invalidEndpoints := []string{
"/v1",
"https:///v1",
"ftp://provider.example/v1",
"https://user@provider.example/v1",
"https://provider.example/v1?mode=chat",
"https://provider.example/v1#chat",
"https://sensitive-endpoint.example/%zz",
}
for _, endpoint := range invalidEndpoints {
t.Run(endpoint, func(t *testing.T) {
transportCalls := 0
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://configured.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
transportCalls++
return nil, errors.New("transport must not be called")
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Endpoint: endpoint},
})
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
if strings.Contains(err.Error(), endpoint) {
t.Fatalf("error exposed selected endpoint %q: %v", endpoint, err)
}
if transportCalls != 0 {
t.Fatalf("transport calls = %d, want 0", transportCalls)
}
})
}
}
func checkRequiredEndpoint(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "",
Model: "m",
})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{},
})
if err == nil {
t.Fatal("expected endpoint-required error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
}