1034 lines
34 KiB
Go
1034 lines
34 KiB
Go
package llm
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"math"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
|
)
|
|
|
|
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 TestOpenAICompatibleClientExplicitZeroTimeoutDisablesClientTimeout(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
time.Sleep(20 * time.Millisecond)
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
|
BaseURL: ts.URL + "/v1",
|
|
Timeout: time.Nanosecond,
|
|
})
|
|
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", TimeoutSeconds: 0},
|
|
TargetPresence: domain.ExecutionTargetPresence{TimeoutSeconds: true},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("expected explicit zero timeout to disable client timeout, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientOmittedTimeoutUsesClientTimeout(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
time.Sleep(20 * time.Millisecond)
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
|
BaseURL: ts.URL + "/v1",
|
|
Timeout: time.Nanosecond,
|
|
})
|
|
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", TimeoutSeconds: 0},
|
|
})
|
|
if err == nil {
|
|
t.Fatal("expected omitted timeout to use client timeout")
|
|
}
|
|
if !errors.Is(err, ErrRequestFailed) {
|
|
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
|
}
|
|
}
|
|
|
|
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) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
_, _ = w.Write([]byte(`{"error":"bad request payload"}`))
|
|
}))
|
|
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(), "400") || !strings.Contains(err.Error(), "bad request payload") {
|
|
t.Fatalf("expected status/body details, 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 TestOpenAICompatibleClientTimeout(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
time.Sleep(250 * time.Millisecond)
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
|
BaseURL: ts.URL + "/v1",
|
|
Model: "m",
|
|
Timeout: 50 * time.Millisecond,
|
|
})
|
|
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 timeout error")
|
|
}
|
|
if !errors.Is(err, ErrRequestFailed) {
|
|
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientRequestTimeoutOverride(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
time.Sleep(100 * time.Millisecond)
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
|
BaseURL: ts.URL + "/v1",
|
|
Model: "m",
|
|
Timeout: 50 * time.Millisecond,
|
|
})
|
|
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{TimeoutSeconds: 1},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("expected request-level timeout override to succeed, got %v", err)
|
|
}
|
|
if resp.Content != "ok" {
|
|
t.Fatalf("expected response content ok, got %q", resp.Content)
|
|
}
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|