From e291b8bfe905f79911001811cc0b15d54c0013de Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 11 Aug 2026 23:58:31 +0000 Subject: [PATCH] Consolidate provider transport test scaffolding --- internal/llm/openai_compatible_client_test.go | 1120 ++++++++--------- 1 file changed, 560 insertions(+), 560 deletions(-) diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go index d70ca37..b3e8574 100644 --- a/internal/llm/openai_compatible_client_test.go +++ b/internal/llm/openai_compatible_client_test.go @@ -11,6 +11,7 @@ import ( "net/url" "strconv" "strings" + "sync" "testing" "time" @@ -22,6 +23,126 @@ var ( 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 @@ -236,42 +357,36 @@ func TestNewOpenAICompatibleClientTreatsSuppliedNegativeTimeoutAsUnset(t *testin } } -func TestOpenAICompatibleClientGenerateSuccess(t *testing.T) { - type observedRequest struct { - Authorization string - Body map[string]any +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}, } - obs := &observedRequest{} - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - obs.Authorization = r.Header.Get("Authorization") - if r.URL.Path != "/v1/chat/completions" { - t.Fatalf("unexpected path: %s", r.URL.Path) - } - if ct := r.Header.Get("Content-Type"); ct != "application/json" { - t.Fatalf("unexpected content type: %s", ct) - } + for _, tc := range tests { + t.Run(tc.name, tc.run) + } +} - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&obs.Body); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ +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} -}`)) - })) - defer ts.Close() +}`) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ - BaseURL: ts.URL + "/v1", + client := newProviderClient(t, provider, OpenAICompatibleConfig{ Timeout: 2 * time.Second, }) - if err != nil { - t.Fatalf("unexpected constructor error: %v", err) - } t.Setenv("PROMPTKIT_TEST_API_KEY", "secret-key") resp, err := client.Generate(context.Background(), domain.GenerateRequest{ @@ -316,109 +431,144 @@ func TestOpenAICompatibleClientGenerateSuccess(t *testing.T) { t.Fatalf("expected absent cache usage fields to remain zero, got %+v", resp.Usage) } - if obs.Authorization != "Bearer secret-key" { - t.Fatalf("unexpected Authorization header: %q", obs.Authorization) + request := provider.lastRequest(t) + if request.method != http.MethodPost { + t.Fatalf("method = %q, want POST", request.method) } - if got, ok := obs.Body["model"].(string); !ok || got != "gpt-test" { - t.Fatalf("unexpected model payload: %#v", obs.Body["model"]) + if request.url != "/v1/chat/completions" { + t.Fatalf("URL = %q, want /v1/chat/completions", request.url) } - if got, ok := obs.Body["temperature"].(float64); !ok || got != 0.4 { - t.Fatalf("unexpected temperature payload: %#v", obs.Body["temperature"]) + if got := request.header.Get("Content-Type"); got != "application/json" { + t.Fatalf("Content-Type = %q, want application/json", got) } - if got, ok := obs.Body["max_tokens"].(float64); !ok || got != 123 { - t.Fatalf("unexpected max_tokens payload: %#v", obs.Body["max_tokens"]) - } - if got, ok := obs.Body["top_p"].(float64); !ok || got != 0.7 { - t.Fatalf("unexpected top_p payload: %#v", obs.Body["top_p"]) - } - if got, ok := obs.Body["service_tier"].(string); !ok || got != "priority" { - t.Fatalf("unexpected service_tier payload: %#v", obs.Body["service_tier"]) + if got := request.header.Get("Authorization"); got != "Bearer secret-key" { + t.Fatalf("Authorization = %q, want Bearer secret-key", got) } - msgs, ok := obs.Body["messages"].([]any) - if !ok || len(msgs) != 2 { - t.Fatalf("unexpected messages payload: %#v", obs.Body["messages"]) + var model string + if !request.decodeField(t, "model", &model) || model != "gpt-test" { + t.Fatalf("model = %q, want gpt-test", model) } - msg0 := msgs[0].(map[string]any) - if msg0["role"] != "system" || msg0["content"] != "You are helpful." { - t.Fatalf("unexpected first message: %#v", msg0) + var temperature float64 + if !request.decodeField(t, "temperature", &temperature) || temperature != 0.4 { + t.Fatalf("temperature = %v, want 0.4", temperature) } - msg1 := msgs[1].(map[string]any) - if msg1["role"] != "user" || msg1["content"] != "Say hello" { - t.Fatalf("unexpected second message: %#v", msg1) + 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) } - responseFormat, ok := obs.Body["response_format"].(map[string]any) - if !ok { - t.Fatalf("expected response_format payload, got %#v", obs.Body["response_format"]) + var messages []struct { + Role string `json:"role"` + Content string `json:"content"` } - if responseFormat["type"] != "json_schema" { - t.Fatalf("expected response_format.type=json_schema, got %#v", responseFormat["type"]) + if !request.decodeField(t, "messages", &messages) || len(messages) != 2 { + t.Fatalf("messages = %#v, want two entries", messages) } - jsonSchema, ok := responseFormat["json_schema"].(map[string]any) - if !ok { - t.Fatalf("expected response_format.json_schema map, got %#v", responseFormat["json_schema"]) + if messages[0].Role != "system" || messages[0].Content != "You are helpful." { + t.Fatalf("unexpected first message: %#v", messages[0]) } - if jsonSchema["name"] != "weather_schema" { - t.Fatalf("expected json_schema.name weather_schema, got %#v", jsonSchema["name"]) + if messages[1].Role != "user" || messages[1].Content != "Say hello" { + t.Fatalf("unexpected second message: %#v", messages[1]) } - if jsonSchema["strict"] != true { - t.Fatalf("expected json_schema.strict=true, got %#v", jsonSchema["strict"]) + + 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 _, ok := jsonSchema["schema"].(map[string]any); !ok { - t.Fatalf("expected json_schema.schema object, got %#v", jsonSchema["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 TestOpenAICompatibleClientDirectAPIKeyPreferredOverEnv(t *testing.T) { - const directKey = "direct-llm-key" - t.Setenv("PROMPTKIT_TEST_API_KEY", "env-key") - - var gotAuth string - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotAuth = r.Header.Get("Authorization") - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - 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{ - Model: "model", - APIKeyEnv: "PROMPTKIT_TEST_API_KEY", - APIKey: directKey, +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, }, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) } - if gotAuth != "Bearer "+directKey { - t.Fatalf("unexpected Authorization header: %q", gotAuth) + + 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 TestOpenAICompatibleClientSerializesCacheControlledMessageAsContentBlock(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() +func checkCacheControlledMessageMapping(t *testing.T) { + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ + _, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{ { Role: "system", @@ -436,59 +586,55 @@ func TestOpenAICompatibleClientSerializesCacheControlledMessageAsContentBlock(t t.Fatalf("expected no error, got %v", err) } + request := provider.lastRequest(t) for _, forbidden := range []string{"cache_control", "extra_params"} { - if _, exists := observedBody[forbidden]; exists { - t.Fatalf("expected top-level %s to be omitted, got %#v", forbidden, observedBody[forbidden]) + if _, exists := request.fields[forbidden]; exists { + t.Fatalf("expected top-level %s to be omitted", forbidden) } } - msgs, ok := observedBody["messages"].([]any) - if !ok || len(msgs) != 2 { - t.Fatalf("unexpected messages payload: %#v", observedBody["messages"]) + var messages []struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` } - msg0 := msgs[0].(map[string]any) - if msg0["role"] != "system" { - t.Fatalf("unexpected first message role: %#v", msg0["role"]) + if !request.decodeField(t, "messages", &messages) || len(messages) != 2 { + t.Fatalf("messages = %#v, want two entries", messages) } - contentBlocks, ok := msg0["content"].([]any) - if !ok || len(contentBlocks) != 1 { - t.Fatalf("expected first message content block array, got %#v", msg0["content"]) + if messages[0].Role != "system" { + t.Fatalf("first message role = %q, want system", messages[0].Role) } - block := contentBlocks[0].(map[string]any) - if block["type"] != "text" || block["text"] != "Stable instructions." { + 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) } - cacheControl, ok := block["cache_control"].(map[string]any) - if !ok { - t.Fatalf("expected cache_control on content block, got %#v", block) + if block.CacheControl.Type != string(domain.CacheControlEphemeral) || block.CacheControl.TTL != "1h" { + t.Fatalf("unexpected cache control: %#v", block.CacheControl) } - if cacheControl["type"] != string(domain.CacheControlEphemeral) || cacheControl["ttl"] != "1h" { - t.Fatalf("unexpected cache_control payload: %#v", cacheControl) + var ordinaryContent string + if err := json.Unmarshal(messages[1].Content, &ordinaryContent); err != nil { + t.Fatalf("decode ordinary message content: %v", err) } - - msg1 := msgs[1].(map[string]any) - if msg1["role"] != "user" || msg1["content"] != "Dynamic request." { - t.Fatalf("expected uncached message to keep string content, got %#v", msg1) + if messages[1].Role != "user" || ordinaryContent != "Dynamic request." { + t.Fatalf("unexpected ordinary message: role=%q content=%q", messages[1].Role, ordinaryContent) } } -func TestOpenAICompatibleClientOmitsEmptyCacheControlTTL(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() +func checkEmptyCacheControlTTLOmission(t *testing.T) { + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ + _, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{ { Role: "system", @@ -504,83 +650,59 @@ func TestOpenAICompatibleClientOmitsEmptyCacheControlTTL(t *testing.T) { t.Fatalf("expected no error, got %v", err) } - msgs := observedBody["messages"].([]any) - msg0 := msgs[0].(map[string]any) - contentBlocks := msg0["content"].([]any) - block := contentBlocks[0].(map[string]any) - cacheControl := block["cache_control"].(map[string]any) - if cacheControl["type"] != string(domain.CacheControlEphemeral) { - t.Fatalf("unexpected cache_control type: %#v", cacheControl) + 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 TestOpenAICompatibleClientSerializesSessionID(t *testing.T) { - var observedBody map[string]any - var observedSessionHeader string - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - observedSessionHeader = r.Header.Get("x-session-id") - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) +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: " "}, } - _, err = client.Generate(context.Background(), domain.GenerateRequest{ - Prompt: domain.RenderedPrompt{ - SessionID: " session-123 ", - Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}, - }, - Target: domain.ExecutionTarget{Model: "model"}, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - if observedBody["session_id"] != "session-123" { - t.Fatalf("expected top-level session_id, got %#v", observedBody["session_id"]) - } - if observedSessionHeader != "" { - t.Fatalf("did not expect x-session-id header, got %q", observedSessionHeader) - } -} + 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 -func TestOpenAICompatibleClientOmitsEmptySessionID(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ - Prompt: domain.RenderedPrompt{ - SessionID: " ", - Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}, - }, - Target: domain.ExecutionTarget{Model: "model"}, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - if _, exists := observedBody["session_id"]; exists { - t.Fatalf("expected empty session_id to be omitted, got %#v", observedBody["session_id"]) + 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) + } + }) } } @@ -607,9 +729,26 @@ func TestOpenAICompatibleClientRejectsTooLongSessionID(t *testing.T) { } } -func TestOpenAICompatibleClientParsesCacheUsage(t *testing.T) { - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{ +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, @@ -618,14 +757,8 @@ func TestOpenAICompatibleClientParsesCacheUsage(t *testing.T) { "prompt_tokens_details": {"cached_tokens": 80}, "cache_write_tokens": 60 } -}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "model"}) - if err != nil { - t.Fatal(err) - } +}`) + 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"}}}, @@ -641,54 +774,28 @@ func TestOpenAICompatibleClientParsesCacheUsage(t *testing.T) { } } -func TestOpenAICompatibleClientOmitsResponseFormatWhenNoStructuredOutput(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() +func checkOptionalResponseFieldOmission(t *testing.T) { + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - 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{Model: "model"}, - }) + _, err := client.Generate(context.Background(), ordinaryGenerateRequest()) if err != nil { t.Fatalf("expected no error, got %v", err) } - if _, exists := observedBody["response_format"]; exists { - t.Fatalf("expected response_format omitted, got %#v", observedBody["response_format"]) + fields := provider.lastRequest(t).fields + if _, exists := fields["response_format"]; exists { + t.Fatal("expected response_format omitted") } - if _, exists := observedBody["service_tier"]; exists { - t.Fatalf("expected service_tier omitted, got %#v", observedBody["service_tier"]) + if _, exists := fields["service_tier"]; exists { + t.Fatal("expected service_tier omitted") } } -func TestOpenAICompatibleClientSerializesReasoningEffortAndExtraParams(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() +func checkReasoningAndExtraParameterMapping(t *testing.T) { + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ + _, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, Target: domain.ExecutionTarget{ Model: "model", @@ -705,132 +812,119 @@ func TestOpenAICompatibleClientSerializesReasoningEffortAndExtraParams(t *testin if err != nil { t.Fatalf("expected no error, got %v", err) } - if observedBody["reasoning_effort"] != "high" { - t.Fatalf("expected reasoning_effort high, got %#v", observedBody["reasoning_effort"]) + + request := provider.lastRequest(t) + var reasoningEffort string + if !request.decodeField(t, "reasoning_effort", &reasoningEffort) || reasoningEffort != "high" { + t.Fatalf("reasoning_effort = %q, want high", reasoningEffort) } - if observedBody["string_value"] != "on" { - t.Fatalf("unexpected string extra param: %#v", observedBody["string_value"]) + var stringValue string + if !request.decodeField(t, "string_value", &stringValue) || stringValue != "on" { + t.Fatalf("string_value = %q, want on", stringValue) } - if observedBody["number_value"] != float64(42) { - t.Fatalf("unexpected number extra param: %#v", observedBody["number_value"]) + var numberValue int + if !request.decodeField(t, "number_value", &numberValue) || numberValue != 42 { + t.Fatalf("number_value = %d, want 42", numberValue) } - if observedBody["boolean_value"] != true { - t.Fatalf("unexpected boolean extra param: %#v", observedBody["boolean_value"]) + var booleanValue bool + if !request.decodeField(t, "boolean_value", &booleanValue) || !booleanValue { + t.Fatalf("boolean_value = %v, want true", booleanValue) } - objectValue, ok := observedBody["object_value"].(map[string]any) - if !ok || objectValue["nested"] != "value" || objectValue["count"] != float64(2) { - t.Fatalf("unexpected object extra param: %#v", observedBody["object_value"]) + var objectValue struct { + Nested string `json:"nested"` + Count int `json:"count"` } - if _, exists := observedBody["extra_params"]; exists { - t.Fatalf("expected extra_params wrapper omitted, got %#v", observedBody["extra_params"]) + if !request.decodeField(t, "object_value", &objectValue) || objectValue.Nested != "value" || objectValue.Count != 2 { + t.Fatalf("unexpected object_value: %#v", objectValue) } - arrayValue, ok := observedBody["array_value"].([]any) - if !ok || len(arrayValue) != 3 || arrayValue[0] != "first" || arrayValue[1] != float64(3) || arrayValue[2] != false { - t.Fatalf("unexpected array extra param: %#v", observedBody["array_value"]) + 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 TestOpenAICompatibleClientOmitsReasoningEffortWhenUnset(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - 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{Model: "model"}, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - if _, exists := observedBody["reasoning_effort"]; exists { - t.Fatalf("expected reasoning_effort omitted, got %#v", observedBody["reasoning_effort"]) - } - if _, exists := observedBody["extra_params"]; exists { - t.Fatalf("expected extra_params wrapper omitted, got %#v", observedBody["extra_params"]) - } -} - -func TestOpenAICompatibleClientSerializesExplicitZeroNumericOverrides(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - 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{Model: "model"}, - TargetPresence: domain.ExecutionTargetPresence{ - Temperature: true, - MaxTokens: true, - TopP: true, +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"}, }, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) } - if observedBody["temperature"] != float64(0) { - t.Fatalf("expected explicit zero temperature, got %#v", observedBody["temperature"]) - } - if observedBody["max_tokens"] != float64(0) { - t.Fatalf("expected explicit zero max_tokens, got %#v", observedBody["max_tokens"]) - } - if observedBody["top_p"] != float64(0) { - t.Fatalf("expected explicit zero top_p, got %#v", observedBody["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 TestOpenAICompatibleClientOmitsImplicitZeroNumericFields(t *testing.T) { - var observedBody map[string]any - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defer r.Body.Close() - if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil { - t.Fatalf("failed to decode request body: %v", err) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) +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}, } - _, err = client.Generate(context.Background(), domain.GenerateRequest{ - Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, - Target: domain.ExecutionTarget{Model: "model"}, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - for _, field := range []string{"temperature", "max_tokens", "top_p"} { - if _, exists := observedBody[field]; exists { - t.Fatalf("expected implicit zero field %q to be omitted, got body %#v", field, observedBody) - } + for _, tc := range tests { + t.Run(tc.name, tc.run) } } -func TestOpenAICompatibleClientOmittedTimeoutUsesClientTimeout(t *testing.T) { +func checkOmittedGenerationTimeout(t *testing.T) { transport := &deadlineCapturingTransport{} clientTimeout := 5 * time.Second @@ -899,19 +993,10 @@ func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *te for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - called := false - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - called = true - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ + _, 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}, }) @@ -924,117 +1009,58 @@ func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *te if !strings.Contains(err.Error(), tc.want) { t.Fatalf("expected error to contain %q, got %v", tc.want, err) } - if called { - t.Fatal("provider should not be called for invalid extra_params") + if calls := provider.requestCount(); calls != 0 { + t.Fatalf("provider calls = %d, want 0", calls) } }) } } -func TestOpenAICompatibleClientNoAuthorizationHeaderWhenNoAPIKey(t *testing.T) { - hadAuth := false - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - hadAuth = r.Header.Get("Authorization") != "" - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() +func checkConfiguredModelFallback(t *testing.T) { + provider := newRecordingProvider(t) + client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "default-model"}) - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) - 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{Model: "model"}, - }) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - if hadAuth { - t.Fatal("did not expect Authorization header") - } -} - -func TestOpenAICompatibleClientAPIKeyEnvMissing(t *testing.T) { - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "model"}) - 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{APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, - }) - if err == nil { - t.Fatal("expected missing API key env error") - } - if !errors.Is(err, ErrInvalidRequest) { - t.Fatalf("expected ErrInvalidRequest, got %v", err) - } -} - -func TestOpenAICompatibleClientModelFallbackFromConfig(t *testing.T) { - gotModel := "" - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var body map[string]any - _ = json.NewDecoder(r.Body).Decode(&body) - if m, ok := body["model"].(string); ok { - gotModel = m - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "default-model"}) - if err != nil { - t.Fatal(err) - } - - _, err = client.Generate(context.Background(), domain.GenerateRequest{ + _, 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) } - if gotModel != "default-model" { - t.Fatalf("expected default model, got %q", gotModel) + var model string + if !provider.lastRequest(t).decodeField(t, "model", &model) || model != "default-model" { + t.Fatalf("model = %q, want default-model", model) } } -func TestOpenAICompatibleClientEndpointOverride(t *testing.T) { - defaultHit := false - overrideHit := false - - defaultServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - defaultHit = true - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"default"}}]}`)) - })) - defer defaultServer.Close() - - overrideServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - overrideHit = true - if r.URL.Path != "/v1/chat/completions" { - t.Fatalf("unexpected path: %s", r.URL.Path) - } - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"override"}}]}`)) - })) - defer overrideServer.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: defaultServer.URL + "/v1", Model: "m"}) - if err != nil { - t.Fatal(err) +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: overrideServer.URL + "/v1"}, + Target: domain.ExecutionTarget{Endpoint: selectedProvider.endpoint("/v1")}, }) if err != nil { t.Fatalf("expected no error, got %v", err) @@ -1042,89 +1068,63 @@ func TestOpenAICompatibleClientEndpointOverride(t *testing.T) { if resp.Content != "override" { t.Fatalf("expected override response, got %q", resp.Content) } - if defaultHit { - t.Fatal("default endpoint should not have been called") + if got := configuredProvider.requestCount(); got != 0 { + t.Fatalf("configured provider calls = %d, want 0", got) } - if !overrideHit { - t.Fatal("override endpoint should have been called") + 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 TestOpenAICompatibleClientNon2xxError(t *testing.T) { +func checkCommonResponseFailures(t *testing.T) { const sensitiveBody = `provider-secret-fragment request_payload_details` - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":"` + sensitiveBody + `"}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "m"}) - if err != nil { - t.Fatal(err) + 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}, } - _, err = client.Generate(context.Background(), domain.GenerateRequest{ - Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, - }) - if err == nil { - t.Fatal("expected non-2xx error") - } - if !errors.Is(err, ErrUnexpectedStatus) { - t.Fatalf("expected ErrUnexpectedStatus, got %v", err) - } - if !strings.Contains(err.Error(), "status=400") { - t.Fatalf("expected status detail, got %v", err) - } - if strings.Contains(err.Error(), sensitiveBody) { - t.Fatalf("expected provider response body to be redacted, got %v", err) + 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 TestOpenAICompatibleClientMalformedResponseInvalidJSON(t *testing.T) { - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{not valid json`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", 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"}}}, - }) - if err == nil { - t.Fatal("expected malformed response error") - } - if !errors.Is(err, ErrMalformedResponse) { - t.Fatalf("expected ErrMalformedResponse, got %v", err) - } -} - -func TestOpenAICompatibleClientMalformedResponseMissingChoices(t *testing.T) { - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{"choices": []}`)) - })) - defer ts.Close() - - client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", 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"}}}, - }) - if err == nil { - t.Fatal("expected malformed response error") - } - if !errors.Is(err, ErrMalformedResponse) { - t.Fatalf("expected ErrMalformedResponse, got %v", err) - } -} - -func TestOpenAICompatibleClientBoundsSuccessfulResponseBodies(t *testing.T) { +func checkSuccessfulResponseByteBoundary(t *testing.T) { tests := []struct { name string size int64 @@ -1201,7 +1201,7 @@ func TestOpenAICompatibleClientBoundsSuccessfulResponseBodies(t *testing.T) { } } -func TestOpenAICompatibleClientRejectsContinuingOversizedResponse(t *testing.T) { +func checkContinuingOversizedResponse(t *testing.T) { prefix := successResponsePrefix + responseContentMarker continuation := &guardedRepeatingReader{ value: 'x', @@ -1265,7 +1265,7 @@ func TestOpenAICompatibleClientRejectsContinuingOversizedResponse(t *testing.T) } } -func TestOpenAICompatibleClientRequiresSingleResponseDocument(t *testing.T) { +func checkSingleResponseDocument(t *testing.T) { validResponse := `{"choices":[{"message":{"content":"ok"}}]}` tests := []struct { name string @@ -1323,7 +1323,7 @@ func TestOpenAICompatibleClientRequiresSingleResponseDocument(t *testing.T) { } } -func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T) { +func checkGenerationTimeoutPrecedence(t *testing.T) { transport := &deadlineCapturingTransport{err: context.DeadlineExceeded} generationTimeout := 2 * time.Second @@ -1362,7 +1362,7 @@ func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T assertDeadlineNear(t, transport.deadline, before, after, generationTimeout) } -func TestOpenAICompatibleClientCallerDeadlineTakesPrecedence(t *testing.T) { +func checkCallerDeadlinePrecedence(t *testing.T) { transport := &deadlineCapturingTransport{err: context.DeadlineExceeded} client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "http://example.com/v1", @@ -1403,7 +1403,7 @@ func TestOpenAICompatibleClientCallerDeadlineTakesPrecedence(t *testing.T) { } } -func TestOpenAICompatibleClientCancellationReturnsRequestFailure(t *testing.T) { +func checkCallerCancellationIdentity(t *testing.T) { client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "http://example.com/v1", Model: "m", @@ -1432,7 +1432,7 @@ func TestOpenAICompatibleClientCancellationReturnsRequestFailure(t *testing.T) { } } -func TestOpenAICompatibleClientExpiredCallerDeadlineReturnsRequestFailure(t *testing.T) { +func checkExpiredCallerDeadlineIdentity(t *testing.T) { client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "http://example.com/v1", Model: "m", @@ -1454,7 +1454,7 @@ func TestOpenAICompatibleClientExpiredCallerDeadlineReturnsRequestFailure(t *tes } } -func TestOpenAICompatibleClientWholeRequestTimeoutReturnsRequestFailure(t *testing.T) { +func checkWholeRequestTimeoutIdentity(t *testing.T) { client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "http://example.com/v1", Model: "m", @@ -1475,7 +1475,7 @@ func TestOpenAICompatibleClientWholeRequestTimeoutReturnsRequestFailure(t *testi } } -func TestOpenAICompatibleClientTransportFailurePreservesCauseWithoutSensitiveText(t *testing.T) { +func checkTransportFailureIdentityAndRedaction(t *testing.T) { transportCause := errors.New("transport diagnostic") const ( endpoint = "http://sensitive-endpoint.example/private" @@ -1551,7 +1551,7 @@ func TestOpenAICompatibleClientRejectsInvalidExecutionSettings(t *testing.T) { } } -func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) { +func checkEmptyConfiguredEndpoint(t *testing.T) { var selectedURL string client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "", @@ -1587,7 +1587,7 @@ func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) { } } -func TestOpenAICompatibleClientComposesCompletionURL(t *testing.T) { +func checkCompletionURLComposition(t *testing.T) { tests := []struct { name string baseURL string @@ -1633,7 +1633,7 @@ func TestOpenAICompatibleClientComposesCompletionURL(t *testing.T) { } } -func TestOpenAICompatibleClientRejectsInvalidSelectedEndpointBeforeTransport(t *testing.T) { +func checkInvalidSelectedEndpointRejection(t *testing.T) { invalidEndpoints := []string{ "/v1", "https:///v1", @@ -1675,7 +1675,7 @@ func TestOpenAICompatibleClientRejectsInvalidSelectedEndpointBeforeTransport(t * } } -func TestOpenAICompatibleClientRequiresEndpointWhenUnsetEverywhere(t *testing.T) { +func checkRequiredEndpoint(t *testing.T) { client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "", Model: "m",