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 wantAuthorization string wantError 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 ", }, wantAuthorization: "Bearer direct-llm-key", wantCallCount: 1, }, { name: "environment key supplies authorization", configureEnv: func(t *testing.T) { t.Setenv("PROMPTKIT_TEST_API_KEY", " env-key ") }, target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_TEST_API_KEY"}, wantAuthorization: "Bearer env-key", wantCallCount: 1, }, { name: "no key omits authorization", wantCallCount: 1, }, { name: "optional missing environment omits authorization", configureEnv: func(t *testing.T) { t.Setenv("PROMPTKIT_MISSING_KEY", "") }, target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, wantCallCount: 1, }, { name: "required missing environment fails before transport", configureEnv: func(t *testing.T) { t.Setenv("PROMPTKIT_MISSING_KEY", "") }, target: domain.ExecutionTarget{ APIKeyEnv: "PROMPTKIT_MISSING_KEY", APIKeyRequired: true, }, wantError: ErrInvalidRequest, }, { name: "required target without source fails before transport", target: domain.ExecutionTarget{APIKeyRequired: true}, wantError: 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.wantError != nil { if !errors.Is(err, tc.wantError) { t.Fatalf("error = %v, want %v", err, tc.wantError) } } 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 { values := provider.lastRequest(t).header.Values("Authorization") if tc.wantAuthorization == "" { if len(values) != 0 { t.Fatalf("Authorization values = %q, want absent", values) } } else if len(values) != 1 || values[0] != tc.wantAuthorization { t.Fatalf("Authorization values = %q, want [%q]", values, tc.wantAuthorization) } } }) } } 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: "content presence", run: checkContentPresence}, {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 checkContentPresence(t *testing.T) { tests := []struct { name string content string }{ {name: "explicit empty string", content: ""}, {name: "whitespace string", content: " \n\t "}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { provider := newRecordingProvider(t) provider.respond(http.StatusOK, `{ "choices": [{"message": {"content": `+strconv.Quote(tc.content)+`}}], "usage": { "prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30, "prompt_tokens_details": {"cached_tokens": 4}, "cache_write_tokens": 5 } }`) client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "model"}) response, err := client.Generate(context.Background(), ordinaryGenerateRequest()) if err != nil { t.Fatalf("generate: %v", err) } if response == nil { t.Fatal("expected response") } if response.Content != tc.content { t.Fatalf("content = %q, want %q", response.Content, tc.content) } if response.Usage != (domain.TokenUsage{ PromptTokens: 10, CompletionTokens: 20, TotalTokens: 30, CachedTokens: 4, CacheWriteTokens: 5, }) { t.Fatalf("usage = %+v", response.Usage) } }) } } 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}, {name: "missing content", statusCode: http.StatusOK, body: `{"choices": [{"message": {}}]}`, wantErr: ErrMalformedResponse}, {name: "null content", statusCode: http.StatusOK, body: `{"choices": [{"message": {"content": null}}]}`, wantErr: ErrMalformedResponse}, {name: "non-string content", statusCode: http.StatusOK, body: `{"choices": [{"message": {"content": 1}}]}`, 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.statusCode < http.StatusOK || tc.statusCode >= http.StatusMultipleChoices { var providerHTTPError *ProviderHTTPError if !errors.As(err, &providerHTTPError) { t.Fatalf("error = %T, want *ProviderHTTPError", err) } if got := providerHTTPError.StatusCode(); got != tc.statusCode { t.Fatalf("provider status = %d, want %d", got, tc.statusCode) } } 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) } }