From 7e94ab133bd73b3c91894c13aa610c741e8ba00f Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 28 Jul 2026 04:29:17 +0000 Subject: [PATCH] Add OpenAI-compatible model client --- docs/integrations/openai-compatible-chat.md | 86 ++ docs/internal/llm.md | 53 + docs/internal/overview.md | 7 +- docs/policy/architecture.md | 13 +- internal/llm/client.go | 11 + internal/llm/openai_compatible_client.go | 385 ++++++ internal/llm/openai_compatible_client_test.go | 1192 +++++++++++++++++ 7 files changed, 1739 insertions(+), 8 deletions(-) create mode 100644 docs/integrations/openai-compatible-chat.md create mode 100644 docs/internal/llm.md create mode 100644 internal/llm/client.go create mode 100644 internal/llm/openai_compatible_client.go create mode 100644 internal/llm/openai_compatible_client_test.go diff --git a/docs/integrations/openai-compatible-chat.md b/docs/integrations/openai-compatible-chat.md new file mode 100644 index 0000000..a6985bd --- /dev/null +++ b/docs/integrations/openai-compatible-chat.md @@ -0,0 +1,86 @@ +# OpenAI-Compatible Chat Integration + +## Purpose + +This document defines the outbound HTTP behavior implemented by Promptkit's +internal OpenAI-compatible model client. The +[internal model-client document](../internal/llm.md) owns implementation flow, +errors, and test ownership. The client is not yet available through a usable +public Promptkit engine. + +## Endpoint And Method + +Generation sends an HTTP `POST` with `Content-Type: application/json`. +A non-empty endpoint from the execution target overrides the client's +configured base URL. After trailing slashes are removed, +`/chat/completions` is appended. Generation fails before sending when neither +source supplies an endpoint. + +## Authentication + +A non-empty API key supplied directly on the execution target takes +precedence. Otherwise, when an API-key environment-variable name is supplied, +the client reads that variable and requires a non-empty value. The selected +key is sent as `Authorization: Bearer `. No authorization header is sent +when neither mechanism is configured. + +## Request Body + +The request body always contains `model` and `messages`. The execution +target's model takes precedence over the client's configured model, and one +must be available. + +Each ordinary message contains its `role` and string `content`. A +cache-controlled message instead uses a text content block containing `type`, +`text`, and `cache_control`; an empty cache-control TTL is omitted. + +A non-empty session ID is trimmed, checked against the internal domain limit, +and sent as top-level `session_id`. It is not sent as a session header. + +The client conditionally includes: + +- `temperature`, `max_tokens`, and `top_p` when non-zero or explicitly + present; +- non-empty `service_tier` and `reasoning_effort`; and +- `response_format` for JSON Schema structured output, including its name, + strict flag, and schema document. + +Extra parameters are merged directly into the top-level body after JSON +serialization is verified. Empty keys and collisions with these reserved +fields are rejected before any provider call: + +- `model` +- `session_id` +- `messages` +- `temperature` +- `max_tokens` +- `top_p` +- `service_tier` +- `reasoning_effort` +- `response_format` + +## Response Handling + +Any 2xx response is decoded as an OpenAI-compatible chat response. The client +returns the first choice's non-empty message content and maps prompt, +completion, total, cached, and cache-write token counts. + +Invalid JSON, absent choices, and empty first-choice content are malformed +responses. For a non-2xx status, the error includes the status code but never +the provider response body. + +## Timeout And Cancellation + +Timeouts are layered: + +- the caller context remains the outer cancellation boundary; +- a positive generation timeout adds a request context deadline; +- zero adds no generation-specific deadline; +- a negative generation timeout is invalid; and +- the cloned `http.Client` supplies the whole-request transport cap, retaining + a positive supplied-client timeout or applying the configured/default + timeout when the supplied value is not positive. + +The earliest applicable caller, generation, or transport deadline controls the +request. Constructing the internal client does not mutate a supplied +`http.Client`. diff --git a/docs/internal/llm.md b/docs/internal/llm.md new file mode 100644 index 0000000..49b36cc --- /dev/null +++ b/docs/internal/llm.md @@ -0,0 +1,53 @@ +# Internal Model Client + +## Purpose + +This document describes Promptkit's internal model-client implementation. The +[architecture policy](../policy/architecture.md) owns the library boundary, +and the +[OpenAI-compatible chat integration](../integrations/openai-compatible-chat.md) +owns the observable outbound HTTP contract. + +The client is implemented only under `internal/llm`. The root package does not +yet assemble it into a usable public engine. + +## Components And Flow + +`Client` is the provider-neutral generation boundary consumed by later +orchestration. `OpenAICompatibleClient` is the built-in implementation. It +uses internal domain values for rendered prompts, execution targets, +structured output, responses, and token usage. + +Construction validates the configured base URL and clones any supplied +`http.Client` so Promptkit can apply its timeout default without mutating the +caller's client. Generation then: + +1. validates request-level timeout and endpoint requirements; +2. maps the internal request into the OpenAI-compatible chat payload; +3. validates and merges extra parameters; +4. resolves authentication; +5. performs the outbound request under the applicable deadlines; and +6. decodes the first response choice and token usage. + +The implementation has no retry loop, tool-call support, provider catalog, +inbound HTTP behavior, or durable session store. + +## Failure Categories + +The package preserves distinct error identities for invalid client +configuration, invalid generation requests, request execution failures, +non-success provider statuses, and malformed successful responses. Provider +response bodies are not included in non-success errors. + +Caller cancellation and deadline failures during the outbound request are +reported as request execution failures. The future runner can classify these +identities without depending on HTTP status mapping. + +## Test Ownership + +The +[OpenAI-compatible client tests](../../internal/llm/openai_compatible_client_test.go) +own configuration, client cloning, deterministic deadline precedence, +authentication, request and response mapping, malformed data, error identity, +cancellation, and response-body suppression. They use local test servers and +test transports; the default suite makes no live or paid provider requests. diff --git a/docs/internal/overview.md b/docs/internal/overview.md index 924d49e..ec1b608 100644 --- a/docs/internal/overview.md +++ b/docs/internal/overview.md @@ -21,10 +21,11 @@ contributor workflow and validation. | `internal/prompt` | Renders prompt messages from Go templates with artifact, variable, session, and cache-control data. | [Go-template renderer](../../internal/prompt/go_renderer.go) | | `internal/artifact` | Resolves ordinary inline and unrestricted caller-selected file references into copied artifacts with metadata and hashes. | [Internal sources and validation](sources.md) | | `internal/validate` | Validates basic, JSON, and JSON Schema output using operating-system filesystem or `fs.FS` schema sources. | [Internal sources and validation](sources.md) | +| `internal/llm` | Defines the internal generation boundary and implements outbound OpenAI-compatible chat requests, response decoding, authentication, and deadline handling. | [Internal model client](llm.md) | -These packages provide the internal model, source, and rendering foundation. -Model clients, orchestration, and a usable public engine are not implemented in -Promptkit yet. +These packages provide the internal model, source, rendering, validation, and +model-client foundation. Orchestration and a usable public engine are not +implemented in Promptkit yet. ## Maintenance diff --git a/docs/policy/architecture.md b/docs/policy/architecture.md index 4bf5ec3..3fe15bb 100644 --- a/docs/policy/architecture.md +++ b/docs/policy/architecture.md @@ -21,7 +21,7 @@ The implemented internal components consist of: - `internal/domain`, which owns framework data values shared by later internal components; - `internal/defaults`, which owns application-neutral framework defaults and - constructs the default execution target; and + constructs the default execution target; - `internal/filecatalog`, which discovers YAML files and provides source-path helpers for filesystem and `fs.FS` consumers; - `internal/promptdef`, which loads and validates prompt definitions from @@ -32,17 +32,20 @@ The implemented internal components consist of: catalog; - `internal/prompt`, which renders prompt messages from Go templates; - `internal/artifact`, which resolves ordinary inline and unrestricted - caller-selected file references; and + caller-selected file references; - `internal/validate`, which validates basic, JSON, and JSON Schema output - using filesystem and `fs.FS` schema sources. + using filesystem and `fs.FS` schema sources; and +- `internal/llm`, which defines the provider-neutral generation boundary and + implements outbound OpenAI-compatible chat requests. The defaults and renderer depend on the domain model. Prompt-definition and profile repositories use the domain model, file catalog, and YAML decoder. The built-in profile repository supplies an embedded `fs.FS` to the profile package. Artifact reading uses the domain model and application-neutral defaults. Validation uses the domain model, file catalog, and JSON Schema -implementation. Model clients, orchestration, and the public engine have not -yet been extracted. +implementation. The model client uses the domain model, application-neutral +defaults, and an injected or standard-library HTTP client. Orchestration and +the public engine have not yet been extracted. Future framework extraction must follow this dependency direction: diff --git a/internal/llm/client.go b/internal/llm/client.go new file mode 100644 index 0000000..6f24f93 --- /dev/null +++ b/internal/llm/client.go @@ -0,0 +1,11 @@ +package llm + +import ( + "context" + "gitea.maximumdirect.net/eric/promptkit/internal/domain" +) + +// Client executes a rendered prompt against an LLM endpoint. +type Client interface { + Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) +} diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go new file mode 100644 index 0000000..2e9f230 --- /dev/null +++ b/internal/llm/openai_compatible_client.go @@ -0,0 +1,385 @@ +package llm + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "os" + "strings" + "time" + "unicode/utf8" + + "gitea.maximumdirect.net/eric/promptkit/internal/defaults" + "gitea.maximumdirect.net/eric/promptkit/internal/domain" +) + +var ( + ErrInvalidConfig = errors.New("invalid llm client configuration") + ErrInvalidRequest = errors.New("invalid generate request") + ErrRequestFailed = errors.New("llm request failed") + ErrUnexpectedStatus = errors.New("llm returned non-success status") + ErrMalformedResponse = errors.New("malformed llm response") +) + +type OpenAICompatibleConfig struct { + BaseURL string + Model string + Timeout time.Duration + HTTPClient *http.Client +} + +type OpenAICompatibleClient struct { + baseURL string + defaultModel string + httpClient *http.Client +} + +func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { + baseURL := strings.TrimSpace(cfg.BaseURL) + if baseURL != "" { + if _, err := url.ParseRequestURI(baseURL); err != nil { + return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) + } + } + + timeout := cfg.Timeout + if timeout <= 0 { + timeout = defaults.LLMRequestTimeoutDefault + } + + var client *http.Client + if cfg.HTTPClient != nil { + cloned := *cfg.HTTPClient + if cloned.Timeout <= 0 { + cloned.Timeout = timeout + } + client = &cloned + } else { + client = &http.Client{Timeout: timeout} + } + + return &OpenAICompatibleClient{ + baseURL: strings.TrimRight(baseURL, "/"), + defaultModel: cfg.Model, + httpClient: client, + }, nil +} + +func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) { + if req.Target.TimeoutSeconds < 0 { + return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest) + } + + endpoint := strings.TrimSpace(req.Target.Endpoint) + if endpoint == "" { + endpoint = c.baseURL + } + if endpoint == "" { + return nil, fmt.Errorf("%w: endpoint is required", ErrInvalidRequest) + } + endpoint = strings.TrimRight(endpoint, "/") + defaults.OpenAIChatCompletionsPath + + wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) + } + + wirePayload, err := openAIChatRequestPayload(wireReq) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) + } + + payload, err := json.Marshal(wirePayload) + if err != nil { + return nil, fmt.Errorf("%w: failed to encode request: %v", ErrRequestFailed, err) + } + + requestContext := ctx + if req.Target.TimeoutSeconds > 0 { + var cancel context.CancelFunc + requestContext, cancel = context.WithTimeout( + ctx, + time.Duration(req.Target.TimeoutSeconds)*time.Second, + ) + defer cancel() + } + + httpReq, err := http.NewRequestWithContext(requestContext, http.MethodPost, endpoint, bytes.NewReader(payload)) + if err != nil { + return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err) + } + httpReq.Header.Set("Content-Type", "application/json") + if apiKey := strings.TrimSpace(req.Target.APIKey); apiKey != "" { + httpReq.Header.Set("Authorization", "Bearer "+apiKey) + } else if envName := strings.TrimSpace(req.Target.APIKeyEnv); envName != "" { + apiKey := strings.TrimSpace(os.Getenv(envName)) + if apiKey == "" { + return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName) + } + httpReq.Header.Set("Authorization", "Bearer "+apiKey) + } + + httpClient := c.httpClient + if httpClient == nil { + httpClient = &http.Client{Timeout: defaults.LLMRequestTimeoutDefault} + } + + httpResp, err := httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err) + } + defer httpResp.Body.Close() + + if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { + _, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096)) + return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode) + } + + var wireResp openAIChatResponse + if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil { + return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err) + } + + if len(wireResp.Choices) == 0 { + return nil, fmt.Errorf("%w: no choices returned", ErrMalformedResponse) + } + content := wireResp.Choices[0].Message.Content + if content == "" { + return nil, fmt.Errorf("%w: first choice has empty message content", ErrMalformedResponse) + } + + return &domain.GenerateResponse{ + Content: content, + Usage: domain.TokenUsage{ + PromptTokens: wireResp.Usage.PromptTokens, + CompletionTokens: wireResp.Usage.CompletionTokens, + TotalTokens: wireResp.Usage.TotalTokens, + CachedTokens: wireResp.Usage.PromptTokensDetails.CachedTokens, + CacheWriteTokens: wireResp.Usage.CacheWriteTokens, + }, + }, nil +} + +func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) { + model := strings.TrimSpace(req.Target.Model) + if model == "" { + model = strings.TrimSpace(defaultModel) + } + if model == "" { + return openAIChatRequest{}, errors.New("model is required") + } + + wireReq := openAIChatRequest{ + Model: model, + } + if sessionID := strings.TrimSpace(req.Prompt.SessionID); sessionID != "" { + if n := utf8.RuneCountInString(sessionID); n > domain.SessionIDMaxLength { + return openAIChatRequest{}, fmt.Errorf("session_id length %d exceeds maximum %d", n, domain.SessionIDMaxLength) + } + wireReq.SessionID = sessionID + } + + wireReq.Messages = make([]openAIChatRequestMessage, 0, len(req.Prompt.Messages)) + for _, msg := range req.Prompt.Messages { + wireReq.Messages = append(wireReq.Messages, openAIChatRequestMessageFromRenderedMessage(msg)) + } + + if req.Target.Temperature != 0 || req.TargetPresence.Temperature { + wireReq.Temperature = &req.Target.Temperature + } + if req.Target.MaxTokens != 0 || req.TargetPresence.MaxTokens { + wireReq.MaxTokens = &req.Target.MaxTokens + } + if req.Target.TopP != 0 || req.TargetPresence.TopP { + wireReq.TopP = &req.Target.TopP + } + if strings.TrimSpace(req.Target.ServiceTier) != "" { + wireReq.ServiceTier = req.Target.ServiceTier + } + if strings.TrimSpace(req.Target.ReasoningEffort) != "" { + wireReq.ReasoningEffort = req.Target.ReasoningEffort + } + if len(req.Target.ExtraParams) > 0 { + wireReq.ExtraParams = req.Target.ExtraParams + } + if req.StructuredOutput != nil { + responseFormat, err := toOpenAIResponseFormat(req.StructuredOutput) + if err != nil { + return openAIChatRequest{}, err + } + wireReq.ResponseFormat = responseFormat + } + + return wireReq, nil +} + +type openAIChatRequest struct { + Model string `json:"model"` + SessionID string `json:"session_id,omitempty"` + Messages []openAIChatRequestMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + ServiceTier string `json:"service_tier,omitempty"` + ReasoningEffort string `json:"reasoning_effort,omitempty"` + ResponseFormat *openAIResponseFormat `json:"response_format,omitempty"` + ExtraParams map[string]any `json:"-"` +} + +func openAIChatRequestPayload(req openAIChatRequest) (map[string]any, error) { + out := map[string]any{ + "model": req.Model, + "messages": req.Messages, + } + if req.SessionID != "" { + out["session_id"] = req.SessionID + } + if req.Temperature != nil { + out["temperature"] = *req.Temperature + } + if req.MaxTokens != nil { + out["max_tokens"] = *req.MaxTokens + } + if req.TopP != nil { + out["top_p"] = *req.TopP + } + if req.ServiceTier != "" { + out["service_tier"] = req.ServiceTier + } + if req.ReasoningEffort != "" { + out["reasoning_effort"] = req.ReasoningEffort + } + if req.ResponseFormat != nil { + out["response_format"] = req.ResponseFormat + } + + for key, value := range req.ExtraParams { + if key == "" { + return nil, errors.New("extra_params key must not be empty") + } + if _, reserved := reservedOpenAIChatRequestFields[key]; reserved { + return nil, fmt.Errorf("extra_params key %q collides with reserved request field", key) + } + if _, err := json.Marshal(value); err != nil { + return nil, fmt.Errorf("extra_params.%s must be JSON-serializable: %w", key, err) + } + out[key] = value + } + + return out, nil +} + +var reservedOpenAIChatRequestFields = map[string]struct{}{ + "model": {}, + "session_id": {}, + "messages": {}, + "temperature": {}, + "max_tokens": {}, + "top_p": {}, + "service_tier": {}, + "reasoning_effort": {}, + "response_format": {}, +} + +type openAIChatRequestMessage struct { + Role string `json:"role"` + Content any `json:"content"` +} + +type openAIChatTextContentBlock struct { + Type string `json:"type"` + Text string `json:"text"` + CacheControl *openAICacheControl `json:"cache_control,omitempty"` +} + +type openAICacheControl struct { + Type string `json:"type"` + TTL string `json:"ttl,omitempty"` +} + +type openAIChatResponseMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type openAIChatResponse struct { + Choices []struct { + Message openAIChatResponseMessage `json:"message"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + PromptTokensDetails struct { + CachedTokens int `json:"cached_tokens"` + } `json:"prompt_tokens_details"` + CacheWriteTokens int `json:"cache_write_tokens"` + } `json:"usage"` +} + +type openAIResponseFormat struct { + Type string `json:"type"` + JSONSchema *openAIJSONSchemaEnvelope `json:"json_schema,omitempty"` +} + +type openAIJSONSchemaEnvelope struct { + Name string `json:"name"` + Strict bool `json:"strict"` + Schema any `json:"schema"` +} + +func openAIChatRequestMessageFromRenderedMessage(msg domain.RenderedMessage) openAIChatRequestMessage { + wireMsg := openAIChatRequestMessage{ + Role: msg.Role, + Content: msg.Content, + } + if msg.CacheControl == nil { + return wireMsg + } + + wireMsg.Content = []openAIChatTextContentBlock{ + { + Type: "text", + Text: msg.Content, + CacheControl: &openAICacheControl{ + Type: string(msg.CacheControl.Type), + TTL: msg.CacheControl.TTL, + }, + }, + } + return wireMsg +} + +func toOpenAIResponseFormat(spec *domain.StructuredOutputSpec) (*openAIResponseFormat, error) { + if spec == nil { + return nil, nil + } + + switch spec.Type { + case domain.StructuredOutputJSONSchema: + if spec.JSONSchema == nil { + return nil, errors.New("json_schema structured output requires schema payload") + } + if strings.TrimSpace(spec.JSONSchema.Name) == "" { + return nil, errors.New("json_schema structured output requires non-empty schema name") + } + if spec.JSONSchema.Schema == nil { + return nil, errors.New("json_schema structured output requires schema document") + } + return &openAIResponseFormat{ + Type: "json_schema", + JSONSchema: &openAIJSONSchemaEnvelope{ + Name: spec.JSONSchema.Name, + Strict: spec.JSONSchema.Strict, + Schema: spec.JSONSchema.Schema, + }, + }, nil + default: + return nil, fmt.Errorf("unsupported structured output type %q", spec.Type) + } +} diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go new file mode 100644 index 0000000..d363837 --- /dev/null +++ b/internal/llm/openai_compatible_client_test.go @@ -0,0 +1,1192 @@ +package llm + +import ( + "context" + "encoding/json" + "errors" + "math" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "gitea.maximumdirect.net/eric/promptkit/internal/domain" +) + +var errTransportStopped = errors.New("transport stopped after request inspection") + +type deadlineCapturingTransport struct { + deadline time.Time + hasDeadline bool +} + +func (t *deadlineCapturingTransport) RoundTrip(req *http.Request) (*http.Response, error) { + t.deadline, t.hasDeadline = req.Context().Deadline() + return nil, errTransportStopped +} + +type contextErrorTransport struct{} + +func (contextErrorTransport) RoundTrip(req *http.Request) (*http.Response, error) { + return nil, req.Context().Err() +} + +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) { + _, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "://invalid", + }) + if err == nil { + t.Fatal("expected invalid configuration error") + } + 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 TestOpenAICompatibleClientGenerateSuccess(t *testing.T) { + type observedRequest struct { + Authorization string + Body map[string]any + } + 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) + } + + 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(`{ + "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", + Timeout: 2 * time.Second, + }) + if err != nil { + t.Fatalf("unexpected constructor error: %v", err) + } + t.Setenv("SCRIPTORIUM_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: "SCRIPTORIUM_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) + } + + if obs.Authorization != "Bearer secret-key" { + t.Fatalf("unexpected Authorization header: %q", obs.Authorization) + } + if got, ok := obs.Body["model"].(string); !ok || got != "gpt-test" { + t.Fatalf("unexpected model payload: %#v", obs.Body["model"]) + } + if got, ok := obs.Body["temperature"].(float64); !ok || got != 0.4 { + t.Fatalf("unexpected temperature payload: %#v", obs.Body["temperature"]) + } + 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"]) + } + + msgs, ok := obs.Body["messages"].([]any) + if !ok || len(msgs) != 2 { + t.Fatalf("unexpected messages payload: %#v", obs.Body["messages"]) + } + msg0 := msgs[0].(map[string]any) + if msg0["role"] != "system" || msg0["content"] != "You are helpful." { + t.Fatalf("unexpected first message: %#v", msg0) + } + msg1 := msgs[1].(map[string]any) + if msg1["role"] != "user" || msg1["content"] != "Say hello" { + t.Fatalf("unexpected second message: %#v", msg1) + } + + responseFormat, ok := obs.Body["response_format"].(map[string]any) + if !ok { + t.Fatalf("expected response_format payload, got %#v", obs.Body["response_format"]) + } + if responseFormat["type"] != "json_schema" { + t.Fatalf("expected response_format.type=json_schema, got %#v", responseFormat["type"]) + } + jsonSchema, ok := responseFormat["json_schema"].(map[string]any) + if !ok { + t.Fatalf("expected response_format.json_schema map, got %#v", responseFormat["json_schema"]) + } + if jsonSchema["name"] != "weather_schema" { + t.Fatalf("expected json_schema.name weather_schema, got %#v", jsonSchema["name"]) + } + if jsonSchema["strict"] != true { + t.Fatalf("expected json_schema.strict=true, got %#v", jsonSchema["strict"]) + } + if _, ok := jsonSchema["schema"].(map[string]any); !ok { + t.Fatalf("expected json_schema.schema object, got %#v", jsonSchema["schema"]) + } +} + +func TestOpenAICompatibleClientDirectAPIKeyPreferredOverEnv(t *testing.T) { + const directKey = "direct-llm-key" + t.Setenv("SCRIPTORIUM_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: "SCRIPTORIUM_TEST_API_KEY", + APIKey: directKey, + }, + }) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if gotAuth != "Bearer "+directKey { + t.Fatalf("unexpected Authorization header: %q", gotAuth) + } +} + +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() + + 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: "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) + } + + 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]) + } + } + + msgs, ok := observedBody["messages"].([]any) + if !ok || len(msgs) != 2 { + t.Fatalf("unexpected messages payload: %#v", observedBody["messages"]) + } + msg0 := msgs[0].(map[string]any) + if msg0["role"] != "system" { + t.Fatalf("unexpected first message role: %#v", msg0["role"]) + } + contentBlocks, ok := msg0["content"].([]any) + if !ok || len(contentBlocks) != 1 { + t.Fatalf("expected first message content block array, got %#v", msg0["content"]) + } + block := contentBlocks[0].(map[string]any) + 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 cacheControl["type"] != string(domain.CacheControlEphemeral) || cacheControl["ttl"] != "1h" { + t.Fatalf("unexpected cache_control payload: %#v", cacheControl) + } + + 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) + } +} + +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() + + 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: "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) + } + + 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) + } + 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) + } + + _, 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) + } +} + +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"]) + } +} + +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 TestOpenAICompatibleClientParsesCacheUsage(t *testing.T) { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{ + "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 + } +}`)) + })) + defer ts.Close() + + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "model"}) + if err != nil { + t.Fatal(err) + } + + 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 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() + + 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["response_format"]; exists { + t.Fatalf("expected response_format omitted, got %#v", observedBody["response_format"]) + } + if _, exists := observedBody["service_tier"]; exists { + t.Fatalf("expected service_tier omitted, got %#v", observedBody["service_tier"]) + } +} + +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() + + 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", + 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) + } + if observedBody["reasoning_effort"] != "high" { + t.Fatalf("expected reasoning_effort high, got %#v", observedBody["reasoning_effort"]) + } + if observedBody["string_value"] != "on" { + t.Fatalf("unexpected string extra param: %#v", observedBody["string_value"]) + } + if observedBody["number_value"] != float64(42) { + t.Fatalf("unexpected number extra param: %#v", observedBody["number_value"]) + } + if observedBody["boolean_value"] != true { + t.Fatalf("unexpected boolean extra param: %#v", observedBody["boolean_value"]) + } + 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"]) + } + if _, exists := observedBody["extra_params"]; exists { + t.Fatalf("expected extra_params wrapper omitted, got %#v", observedBody["extra_params"]) + } + 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"]) + } +} + +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, + }, + }) + 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"]) + } +} + +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) + } + + _, 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) + } + } +} + +func TestOpenAICompatibleClientOmittedTimeoutUsesClientTimeout(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 !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) { + 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() + + 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", 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 called { + t.Fatal("provider should not be called for invalid extra_params") + } + }) + } +} + +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() + + 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: "SCRIPTORIUM_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{ + 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) + } +} + +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) + } + + 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"}, + }) + 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 defaultHit { + t.Fatal("default endpoint should not have been called") + } + if !overrideHit { + t.Fatal("override endpoint should have been called") + } +} + +func TestOpenAICompatibleClientNon2xxError(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) + } + + _, 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) + } +} + +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 TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T) { + transport := &deadlineCapturingTransport{} + 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 !transport.hasDeadline { + t.Fatal("expected generation timeout to set a transport deadline") + } + assertDeadlineNear(t, transport.deadline, before, after, generationTimeout) +} + +func TestOpenAICompatibleClientCallerDeadlineTakesPrecedence(t *testing.T) { + transport := &deadlineCapturingTransport{} + 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 !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 TestOpenAICompatibleClientCancellationReturnsRequestFailure(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 !strings.Contains(err.Error(), context.Canceled.Error()) { + t.Fatalf("expected cancellation detail, got %v", err) + } +} + +func TestOpenAICompatibleClientNegativeTimeoutRejected(t *testing.T) { + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "http://example.com/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"}}}, + Target: domain.ExecutionTarget{TimeoutSeconds: -1}, + }) + if err == nil { + t.Fatal("expected invalid request error") + } + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } +} + +func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) { + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "", + Model: "m", + }) + if err != nil { + t.Fatalf("expected empty configured base URL to be allowed, got %v", err) + } + + _, err = client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + Target: domain.ExecutionTarget{Endpoint: "http://localhost:9999/v1"}, + }) + if err == nil { + t.Fatal("expected request failure due to unreachable endpoint") + } + if !errors.Is(err, ErrRequestFailed) { + t.Fatalf("expected ErrRequestFailed with request endpoint override, got %v", err) + } +} + +func TestOpenAICompatibleClientRequiresEndpointWhenUnsetEverywhere(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) + } +}