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