diff --git a/internal/providers/azure/azure_test.go b/internal/providers/azure/azure_test.go index 6cc178bfd..7662651fa 100644 --- a/internal/providers/azure/azure_test.go +++ b/internal/providers/azure/azure_test.go @@ -2,133 +2,96 @@ package azure import ( "context" + "io" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_UsesAzureAuthAndDefaultAPIVersion(t *testing.T) { - var gotPath string - var gotAPIVersion string - var gotAPIKey string - var gotAuthorization string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAPIVersion = r.URL.Query().Get("api-version") - gotAPIKey = r.Header.Get("api-key") - gotAuthorization = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "gpt-4o", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "hello"}, - "finish_reason": "stop" - }] - }`)) - })) - defer server.Close() +const chatCompletionJSON = `{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop" + }] +}` + +// newTestProvider points a provider at the given upstream base URL. +func newTestProvider(client *http.Client, baseURL string) *Provider { + provider := NewWithHTTPClient("test-api-key", client, llmclient.Hooks{}) + provider.SetBaseURL(baseURL) + return provider +} - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) +func TestChatCompletion_UsesAzureAuthAndDefaultAPIVersion(t *testing.T) { + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.Client(), server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "gpt-4o", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: "gpt-4o", + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAPIKey != "test-api-key" { - t.Fatalf("api-key = %q, want test-api-key", gotAPIKey) - } - if gotAuthorization != "" { - t.Fatalf("authorization = %q, want empty", gotAuthorization) - } - if gotAPIVersion != defaultAPIVersion { - t.Fatalf("api-version = %q, want %q", gotAPIVersion, defaultAPIVersion) - } + require.NoError(t, err) + + sent := capture.Last(t) + assert.Equal(t, "/chat/completions", sent.Path) + assert.Equal(t, "test-api-key", sent.Header.Get("api-key")) + assert.Empty(t, sent.Header.Get("Authorization")) + assert.Equal(t, defaultAPIVersion, sent.Query.Get("api-version")) } func TestSetAPIVersion_OverridesDefault(t *testing.T) { - var gotAPIVersion string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotAPIVersion = r.URL.Query().Get("api-version") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "gpt-4o", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "hello"}, - "finish_reason": "stop" - }] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.Client(), server.URL) provider.SetAPIVersion("2025-04-01-preview") _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "gpt-4o", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: "gpt-4o", + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotAPIVersion != "2025-04-01-preview" { - t.Fatalf("api-version = %q, want 2025-04-01-preview", gotAPIVersion) - } + require.NoError(t, err) + assert.Equal(t, "2025-04-01-preview", capture.Last(t).Query.Get("api-version")) } func TestListModels_UsesAzureOpenAIPath(t *testing.T) { - var gotPath string - var gotAPIVersion string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAPIVersion = r.URL.Query().Get("api-version") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) + provider := newTestProvider(server.Client(), server.URL) _, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotPath != "/openai/models" { - t.Fatalf("path = %q, want /openai/models", gotPath) - } - if gotAPIVersion != defaultAPIVersion { - t.Fatalf("api-version = %q, want %q", gotAPIVersion, defaultAPIVersion) - } + require.NoError(t, err) + + sent := capture.Last(t) + assert.Equal(t, "/openai/models", sent.Path) + assert.Equal(t, defaultAPIVersion, sent.Query.Get("api-version")) } -func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { - tests := []struct { +// batchEndpointCases exercises every batch surface; each case names the Azure +// path and method the shared OpenAI batch adapter must produce. +func batchEndpointCases() []struct { + name string + call func(*Provider) error + wantPath string + wantMethod string + responseBody string +} { + const batchJSON = `{ + "id":"batch_123", + "object":"batch", + "endpoint":"/v1/chat/completions", + "status":"validating", + "created_at":1677652288, + "request_counts":{"total":1,"completed":0,"failed":0} + }` + return []struct { name string call func(*Provider) error wantPath string @@ -145,16 +108,9 @@ func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { }) return err }, - wantPath: "/openai/batches", - wantMethod: http.MethodPost, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"validating", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, + wantPath: "/openai/batches", + wantMethod: http.MethodPost, + responseBody: batchJSON, }, { name: "get", @@ -162,16 +118,9 @@ func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { _, err := p.GetBatch(context.Background(), "batch_123") return err }, - wantPath: "/openai/batches/batch_123", - wantMethod: http.MethodGet, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"validating", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, + wantPath: "/openai/batches/batch_123", + wantMethod: http.MethodGet, + responseBody: batchJSON, }, { name: "list", @@ -179,13 +128,9 @@ func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { _, err := p.ListBatches(context.Background(), 10, "batch_122") return err }, - wantPath: "/openai/batches", - wantMethod: http.MethodGet, - responseBody: `{ - "object":"list", - "data":[], - "has_more":false - }`, + wantPath: "/openai/batches", + wantMethod: http.MethodGet, + responseBody: `{"object":"list","data":[],"has_more":false}`, }, { name: "cancel", @@ -193,226 +138,75 @@ func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { _, err := p.CancelBatch(context.Background(), "batch_123") return err }, - wantPath: "/openai/batches/batch_123/cancel", - wantMethod: http.MethodPost, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"cancelling", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, + wantPath: "/openai/batches/batch_123/cancel", + wantMethod: http.MethodPost, + responseBody: batchJSON, }, } +} - for _, tt := range tests { +func TestBatchEndpoints_UseAzureOpenAIPaths(t *testing.T) { + for _, tt := range batchEndpointCases() { t.Run(tt.name, func(t *testing.T) { - var gotPath string - var gotMethod string - var gotAPIVersion string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotMethod = r.Method - gotAPIVersion = r.URL.Query().Get("api-version") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, tt.responseBody) + provider := newTestProvider(server.Client(), server.URL) - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + require.NoError(t, tt.call(provider)) - if err := tt.call(provider); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } - if gotMethod != tt.wantMethod { - t.Fatalf("method = %q, want %q", gotMethod, tt.wantMethod) - } - if gotAPIVersion != defaultAPIVersion { - t.Fatalf("api-version = %q, want %q", gotAPIVersion, defaultAPIVersion) - } + sent := capture.Last(t) + assert.Equal(t, tt.wantPath, sent.Path) + assert.Equal(t, tt.wantMethod, sent.Method) + assert.Equal(t, defaultAPIVersion, sent.Query.Get("api-version")) }) } } func TestListModels_UsesAzureResourceRootForDeploymentScopedBaseURL(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/openai/deployments/gpt-4o") + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) + provider := newTestProvider(server.Client(), server.URL+"/openai/deployments/gpt-4o") _, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotPath != "/openai/models" { - t.Fatalf("path = %q, want /openai/models", gotPath) - } + require.NoError(t, err) + assert.Equal(t, "/openai/models", capture.Last(t).Path) } func TestBatchEndpoints_UseAzureResourceRootForDeploymentScopedBaseURL(t *testing.T) { - tests := []struct { - name string - call func(*Provider) error - wantPath string - wantMethod string - responseBody string - }{ - { - name: "create", - call: func(p *Provider) error { - _, err := p.CreateBatch(context.Background(), &core.BatchRequest{ - InputFileID: "file-123", - Endpoint: "/v1/chat/completions", - CompletionWindow: "24h", - }) - return err - }, - wantPath: "/openai/batches", - wantMethod: http.MethodPost, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"validating", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, - }, - { - name: "get", - call: func(p *Provider) error { - _, err := p.GetBatch(context.Background(), "batch_123") - return err - }, - wantPath: "/openai/batches/batch_123", - wantMethod: http.MethodGet, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"validating", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, - }, - { - name: "list", - call: func(p *Provider) error { - _, err := p.ListBatches(context.Background(), 10, "batch_122") - return err - }, - wantPath: "/openai/batches", - wantMethod: http.MethodGet, - responseBody: `{ - "object":"list", - "data":[], - "has_more":false - }`, - }, - { - name: "cancel", - call: func(p *Provider) error { - _, err := p.CancelBatch(context.Background(), "batch_123") - return err - }, - wantPath: "/openai/batches/batch_123/cancel", - wantMethod: http.MethodPost, - responseBody: `{ - "id":"batch_123", - "object":"batch", - "endpoint":"/v1/chat/completions", - "status":"cancelling", - "created_at":1677652288, - "request_counts":{"total":1,"completed":0,"failed":0} - }`, - }, - } - - for _, tt := range tests { + for _, tt := range batchEndpointCases() { t.Run(tt.name, func(t *testing.T) { - var gotPath string - var gotMethod string + server, capture := providertest.JSONServer(t, http.StatusOK, tt.responseBody) + provider := newTestProvider(server.Client(), server.URL+"/openai/deployments/gpt-4o") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotMethod = r.Method - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + require.NoError(t, tt.call(provider)) - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/openai/deployments/gpt-4o") - - if err := tt.call(provider); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } - if gotMethod != tt.wantMethod { - t.Fatalf("method = %q, want %q", gotMethod, tt.wantMethod) - } + sent := capture.Last(t) + assert.Equal(t, tt.wantPath, sent.Path) + assert.Equal(t, tt.wantMethod, sent.Method) }) } } func TestGetBatchResults_UsesAzureResourceRootForDeploymentScopedBaseURL(t *testing.T) { - var gotPaths []string - var gotVersions []string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPaths = append(gotPaths, r.URL.Path) - gotVersions = append(gotVersions, r.URL.Query().Get("api-version")) - - switch r.URL.Path { - case "/openai/batches/batch_1": + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/openai/batches/batch_1": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"batch_1","status":"completed","output_file_id":"file_1","endpoint":"/v1/chat/completions"}`)) - case "/openai/files/file_1/content": + _, _ = io.WriteString(w, `{"id":"batch_1","status":"completed","output_file_id":"file_1","endpoint":"/v1/chat/completions"}`) + }, + "/openai/files/file_1/content": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/jsonl") - _, _ = w.Write([]byte(`{"custom_id":"ok-1","response":{"status_code":200,"url":"/v1/chat/completions","body":{"id":"resp-1","model":"gpt-4o-mini"}}}`)) - default: - http.NotFound(w, r) - } - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/openai/deployments/gpt-4o") + _, _ = io.WriteString(w, `{"custom_id":"ok-1","response":{"status_code":200,"url":"/v1/chat/completions","body":{"id":"resp-1","model":"gpt-4o-mini"}}}`) + }, + }) + provider := newTestProvider(server.Client(), server.URL+"/openai/deployments/gpt-4o") resp, err := provider.GetBatchResults(context.Background(), "batch_1") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if resp.BatchID != "batch_1" { - t.Fatalf("BatchID = %q, want batch_1", resp.BatchID) - } - if len(gotPaths) != 2 { - t.Fatalf("saw %d requests, want 2", len(gotPaths)) - } - if gotPaths[0] != "/openai/batches/batch_1" { - t.Fatalf("first path = %q, want /openai/batches/batch_1", gotPaths[0]) - } - if gotPaths[1] != "/openai/files/file_1/content" { - t.Fatalf("second path = %q, want /openai/files/file_1/content", gotPaths[1]) - } - for i, gotVersion := range gotVersions { - if gotVersion != defaultAPIVersion { - t.Fatalf("request %d api-version = %q, want %q", i, gotVersion, defaultAPIVersion) - } + require.NoError(t, err) + assert.Equal(t, "batch_1", resp.BatchID) + + requests := capture.All() + require.Len(t, requests, 2) + assert.Equal(t, "/openai/batches/batch_1", requests[0].Path) + assert.Equal(t, "/openai/files/file_1/content", requests[1].Path) + for i, sent := range requests { + assert.Equal(t, defaultAPIVersion, sent.Query.Get("api-version"), "request %d", i) } } diff --git a/internal/providers/azure/realtime_test.go b/internal/providers/azure/realtime_test.go index 67ab94769..0d0558b0c 100644 --- a/internal/providers/azure/realtime_test.go +++ b/internal/providers/azure/realtime_test.go @@ -3,11 +3,12 @@ package azure import ( "context" "net/url" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestRealtimeTarget(t *testing.T) { @@ -19,30 +20,19 @@ func TestRealtimeTarget(t *testing.T) { }, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "gpt-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) u, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("parse target url: %v", err) - } - if u.Scheme != "wss" || u.Host != "myres.openai.azure.com" || u.Path != "/openai/realtime" { - t.Errorf("endpoint = %q, want wss://myres.openai.azure.com/openai/realtime", target.URL) - } - if got := u.Query().Get("deployment"); got != "gpt-realtime" { - t.Errorf("deployment = %q, want gpt-realtime", got) - } - if got := u.Query().Get("api-version"); got != "2025-04-01-preview" { - t.Errorf("api-version = %q, want 2025-04-01-preview", got) - } + require.NoError(t, err) + assert.Equal(t, "wss", u.Scheme) + assert.Equal(t, "myres.openai.azure.com", u.Host) + assert.Equal(t, "/openai/realtime", u.Path) + assert.Equal(t, "gpt-realtime", u.Query().Get("deployment")) + assert.Equal(t, "2025-04-01-preview", u.Query().Get("api-version")) + // Azure authenticates with the api-key header, not Bearer. - if got := target.Headers.Get("api-key"); got != apiKey { - t.Errorf("api-key = %q, want %q", got, apiKey) - } - if target.Headers.Get("Authorization") != "" { - t.Error("Authorization header must not be set for Azure (uses api-key)") - } + assert.Equal(t, apiKey, target.Headers.Get("api-key")) + assert.Empty(t, target.Headers.Get("Authorization")) } func TestRealtimeTargetStripsExistingOpenAIPath(t *testing.T) { @@ -53,16 +43,11 @@ func TestRealtimeTargetStripsExistingOpenAIPath(t *testing.T) { } { p := New(providers.ProviderConfig{APIKey: "k", BaseURL: base}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("base %q: unexpected error: %v", base, err) - } + require.NoError(t, err) + u, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("base %q: parse target url: %v", base, err) - } - if u.Path != "/openai/realtime" { - t.Errorf("base %q: path = %q, want /openai/realtime", base, u.Path) - } + require.NoError(t, err, "base %q", base) + assert.Equal(t, "/openai/realtime", u.Path, "base %q", base) } } @@ -72,19 +57,15 @@ func TestRealtimeTargetOmitsAuthWhenNoKey(t *testing.T) { BaseURL: "https://myres.openai.azure.com", }, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if _, present := target.Headers["Api-Key"]; present { - t.Error("api-key header should be absent when no key is configured") - } + require.NoError(t, err) + _, present := target.Headers["Api-Key"] + assert.False(t, present) } func TestRealtimeTargetMissingModel(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k", BaseURL: "https://myres.openai.azure.com"}, providers.ProviderOptions{}).(*Provider) - if _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } func TestRealtimeCallTarget(t *testing.T) { @@ -99,22 +80,13 @@ func TestRealtimeCallTarget(t *testing.T) { } { p := New(providers.ProviderConfig{APIKey: apiKey, BaseURL: base}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: "gpt-realtime"}) - if err != nil { - t.Fatalf("base %q: unexpected error: %v", base, err) - } - if target.URL != "https://myres.openai.azure.com/openai/v1/realtime/calls" { - t.Errorf("base %q: url = %q, want the GA calls endpoint", base, target.URL) - } + require.NoError(t, err) + assert.Equal(t, "https://myres.openai.azure.com/openai/v1/realtime/calls", target.URL, "base %q", base) + // The GA v1 surface takes no api-version parameter. - if strings.Contains(target.URL, "api-version") { - t.Errorf("base %q: url = %q, want no api-version on the GA surface", base, target.URL) - } - if got := target.Headers.Get("api-key"); got != apiKey { - t.Errorf("base %q: api-key = %q, want %q", base, got, apiKey) - } - if target.Headers.Get("Authorization") != "" { - t.Errorf("base %q: Authorization must not be set for Azure (uses api-key)", base) - } + assert.NotContains(t, target.URL, "api-version", "base %q", base) + assert.Equal(t, apiKey, target.Headers.Get("api-key"), "base %q", base) + assert.Empty(t, target.Headers.Get("Authorization"), "base %q: Azure uses api-key, not Authorization", base) } } @@ -125,19 +97,14 @@ func TestRealtimeClientSecretTarget(t *testing.T) { }, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeClientSecretTarget(context.Background(), &core.RealtimeRequest{Model: "gpt-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if target.URL != "https://myres.openai.azure.com/openai/v1/realtime/client_secrets" { - t.Errorf("url = %q, want the GA client secrets endpoint", target.URL) - } + require.NoError(t, err) + assert.Equal(t, "https://myres.openai.azure.com/openai/v1/realtime/client_secrets", target.URL) } func TestRealtimeCallTargetMissingModel(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k", BaseURL: "https://myres.openai.azure.com"}, providers.ProviderOptions{}).(*Provider) - if _, err := p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + _, err := p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } func TestRealtimeTargetAttachesByCallID(t *testing.T) { @@ -147,21 +114,16 @@ func TestRealtimeTargetAttachesByCallID(t *testing.T) { }, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "gpt-realtime", CallID: "rtc_3"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) + u, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("parse target url: %v", err) - } - if u.Scheme != "wss" || u.Host != "myres.openai.azure.com" || u.Path != "/openai/v1/realtime" { - t.Errorf("endpoint = %q, want wss://myres.openai.azure.com/openai/v1/realtime", target.URL) - } - if got := u.Query().Get("call_id"); got != "rtc_3" { - t.Errorf("call_id = %q, want rtc_3", got) - } + require.NoError(t, err) + assert.Equal(t, "wss", u.Scheme) + assert.Equal(t, "myres.openai.azure.com", u.Host) + assert.Equal(t, "/openai/v1/realtime", u.Path) + assert.Equal(t, "rtc_3", u.Query().Get("call_id")) + // The GA attach surface takes neither api-version nor deployment. - if u.Query().Has("api-version") || u.Query().Has("deployment") { - t.Errorf("query = %q, want only call_id on the GA attach surface", u.RawQuery) - } + assert.False(t, u.Query().Has("api-version")) + assert.False(t, u.Query().Has("deployment"), "query = %q, want only call_id", u.RawQuery) } diff --git a/internal/providers/bailian/bailian_test.go b/internal/providers/bailian/bailian_test.go index 84f1d50fd..284a604f5 100644 --- a/internal/providers/bailian/bailian_test.go +++ b/internal/providers/bailian/bailian_test.go @@ -6,75 +6,55 @@ import ( "errors" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_SendsBearerAuthAndCorrectPath(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-bailian", - "created":1677652288, - "model":"qwen3-max", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":5,"completion_tokens":10,"total_tokens":15} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("bailian-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "qwen3-max", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +var ( + _ core.NativeBatchProvider = (*Provider)(nil) + _ core.NativeFileProvider = (*Provider)(nil) +) + +const chatCompletionJSON = `{ + "id":"chatcmpl-bailian", + "created":1677652288, + "model":"qwen3-max", + "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}], + "usage":{"prompt_tokens":5,"completion_tokens":10,"total_tokens":15} +}` + +const upstreamErrorJSON = `{"error":{"message":"bad request","type":"invalid_request_error"}}` + +// newTestProvider points a provider at the given upstream URL. +func newTestProvider(baseURL string) *Provider { + p := NewWithHTTPClient("key", nil, llmclient.Hooks{}) + p.SetBaseURL(baseURL) + return p +} + +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "bailian", + DefaultBaseURL: defaultBaseURL, + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + p := NewWithHTTPClient(apiKey, client, hooks) + p.SetBaseURL(baseURL) + return p }, + Embeddings: true, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "qwen3-max" { - t.Fatalf("resp.Model = %q, want qwen3-max", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer bailian-key" { - t.Fatalf("authorization = %q, want Bearer bailian-key", gotAuth) - } } func TestChatCompletion_MaxTokensMapping(t *testing.T) { - var gotBody []byte - var readErr error - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-bailian", - "created":1677652288, - "model":"qwen3-max", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":5,"completion_tokens":10,"total_tokens":15} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) maxTokens := 4096 _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ @@ -82,316 +62,61 @@ func TestChatCompletion_MaxTokensMapping(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hi"}}, MaxTokens: &maxTokens, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } + require.NoError(t, err) - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_tokens"]; exists { - t.Fatal("sent body should NOT contain max_tokens (Bailian deprecated it)") - } - mct, exists := sentBody["max_completion_tokens"] - if !exists { - t.Fatal("sent body should contain max_completion_tokens") - } - if mct.(float64) != 4096 { - t.Fatalf("max_completion_tokens = %v, want 4096", mct) - } + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "max_tokens") + assert.Equal(t, float64(4096), sent["max_completion_tokens"]) } func TestChatCompletion_NoMaxTokensMapping(t *testing.T) { - var gotBody []byte - var readErr error - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-bailian", - "created":1677652288, - "model":"qwen3-max", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":5,"completion_tokens":10,"total_tokens":15} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } - - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_completion_tokens"]; exists { - t.Fatal("sent body should NOT contain max_completion_tokens when request had no max_tokens") - } + require.NoError(t, err) + assert.NotContains(t, capture.Last(t).JSON(t), "max_completion_tokens") } func TestStreamChatCompletion_MaxTokensMapping(t *testing.T) { - var gotBody []byte - var readErr error - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.Header().Set("Content-Type", "text/event-stream") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.SSEServer(t, "data: [DONE]\n\n") + provider := newTestProvider(server.URL) maxTokens := 2048 - _, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ + stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3-flash", Messages: []core.Message{{Role: "user", Content: "hi"}}, MaxTokens: &maxTokens, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } - - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_tokens"]; exists { - t.Fatal("sent body should NOT contain max_tokens for streaming either") - } - mct, exists := sentBody["max_completion_tokens"] - if !exists { - t.Fatal("sent body should contain max_completion_tokens for streaming") - } - if mct.(float64) != 2048 { - t.Fatalf("max_completion_tokens = %v, want 2048", mct) - } -} - -func TestStreamChatCompletion_ReturnsSSE(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hi\"},\"finish_reason\":null}]}\n\n")) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ - Model: "qwen3-flash", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() - body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("failed to read stream: %v", err) - } - if !strings.Contains(string(body), "[DONE]") { - t.Fatalf("stream should contain [DONE], got: %s", string(body)) - } -} - -func TestListModels_ReturnsModels(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[ - {"id":"qwen3-max","object":"model","owned_by":"alibaba"}, - {"id":"qwen3-plus","object":"model","owned_by":"alibaba"}, - {"id":"qwen3-flash","object":"model","owned_by":"alibaba"} - ] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 3 { - t.Fatalf("got %d models, want 3", len(resp.Data)) - } -} - -func TestEmbeddings_SendsRequest(t *testing.T) { - var gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], - "model":"text-embedding-v3", - "usage":{"prompt_tokens":2,"total_tokens":2} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: "text-embedding-v3", - Input: "test", - }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if gotPath != "/embeddings" { - t.Fatalf("path = %q, want /embeddings", gotPath) - } -} - -func TestResponsesViaChat_DelegatesToChat(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-bailian", - "created":1677652288, - "model":"qwen3-max", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":5,"completion_tokens":5,"total_tokens":10} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ - Model: "qwen3-max", - Input: "hello", - }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if resp.Status != "completed" { - t.Fatalf("status = %q, want completed", resp.Status) - } -} - -func TestDefaultBaseURL(t *testing.T) { - provider := NewWithHTTPClient("key", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("expected non-nil provider") - } - // Verify the registration exposes the correct default base URL - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Fatalf("Registration.DefaultBaseURL = %q, want %q", - Registration.Discovery.DefaultBaseURL, defaultBaseURL) - } - // Verify the provider actually uses the default base URL by checking - // that a ChatCompletion request hits the correct host. - var gotHost string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotHost = r.Host - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-1", - "created":1, - "model":"qwen3-max", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2} - }`)) - })) - defer server.Close() - - // Override the base URL to our test server to capture the request, - // but first verify the provider's default is the expected DashScope URL. - provider.SetBaseURL(server.URL) - - _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "qwen3-max", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - // The request should have reached our test server, confirming SetBaseURL works. - if gotHost == "" { - t.Fatal("expected a non-empty host in the request") - } -} - -func TestProvider_ExposesBatchAndFileInterfaces(t *testing.T) { - provider := NewWithHTTPClient("key", nil, llmclient.Hooks{}) - if _, ok := any(provider).(core.NativeBatchProvider); !ok { - t.Fatal("bailian should implement native batch") - } - if _, ok := any(provider).(core.NativeFileProvider); !ok { - t.Fatal("bailian should implement native file") - } -} - -func TestRegistration_TypeAndBaseURL(t *testing.T) { - if Registration.Type != "bailian" { - t.Fatalf("Registration.Type = %q, want bailian", Registration.Type) - } - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Fatalf("DefaultBaseURL mismatch") - } + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "max_tokens") + assert.Equal(t, float64(2048), sent["max_completion_tokens"]) } func TestPassthrough_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"ok":true}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"ok":true}`) + provider := newTestProvider(server.URL) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/chat/completions", Body: io.NopCloser(strings.NewReader(`{}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - if resp.StatusCode != http.StatusOK { - t.Fatalf("StatusCode = %d, want 200", resp.StatusCode) - } + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) } func TestPassthrough_NilRequest(t *testing.T) { provider := NewWithHTTPClient("key", nil, llmclient.Hooks{}) - - if _, err := provider.Passthrough(context.Background(), nil); err == nil { - t.Fatal("expected error for nil passthrough request") - } + _, err := provider.Passthrough(context.Background(), nil) + require.Error(t, err) } func TestPassthrough_ReadError(t *testing.T) { @@ -403,75 +128,40 @@ func TestPassthrough_ReadError(t *testing.T) { Endpoint: "/chat/completions", Body: errReadCloser{err: readErr}, }) - if !errors.Is(err, readErr) { - t.Fatalf("Passthrough() error = %v, want %v", err, readErr) - } + require.ErrorIs(t, err, readErr) } func TestAdaptBailianRequest_Nil(t *testing.T) { r, err := adaptChatRequest(nil) - if err != nil { - t.Fatalf("adaptChatRequest(nil) error = %v", err) - } - if r != nil { - t.Fatal("expected nil") - } + require.NoError(t, err) + assert.Nil(t, r) } func TestAdaptBailianRequest_NoMaxTokens(t *testing.T) { req := &core.ChatRequest{Model: "qwen3-max"} r, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if r.MaxTokens != nil { - t.Fatal("should not set max_completion_tokens when request had none") - } - -} -func TestNew_UsesRegistrationAndDefaultBaseURL(t *testing.T) { - provider := New(providers.ProviderConfig{ - APIKey: "reg-key", - }, providers.ProviderOptions{}) - if provider == nil { - t.Fatal("New() returned nil") - } - // Verify the provider constructed via registration uses the expected base URL - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Fatalf("Registration.DefaultBaseURL = %q, want %q", - Registration.Discovery.DefaultBaseURL, defaultBaseURL) - } + require.NoError(t, err) + assert.Nil(t, r.MaxTokens) } func TestStreamResponses_DelegatesToChat(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n")) - _, _ = w.Write([]byte("data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n")) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.SSEServer(t, + "data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n"+ + "data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n"+ + "data: [DONE]\n\n") + provider := newTestProvider(server.URL) stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "qwen3-max", Input: "hello", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("failed to read stream: %v", err) - } - if len(body) == 0 { - t.Fatal("expected non-empty stream body") - } + require.NoError(t, err) + assert.NotEmpty(t, body) + assert.Equal(t, "/chat/completions", capture.Last(t).Path) } func TestAdaptBailianRequest_PreservesOtherFields(t *testing.T) { @@ -482,18 +172,10 @@ func TestAdaptBailianRequest_PreservesOtherFields(t *testing.T) { MaxTokens: &maxTokens, } r, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if r == req { - t.Fatal("should return a clone, not the original") - } - if r.Model != "qwen3-max" { - t.Fatalf("model = %q", r.Model) - } - if r.MaxTokens != nil { - t.Fatal("MaxTokens should be nil in the clone") - } + require.NoError(t, err) + assert.NotSame(t, req, r) + assert.Equal(t, "qwen3-max", r.Model) + assert.Nil(t, r.MaxTokens) } func TestAdaptBailianRequest_RespectsExistingMaxCompletionTokens(t *testing.T) { @@ -508,60 +190,35 @@ func TestAdaptBailianRequest_RespectsExistingMaxCompletionTokens(t *testing.T) { ExtraFields: extra, } r, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if r == req { - t.Fatal("should return a clone") - } - if r.MaxTokens != nil { - t.Fatal("MaxTokens should be nil") - } + require.NoError(t, err) + assert.NotSame(t, req, r) + assert.Nil(t, r.MaxTokens) body, err := json.Marshal(r) - if err != nil { - t.Fatalf("failed to marshal adapted request: %v", err) - } + require.NoError(t, err) + var raw map[string]json.RawMessage - if err := json.Unmarshal(body, &raw); err != nil { - t.Fatalf("failed to unmarshal body: %v", err) - } - if _, exists := raw["max_completion_tokens"]; !exists { - t.Fatal("max_completion_tokens should exist") - } + require.NoError(t, json.Unmarshal(body, &raw)) + assert.NotContains(t, raw, "max_tokens") + require.Contains(t, raw, "max_completion_tokens") + var mct int - if err := json.Unmarshal(raw["max_completion_tokens"], &mct); err != nil { - t.Fatalf("failed to unmarshal max_completion_tokens: %v", err) - } - if mct != 200 { - t.Fatalf("max_completion_tokens = %d, want 200", mct) - } - if _, exists := raw["max_tokens"]; exists { - t.Fatal("max_tokens should NOT exist in output") - } + require.NoError(t, json.Unmarshal(raw["max_completion_tokens"], &mct)) + assert.Equal(t, 200, mct) } func TestChatCompletion_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":{"message":"bad request","type":"invalid_request_error"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusBadRequest, upstreamErrorJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err == nil { - t.Fatal("expected error from upstream 400") - } + require.Error(t, err) } func TestChatCompletion_TransportFailure(t *testing.T) { - // Use a stub RoundTripper that always returns an error errTransport := errors.New("simulated transport failure") provider := NewWithHTTPClient("key", &http.Client{ Transport: roundTripperFunc(func(*http.Request) (*http.Response, error) { @@ -573,289 +230,154 @@ func TestChatCompletion_TransportFailure(t *testing.T) { Model: "qwen3-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err == nil { - t.Fatal("expected transport error") - } + require.Error(t, err) } func TestStreamChatCompletion_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte(`{"error":{"message":"unauthorized","type":"authentication_error"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusUnauthorized, `{"error":{"message":"unauthorized","type":"authentication_error"}}`) + provider := newTestProvider(server.URL) _, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err == nil { - t.Fatal("expected error from upstream 401") - } + require.Error(t, err) } func TestResponses_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":{"message":"bad request","type":"invalid_request_error"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusBadRequest, upstreamErrorJSON) + provider := newTestProvider(server.URL) _, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "qwen3-max", Input: "hello", }) - if err == nil { - t.Fatal("expected error from upstream 400") - } + require.Error(t, err) } func TestEmbeddings_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":{"message":"bad request","type":"invalid_request_error"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusBadRequest, upstreamErrorJSON) + provider := newTestProvider(server.URL) _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "text-embedding-v3", Input: "test", }) - if err == nil { - t.Fatal("expected error from upstream 400") - } + require.Error(t, err) } func TestPassthrough_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusInternalServerError) - _, _ = w.Write([]byte(`{"error":"internal"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusInternalServerError, `{"error":"internal"}`) + provider := newTestProvider(server.URL) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/chat/completions", Body: io.NopCloser(strings.NewReader(`{}`)), }) - if err != nil { - t.Fatalf("Passthrough() should not return error on non-2xx: %v", err) - } - if resp.StatusCode != http.StatusInternalServerError { - t.Fatalf("StatusCode = %d, want 500", resp.StatusCode) - } + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusInternalServerError, resp.StatusCode) } func TestListModels_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusInternalServerError) - _, _ = w.Write([]byte(`{"error":"internal"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusInternalServerError, `{"error":"internal"}`) + provider := newTestProvider(server.URL) _, err := provider.ListModels(context.Background()) - if err == nil { - t.Fatal("expected error from upstream 500") - } + require.Error(t, err) } func TestCreateBatch_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"batch-bailian-1","object":"batch","status":"validating"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"batch-bailian-1","object":"batch","status":"validating"}`) + provider := newTestProvider(server.URL) resp, err := provider.CreateBatch(context.Background(), &core.BatchRequest{ InputFileID: "file-1", Endpoint: "/v1/chat/completions", }) - if err != nil { - t.Fatalf("CreateBatch() error = %v", err) - } - if resp.ID != "batch-bailian-1" { - t.Fatalf("batch id = %q", resp.ID) - } + require.NoError(t, err) + assert.Equal(t, "batch-bailian-1", resp.ID) } func TestGetBatch_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"batch-bailian-1","object":"batch","status":"completed"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"batch-bailian-1","object":"batch","status":"completed"}`) + provider := newTestProvider(server.URL) resp, err := provider.GetBatch(context.Background(), "batch-bailian-1") - if err != nil { - t.Fatalf("GetBatch() error = %v", err) - } - if resp.ID != "batch-bailian-1" { - t.Fatalf("batch id = %q", resp.ID) - } + require.NoError(t, err) + assert.Equal(t, "batch-bailian-1", resp.ID) } func TestListBatches_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) + provider := newTestProvider(server.URL) _, err := provider.ListBatches(context.Background(), 10, "") - if err != nil { - t.Fatalf("ListBatches() error = %v", err) - } + require.NoError(t, err) } func TestCancelBatch_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"batch-bailian-1","object":"batch","status":"cancelling"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"batch-bailian-1","object":"batch","status":"cancelling"}`) + provider := newTestProvider(server.URL) resp, err := provider.CancelBatch(context.Background(), "batch-bailian-1") - if err != nil { - t.Fatalf("CancelBatch() error = %v", err) - } - if resp.ID != "batch-bailian-1" { - t.Fatalf("batch id = %q", resp.ID) - } + require.NoError(t, err) + assert.Equal(t, "batch-bailian-1", resp.ID) } func TestCreateFile_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"file-1","object":"file","purpose":"batch","bytes":100}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"file-1","object":"file","purpose":"batch","bytes":100}`) + provider := newTestProvider(server.URL) resp, err := provider.CreateFile(context.Background(), &core.FileCreateRequest{ Content: []byte("data"), Purpose: "batch", }) - if err != nil { - t.Fatalf("CreateFile() error = %v", err) - } - if resp.Provider != "bailian" { - t.Fatalf("provider = %q, want bailian", resp.Provider) - } + require.NoError(t, err) + assert.Equal(t, "bailian", resp.Provider) } func TestDeleteFile_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"file-1","object":"file","deleted":true}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"file-1","object":"file","deleted":true}`) + provider := newTestProvider(server.URL) resp, err := provider.DeleteFile(context.Background(), "file-1") - if err != nil { - t.Fatalf("DeleteFile() error = %v", err) - } - if !resp.Deleted { - t.Fatal("expected deleted=true") - } + require.NoError(t, err) + assert.True(t, resp.Deleted) } func TestListFiles_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) + provider := newTestProvider(server.URL) _, err := provider.ListFiles(context.Background(), "batch", 10, "") - if err != nil { - t.Fatalf("ListFiles() error = %v", err) - } + require.NoError(t, err) } func TestGetFile_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"file-1","object":"file","purpose":"batch"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"file-1","object":"file","purpose":"batch"}`) + provider := newTestProvider(server.URL) resp, err := provider.GetFile(context.Background(), "file-1") - if err != nil { - t.Fatalf("GetFile() error = %v", err) - } - if resp.Provider != "bailian" { - t.Fatalf("provider = %q, want bailian", resp.Provider) - } + require.NoError(t, err) + assert.Equal(t, "bailian", resp.Provider) } func TestGetFileContent_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"text":"content"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"text":"content"}`) + provider := newTestProvider(server.URL) _, err := provider.GetFileContent(context.Background(), "file-1") - if err != nil { - t.Fatalf("GetFileContent() error = %v", err) - } + require.NoError(t, err) } func TestGetBatchResults_Delegates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"batch-1","output_file_id":"file-out-1"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"id":"batch-1","output_file_id":"file-out-1"}`) + provider := newTestProvider(server.URL) _, err := provider.GetBatchResults(context.Background(), "batch-1") - if err != nil { - t.Fatalf("GetBatchResults() error = %v", err) - } + require.NoError(t, err) } // roundTripperFunc adapts a function to the http.RoundTripper interface. @@ -878,120 +400,50 @@ func (r errReadCloser) Close() error { } func TestPassthrough_MaxTokensMapping(t *testing.T) { - var gotBody []byte - var readErr error - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"id":"chatcmpl-bailian","model":"qwen3-max"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"chatcmpl-bailian","model":"qwen3-max"}`) + provider := newTestProvider(server.URL) - _, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/chat/completions", Body: io.NopCloser(strings.NewReader(`{"model":"qwen3-max","messages":[{"role":"user","content":"hi"}],"max_tokens":4096}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } + require.NoError(t, err) + defer resp.Body.Close() - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_tokens"]; exists { - t.Fatal("passthrough body should NOT contain max_tokens") - } - mct, exists := sentBody["max_completion_tokens"] - if !exists { - t.Fatal("passthrough body should contain max_completion_tokens") - } - if mct.(float64) != 4096 { - t.Fatalf("max_completion_tokens = %v, want 4096", mct) - } + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "max_tokens") + assert.Equal(t, float64(4096), sent["max_completion_tokens"]) } func TestPassthrough_PreservesExistingMaxCompletionTokens(t *testing.T) { - var gotBody []byte - var readErr error + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"chatcmpl-bailian","model":"qwen3-max"}`) + provider := newTestProvider(server.URL) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"id":"chatcmpl-bailian","model":"qwen3-max"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - _, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/chat/completions", Body: io.NopCloser(strings.NewReader(`{"model":"qwen3-max","messages":[{"role":"user","content":"hi"}],"max_tokens":100,"max_completion_tokens":200}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } + require.NoError(t, err) + defer resp.Body.Close() - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_tokens"]; exists { - t.Fatal("passthrough body should NOT contain max_tokens when max_completion_tokens already set") - } - mct, exists := sentBody["max_completion_tokens"] - if !exists { - t.Fatal("passthrough body should contain max_completion_tokens") - } - if mct.(float64) != 200 { - t.Fatalf("max_completion_tokens = %v, want 200 (explicit value should win)", mct) - } + // The explicit max_completion_tokens wins over the mapped max_tokens. + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "max_tokens") + assert.Equal(t, float64(200), sent["max_completion_tokens"]) } func TestPassthrough_NoMaxTokens(t *testing.T) { - var gotBody []byte - var readErr error - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, readErr = io.ReadAll(r.Body) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"id":"chatcmpl-bailian","model":"qwen3-max"}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"chatcmpl-bailian","model":"qwen3-max"}`) + provider := newTestProvider(server.URL) - provider := NewWithHTTPClient("key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - _, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/chat/completions", Body: io.NopCloser(strings.NewReader(`{"model":"qwen3-max","messages":[{"role":"user","content":"hi"}]}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - if readErr != nil { - t.Fatalf("reading request body: %v", readErr) - } - - var sentBody map[string]any - if err := json.Unmarshal(gotBody, &sentBody); err != nil { - t.Fatalf("unmarshal sent body: %v", err) - } - if _, exists := sentBody["max_completion_tokens"]; exists { - t.Fatal("passthrough body should NOT contain max_completion_tokens when request had no max_tokens") - } + require.NoError(t, err) + defer resp.Body.Close() + assert.NotContains(t, capture.Last(t).JSON(t), "max_completion_tokens") } diff --git a/internal/providers/bailian/realtime_test.go b/internal/providers/bailian/realtime_test.go index 13eff6900..0d97f348f 100644 --- a/internal/providers/bailian/realtime_test.go +++ b/internal/providers/bailian/realtime_test.go @@ -8,6 +8,8 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestRealtimeURL(t *testing.T) { @@ -27,24 +29,17 @@ func TestRealtimeURL(t *testing.T) { t.Run(tt.name, func(t *testing.T) { got, err := realtimeURL(tt.baseURL, tt.model) if tt.wantErr { - if err == nil { - t.Fatalf("expected error, got %q", got) - } + require.Error(t, err) + return } - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) + u, parseErr := url.Parse(got) - if parseErr != nil { - t.Fatalf("invalid URL: %v", parseErr) - } - if base := u.Scheme + "://" + u.Host + u.Path; base != tt.wantBase { - t.Errorf("base = %q, want %q", base, tt.wantBase) - } - if u.Query().Get("model") != tt.model { - t.Errorf("model = %q, want %q", u.Query().Get("model"), tt.model) - } + require.NoError(t, parseErr) + base := u.Scheme + "://" + u.Host + u.Path + assert.Equal(t, tt.wantBase, base) + assert.Equal(t, tt.model, u.Query().Get("model")) }) } } @@ -54,30 +49,20 @@ func TestRealtimeTarget(t *testing.T) { p := New(providers.ProviderConfig{APIKey: apiKey}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "qwen3-omni-flash-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://dashscope.aliyuncs.com/api-ws/v1/realtime?") { - t.Errorf("url = %q, want DashScope realtime endpoint", target.URL) - } - if got := target.Headers.Get("Authorization"); got != "Bearer "+apiKey { - t.Errorf("Authorization = %q, want bearer with key", got) - } - - if _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://dashscope.aliyuncs.com/api-ws/v1/realtime?"), "url = %q, want DashScope realtime endpoint", target.URL) + got := target.Headers.Get("Authorization") + assert.Equal(t, "Bearer "+apiKey, got) + _, err = p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } func TestRealtimeTargetOmitsAuthWhenNoKey(t *testing.T) { p := New(providers.ProviderConfig{APIKey: ""}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if _, present := target.Headers["Authorization"]; present { - t.Error("Authorization header should be absent when no API key is configured") - } + require.NoError(t, err) + _, present := target.Headers["Authorization"] + assert.False(t, present) } func TestRealtimeTargetFollowsSetBaseURL(t *testing.T) { @@ -85,10 +70,6 @@ func TestRealtimeTargetFollowsSetBaseURL(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) p.SetBaseURL("https://dashscope-intl.aliyuncs.com/compatible-mode/v1") target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://dashscope-intl.aliyuncs.com/api-ws/v1/realtime") { - t.Errorf("url = %q, want the SetBaseURL region host", target.URL) - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://dashscope-intl.aliyuncs.com/api-ws/v1/realtime"), "url = %q, want the SetBaseURL region host", target.URL) } diff --git a/internal/providers/bedrock/bedrock_test.go b/internal/providers/bedrock/bedrock_test.go index 01b4a0025..81e3e0161 100644 --- a/internal/providers/bedrock/bedrock_test.go +++ b/internal/providers/bedrock/bedrock_test.go @@ -3,7 +3,6 @@ package bedrock import ( "context" "encoding/json" - "errors" "io" "net/http" "strings" @@ -12,6 +11,8 @@ import ( awssdk "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/service/bedrockruntime" brtypes "github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" @@ -31,9 +32,8 @@ func TestCallObservationCoversBedrockSDKAndFirstChunk(t *testing.T) { ends = append(ends, info) }, OnStreamFirstChunk: func(ctx context.Context, info llmclient.ResponseInfo) { - if ctx.Value(contextKey{}) != true { - t.Error("first chunk hook did not receive derived context") - } + started, _ := ctx.Value(contextKey{}).(bool) + assert.True(t, started, "first-chunk hook must see the OnRequestStart context") chunks = append(chunks, info) }, }} @@ -41,18 +41,15 @@ func TestCallObservationCoversBedrockSDKAndFirstChunk(t *testing.T) { observation := p.beginCallObservation(t.Context(), "anthropic.claude", true) observation.end(http.StatusOK, nil) stream := observedStream(io.NopCloser(strings.NewReader("data: first\n\n")), observation) - if len(chunks) != 0 { - t.Fatal("first chunk hook fired before stream read") - } - if _, err := io.ReadAll(stream); err != nil { - t.Fatal(err) - } - if len(starts) != 1 || len(ends) != 1 || len(chunks) != 1 { - t.Fatalf("hook counts = start:%d end:%d chunk:%d, want 1/1/1", len(starts), len(ends), len(chunks)) - } - if starts[0].Operation != llmclient.OperationChat || starts[0].Endpoint != converseEndpoint || !starts[0].Stream { - t.Fatalf("start info = %+v, want streaming Bedrock chat", starts[0]) - } + require.Empty(t, chunks) + _, err := io.ReadAll(stream) + require.NoError(t, err) + require.Len(t, starts, 1) + require.Len(t, ends, 1) + require.Len(t, chunks, 1) + require.Equal(t, llmclient.OperationChat, starts[0].Operation) + require.Equal(t, converseEndpoint, starts[0].Endpoint) + require.True(t, starts[0].Stream) } func TestParseBaseURL(t *testing.T) { @@ -72,12 +69,8 @@ func TestParseBaseURL(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { region, endpoint := parseBaseURL(tc.in) - if region != tc.wantRegion { - t.Errorf("region = %q, want %q", region, tc.wantRegion) - } - if endpoint != tc.wantEndpoint { - t.Errorf("endpoint = %q, want %q", endpoint, tc.wantEndpoint) - } + assert.Equal(t, tc.wantRegion, region) + assert.Equal(t, tc.wantEndpoint, endpoint) }) } } @@ -116,12 +109,8 @@ func TestPlaneEndpoint(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - if got := runtimePlaneEndpoint(tc.in); got != tc.wantRuntime { - t.Errorf("runtimePlaneEndpoint(%q) = %q, want %q", tc.in, got, tc.wantRuntime) - } - if got := controlPlaneEndpoint(tc.in); got != tc.wantControl { - t.Errorf("controlPlaneEndpoint(%q) = %q, want %q", tc.in, got, tc.wantControl) - } + assert.Equal(t, tc.wantRuntime, runtimePlaneEndpoint(tc.in), "runtimePlaneEndpoint(%q)", tc.in) + assert.Equal(t, tc.wantControl, controlPlaneEndpoint(tc.in), "controlPlaneEndpoint(%q)", tc.in) }) } } @@ -144,10 +133,7 @@ func TestMapStopReason(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - got := mapStopReason(tc.reason, tc.hasTools) - if got != tc.want { - t.Errorf("got %q, want %q", got, tc.want) - } + assert.Equal(t, tc.want, mapStopReason(tc.reason, tc.hasTools)) }) } } @@ -166,33 +152,16 @@ func TestBuildConverseParts_BasicRequest(t *testing.T) { } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if awssdk.ToString(parts.modelID) != req.Model { - t.Errorf("modelID = %q, want %q", awssdk.ToString(parts.modelID), req.Model) - } - if len(parts.system) != 1 { - t.Fatalf("expected 1 system block, got %d", len(parts.system)) - } - if got := parts.system[0].(*brtypes.SystemContentBlockMemberText).Value; got != "You are concise" { - t.Errorf("system text = %q", got) - } - if len(parts.messages) != 1 { - t.Fatalf("expected 1 message, got %d", len(parts.messages)) - } - if parts.messages[0].Role != brtypes.ConversationRoleUser { - t.Errorf("role = %q", parts.messages[0].Role) - } - if parts.infCfg == nil { - t.Fatal("inference config should be set") - } - if awssdk.ToInt32(parts.infCfg.MaxTokens) != int32(maxTokens) { - t.Errorf("max tokens = %d", awssdk.ToInt32(parts.infCfg.MaxTokens)) - } - if awssdk.ToFloat32(parts.infCfg.Temperature) != float32(temp) { - t.Errorf("temperature = %v", awssdk.ToFloat32(parts.infCfg.Temperature)) - } + require.NoError(t, err) + assert.Equal(t, req.Model, awssdk.ToString(parts.modelID)) + require.Len(t, parts.system, 1) + got := parts.system[0].(*brtypes.SystemContentBlockMemberText).Value + assert.Equal(t, "You are concise", got) + require.Len(t, parts.messages, 1) + assert.Equal(t, brtypes.ConversationRoleUser, parts.messages[0].Role) + require.NotNil(t, parts.infCfg) + assert.Equal(t, int32(maxTokens), awssdk.ToInt32(parts.infCfg.MaxTokens)) + assert.Equal(t, float32(temp), awssdk.ToFloat32(parts.infCfg.Temperature)) } func TestBuildConverseParts_MaxCompletionTokensFallback(t *testing.T) { @@ -204,15 +173,10 @@ func TestBuildConverseParts_MaxCompletionTokensFallback(t *testing.T) { }), } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if parts.infCfg == nil { - t.Fatal("inference config should be set when max_completion_tokens is provided via ExtraFields") - } - if got := awssdk.ToInt32(parts.infCfg.MaxTokens); got != 256 { - t.Errorf("max tokens = %d, want 256", got) - } + require.NoError(t, err) + require.NotNil(t, parts.infCfg) + got := awssdk.ToInt32(parts.infCfg.MaxTokens) + assert.Equal(t, int32(256), got) } func TestBuildConverseParts_MaxTokensWinsOverFallback(t *testing.T) { @@ -226,23 +190,18 @@ func TestBuildConverseParts_MaxTokensWinsOverFallback(t *testing.T) { }), } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if got := awssdk.ToInt32(parts.infCfg.MaxTokens); got != 128 { - t.Errorf("max tokens = %d, want 128 (max_tokens should take precedence)", got) - } + require.NoError(t, err) + got := awssdk.ToInt32(parts.infCfg.MaxTokens) + assert.Equal(t, int32(128), got) } func TestBuildConverseParts_RejectsEmptyModel(t *testing.T) { _, err := buildConverseParts(&core.ChatRequest{Messages: []core.Message{{Role: "user", Content: "hi"}}}) - if err == nil { - t.Fatal("expected error for missing model") - } + require.Error(t, err) + var ge *core.GatewayError - if !errors.As(err, &ge) || ge.Type != core.ErrorTypeInvalidRequest { - t.Fatalf("expected invalid_request_error, got %v", err) - } + require.ErrorAs(t, err, &ge) + require.Equal(t, core.ErrorTypeInvalidRequest, ge.Type) } func TestBuildConverseParts_MergesParallelToolResults(t *testing.T) { @@ -266,28 +225,20 @@ func TestBuildConverseParts_MergesParallelToolResults(t *testing.T) { }, } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } + require.NoError(t, err) + // Expect: user(text), assistant(2 tool_use), user(2 tool_result) — three messages. - if len(parts.messages) != 3 { - t.Fatalf("expected 3 messages, got %d: %+v", len(parts.messages), parts.messages) - } + require.Len(t, parts.messages, 3) + last := parts.messages[2] - if last.Role != brtypes.ConversationRoleUser { - t.Fatalf("merged tool results must be user-role, got %q", last.Role) - } - if len(last.Content) != 2 { - t.Fatalf("expected 2 ToolResult blocks in merged message, got %d", len(last.Content)) - } + require.Equal(t, brtypes.ConversationRoleUser, last.Role) + require.Len(t, last.Content, 2) + for i, want := range []string{"call_1", "call_2"} { tr, ok := last.Content[i].(*brtypes.ContentBlockMemberToolResult) - if !ok { - t.Fatalf("block %d not a ToolResult: %T", i, last.Content[i]) - } - if got := awssdk.ToString(tr.Value.ToolUseId); got != want { - t.Errorf("block %d ToolUseId = %q, want %q", i, got, want) - } + require.True(t, ok, "block %d not a ToolResult: %T", i, last.Content[i]) + got := awssdk.ToString(tr.Value.ToolUseId) + assert.Equal(t, want, got) } } @@ -300,15 +251,11 @@ func TestBuildConverseParts_TopPFromExtraFields(t *testing.T) { }), } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if parts.infCfg == nil || parts.infCfg.TopP == nil { - t.Fatal("top_p was not forwarded to InferenceConfiguration.TopP") - } - if got := awssdk.ToFloat32(parts.infCfg.TopP); got != 0.7 { - t.Errorf("top_p = %v, want 0.7", got) - } + require.NoError(t, err) + require.NotNil(t, parts.infCfg) + require.NotNil(t, parts.infCfg.TopP) + got := awssdk.ToFloat32(parts.infCfg.TopP) + assert.Equal(t, float32(0.7), got) } func TestBuildConverseParts_TopPFromTypedField(t *testing.T) { @@ -319,15 +266,11 @@ func TestBuildConverseParts_TopPFromTypedField(t *testing.T) { TopP: &topP, } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if parts.infCfg == nil || parts.infCfg.TopP == nil { - t.Fatal("typed top_p was not forwarded to InferenceConfiguration.TopP") - } - if got := awssdk.ToFloat32(parts.infCfg.TopP); got != 0.8 { - t.Errorf("top_p = %v, want 0.8", got) - } + require.NoError(t, err) + require.NotNil(t, parts.infCfg) + require.NotNil(t, parts.infCfg.TopP) + got := awssdk.ToFloat32(parts.infCfg.TopP) + assert.Equal(t, float32(0.8), got) } func TestBuildConverseParts_TypedTopPWinsOverExtraFields(t *testing.T) { @@ -341,15 +284,11 @@ func TestBuildConverseParts_TypedTopPWinsOverExtraFields(t *testing.T) { }), } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if parts.infCfg == nil || parts.infCfg.TopP == nil { - t.Fatal("typed top_p was not forwarded to InferenceConfiguration.TopP") - } - if got := awssdk.ToFloat32(parts.infCfg.TopP); got != 0.8 { - t.Errorf("top_p = %v, want typed value 0.8", got) - } + require.NoError(t, err) + require.NotNil(t, parts.infCfg) + require.NotNil(t, parts.infCfg.TopP) + got := awssdk.ToFloat32(parts.infCfg.TopP) + assert.Equal(t, float32(0.8), got) } func TestBuildConverseParts_RejectsMaxTokensOverflow(t *testing.T) { @@ -360,13 +299,11 @@ func TestBuildConverseParts_RejectsMaxTokensOverflow(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hi"}}, } _, err := buildConverseParts(req) - if err == nil { - t.Fatal("expected invalid_request_error for oversized max_tokens") - } + require.Error(t, err) + var ge *core.GatewayError - if !errors.As(err, &ge) || ge.Type != core.ErrorTypeInvalidRequest { - t.Fatalf("expected invalid_request_error, got %v", err) - } + require.ErrorAs(t, err, &ge) + require.Equal(t, core.ErrorTypeInvalidRequest, ge.Type) } func TestBuildConverseParts_ToolResultBatchesDoNotAliasAcrossTurns(t *testing.T) { @@ -391,9 +328,7 @@ func TestBuildConverseParts_ToolResultBatchesDoNotAliasAcrossTurns(t *testing.T) }, } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } + require.NoError(t, err) collectIDs := func(content []brtypes.ContentBlock) []string { var ids []string @@ -419,12 +354,8 @@ func TestBuildConverseParts_ToolResultBatchesDoNotAliasAcrossTurns(t *testing.T) secondBatch = ids } } - if got, want := firstBatch, []string{"c1", "c2"}; !equalStrings(got, want) { - t.Errorf("first turn tool result IDs = %v, want %v (aliasing bug overwrote them)", got, want) - } - if got, want := secondBatch, []string{"c3"}; !equalStrings(got, want) { - t.Errorf("second turn tool result IDs = %v, want %v", got, want) - } + assert.Equal(t, []string{"c1", "c2"}, firstBatch, "first turn tool result IDs (aliasing bug overwrote them)") + assert.Equal(t, []string{"c3"}, secondBatch, "second turn tool result IDs") } func TestBuildConverseParts_MergesUserTextAfterToolResult(t *testing.T) { @@ -444,42 +375,20 @@ func TestBuildConverseParts_MergesUserTextAfterToolResult(t *testing.T) { }, } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if len(parts.messages) != 3 { - t.Fatalf("expected 3 turns (user, asst, merged-user), got %d", len(parts.messages)) - } + require.NoError(t, err) + require.Len(t, parts.messages, 3) + last := parts.messages[2] - if last.Role != brtypes.ConversationRoleUser { - t.Fatalf("last role = %q, want user", last.Role) - } + require.Equal(t, brtypes.ConversationRoleUser, last.Role) + // Expect [ToolResult, Text] in the merged user message. - if len(last.Content) != 2 { - t.Fatalf("merged user message should have 2 blocks, got %d", len(last.Content)) - } - if _, ok := last.Content[0].(*brtypes.ContentBlockMemberToolResult); !ok { - t.Errorf("first block should be ToolResult, got %T", last.Content[0]) - } - tb, ok := last.Content[1].(*brtypes.ContentBlockMemberText) - if !ok { - t.Fatalf("second block should be Text, got %T", last.Content[1]) - } - if tb.Value != "thanks!" { - t.Errorf("merged text = %q, want %q", tb.Value, "thanks!") - } -} + require.Len(t, last.Content, 2) + _, ok := last.Content[0].(*brtypes.ContentBlockMemberToolResult) + assert.True(t, ok, "first block should be ToolResult, got %T", last.Content[0]) -func equalStrings(a, b []string) bool { - if len(a) != len(b) { - return false - } - for i := range a { - if a[i] != b[i] { - return false - } - } - return true + tb, ok := last.Content[1].(*brtypes.ContentBlockMemberText) + require.True(t, ok, "second block should be Text, got %T", last.Content[1]) + assert.Equal(t, "thanks!", tb.Value) } func TestBuildConverseParts_AssistantToolCallsRoundtrip(t *testing.T) { @@ -502,36 +411,24 @@ func TestBuildConverseParts_AssistantToolCallsRoundtrip(t *testing.T) { }, } parts, err := buildConverseParts(req) - if err != nil { - t.Fatalf("buildConverseParts: %v", err) - } - if len(parts.messages) != 3 { - t.Fatalf("expected 3 messages, got %d", len(parts.messages)) - } + require.NoError(t, err) + require.Len(t, parts.messages, 3) + // Assistant message should carry a ToolUse content block asst := parts.messages[1] - if asst.Role != brtypes.ConversationRoleAssistant { - t.Fatalf("expected assistant role, got %q", asst.Role) - } + require.Equal(t, brtypes.ConversationRoleAssistant, asst.Role) + tu, ok := asst.Content[0].(*brtypes.ContentBlockMemberToolUse) - if !ok { - t.Fatalf("expected tool use block, got %T", asst.Content[0]) - } - if awssdk.ToString(tu.Value.ToolUseId) != "tool_call_1" { - t.Errorf("tool use id = %q", awssdk.ToString(tu.Value.ToolUseId)) - } + require.True(t, ok, "expected tool use block, got %T", asst.Content[0]) + assert.Equal(t, "tool_call_1", awssdk.ToString(tu.Value.ToolUseId)) + // Tool result message must be sent as user role with ContentBlockMemberToolResult toolMsg := parts.messages[2] - if toolMsg.Role != brtypes.ConversationRoleUser { - t.Fatalf("expected tool result to use user role, got %q", toolMsg.Role) - } + require.Equal(t, brtypes.ConversationRoleUser, toolMsg.Role) + tr, ok := toolMsg.Content[0].(*brtypes.ContentBlockMemberToolResult) - if !ok { - t.Fatalf("expected tool result block, got %T", toolMsg.Content[0]) - } - if awssdk.ToString(tr.Value.ToolUseId) != "tool_call_1" { - t.Errorf("tool result id = %q", awssdk.ToString(tr.Value.ToolUseId)) - } + require.True(t, ok, "expected tool result block, got %T", toolMsg.Content[0]) + assert.Equal(t, "tool_call_1", awssdk.ToString(tr.Value.ToolUseId)) } func TestConvertTools_ToolChoiceNormalization(t *testing.T) { @@ -564,18 +461,14 @@ func TestConvertTools_ToolChoiceNormalization(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { cfg, err := convertTools(tools, tc.choice) - if err != nil { - t.Fatalf("convertTools: %v", err) - } + require.NoError(t, err) + if tc.wantNil { - if cfg != nil { - t.Fatalf("expected nil ToolConfiguration when tool_choice is none, got %+v", cfg) - } + require.Nil(t, cfg) return } - if cfg == nil { - t.Fatal("expected ToolConfiguration") - } + require.NotNil(t, cfg) + gotName := "" switch cfg.ToolChoice.(type) { case *brtypes.ToolChoiceMemberAuto: @@ -585,9 +478,7 @@ func TestConvertTools_ToolChoiceNormalization(t *testing.T) { case *brtypes.ToolChoiceMemberTool: gotName = "Tool" } - if gotName != tc.wantChoice { - t.Errorf("got %s, want %s", gotName, tc.wantChoice) - } + assert.Equal(t, tc.wantChoice, gotName) }) } } @@ -611,21 +502,14 @@ func TestConvertConverseOutput_TextAndUsage(t *testing.T) { Usage: usage, } resp := convertConverseOutput("anthropic.claude-3-5-haiku-20241022-v1:0", out) - if resp.Provider != providerName { - t.Errorf("provider = %q", resp.Provider) - } - if len(resp.Choices) != 1 { - t.Fatalf("expected 1 choice, got %d", len(resp.Choices)) - } - if resp.Choices[0].FinishReason != "stop" { - t.Errorf("finish_reason = %q", resp.Choices[0].FinishReason) - } - if got := core.ExtractTextContent(resp.Choices[0].Message.Content); got != "Hello there" { - t.Errorf("content = %q", got) - } - if resp.Usage.PromptTokens != 10 || resp.Usage.CompletionTokens != 20 || resp.Usage.TotalTokens != 30 { - t.Errorf("usage = %+v", resp.Usage) - } + assert.Equal(t, providerName, resp.Provider) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "stop", resp.Choices[0].FinishReason) + got := core.ExtractTextContent(resp.Choices[0].Message.Content) + assert.Equal(t, "Hello there", got) + assert.Equal(t, 10, resp.Usage.PromptTokens) + assert.Equal(t, 20, resp.Usage.CompletionTokens) + assert.Equal(t, 30, resp.Usage.TotalTokens) } func TestConvertConverseOutput_ToolUseRoundtripsArguments(t *testing.T) { @@ -648,23 +532,16 @@ func TestConvertConverseOutput_ToolUseRoundtripsArguments(t *testing.T) { Usage: &brtypes.TokenUsage{InputTokens: awssdk.Int32(1), OutputTokens: awssdk.Int32(1), TotalTokens: awssdk.Int32(2)}, } resp := convertConverseOutput("model", out) - if resp.Choices[0].FinishReason != "tool_calls" { - t.Fatalf("finish_reason = %q", resp.Choices[0].FinishReason) - } + require.Equal(t, "tool_calls", resp.Choices[0].FinishReason) + calls := resp.Choices[0].Message.ToolCalls - if len(calls) != 1 { - t.Fatalf("expected 1 tool call, got %d", len(calls)) - } - if calls[0].Function.Name != "get_weather" { - t.Errorf("name = %q", calls[0].Function.Name) - } + require.Len(t, calls, 1) + assert.Equal(t, "get_weather", calls[0].Function.Name) + var args map[string]any - if err := json.Unmarshal([]byte(calls[0].Function.Arguments), &args); err != nil { - t.Fatalf("arguments not valid JSON: %v (%s)", err, calls[0].Function.Arguments) - } - if args["city"] != "Paris" { - t.Errorf("args = %v", args) - } + err := json.Unmarshal([]byte(calls[0].Function.Arguments), &args) + require.NoError(t, err, "arguments not valid JSON: %s", calls[0].Function.Arguments) + assert.Equal(t, "Paris", args["city"]) } // TestNew_BearerTokenOnly ensures the provider initializes when the only AWS @@ -685,61 +562,39 @@ func TestNew_BearerTokenOnly(t *testing.T) { t.Setenv("AWS_SHARED_CREDENTIALS_FILE", "/dev/null") p := New(providers.ProviderConfig{BaseURL: "us-east-1"}, providers.ProviderOptions{}).(*Provider) - if p.configErr != nil { - t.Fatalf("configErr = %v, want nil", p.configErr) - } - if p.runtime == nil { - t.Fatal("runtime client should be constructed") - } - if p.region != "us-east-1" { - t.Errorf("region = %q, want us-east-1", p.region) - } - if err := p.ready(); err != nil { - t.Errorf("ready() = %v, want nil", err) - } + require.NoError(t, p.configErr) + require.NotNil(t, p.runtime) + assert.Equal(t, "us-east-1", p.region) + err := p.ready() + assert.NoError(t, err) } func TestRegistration(t *testing.T) { - if Registration.Type != providerName { - t.Errorf("Registration.Type = %q, want %q", Registration.Type, providerName) - } - if !Registration.Discovery.AllowAPIKeyless { - t.Error("Registration.Discovery.AllowAPIKeyless should be true") - } - if Registration.New == nil { - t.Fatal("Registration.New should not be nil") - } + assert.Equal(t, providerName, Registration.Type) + assert.True(t, Registration.Discovery.AllowAPIKeyless) + require.NotNil(t, Registration.New) } func TestStreamConverter_FormatChunkContent(t *testing.T) { sc := newOpenAIStream(nil, "test-model") chunk := sc.formatChunk(map[string]any{"content": "Hi"}, nil, nil) - if !strings.HasPrefix(chunk, "data: ") || !strings.HasSuffix(chunk, "\n\n") { - t.Fatalf("malformed SSE framing: %q", chunk) - } + require.True(t, strings.HasPrefix(chunk, "data: ")) + require.True(t, strings.HasSuffix(chunk, "\n\n"), "malformed SSE framing: %q", chunk) + payload := strings.TrimSuffix(strings.TrimPrefix(chunk, "data: "), "\n\n") var parsed map[string]any - if err := json.Unmarshal([]byte(payload), &parsed); err != nil { - t.Fatalf("payload not JSON: %v", err) - } - if parsed["object"] != "chat.completion.chunk" { - t.Errorf("object = %v", parsed["object"]) - } - if parsed["model"] != "test-model" { - t.Errorf("model = %v", parsed["model"]) - } - if parsed["provider"] != providerName { - t.Errorf("provider = %v", parsed["provider"]) - } + err := json.Unmarshal([]byte(payload), &parsed) + require.NoError(t, err) + assert.Equal(t, "chat.completion.chunk", parsed["object"]) + assert.Equal(t, "test-model", parsed["model"]) + assert.Equal(t, providerName, parsed["provider"]) + choices, _ := parsed["choices"].([]any) - if len(choices) != 1 { - t.Fatalf("expected 1 choice, got %d", len(choices)) - } + require.Len(t, choices, 1) + choice := choices[0].(map[string]any) delta := choice["delta"].(map[string]any) - if delta["content"] != "Hi" { - t.Errorf("delta.content = %v", delta["content"]) - } + assert.Equal(t, "Hi", delta["content"]) } // TestStreamConverter_DeferredFinishWithUsage asserts that messageStop alone @@ -751,12 +606,9 @@ func TestStreamConverter_DeferredFinishWithUsage(t *testing.T) { sc.handleEvent(&brtypes.ConverseStreamOutputMemberMessageStop{ Value: brtypes.MessageStopEvent{StopReason: brtypes.StopReasonEndTurn}, }) - if len(sc.buf) != 0 { - t.Fatalf("messageStop should not emit yet, got %q", string(sc.buf)) - } - if !sc.havePendingStop || sc.finishSent { - t.Fatalf("expected pending finish, sent=%v pending=%v", sc.finishSent, sc.havePendingStop) - } + require.Empty(t, sc.buf) + require.True(t, sc.havePendingStop) + require.False(t, sc.finishSent) sc.handleEvent(&brtypes.ConverseStreamOutputMemberMetadata{ Value: brtypes.ConverseStreamMetadataEvent{ @@ -767,27 +619,20 @@ func TestStreamConverter_DeferredFinishWithUsage(t *testing.T) { }, }, }) - if !sc.finishSent { - t.Fatal("metadata should have flushed the finish chunk") - } + require.True(t, sc.finishSent) payload := strings.TrimSuffix(strings.TrimPrefix(string(sc.buf), "data: "), "\n\n") var parsed map[string]any - if err := json.Unmarshal([]byte(payload), &parsed); err != nil { - t.Fatalf("payload not JSON: %v (%q)", err, payload) - } + err := json.Unmarshal([]byte(payload), &parsed) + require.NoError(t, err, "payload not JSON: %q", payload) + choices := parsed["choices"].([]any) choice := choices[0].(map[string]any) - if choice["finish_reason"] != "stop" { - t.Errorf("finish_reason = %v, want stop", choice["finish_reason"]) - } + assert.Equal(t, "stop", choice["finish_reason"]) + usage, ok := parsed["usage"].(map[string]any) - if !ok { - t.Fatalf("usage missing from finish chunk: %v", parsed) - } - if usage["total_tokens"].(float64) != 24 { - t.Errorf("total_tokens = %v, want 24", usage["total_tokens"]) - } + require.True(t, ok, "usage missing from finish chunk: %v", parsed) + assert.Equal(t, float64(24), usage["total_tokens"]) } // TestStreamConverter_DeferredFinishWithoutMetadata asserts that we still @@ -798,25 +643,20 @@ func TestStreamConverter_DeferredFinishWithoutMetadata(t *testing.T) { sc.handleEvent(&brtypes.ConverseStreamOutputMemberMessageStop{ Value: brtypes.MessageStopEvent{StopReason: brtypes.StopReasonMaxTokens}, }) - if sc.finishSent { - t.Fatal("finish should still be deferred") - } + require.False(t, sc.finishSent) + sc.flushFinish() - if !sc.finishSent { - t.Fatal("flushFinish should have sent the chunk") - } + require.True(t, sc.finishSent) + payload := strings.TrimSuffix(strings.TrimPrefix(string(sc.buf), "data: "), "\n\n") var parsed map[string]any - if err := json.Unmarshal([]byte(payload), &parsed); err != nil { - t.Fatalf("payload not JSON: %v", err) - } - if _, ok := parsed["usage"]; ok { - t.Errorf("usage should be absent when metadata never arrived") - } + err := json.Unmarshal([]byte(payload), &parsed) + require.NoError(t, err) + _, ok := parsed["usage"] + assert.False(t, ok) + choice := parsed["choices"].([]any)[0].(map[string]any) - if choice["finish_reason"] != "length" { - t.Errorf("finish_reason = %v, want length", choice["finish_reason"]) - } + assert.Equal(t, "length", choice["finish_reason"]) } func TestStreamConverter_FormatChunkUsage(t *testing.T) { @@ -828,38 +668,28 @@ func TestStreamConverter_FormatChunkUsage(t *testing.T) { }) payload := strings.TrimSuffix(strings.TrimPrefix(chunk, "data: "), "\n\n") var parsed map[string]any - if err := json.Unmarshal([]byte(payload), &parsed); err != nil { - t.Fatalf("payload not JSON: %v", err) - } + err := json.Unmarshal([]byte(payload), &parsed) + require.NoError(t, err) + usage, ok := parsed["usage"].(map[string]any) - if !ok { - t.Fatalf("usage missing: %v", parsed) - } - if usage["total_tokens"].(float64) != 10 { - t.Errorf("total_tokens = %v", usage["total_tokens"]) - } + require.True(t, ok, "usage missing: %v", parsed) + assert.Equal(t, float64(10), usage["total_tokens"]) } func TestGatewayCachePointUsesJSONBooleanAndNeverCreatesEmptyMessage(t *testing.T) { fields := core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ core.GatewayCachePointField: json.RawMessage(" true\n"), }) - if !isGatewayCachePoint(fields) { - t.Fatal("formatted JSON true was not recognized") - } + require.True(t, isGatewayCachePoint(fields)) + falseFields := core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ core.GatewayCachePointField: json.RawMessage(`false`), }) - if isGatewayCachePoint(falseFields) { - t.Fatal("JSON false was recognized as a cache point") - } + require.False(t, isGatewayCachePoint(falseFields)) + _, messages, err := convertMessages([]core.Message{{Role: "assistant", ExtraFields: fields}}) - if err != nil { - t.Fatal(err) - } - if len(messages) != 0 { - t.Fatalf("empty assistant created a cache-only message: %+v", messages) - } + require.NoError(t, err) + require.Empty(t, messages) } type testBedrockAPIError struct{ code, message string } @@ -869,12 +699,9 @@ func (e testBedrockAPIError) ErrorCode() string { return e.code } func (e testBedrockAPIError) ErrorMessage() string { return e.message } func TestCachePointFallbackIsNarrowAndLossless(t *testing.T) { - if !isCachePointValidationError(testBedrockAPIError{"ValidationException", "cache point below minimum tokens"}) { - t.Fatal("cache-point validation error was not recognized") - } - if isCachePointValidationError(testBedrockAPIError{"ValidationException", "invalid tool schema"}) { - t.Fatal("unrelated validation error would trigger a retry") - } + require.True(t, isCachePointValidationError(testBedrockAPIError{"ValidationException", "cache point below minimum tokens"})) + require.False(t, isCachePointValidationError(testBedrockAPIError{"ValidationException", "invalid tool schema"})) + parts := converseParts{ system: []brtypes.SystemContentBlock{ &brtypes.SystemContentBlockMemberText{Value: "system"}, @@ -886,12 +713,11 @@ func TestCachePointFallbackIsNarrowAndLossless(t *testing.T) { }}}, } clean := withoutCachePoints(parts) - if partsHaveCachePoints(clean) || len(clean.system) != 1 || len(clean.messages) != 1 || len(clean.messages[0].Content) != 1 { - t.Fatalf("cache-point removal damaged request content: %+v", clean) - } - if !partsHaveCachePoints(parts) { - t.Fatal("cache-point removal mutated original parts") - } + require.False(t, partsHaveCachePoints(clean)) + require.Len(t, clean.system, 1) + require.Len(t, clean.messages, 1) + require.Len(t, clean.messages[0].Content, 1, "cache-point removal damaged request content") + require.True(t, partsHaveCachePoints(parts)) } func TestStreamConverter_FormatChunkForwardsCacheUsage(t *testing.T) { @@ -919,13 +745,12 @@ func TestStreamConverter_FormatChunkForwardsCacheUsage(t *testing.T) { }) payload := strings.TrimSuffix(strings.TrimPrefix(chunk, "data: "), "\n\n") var parsed map[string]any - if err := json.Unmarshal([]byte(payload), &parsed); err != nil { - t.Fatalf("payload not JSON: %v", err) - } + err := json.Unmarshal([]byte(payload), &parsed) + require.NoError(t, err) + usage, ok := parsed["usage"].(map[string]any) - if !ok { - t.Fatalf("usage missing: %v", parsed) - } + require.True(t, ok, "usage missing: %v", parsed) + assertCacheKey(t, usage, "cache_read_input_tokens", tc.wantRead) assertCacheKey(t, usage, "cache_creation_input_tokens", tc.wantCreate) }) @@ -936,20 +761,11 @@ func TestStreamConverter_FormatChunkForwardsCacheUsage(t *testing.T) { // absent, otherwise the key must be present with that value. func assertCacheKey(t *testing.T, usage map[string]any, key string, want int) { t.Helper() - got, ok := usage[key] if want < 0 { - if ok { - t.Errorf("%s = %v, want absent", key, got) - } - return - } - if !ok { - t.Errorf("%s missing, want %d", key, want) + assert.NotContains(t, usage, key) return } - if got.(float64) != float64(want) { - t.Errorf("%s = %v, want %d", key, got, want) - } + assert.Equal(t, float64(want), usage[key], "%s", key) } func TestBedrockUsageExtrasCacheKeys(t *testing.T) { @@ -987,20 +803,14 @@ func TestBedrockUsageExtrasCacheKeys(t *testing.T) { CacheWriteInputTokens: tc.write, }) if len(tc.wantKeys) == 0 { - if out != nil { - t.Fatalf("extras = %v, want nil when no cache counters are set", out) - } + require.Nil(t, out) return } for key, want := range tc.wantKeys { - if out[key] != want { - t.Errorf("%s = %v, want %d", key, out[key], want) - } + assert.Equal(t, want, out[key]) } for _, key := range tc.absentKeys { - if got, ok := out[key]; ok { - t.Errorf("%s = %v, want absent", key, got) - } + assert.NotContains(t, out, key) } }) } diff --git a/internal/providers/bedrockmantle/bedrock_mantle_test.go b/internal/providers/bedrockmantle/bedrock_mantle_test.go index 60264f893..962419098 100644 --- a/internal/providers/bedrockmantle/bedrock_mantle_test.go +++ b/internal/providers/bedrockmantle/bedrock_mantle_test.go @@ -10,52 +10,38 @@ import ( "testing" "github.com/aws/aws-sdk-go-v2/aws" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) func TestResponsesUsesOpenAIPathForGPT56(t *testing.T) { - var path, authorization string - var body map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - path = r.URL.Path - authorization = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&body); err != nil { - t.Errorf("decode request: %v", err) - } - w.Header().Set("Content-Type", "application/json") - _, _ = io.WriteString(w, `{"id":"resp_1","object":"response","model":"openai.gpt-5.6-sol","status":"completed","output":[]}`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"resp_1","object":"response","model":"openai.gpt-5.6-sol","status":"completed","output":[]}`) p := testProvider(t, server, modeAuto, providers.NewKeyring("secret", "second-secret"), nil) var req core.ResponsesRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"openai.gpt-5.6-sol", "input":"hello", "previous_response_id":"resp_previous", "custom_bedrock_option":true - }`), &req); err != nil { - t.Fatal(err) - } + }`), &req) + require.NoError(t, err) resp, err := p.Responses(context.Background(), &req) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if resp.ID != "resp_1" { - t.Errorf("response ID = %q, want resp_1", resp.ID) - } - if path != "/openai/v1/responses" { - t.Errorf("path = %q, want /openai/v1/responses", path) - } - if authorization != "Bearer secret" { - t.Errorf("Authorization = %q, want first bearer token", authorization) - } - if body["previous_response_id"] != "resp_previous" || body["custom_bedrock_option"] != true { - t.Errorf("request body did not preserve Responses fields: %#v", body) - } + require.NoError(t, err) + assert.Equal(t, "resp_1", resp.ID) + + sent := capture.Last(t) + assert.Equal(t, "/openai/v1/responses", sent.Path) + assert.Equal(t, "Bearer secret", sent.Header.Get("Authorization")) + body := sent.JSON(t) + assert.Equal(t, "resp_previous", body["previous_response_id"]) + custom, _ := body["custom_bedrock_option"].(bool) + assert.True(t, custom, "request body did not preserve Responses fields: %#v", body) } func TestMantleEndpointRouting(t *testing.T) { @@ -96,86 +82,50 @@ func TestMantleEndpointRouting(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var path string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - path = r.URL.Path + server, capture := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") - switch { - case strings.HasSuffix(path, "/chat/completions"): + if strings.HasSuffix(r.URL.Path, "/chat/completions") { _, _ = io.WriteString(w, `{"id":"chat_1","object":"chat.completion","choices":[]}`) - default: - _, _ = io.WriteString(w, `{"id":"resp_1","object":"response","status":"completed","output":[]}`) + return } - })) - defer server.Close() + _, _ = io.WriteString(w, `{"id":"resp_1","object":"response","status":"completed","output":[]}`) + }) p := testProvider(t, server, tt.mode, providers.NewKeyring("secret"), nil) - if err := tt.call(context.Background(), p); err != nil { - t.Fatalf("request error = %v", err) - } - if path != tt.wantPath { - t.Errorf("path = %q, want %q", path, tt.wantPath) - } + require.NoError(t, tt.call(context.Background(), p)) + assert.Equal(t, tt.wantPath, capture.Last(t).Path) }) } } func TestListModelsAlwaysUsesCatalogPath(t *testing.T) { - var path string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - path = r.URL.Path - _, _ = io.WriteString(w, `{"data":[{"id":"openai.gpt-5.6-luna"}]}`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":[{"id":"openai.gpt-5.6-luna"}]}`) p := testProvider(t, server, modeOpenAI, providers.NewKeyring("secret"), nil) models, err := p.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if path != "/v1/models" { - t.Errorf("path = %q, want /v1/models", path) - } - if models.Object != "list" || len(models.Data) != 1 || models.Data[0].Object != "model" { - t.Errorf("models were not normalized: %#v", models) - } + require.NoError(t, err) + assert.Equal(t, "/v1/models", capture.Last(t).Path) + assert.Equal(t, "list", models.Object) + require.Len(t, models.Data, 1) + assert.Equal(t, "model", models.Data[0].Object, "models were not normalized") } func TestStreamResponsesUsesOpenAIPath(t *testing.T) { - var path string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - path = r.URL.Path - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n") - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n") p := testProvider(t, server, modeAuto, providers.NewKeyring("secret"), nil) stream, err := p.StreamResponses(context.Background(), &core.ResponsesRequest{Model: "openai.gpt-5.6-luna", Input: "hello"}) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer func() { _ = stream.Close() }() + body, err := io.ReadAll(stream) - if err != nil { - t.Fatal(err) - } - if path != "/openai/v1/responses" { - t.Errorf("path = %q, want /openai/v1/responses", path) - } - if !strings.Contains(string(body), "[DONE]") { - t.Errorf("stream = %q, want terminal [DONE]", body) - } + require.NoError(t, err) + assert.Equal(t, "/openai/v1/responses", capture.Last(t).Path) + assert.Contains(t, string(body), "[DONE]", "stream must end with the terminal marker") } func TestSigV4AuthenticationUsesBedrockService(t *testing.T) { - var authorization, securityToken string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - authorization = r.Header.Get("Authorization") - securityToken = r.Header.Get("X-Amz-Security-Token") - _, _ = io.WriteString(w, `{"data":[]}`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":[]}`) provider := aws.CredentialsProviderFunc(func(context.Context) (aws.Credentials, error) { return aws.Credentials{ @@ -185,50 +135,39 @@ func TestSigV4AuthenticationUsesBedrockService(t *testing.T) { }, nil }) p := testProvider(t, server, modeAuto, nil, provider) - if _, err := p.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if !strings.HasPrefix(authorization, "AWS4-HMAC-SHA256 ") { - t.Fatalf("Authorization = %q, want SigV4", authorization) - } - if !strings.Contains(authorization, "Credential=AKID/") || !strings.Contains(authorization, "/us-east-1/bedrock/aws4_request") { - t.Errorf("Authorization has wrong credential scope: %q", authorization) - } - if securityToken != "session-token" { - t.Errorf("X-Amz-Security-Token = %q", securityToken) - } + _, err := p.ListModels(context.Background()) + require.NoError(t, err) + + sent := capture.Last(t) + authorization := sent.Header.Get("Authorization") + require.True(t, strings.HasPrefix(authorization, "AWS4-HMAC-SHA256 "), "Authorization = %q, want SigV4", authorization) + assert.Contains(t, authorization, "Credential=AKID/") + assert.Contains(t, authorization, "/us-east-1/bedrock/aws4_request") + assert.Equal(t, "session-token", sent.Header.Get("X-Amz-Security-Token")) } func TestBearerAuthenticationRotatesKeys(t *testing.T) { - var authorizations []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - authorizations = append(authorizations, r.Header.Get("Authorization")) - _, _ = io.WriteString(w, `{"data":[]}`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":[]}`) p := testProvider(t, server, modeAuto, providers.NewKeyring("first", "second"), nil) for range 2 { - if _, err := p.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } - } - if len(authorizations) != 2 || authorizations[0] != "Bearer first" || authorizations[1] != "Bearer second" { - t.Errorf("Authorization headers = %v", authorizations) + _, err := p.ListModels(context.Background()) + require.NoError(t, err) } + requests := capture.All() + require.Len(t, requests, 2) + assert.Equal(t, "Bearer first", requests[0].Header.Get("Authorization")) + assert.Equal(t, "Bearer second", requests[1].Header.Get("Authorization")) } func TestProviderDoesNotAdvertiseUnsupportedOpenAISurfaces(t *testing.T) { p := &Provider{} - if _, ok := any(p).(core.NativeResponseLifecycleProvider); ok { - t.Error("Bedrock Mantle unexpectedly implements response lifecycle APIs") - } - if _, ok := any(p).(core.NativeBatchProvider); ok { - t.Error("Bedrock Mantle unexpectedly implements batch APIs") - } - if _, ok := any(p).(core.NativeFileProvider); ok { - t.Error("Bedrock Mantle unexpectedly implements file APIs") - } + _, ok := any(p).(core.NativeResponseLifecycleProvider) + assert.False(t, ok) + _, ok = any(p).(core.NativeBatchProvider) + assert.False(t, ok) + _, ok = any(p).(core.NativeFileProvider) + assert.False(t, ok) } func testProvider(t *testing.T, server *httptest.Server, mode string, keys *providers.Keyring, credentialsProvider aws.CredentialsProvider) *Provider { diff --git a/internal/providers/bedrockmantle/config_test.go b/internal/providers/bedrockmantle/config_test.go index d96544b7c..8ea642747 100644 --- a/internal/providers/bedrockmantle/config_test.go +++ b/internal/providers/bedrockmantle/config_test.go @@ -1,6 +1,11 @@ package bedrockmantle -import "testing" +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) func TestResolveEndpoint(t *testing.T) { tests := []struct { @@ -42,12 +47,10 @@ func TestResolveEndpoint(t *testing.T) { t.Setenv("AWS_REGION", "") t.Setenv("AWS_DEFAULT_REGION", "") got, err := resolveEndpoint(tt.baseURL, tt.apiMode) - if err != nil { - t.Fatalf("resolveEndpoint() error = %v", err) - } - if got.baseURL != tt.wantURL || got.region != tt.wantRegion || got.mode != tt.wantMode { - t.Errorf("resolveEndpoint() = %+v, want URL %q, region %q, mode %q", got, tt.wantURL, tt.wantRegion, tt.wantMode) - } + require.NoError(t, err) + assert.Equal(t, tt.wantURL, got.baseURL) + assert.Equal(t, tt.wantRegion, got.region) + assert.Equal(t, tt.wantMode, got.mode) }) } } @@ -62,17 +65,15 @@ func TestResolveEndpointRejectsInvalidConfiguration(t *testing.T) { {baseURL: "us-east-1", apiMode: "legacy"}, } for _, tt := range tests { - if _, err := resolveEndpoint(tt.baseURL, tt.apiMode); err == nil { - t.Errorf("resolveEndpoint(%q, %q) error = nil", tt.baseURL, tt.apiMode) - } + _, err := resolveEndpoint(tt.baseURL, tt.apiMode) + assert.Error(t, err, "resolveEndpoint(%q, %q)", tt.baseURL, tt.apiMode) } } func TestResolveEndpointRejectsInvalidEnvironmentRegion(t *testing.T) { t.Setenv("BEDROCK_MANTLE_REGION", "not-a-region") - if _, err := resolveEndpoint("", ""); err == nil { - t.Fatal("resolveEndpoint() error = nil") - } + _, err := resolveEndpoint("", "") + require.Error(t, err) } func TestUsesOpenAIPath(t *testing.T) { @@ -89,8 +90,6 @@ func TestUsesOpenAIPath(t *testing.T) { {model: "amazon.nova-2-lite-v1:0", want: false}, } for _, tt := range tests { - if got := usesOpenAIPath(tt.model); got != tt.want { - t.Errorf("usesOpenAIPath(%q) = %v, want %v", tt.model, got, tt.want) - } + assert.Equal(t, tt.want, usesOpenAIPath(tt.model), "usesOpenAIPath(%q)", tt.model) } } diff --git a/internal/providers/chatgpt/chatgpt_test.go b/internal/providers/chatgpt/chatgpt_test.go index 1b28e5262..1fa58acc7 100644 --- a/internal/providers/chatgpt/chatgpt_test.go +++ b/internal/providers/chatgpt/chatgpt_test.go @@ -3,18 +3,18 @@ package chatgpt import ( "context" "encoding/base64" - "errors" "io" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) // codexSSE is a minimal Codex-backend stream: one text delta and the terminal @@ -34,44 +34,23 @@ func tokenWithAccount(t *testing.T, accountID string) string { accountIDClaim: map[string]string{"chatgpt_account_id": accountID}, "exp": 1787235658, }) - if err != nil { - t.Fatalf("marshal claims: %v", err) - } + require.NoError(t, err) + enc := base64.RawURLEncoding.EncodeToString return enc([]byte(`{"alg":"none"}`)) + "." + enc(payload) + ".sig" } func TestRegistration_TypeIsChatGPT(t *testing.T) { - if Registration.Type != "chatgpt" { - t.Errorf("Registration.Type = %q, want %q", Registration.Type, "chatgpt") - } - if Registration.New == nil { - t.Error("Registration.New should not be nil") - } - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Errorf("DefaultBaseURL = %q, want %q", Registration.Discovery.DefaultBaseURL, defaultBaseURL) - } + assert.Equal(t, "chatgpt", Registration.Type) + assert.NotNil(t, Registration.New) + assert.Equal(t, defaultBaseURL, Registration.Discovery.DefaultBaseURL) } // TestStreamResponses_SendsCodexDialect locks the wire contract: the ChatGPT // Codex backend requires stream/store pinned, rejects public Responses // parameters it does not implement, and needs a list-shaped input. func TestStreamResponses_SendsCodexDialect(t *testing.T) { - var gotPath string - var gotHeader http.Header - var gotBody map[string]any - - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotHeader = r.Header.Clone() - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - t.Errorf("decode body: %v", err) - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, codexSSE) - })) - defer srv.Close() - + srv, capture := providertest.SSEServer(t, codexSSE) token := tokenWithAccount(t, "acct-123") provider := NewWithHTTPClient(token, srv.URL, srv.Client(), llmclient.Hooks{}) @@ -90,84 +69,64 @@ func TestStreamResponses_SendsCodexDialect(t *testing.T) { Include: []string{"reasoning.encrypted_content"}, Reasoning: &core.Reasoning{Effort: "low"}, }) - if err != nil { - t.Fatalf("StreamResponses: %v", err) - } + require.NoError(t, err) + defer func() { _ = stream.Close() }() - if _, err := io.ReadAll(stream); err != nil { - t.Fatalf("read stream: %v", err) - } + _, err = io.ReadAll(stream) + require.NoError(t, err) - if gotPath != "/responses" { - t.Errorf("path = %q, want /responses", gotPath) - } - if got := gotHeader.Get("Authorization"); got != "Bearer "+token { - t.Errorf("Authorization header not forwarded") - } - if got := gotHeader.Get("chatgpt-account-id"); got != "acct-123" { - t.Errorf("chatgpt-account-id = %q, want acct-123", got) - } + sent := capture.Last(t) + assert.Equal(t, "/responses", sent.Path) + assert.Equal(t, "Bearer "+token, sent.Header.Get("Authorization")) + assert.Equal(t, "acct-123", sent.Header.Get("chatgpt-account-id")) + + gotBody := sent.JSON(t) + streamed, _ := gotBody["stream"].(bool) + assert.True(t, streamed) + require.Contains(t, gotBody, "store") + stored, _ := gotBody["store"].(bool) + assert.False(t, stored) + assert.Equal(t, "You are Codex.", gotBody["instructions"]) - if gotBody["stream"] != true { - t.Errorf("stream = %v, want true", gotBody["stream"]) - } - if gotBody["store"] != false { - t.Errorf("store = %v, want false", gotBody["store"]) - } - if gotBody["instructions"] != "You are Codex." { - t.Errorf("instructions = %v", gotBody["instructions"]) - } for _, field := range []string{"temperature", "max_output_tokens", "previous_response_id", "truncation", "user", "metadata", "top_p", "service_tier"} { - if _, ok := gotBody[field]; ok { - t.Errorf("%s must not be sent to the Codex backend", field) - } + assert.NotContains(t, gotBody, field, "%s must not be sent to the Codex backend", field) } input, ok := gotBody["input"].([]any) - if !ok || len(input) != 1 { - t.Fatalf("input = %#v, want a one-element list", gotBody["input"]) - } + require.True(t, ok) + require.Len(t, input, 1) + msg, _ := input[0].(map[string]any) - if msg["role"] != "user" || msg["type"] != "message" { - t.Errorf("input[0] = %#v, want a user message", msg) - } + assert.Equal(t, "user", msg["role"]) + assert.Equal(t, "message", msg["type"]) + // The Responses API spells input content "input_text"; core.ContentPart // would have rewritten it to the Chat Completions "text". parts, ok := msg["content"].([]any) - if !ok || len(parts) != 1 { - t.Fatalf("content = %#v, want one part", msg["content"]) - } + require.True(t, ok) + require.Len(t, parts, 1) + part, _ := parts[0].(map[string]any) - if part["type"] != "input_text" || part["text"] != "Reply with exactly ok" { - t.Errorf("content[0] = %#v, want an input_text part", part) - } + assert.Equal(t, "input_text", part["type"]) + assert.Equal(t, "Reply with exactly ok", part["text"]) } // TestResponses_CollapsesUpstreamStream covers the non-streaming path: the // backend refuses stream:false, so GoModel streams and returns the final object. func TestResponses_CollapsesUpstreamStream(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, codexSSE) - })) - defer srv.Close() - + srv, _ := providertest.SSEServer(t, codexSSE) provider := NewWithHTTPClient("token", srv.URL, srv.Client(), llmclient.Hooks{}) resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "gpt-5.6-terra", Input: []core.ResponsesInputElement{{Type: "message", Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("Responses: %v", err) - } - if resp.Status != "completed" || resp.ID != "resp_1" { - t.Errorf("resp = %+v, want completed resp_1", resp) - } - if len(resp.Output) != 1 || len(resp.Output[0].Content) != 1 || resp.Output[0].Content[0].Text != "ok" { - t.Errorf("output = %+v, want a single 'ok' text item", resp.Output) - } - if resp.Usage == nil || resp.Usage.TotalTokens != 5 { - t.Errorf("usage = %+v, want total_tokens 5", resp.Usage) - } + require.NoError(t, err) + assert.Equal(t, "completed", resp.Status) + assert.Equal(t, "resp_1", resp.ID) + require.Len(t, resp.Output, 1) + require.Len(t, resp.Output[0].Content, 1) + assert.Equal(t, "ok", resp.Output[0].Content[0].Text) + require.NotNil(t, resp.Usage) + assert.Equal(t, 5, resp.Usage.TotalTokens) } // TestResponses_TruncatedStreamIsAnError guards the non-streaming path against @@ -194,20 +153,11 @@ func TestResponses_TruncatedStreamIsAnError(t *testing.T) { } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, tc.body) - })) - defer srv.Close() - + srv, _ := providertest.SSEServer(t, tc.body) provider := NewWithHTTPClient("token", srv.URL, srv.Client(), llmclient.Hooks{}) - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{Model: "gpt-5.6-terra", Input: "hi"}) - if err == nil { - t.Fatalf("expected an error, got response %+v", resp) - } - if !strings.Contains(err.Error(), tc.want) { - t.Errorf("error = %q, want it to mention %q", err, tc.want) - } + _, err := provider.Responses(context.Background(), &core.ResponsesRequest{Model: "gpt-5.6-terra", Input: "hi"}) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.want) }) } } @@ -221,12 +171,14 @@ func TestResponses_NonSuccessTerminalIsReturnedAsAResponse(t *testing.T) { event string payload string wantStatus string + wantError string // upstream error message the response must carry, "" for none }{ { name: "failed", event: "response.failed", payload: `{"type":"response.failed","response":{"id":"resp_1","object":"response","status":"failed","model":"gpt-5.6-terra","error":{"code":"server_error","message":"boom"}}}`, wantStatus: "failed", + wantError: "boom", }, { name: "incomplete", @@ -237,50 +189,26 @@ func TestResponses_NonSuccessTerminalIsReturnedAsAResponse(t *testing.T) { } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "event: "+tc.event+"\ndata: "+tc.payload+"\n\n") - })) - defer srv.Close() - + srv, _ := providertest.SSEServer(t, "event: "+tc.event+"\ndata: "+tc.payload+"\n\n") provider := NewWithHTTPClient("token", srv.URL, srv.Client(), llmclient.Hooks{}) resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{Model: "gpt-5.6-terra", Input: "hi"}) - if err != nil { - t.Fatalf("Responses: %v", err) - } - if resp.Status != tc.wantStatus { - t.Errorf("status = %q, want %q", resp.Status, tc.wantStatus) + require.NoError(t, err) + assert.Equal(t, tc.wantStatus, resp.Status) + if tc.wantError == "" { + assert.Nil(t, resp.Error) + return } + require.NotNil(t, resp.Error) + assert.Equal(t, tc.wantError, resp.Error.Message) }) } - t.Run("failed carries the upstream error", func(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "event: response.failed\n"+ - `data: {"type":"response.failed","response":{"id":"resp_1","object":"response","status":"failed","model":"gpt-5.6-terra","error":{"code":"server_error","message":"boom"}}}`+"\n\n") - })) - defer srv.Close() - - provider := NewWithHTTPClient("token", srv.URL, srv.Client(), llmclient.Hooks{}) - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{Model: "gpt-5.6-terra", Input: "hi"}) - if err != nil { - t.Fatalf("Responses: %v", err) - } - if resp.Error == nil || resp.Error.Message != "boom" { - t.Errorf("resp.Error = %+v, want the upstream error", resp.Error) - } - }) } func TestStreamResponses_RequiresToken(t *testing.T) { provider := NewWithHTTPClient("", "http://example.invalid", http.DefaultClient, llmclient.Hooks{}) _, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{Model: "gpt-5.6-terra", Input: "hi"}) - if err == nil { - t.Fatal("expected an authentication error without a token") - } - if !strings.Contains(err.Error(), "CHATGPT_API_KEY") { - t.Errorf("error = %q, want it to name CHATGPT_API_KEY", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "CHATGPT_API_KEY") } func TestListModels(t *testing.T) { @@ -296,16 +224,10 @@ func TestListModels(t *testing.T) { t.Run(tc.name, func(t *testing.T) { provider := New(providers.ProviderConfig{APIKey: "token"}, providers.ProviderOptions{Models: tc.configured}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels: %v", err) - } - if len(resp.Data) != len(tc.want) { - t.Fatalf("got %d models, want %d", len(resp.Data), len(tc.want)) - } + require.NoError(t, err) + require.Len(t, resp.Data, len(tc.want)) for i, model := range resp.Data { - if model.ID != tc.want[i] { - t.Errorf("model[%d] = %q, want %q", i, model.ID, tc.want[i]) - } + assert.Equal(t, tc.want[i], model.ID, "model[%d]", i) } }) } @@ -332,21 +254,16 @@ func TestUnsupportedSurfaces(t *testing.T) { for name, call := range calls { t.Run(name, func(t *testing.T) { err := call() - if err == nil { - t.Fatal("expected an unsupported-surface error") - } + require.Error(t, err) + var gatewayErr *core.GatewayError - if !errors.As(err, &gatewayErr) { - t.Fatalf("error = %T, want *core.GatewayError", err) - } - if gatewayErr.StatusCode != http.StatusNotImplemented { - t.Errorf("status = %d, want %d", gatewayErr.StatusCode, http.StatusNotImplemented) - } + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusNotImplemented, gatewayErr.StatusCode) + // The code is the programmatic half of the contract: callers // branch on it to tell a capability gap from a bad request. - if gatewayErr.Code == nil || *gatewayErr.Code != unsupportedOperationCode { - t.Errorf("code = %v, want %q", gatewayErr.Code, unsupportedOperationCode) - } + require.NotNil(t, gatewayErr.Code) + assert.Equal(t, unsupportedOperationCode, *gatewayErr.Code) }) } } @@ -364,9 +281,7 @@ func TestAccountIDFromToken(t *testing.T) { } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - if got := accountIDFromToken(tc.token); got != tc.want { - t.Errorf("accountIDFromToken() = %q, want %q", got, tc.want) - } + assert.Equal(t, tc.want, accountIDFromToken(tc.token)) }) } } diff --git a/internal/providers/chutes/chutes_test.go b/internal/providers/chutes/chutes_test.go index 5e0416d85..e1e732af3 100644 --- a/internal/providers/chutes/chutes_test.go +++ b/internal/providers/chutes/chutes_test.go @@ -2,170 +2,70 @@ package chutes import ( "context" - "encoding/json" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestNew_ConstructsRegisteredProvider(t *testing.T) { - provider, ok := New(providers.ProviderConfig{ - APIKey: "cpk_test", - BaseURL: "https://chutes.example/v1", - }, providers.ProviderOptions{}).(*Provider) - if !ok || provider.compat == nil { - t.Fatalf("New() = %T, want initialized *Provider", provider) - } - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Fatalf("registration base URL = %q, want %q", Registration.Discovery.DefaultBaseURL, defaultBaseURL) - } +// Chutes wraps the shared OpenAI-compatible adapter and translates Responses +// through chat completions, so the shared contract covers its surface. The +// shared LLM endpoint has no embeddings route, so Embeddings must fail fast +// without an upstream call, and the provider must not advertise native +// batch, file, audio, or response-lifecycle support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "chutes", + DefaultBaseURL: "https://llm.chutes.ai/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) + }, + }) + + provider := NewWithHTTPClient("cpk_test", "", nil, llmclient.Hooks{}) + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + assert.False(t, ok, "provider should not implement core.NativeResponseLifecycleProvider") } func TestSetBaseURL_ChangesRequestTarget(t *testing.T) { - var gotMethod, gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotMethod = r.Method - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) provider := NewWithHTTPClient("cpk_test", "https://unused.example/v1", server.Client(), llmclient.Hooks{}) provider.SetBaseURL(server.URL) - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotMethod != http.MethodGet || gotPath != "/models" { - t.Fatalf("method/path = %q/%q, want GET /models", gotMethod, gotPath) - } -} - -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath, gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-chutes", - "created":1677652288, - "model":"Qwen/Qwen3-32B-TEE", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "Qwen/Qwen3-32B-TEE", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer cpk_test" { - t.Fatalf("authorization = %q, want Bearer cpk_test", gotAuth) - } - if resp.Model != "Qwen/Qwen3-32B-TEE" || resp.Usage.TotalTokens != 4 { - t.Fatalf("response = %+v, want model and usage preserved", resp) - } -} - -func TestStreamChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath, gotAuth string - var gotBody struct { - Model string `json:"model"` - Stream bool `json:"stream"` - } - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-chutes\",\"object\":\"chat.completion.chunk\",\"choices\":[]}\n\ndata: [DONE]\n\n") - })) - defer server.Close() - - provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) - stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ - Model: "Qwen/Qwen3-32B-TEE", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } - defer stream.Close() + _, err := provider.ListModels(context.Background()) + require.NoError(t, err) - body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotPath != "/chat/completions" || gotAuth != "Bearer cpk_test" { - t.Fatalf("request path/auth = %q/%q, want /chat/completions/Bearer cpk_test", gotPath, gotAuth) - } - if gotBody.Model != "Qwen/Qwen3-32B-TEE" || !gotBody.Stream { - t.Fatalf("stream request body = %#v", gotBody) - } - if !strings.Contains(string(body), "data: [DONE]") { - t.Fatalf("stream body = %q, want SSE terminator", body) - } + req := capture.Last(t) + assert.Equal(t, http.MethodGet, req.Method) + assert.Equal(t, "/models", req.Path) } func TestChatCompletion_ReturnsUpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusTooManyRequests) - _, _ = w.Write([]byte(`{"error":{"message":"rate limited","type":"rate_limit_error"}}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusTooManyRequests, `{"error":{"message":"rate limited","type":"rate_limit_error"}}`) provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "Qwen/Qwen3-32B-TEE", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err == nil { - t.Fatal("ChatCompletion() error = nil, want upstream error") - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if gatewayErr.StatusCode != http.StatusTooManyRequests || gatewayErr.Type != core.ErrorTypeRateLimit { - t.Fatalf("gateway error = %+v, want 429 rate_limit_error", gatewayErr) - } + require.Error(t, err) + + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusTooManyRequests, gatewayErr.StatusCode) + assert.Equal(t, core.ErrorTypeRateLimit, gatewayErr.Type) } func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { - var gotURI, gotAuth, gotBeta, gotBody string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotURI = r.URL.RequestURI() - gotAuth = r.Header.Get("Authorization") - gotBeta = r.Header.Get("X-Chutes-Beta") - body, _ := io.ReadAll(r.Body) - gotBody = string(body) - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusAccepted) - _, _ = w.Write([]byte(`{"accepted":true}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusAccepted, `{"accepted":true}`) provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ @@ -177,133 +77,44 @@ func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { "X-Chutes-Beta": {"test"}, }, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() responseBody, err := io.ReadAll(resp.Body) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotURI != "/chat/completions?trace=true" || gotAuth != "Bearer cpk_test" { - t.Fatalf("request URI/auth = %q/%q", gotURI, gotAuth) - } - if gotBeta != "test" || gotBody != `{"model":"Qwen/Qwen3-32B-TEE"}` { - t.Fatalf("request beta/body = %q/%q", gotBeta, gotBody) - } - if resp.StatusCode != http.StatusAccepted || string(responseBody) != `{"accepted":true}` { - t.Fatalf("response status/body = %d/%q", resp.StatusCode, responseBody) - } -} - -func TestResponses_TranslatesToChatCompletions(t *testing.T) { - var gotPath string - var gotBody struct { - Model string `json:"model"` - } - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-chutes", - "created":1677652288, - "model":"Qwen/Qwen3-32B-TEE", - "choices":[{"index":0,"message":{"role":"assistant","content":"translated"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ - Model: "Qwen/Qwen3-32B-TEE", - Input: "hi", - }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotBody.Model != "Qwen/Qwen3-32B-TEE" { - t.Fatalf("request model = %q, want Qwen/Qwen3-32B-TEE", gotBody.Model) - } - if resp.Object != "response" || resp.Status != "completed" { - t.Fatalf("response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) - } + require.NoError(t, err) + assert.Equal(t, http.StatusAccepted, resp.StatusCode) + assert.Equal(t, `{"accepted":true}`, string(responseBody)) + + req := capture.Last(t) + assert.Equal(t, "/chat/completions", req.Path) + assert.Equal(t, "true", req.Query.Get("trace")) + assert.Equal(t, "Bearer cpk_test", req.Header.Get("Authorization")) + assert.Equal(t, "test", req.Header.Get("X-Chutes-Beta")) + assert.Equal(t, `{"model":"Qwen/Qwen3-32B-TEE"}`, string(req.Body)) } func TestStreamResponses_TranslatesToChatCompletions(t *testing.T) { - var gotMethod, gotPath, gotAuth string - var gotBody struct { - Model string `json:"model"` - Stream bool `json:"stream"` - } - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotMethod = r.Method - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-chutes\",\"object\":\"chat.completion.chunk\",\"created\":1677652288,\"model\":\"Qwen/Qwen3-32B-TEE\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"},\"finish_reason\":null}]}\n\ndata: [DONE]\n\n") - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "data: {\"id\":\"chatcmpl-chutes\",\"object\":\"chat.completion.chunk\",\"created\":1677652288,\"model\":\"Qwen/Qwen3-32B-TEE\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"},\"finish_reason\":null}]}\n\ndata: [DONE]\n\n") provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "Qwen/Qwen3-32B-TEE", Input: "hi", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotMethod != http.MethodPost || gotPath != "/chat/completions" || gotAuth != "Bearer cpk_test" { - t.Fatalf("request method/path/auth = %q/%q/%q", gotMethod, gotPath, gotAuth) - } - if gotBody.Model != "Qwen/Qwen3-32B-TEE" || !gotBody.Stream { - t.Fatalf("stream request body = %#v", gotBody) - } - if raw := string(body); !strings.Contains(raw, "response.output_text.delta") || !strings.Contains(raw, "data: [DONE]") { - t.Fatalf("converted stream missing Responses events or done marker: %s", raw) - } -} - -func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { - provider := NewWithHTTPClient("cpk_test", "", nil, llmclient.Hooks{}) - if _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{}); err == nil { - t.Fatal("Embeddings() error = nil, want unsupported error") - } -} - -func TestProvider_DoesNotExposeUnsupportedOptionalInterfaces(t *testing.T) { - provider := NewWithHTTPClient("cpk_test", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("chutes provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("chutes provider should not implement native file provider") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("chutes provider should not implement native response lifecycle provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("chutes provider should not implement audio provider") - } + require.NoError(t, err) + raw := string(body) + assert.Contains(t, raw, "response.output_text.delta") + assert.Contains(t, raw, "data: [DONE]") + + req := capture.Last(t) + assert.Equal(t, http.MethodPost, req.Method) + assert.Equal(t, "/chat/completions", req.Path) + assert.Equal(t, "Bearer cpk_test", req.Header.Get("Authorization")) + sent := req.JSON(t) + assert.Equal(t, "Qwen/Qwen3-32B-TEE", sent["model"]) + assert.Equal(t, true, sent["stream"]) } diff --git a/internal/providers/chutes/models_test.go b/internal/providers/chutes/models_test.go index 1904df1fd..a5b5a6cac 100644 --- a/internal/providers/chutes/models_test.go +++ b/internal/providers/chutes/models_test.go @@ -3,118 +3,97 @@ package chutes import ( "context" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestListModels_PreservesChutesMetadata(t *testing.T) { - var gotMethod, gotPath, gotAuth string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotMethod = r.Method - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"chutes-model-catalog", - "data":[{ - "id":"Qwen/Qwen3.5-397B-A17B-TEE", - "owned_by":"sglang", - "created":1677652288, - "context_length":262144, - "max_output_length":65536, - "input_modalities":["text","image"], - "supported_features":["json_mode","tools","structured_outputs","reasoning"], - "confidential_compute":true, - "pricing":{"prompt":0.45,"completion":3.0,"input_cache_read":0.045} - }] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "object":"chutes-model-catalog", + "data":[{ + "id":"Qwen/Qwen3.5-397B-A17B-TEE", + "owned_by":"sglang", + "created":1677652288, + "context_length":262144, + "max_output_length":65536, + "input_modalities":["text","image"], + "supported_features":["json_mode","tools","structured_outputs","reasoning"], + "confidential_compute":true, + "pricing":{"prompt":0.45,"completion":3.0,"input_cache_read":0.045} + }] + }`) provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotMethod != http.MethodGet || gotPath != "/models" || gotAuth != "Bearer cpk_test" { - t.Fatalf("request method/path/auth = %q/%q/%q, want GET /models/Bearer cpk_test", gotMethod, gotPath, gotAuth) - } - if len(resp.Data) != 1 { - t.Fatalf("len(resp.Data) = %d, want 1", len(resp.Data)) - } - if resp.Object != "list" { - t.Fatalf("resp.Object = %q, want list", resp.Object) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, http.MethodGet, req.Method) + assert.Equal(t, "/models", req.Path) + assert.Equal(t, "Bearer cpk_test", req.Header.Get("Authorization")) + require.Len(t, resp.Data, 1) + assert.Equal(t, "list", resp.Object) + model := resp.Data[0] - if model.Object != "model" { - t.Fatalf("model.Object = %q, want model", model.Object) - } - if model.Metadata == nil || model.Metadata.ContextWindow == nil || *model.Metadata.ContextWindow != 262144 { - t.Fatalf("model context metadata = %+v, want 262144", model.Metadata) - } - if model.Metadata.MaxOutputTokens == nil || *model.Metadata.MaxOutputTokens != 65536 { - t.Fatalf("max output tokens = %+v, want 65536", model.Metadata.MaxOutputTokens) - } - if !model.Metadata.Capabilities["tools"] || !model.Metadata.Capabilities["vision"] || !model.Metadata.Capabilities["confidential_compute"] { - t.Fatalf("capabilities = %v, want tools, vision, and confidential_compute", model.Metadata.Capabilities) - } + assert.Equal(t, "model", model.Object) + require.NotNil(t, model.Metadata) + require.NotNil(t, model.Metadata.ContextWindow) + assert.Equal(t, 262144, *model.Metadata.ContextWindow) + require.NotNil(t, model.Metadata.MaxOutputTokens) + assert.Equal(t, 65536, *model.Metadata.MaxOutputTokens) + assert.True(t, model.Metadata.Capabilities["tools"]) + assert.True(t, model.Metadata.Capabilities["vision"]) + assert.True(t, model.Metadata.Capabilities["confidential_compute"]) + pricing := model.Metadata.Pricing - if pricing == nil || pricing.Currency != "USD" || pricing.InputPerMtok == nil || *pricing.InputPerMtok != 0.45 || - pricing.OutputPerMtok == nil || *pricing.OutputPerMtok != 3.0 || - pricing.CachedInputPerMtok == nil || *pricing.CachedInputPerMtok != 0.045 { - t.Fatalf("pricing = %+v, want Chutes per-MTok USD pricing", pricing) - } + require.NotNil(t, pricing) + assert.Equal(t, "USD", pricing.Currency) + require.NotNil(t, pricing.InputPerMtok) + assert.Equal(t, 0.45, *pricing.InputPerMtok) + require.NotNil(t, pricing.OutputPerMtok) + assert.Equal(t, 3.0, *pricing.OutputPerMtok) + require.NotNil(t, pricing.CachedInputPerMtok) + assert.Equal(t, 0.045, *pricing.CachedInputPerMtok) } func TestListModels_FiltersBlankIDsAndKeepsMinimalModels(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "data":[ - {"id":" "}, - {"id":" minimal-model ","object":"model","owned_by":" chutes ","pricing":{}} - ] - }`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `{ + "data":[ + {"id":" "}, + {"id":" minimal-model ","object":"model","owned_by":" chutes ","pricing":{}} + ] + }`) provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 { - t.Fatalf("models = %+v, want one non-blank model", resp.Data) - } + require.NoError(t, err) + require.Len(t, resp.Data, 1) + model := resp.Data[0] - if model.ID != "minimal-model" || model.Object != "model" || model.OwnedBy != "chutes" { - t.Fatalf("model identity = %+v, want trimmed minimal model", model) - } - if model.Metadata.ContextWindow != nil || model.Metadata.MaxOutputTokens != nil || - model.Metadata.Capabilities != nil || model.Metadata.Pricing != nil { - t.Fatalf("optional metadata = %+v, want omitted zero values", model.Metadata) - } + assert.Equal(t, "minimal-model", model.ID) + assert.Equal(t, "model", model.Object) + assert.Equal(t, "chutes", model.OwnedBy) + assert.Nil(t, model.Metadata.ContextWindow) + assert.Nil(t, model.Metadata.MaxOutputTokens) + assert.Nil(t, model.Metadata.Capabilities) + assert.Nil(t, model.Metadata.Pricing) } func TestListModels_ReturnsUpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusServiceUnavailable) - _, _ = w.Write([]byte(`{"error":{"message":"catalog unavailable"}}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusServiceUnavailable, `{"error":{"message":"catalog unavailable"}}`) provider := NewWithHTTPClient("cpk_test", server.URL, server.Client(), llmclient.Hooks{}) _, err := provider.ListModels(context.Background()) - if err == nil { - t.Fatal("ListModels() error = nil, want upstream error") - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok || gatewayErr.StatusCode != http.StatusServiceUnavailable { - t.Fatalf("error = %#v, want 503 *core.GatewayError", err) - } + require.Error(t, err) + + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusServiceUnavailable, gatewayErr.StatusCode) } func TestModelCapabilities_MapsOptionalModalities(t *testing.T) { @@ -122,23 +101,20 @@ func TestModelCapabilities_MapsOptionalModalities(t *testing.T) { SupportedFeatures: []string{" JSON_Mode ", " "}, InputModalities: []string{"audio", "video", "unknown"}, }) - if !capabilities["json_mode"] || !capabilities["audio"] || !capabilities["video"] { - t.Fatalf("capabilities = %v, want normalized json_mode, audio, and video", capabilities) - } + assert.True(t, capabilities["json_mode"]) + assert.True(t, capabilities["audio"]) + assert.True(t, capabilities["video"]) } func TestModelPricing_HandlesNilAndPartialPrices(t *testing.T) { var absent *modelPricing - if got := absent.toCore(); got != nil { - t.Fatalf("nil pricing = %+v, want nil", got) - } + assert.Nil(t, absent.toCore()) prompt := 0.25 got := (&modelPricing{Prompt: &prompt}).toCore() - if got == nil || got.InputPerMtok == nil || *got.InputPerMtok != prompt { - t.Fatalf("partial pricing = %+v, want input price", got) - } - if got.OutputPerMtok != nil || got.CachedInputPerMtok != nil { - t.Fatalf("partial pricing = %+v, want absent optional prices", got) - } + require.NotNil(t, got) + require.NotNil(t, got.InputPerMtok) + assert.Equal(t, prompt, *got.InputPerMtok) + assert.Nil(t, got.OutputPerMtok) + assert.Nil(t, got.CachedInputPerMtok) } diff --git a/internal/providers/cohere/audio_test.go b/internal/providers/cohere/audio_test.go index 1f2765e81..e4dbdefe5 100644 --- a/internal/providers/cohere/audio_test.go +++ b/internal/providers/cohere/audio_test.go @@ -7,59 +7,58 @@ import ( "mime" "mime/multipart" "net/http" - "net/http/httptest" - "reflect" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestCreateTranscriptionTranslatesMultipartRequest(t *testing.T) { - var ( - partNames []string - fields = map[string]string{} - filename string - audio []byte - ) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v2/audio/transcriptions" { - t.Errorf("path = %q, want /v2/audio/transcriptions", r.URL.Path) +// multipartParts is what the upstream saw in a transcription request: the +// form part names in wire order, the non-file fields, and the file part. +type multipartParts struct { + names []string + fields map[string]string + filename string + audio []byte +} + +func parseMultipart(contentType string, body []byte) (multipartParts, error) { + parts := multipartParts{fields: map[string]string{}} + _, params, err := mime.ParseMediaType(contentType) + if err != nil { + return parts, err + } + reader := multipart.NewReader(bytes.NewReader(body), params["boundary"]) + for { + part, err := reader.NextPart() + if err == io.EOF { + return parts, nil } - if got := r.Header.Get("Authorization"); got != "Bearer test-key" { - t.Errorf("Authorization = %q", got) + if err != nil { + return parts, err } - - _, params, err := mime.ParseMediaType(r.Header.Get("Content-Type")) + data, err := io.ReadAll(part) if err != nil { - t.Fatalf("parse Content-Type: %v", err) + return parts, err } - reader := multipart.NewReader(r.Body, params["boundary"]) - for { - part, nextErr := reader.NextPart() - if nextErr == io.EOF { - break - } - if nextErr != nil { - t.Fatalf("read multipart: %v", nextErr) - } - partNames = append(partNames, part.FormName()) - data, readErr := io.ReadAll(part) - if readErr != nil { - t.Fatalf("read part %q: %v", part.FormName(), readErr) - } - if part.FormName() == "file" { - filename = part.FileName() - audio = data - } else { - fields[part.FormName()] = string(data) - } + parts.names = append(parts.names, part.FormName()) + if part.FormName() == "file" { + parts.filename = part.FileName() + parts.audio = data + continue } + parts.fields[part.FormName()] = string(data) + } +} +func TestCreateTranscriptionTranslatesMultipartRequest(t *testing.T) { + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json; charset=utf-8") _, _ = io.WriteString(w, `{"text":"GoModel routes requests reliably."}`) - })) - defer server.Close() + }) provider := NewWithHTTPClient("test-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ @@ -70,31 +69,25 @@ func TestCreateTranscriptionTranslatesMultipartRequest(t *testing.T) { ResponseFormat: "json", Temperature: "0.2", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } + require.NoError(t, err) - wantPartNames := []string{"model", "language", "temperature", "file"} - if !reflect.DeepEqual(partNames, wantPartNames) { - t.Fatalf("multipart fields = %#v, want %#v", partNames, wantPartNames) - } - wantFields := map[string]string{ + sent := capture.Last(t) + assert.Equal(t, "/v2/audio/transcriptions", sent.Path) + assert.Equal(t, "Bearer test-key", sent.Header.Get("Authorization")) + + parts, err := parseMultipart(sent.Header.Get("Content-Type"), sent.Body) + require.NoError(t, err) + assert.Equal(t, []string{"model", "language", "temperature", "file"}, parts.names) + assert.Equal(t, map[string]string{ "model": "cohere-transcribe-03-2026", "language": "en", "temperature": "0.2", - } - if !reflect.DeepEqual(fields, wantFields) { - t.Fatalf("multipart values = %#v, want %#v", fields, wantFields) - } - if filename != "sample.wav" || string(audio) != "wave-bytes" { - t.Fatalf("file = %q/%q", filename, audio) - } - if resp.ContentType != "application/json; charset=utf-8" { - t.Fatalf("ContentType = %q", resp.ContentType) - } - if string(resp.Data) != `{"text":"GoModel routes requests reliably."}` { - t.Fatalf("Data = %s", resp.Data) - } + }, parts.fields) + assert.Equal(t, "sample.wav", parts.filename) + assert.Equal(t, "wave-bytes", string(parts.audio)) + + assert.Equal(t, "application/json; charset=utf-8", resp.ContentType) + assert.Equal(t, `{"text":"GoModel routes requests reliably."}`, string(resp.Data)) } func TestCreateTranscriptionValidation(t *testing.T) { diff --git a/internal/providers/cohere/cohere_test.go b/internal/providers/cohere/cohere_test.go index a5fd22432..c1a83fe33 100644 --- a/internal/providers/cohere/cohere_test.go +++ b/internal/providers/cohere/cohere_test.go @@ -4,59 +4,48 @@ import ( "context" "io" "net/http" - "net/http/httptest" "strconv" "strings" - "sync" "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) +// sseBody joins SSE lines the way Cohere's stream emits them. +func sseBody(lines ...string) string { + return strings.Join(lines, "\n") +} + func TestChatCompletionTranslatesRequestAndResponse(t *testing.T) { - var captured map[string]any var operation string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v2/chat" { - t.Errorf("path = %q, want /v2/chat", r.URL.Path) - } - if got := r.Header.Get("Authorization"); got != "Bearer test-key" { - t.Errorf("Authorization = %q", got) - } - if got := r.Header.Get("X-Client-Name"); got != "GoModel" { - t.Errorf("X-Client-Name = %q", got) - } - if err := json.NewDecoder(r.Body).Decode(&captured); err != nil { - t.Errorf("decode request: %v", err) + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"cohere-id", + "finish_reason":"TOOL_CALL", + "message":{ + "role":"assistant", + "content":[ + {"type":"thinking","thinking":"check the weather"}, + {"type":"text","text":"I will check."} + ], + "tool_plan":"Use the weather tool.", + "tool_calls":[ + {"id":"call-2","type":"function","function":{"name":"weather","arguments":"{\"city\":\"Warsaw\"}"}} + ], + "citations":[{"start":0,"end":1,"text":"I"}] + }, + "usage":{ + "billed_units":{"input_tokens":10,"output_tokens":3}, + "tokens":{"input_tokens":12,"image_tokens":5,"output_tokens":4}, + "cached_tokens":2 } - w.Header().Set("Content-Type", "application/json") - _, _ = io.WriteString(w, `{ - "id":"cohere-id", - "finish_reason":"TOOL_CALL", - "message":{ - "role":"assistant", - "content":[ - {"type":"thinking","thinking":"check the weather"}, - {"type":"text","text":"I will check."} - ], - "tool_plan":"Use the weather tool.", - "tool_calls":[ - {"id":"call-2","type":"function","function":{"name":"weather","arguments":"{\"city\":\"Warsaw\"}"}} - ], - "citations":[{"start":0,"end":1,"text":"I"}] - }, - "usage":{ - "billed_units":{"input_tokens":10,"output_tokens":3}, - "tokens":{"input_tokens":12,"image_tokens":5,"output_tokens":4}, - "cached_tokens":2 - } - }`) - })) - defer server.Close() + }`) provider := NewWithHTTPClient("test-key", server.URL, server.Client(), llmclient.Hooks{ OnRequestStart: func(ctx context.Context, info llmclient.RequestInfo) context.Context { @@ -115,80 +104,57 @@ func TestChatCompletionTranslatesRequestAndResponse(t *testing.T) { } resp, err := provider.ChatCompletion(context.Background(), req) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + require.NoError(t, err) + + sent := capture.Last(t) + assert.Equal(t, "/v2/chat", sent.Path) + assert.Equal(t, "Bearer test-key", sent.Header.Get("Authorization")) + assert.Equal(t, "GoModel", sent.Header.Get("X-Client-Name")) + assert.Equal(t, llmclient.OperationChat, operation) + + captured := sent.JSON(t) + assert.Equal(t, req.Model, captured["model"]) + require.Contains(t, captured, "stream") + streamed, _ := captured["stream"].(bool) + assert.False(t, streamed) - if captured["model"] != req.Model || captured["stream"] != false { - t.Fatalf("request model/stream = %#v/%#v", captured["model"], captured["stream"]) - } messages := captured["messages"].([]any) - if messages[0].(map[string]any)["role"] != "system" { - t.Fatalf("developer role = %#v, want system", messages[0]) - } - if messages[0].(map[string]any)["content"] != "Be brief." { - t.Fatalf("developer content = %#v, want a Cohere system string", messages[0]) - } + assert.Equal(t, "system", messages[0].(map[string]any)["role"], "developer role maps to system") + assert.Equal(t, "Be brief.", messages[0].(map[string]any)["content"], "developer content becomes a Cohere system string") + userContent := messages[1].(map[string]any)["content"].([]any) image := userContent[1].(map[string]any) - if image["type"] != "image_url" { - t.Fatalf("image content = %#v", image) - } + assert.Equal(t, "image_url", image["type"]) + toolResult := messages[3].(map[string]any)["content"].(string) - if toolResult != `{"result":"sunny"}` { - t.Fatalf("tool result = %q", toolResult) - } - if captured["tool_choice"] != "REQUIRED" { - t.Fatalf("tool_choice = %#v", captured["tool_choice"]) - } - if captured["max_tokens"] != float64(250) { - t.Fatalf("max_tokens = %#v", captured["max_tokens"]) - } - if got := captured["stop_sequences"].([]any)[0]; got != "END" { - t.Fatalf("stop_sequences = %#v", captured["stop_sequences"]) - } - if captured["k"] != float64(20) { - t.Fatalf("k = %#v", captured["k"]) - } - if operation != llmclient.OperationChat { - t.Fatalf("operation = %q, want chat", operation) - } + assert.Equal(t, `{"result":"sunny"}`, toolResult) + assert.Equal(t, "REQUIRED", captured["tool_choice"]) + assert.Equal(t, float64(250), captured["max_tokens"]) + assert.Equal(t, []any{"END"}, captured["stop_sequences"]) + assert.Equal(t, float64(20), captured["k"]) + responseFormat := captured["response_format"].(map[string]any) - if responseFormat["type"] != "json_object" { - t.Fatalf("response_format.type = %#v, want json_object", responseFormat["type"]) - } + assert.Equal(t, "json_object", responseFormat["type"]) jsonSchema, ok := responseFormat["json_schema"].(map[string]any) - if !ok || jsonSchema["type"] != "object" { - t.Fatalf("response_format.json_schema = %#v, want translated schema", responseFormat["json_schema"]) - } - if _, leaked := responseFormat["name"]; leaked { - t.Fatalf("response_format leaked OpenAI wrapper fields: %#v", responseFormat) - } - - if resp.ID != "cohere-id" || resp.Model != req.Model || resp.Provider != "cohere" { - t.Fatalf("response identity = %#v", resp) - } - if resp.Choices[0].FinishReason != "tool_calls" { - t.Fatalf("finish_reason = %q", resp.Choices[0].FinishReason) - } - if resp.Choices[0].Message.Content != "I will check." { - t.Fatalf("content = %#v", resp.Choices[0].Message.Content) - } - if got := rawString(resp.Choices[0].Message.ExtraFields.Lookup("reasoning_content")); got != "check the weather" { - t.Fatalf("reasoning_content = %q", got) - } - if len(resp.Choices[0].Message.ToolCalls) != 1 || - resp.Choices[0].Message.ToolCalls[0].Function.Arguments != `{"city":"Warsaw"}` { - t.Fatalf("tool calls = %#v", resp.Choices[0].Message.ToolCalls) - } - if resp.Usage.PromptTokens != 17 || resp.Usage.CompletionTokens != 4 || resp.Usage.TotalTokens != 21 { - t.Fatalf("usage = %#v", resp.Usage) - } - if resp.Usage.PromptTokensDetails == nil || - resp.Usage.PromptTokensDetails.CachedTokens != 2 || - resp.Usage.PromptTokensDetails.ImageTokens != 5 { - t.Fatalf("prompt token details = %#v", resp.Usage.PromptTokensDetails) - } + require.True(t, ok, "response_format.json_schema = %#v, want translated schema", responseFormat["json_schema"]) + assert.Equal(t, "object", jsonSchema["type"]) + assert.NotContains(t, responseFormat, "name") + + assert.Equal(t, "cohere-id", resp.ID) + assert.Equal(t, req.Model, resp.Model) + assert.Equal(t, "cohere", resp.Provider) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "tool_calls", resp.Choices[0].FinishReason) + assert.Equal(t, "I will check.", resp.Choices[0].Message.Content) + assert.Equal(t, "check the weather", rawString(resp.Choices[0].Message.ExtraFields.Lookup("reasoning_content"))) + require.Len(t, resp.Choices[0].Message.ToolCalls, 1) + assert.Equal(t, `{"city":"Warsaw"}`, resp.Choices[0].Message.ToolCalls[0].Function.Arguments) + assert.Equal(t, 17, resp.Usage.PromptTokens) + assert.Equal(t, 4, resp.Usage.CompletionTokens) + assert.Equal(t, 21, resp.Usage.TotalTokens) + require.NotNil(t, resp.Usage.PromptTokensDetails) + assert.Equal(t, 2, resp.Usage.PromptTokensDetails.CachedTokens) + assert.Equal(t, 5, resp.Usage.PromptTokensDetails.ImageTokens) } func TestChatCompletionReturnsCohereGenerationFailures(t *testing.T) { @@ -203,76 +169,59 @@ func TestChatCompletionReturnsCohereGenerationFailures(t *testing.T) { for _, tt := range tests { t.Run(tt.finishReason, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - _, _ = io.WriteString(w, `{ - "id":"failed-generation", - "finish_reason":"`+tt.finishReason+`", - "message":{"role":"assistant","content":[]}, - "usage":{} - }`) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `{ + "id":"failed-generation", + "finish_reason":"`+tt.finishReason+`", + "message":{"role":"assistant","content":[]}, + "usage":{} + }`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "command-a", Messages: []core.Message{{Role: "user", Content: "hello"}}, }) - if resp != nil { - t.Fatalf("ChatCompletion() response = %#v, want nil", resp) - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("ChatCompletion() error = %#v, want GatewayError", err) - } - if gatewayErr.Type != core.ErrorTypeProvider || gatewayErr.StatusCode != tt.wantStatus { - t.Fatalf("ChatCompletion() error = %#v, want provider status %d", gatewayErr, tt.wantStatus) - } - if !strings.Contains(strings.ToLower(gatewayErr.Message), tt.wantMessage) { - t.Fatalf("ChatCompletion() message = %q, want %q", gatewayErr.Message, tt.wantMessage) - } + require.Nil(t, resp) + + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, core.ErrorTypeProvider, gatewayErr.Type) + assert.Equal(t, tt.wantStatus, gatewayErr.StatusCode) + assert.Contains(t, strings.ToLower(gatewayErr.Message), tt.wantMessage) }) } } func TestStreamChatCompletionConvertsCohereEvents(t *testing.T) { - var streamValue any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var body map[string]any - _ = json.NewDecoder(r.Body).Decode(&body) - streamValue = body["stream"] - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, strings.Join([]string{ - `event: message-start`, - `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant","content":[],"tool_calls":[],"citations":[]}}}`, - ``, - `event: content-delta`, - `data: {"type":"content-delta","index":0,"delta":{"message":{"content":{"thinking":"consider"}}}}`, - ``, - `event: content-delta`, - `data: {"type":"content-delta","index":0,"delta":{"message":{"content":{"text":"Hello"}}}}`, - ``, - `event: tool-plan-delta`, - `data: {"type":"tool-plan-delta","delta":{"message":{"tool_plan":"Use lookup."}}}`, - ``, - `event: tool-call-start`, - `data: {"type":"tool-call-start","index":0,"delta":{"message":{"tool_calls":{"id":"call-1","type":"function","function":{"name":"lookup","arguments":""}}}}}`, - ``, - `event: tool-call-delta`, - `data: {"type":"tool-call-delta","index":0,"delta":{"message":{"tool_calls":{"function":{"arguments":"{\"q\":\"x\"}"}}}}}`, - ``, - `event: citation-start`, - `data: {"type":"citation-start","index":0,"delta":{"message":{"citations":{"start":0,"end":5,"text":"Hello","sources":[{"type":"document","id":"doc-1"}]}}}}`, - ``, - `event: citation-end`, - `data: {"type":"citation-end","index":0}`, - ``, - `event: message-end`, - `data: {"type":"message-end","delta":{"finish_reason":"TOOL_CALL","usage":{"tokens":{"input_tokens":8,"output_tokens":3},"cached_tokens":1}}}`, - ``, - }, "\n")) - })) - defer server.Close() + server, capture := providertest.SSEServer(t, sseBody( + `event: message-start`, + `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant","content":[],"tool_calls":[],"citations":[]}}}`, + ``, + `event: content-delta`, + `data: {"type":"content-delta","index":0,"delta":{"message":{"content":{"thinking":"consider"}}}}`, + ``, + `event: content-delta`, + `data: {"type":"content-delta","index":0,"delta":{"message":{"content":{"text":"Hello"}}}}`, + ``, + `event: tool-plan-delta`, + `data: {"type":"tool-plan-delta","delta":{"message":{"tool_plan":"Use lookup."}}}`, + ``, + `event: tool-call-start`, + `data: {"type":"tool-call-start","index":0,"delta":{"message":{"tool_calls":{"id":"call-1","type":"function","function":{"name":"lookup","arguments":""}}}}}`, + ``, + `event: tool-call-delta`, + `data: {"type":"tool-call-delta","index":0,"delta":{"message":{"tool_calls":{"function":{"arguments":"{\"q\":\"x\"}"}}}}}`, + ``, + `event: citation-start`, + `data: {"type":"citation-start","index":0,"delta":{"message":{"citations":{"start":0,"end":5,"text":"Hello","sources":[{"type":"document","id":"doc-1"}]}}}}`, + ``, + `event: citation-end`, + `data: {"type":"citation-end","index":0}`, + ``, + `event: message-end`, + `data: {"type":"message-end","delta":{"finish_reason":"TOOL_CALL","usage":{"tokens":{"input_tokens":8,"output_tokens":3},"cached_tokens":1}}}`, + ``, + )) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ @@ -280,19 +229,16 @@ func TestStreamChatCompletionConvertsCohereEvents(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hello"}}, StreamOptions: &core.StreamOptions{IncludeUsage: true}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() + body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("read stream: %v", err) - } - output := string(body) + require.NoError(t, err) - if streamValue != true { - t.Fatalf("upstream stream = %#v", streamValue) - } + streamed, _ := capture.Last(t).JSON(t)["stream"].(bool) + assert.True(t, streamed) + + output := string(body) for _, want := range []string{ `"id":"stream-id"`, `"role":"assistant"`, @@ -307,83 +253,59 @@ func TestStreamChatCompletionConvertsCohereEvents(t *testing.T) { `"cached_tokens":1`, "data: [DONE]", } { - if !strings.Contains(output, want) { - t.Errorf("stream missing %q:\n%s", want, output) - } + assert.Contains(t, output, want) } } func TestStreamChatCompletionReturnsGenerationFailure(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, strings.Join([]string{ - `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, - ``, - `data: {"type":"message-end","delta":{"finish_reason":"TIMEOUT"}}`, - ``, - }, "\n")) - })) - defer server.Close() + server, _ := providertest.SSEServer(t, sseBody( + `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, + ``, + `data: {"type":"message-end","delta":{"finish_reason":"TIMEOUT"}}`, + ``, + )) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: "command-a", Messages: []core.Message{{Role: "user", Content: "hello"}}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("read stream: %v", err) - } + require.NoError(t, err) + output := string(body) - if !strings.Contains(output, `"type":"provider_error"`) || - !strings.Contains(output, `"message":"Cohere generation timed out"`) { - t.Fatalf("stream = %s, want provider timeout error", output) - } - if strings.Contains(output, `"finish_reason":"stop"`) { - t.Fatalf("stream fabricated a successful stop finish:\n%s", output) - } + assert.Contains(t, output, `"type":"provider_error"`) + assert.Contains(t, output, `"message":"Cohere generation timed out"`) + assert.NotContains(t, output, `"finish_reason":"stop"`, "stream fabricated a successful stop finish") } func TestStreamResponsesPropagatesGenerationFailure(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, strings.Join([]string{ - `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, - ``, - `data: {"type":"message-end","delta":{"finish_reason":"ERROR","error":"capacity exhausted"}}`, - ``, - }, "\n")) - })) - defer server.Close() + server, _ := providertest.SSEServer(t, sseBody( + `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, + ``, + `data: {"type":"message-end","delta":{"finish_reason":"ERROR","error":"capacity exhausted"}}`, + ``, + )) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "command-a", Input: "hello", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("read stream: %v", err) - } + require.NoError(t, err) + output := string(body) - if !strings.Contains(output, "event: response.failed") || - !strings.Contains(output, `"status":"failed"`) || - !strings.Contains(output, `"message":"capacity exhausted"`) { - t.Fatalf("Responses stream = %s, want response.failed", output) - } - if strings.Contains(output, "event: response.completed") { - t.Fatalf("Responses stream fabricated response.completed:\n%s", output) - } + assert.Contains(t, output, "event: response.failed") + assert.Contains(t, output, `"status":"failed"`) + assert.Contains(t, output, `"message":"capacity exhausted"`) + assert.NotContains(t, output, "event: response.completed", "Responses stream fabricated response.completed") } func TestStreamResponsesPropagatesAdapterFailures(t *testing.T) { @@ -395,22 +317,22 @@ func TestStreamResponsesPropagatesAdapterFailures(t *testing.T) { }{ { name: "malformed event", - body: strings.Join([]string{ + body: sseBody( `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, ``, `data: {not-json}`, ``, - }, "\n"), + ), wantMessage: "failed to parse Cohere stream event", }, { name: "upstream read failure", - body: strings.Join([]string{ + body: sseBody( `data: {"type":"message-start","id":"stream-id","delta":{"message":{"role":"assistant"}}}`, ``, `data: {"type":"content-delta","delta":{"message":{"content":{"text":"partial"}}}}`, ``, - }, "\n"), + ), contentLengthExtra: 100, wantMessage: "failed to read Cohere stream", }, @@ -418,90 +340,64 @@ func TestStreamResponsesPropagatesAdapterFailures(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + // An overstated Content-Length makes the client's read fail + // partway, which must surface as a terminal failure event. + server, _ := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/event-stream") if tt.contentLengthExtra > 0 { w.Header().Set("Content-Length", strconv.Itoa(len(tt.body)+tt.contentLengthExtra)) } _, _ = io.WriteString(w, tt.body) - })) - defer server.Close() + }) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "command-a", Input: "hello", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("read stream: %v", err) - } + require.NoError(t, err) + output := string(body) - if !strings.Contains(output, "event: response.failed") || - !strings.Contains(output, `"status":"failed"`) || - !strings.Contains(output, `"message":"`+tt.wantMessage+`"`) || - !strings.Contains(output, "data: [DONE]") { - t.Fatalf("Responses stream = %s, want terminal failure %q", output, tt.wantMessage) - } - if strings.Contains(output, "event: response.completed") { - t.Fatalf("Responses stream fabricated response.completed:\n%s", output) - } + assert.Contains(t, output, "event: response.failed") + assert.Contains(t, output, `"status":"failed"`) + assert.Contains(t, output, `"message":"`+tt.wantMessage+`"`) + assert.Contains(t, output, "data: [DONE]") + assert.NotContains(t, output, "event: response.completed", "Responses stream fabricated response.completed") }) } } func TestStreamChatCompletionDoesNotTurnIncompleteStreamIntoSuccess(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "data: {\"type\":\"content-delta\",\"delta\":{\"message\":{\"content\":{\"text\":\"partial\"}}}}\n\n") - })) - defer server.Close() + server, _ := providertest.SSEServer(t, "data: {\"type\":\"content-delta\",\"delta\":{\"message\":{\"content\":{\"text\":\"partial\"}}}}\n\n") provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: "command-a", Messages: []core.Message{{Role: "user", Content: "hello"}}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() + body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("read stream: %v", err) - } + require.NoError(t, err) + output := string(body) - if !strings.Contains(output, `"message":"Cohere stream ended before message-end"`) { - t.Fatalf("stream = %s, want incomplete-stream error", output) - } - if strings.Contains(output, `"finish_reason":"stop"`) { - t.Fatalf("stream fabricated a successful stop finish:\n%s", output) - } - if !strings.Contains(output, "data: [DONE]") { - t.Fatalf("stream missing DONE marker:\n%s", output) - } + assert.Contains(t, output, `"message":"Cohere stream ended before message-end"`) + assert.NotContains(t, output, `"finish_reason":"stop"`, "stream fabricated a successful stop finish") + assert.Contains(t, output, "data: [DONE]") } func TestEmbeddingsTranslatesOpenAIShape(t *testing.T) { - var captured map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v2/embed" { - t.Errorf("path = %q, want /v2/embed", r.URL.Path) - } - _ = json.NewDecoder(r.Body).Decode(&captured) - _, _ = io.WriteString(w, `{ - "id":"embed-id", - "response_type":"embeddings_by_type", - "embeddings":{"float":[[0.1,0.2],[0.3,0.4]]}, - "meta":{"billed_units":{"input_tokens":7}} - }`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"embed-id", + "response_type":"embeddings_by_type", + "embeddings":{"float":[[0.1,0.2],[0.3,0.4]]}, + "meta":{"billed_units":{"input_tokens":7}} + }`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) dimensions := 2 @@ -513,78 +409,51 @@ func TestEmbeddingsTranslatesOpenAIShape(t *testing.T) { "truncate": json.RawMessage(`"END"`), }), }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - - if captured["input_type"] != "search_document" { - t.Fatalf("input_type = %#v", captured["input_type"]) - } - if captured["output_dimension"] != float64(2) { - t.Fatalf("output_dimension = %#v", captured["output_dimension"]) - } - if captured["embedding_types"].([]any)[0] != "float" { - t.Fatalf("embedding_types = %#v", captured["embedding_types"]) - } - if captured["truncate"] != "END" { - t.Fatalf("truncate = %#v", captured["truncate"]) - } - if resp.Object != "list" || resp.Model != "embed-v4.0" || resp.Provider != "cohere" { - t.Fatalf("response = %#v", resp) - } - if len(resp.Data) != 2 || string(resp.Data[1].Embedding) != `[0.3,0.4]` { - t.Fatalf("embedding data = %#v", resp.Data) - } - if resp.Usage.PromptTokens != 7 || resp.Usage.TotalTokens != 7 { - t.Fatalf("usage = %#v", resp.Usage) - } + require.NoError(t, err) + + sent := capture.Last(t) + assert.Equal(t, "/v2/embed", sent.Path) + captured := sent.JSON(t) + assert.Equal(t, "search_document", captured["input_type"]) + assert.Equal(t, float64(2), captured["output_dimension"]) + assert.Equal(t, []any{"float"}, captured["embedding_types"]) + assert.Equal(t, "END", captured["truncate"]) + + assert.Equal(t, "list", resp.Object) + assert.Equal(t, "embed-v4.0", resp.Model) + assert.Equal(t, "cohere", resp.Provider) + require.Len(t, resp.Data, 2) + assert.Equal(t, `[0.3,0.4]`, string(resp.Data[1].Embedding)) + assert.Equal(t, 7, resp.Usage.PromptTokens) + assert.Equal(t, 7, resp.Usage.TotalTokens) } func TestEmbeddingsSupportsBase64(t *testing.T) { resp := fromCohereEmbedResponse(&embedResponse{ Embeddings: embedVectors{Base64: []string{"AAAA", "BBBB"}}, }, &core.EmbeddingRequest{Model: "embed-v4.0", EncodingFormat: "base64"}) - if got := string(resp.Data[1].Embedding); got != `"BBBB"` { - t.Fatalf("base64 embedding = %s", got) - } + assert.Equal(t, `"BBBB"`, string(resp.Data[1].Embedding)) } func TestListModelsFiltersUnsupportedEndpointsAndRotatesKeys(t *testing.T) { - var ( - mu sync.Mutex - headers []string - ) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1/models" || r.URL.Query().Get("page_size") != "1000" { - t.Errorf("models URL = %s", r.URL.String()) - } - mu.Lock() - headers = append(headers, r.Header.Get("Authorization")) - mu.Unlock() - _, _ = io.WriteString(w, `{"models":[ - {"name":"command-a","endpoints":["chat"],"context_length":128000}, - {"name":"embed-v4.0","endpoints":["embed"]}, - {"name":"cohere-transcribe-03-2026","endpoints":["transcriptions"]}, - {"name":"rerank-v3.5","endpoints":["rerank"]}, - {"name":"legacy-unknown"} - ]}`) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"models":[ + {"name":"command-a","endpoints":["chat"],"context_length":128000}, + {"name":"embed-v4.0","endpoints":["embed"]}, + {"name":"cohere-transcribe-03-2026","endpoints":["transcriptions"]}, + {"name":"rerank-v3.5","endpoints":["rerank"]}, + {"name":"legacy-unknown"} + ]}`) provider := NewWithHTTPClient("first", server.URL, server.Client(), llmclient.Hooks{}) provider.keys = providers.NewKeyring("first", "second") for range 2 { resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 4 { - t.Fatalf("models = %#v", resp.Data) - } - if resp.Data[0].Metadata == nil || resp.Data[0].Metadata.ContextWindow == nil || - *resp.Data[0].Metadata.ContextWindow != 128000 { - t.Fatalf("model metadata = %#v", resp.Data[0].Metadata) - } + require.NoError(t, err) + require.Len(t, resp.Data, 4) + require.NotNil(t, resp.Data[0].Metadata) + require.NotNil(t, resp.Data[0].Metadata.ContextWindow) + assert.Equal(t, 128000, *resp.Data[0].Metadata.ContextWindow) + wantModes := map[string][]string{ "command-a": {"chat"}, "embed-v4.0": {"embedding"}, @@ -593,24 +462,26 @@ func TestListModelsFiltersUnsupportedEndpointsAndRotatesKeys(t *testing.T) { } for _, model := range resp.Data { want := wantModes[model.ID] - var got []string - if model.Metadata != nil { - got = model.Metadata.Modes - } - if len(got) != len(want) || (len(want) > 0 && got[0] != want[0]) { - t.Errorf("%s Modes = %v, want %v", model.ID, got, want) - } - if len(want) > 0 { - cats := core.CategoriesForModes(want) - if model.Metadata == nil || len(model.Metadata.Categories) != len(cats) || model.Metadata.Categories[0] != cats[0] { - t.Errorf("%s Categories = %+v, want %v", model.ID, model.Metadata, cats) + if len(want) == 0 { + if model.Metadata != nil { + assert.Empty(t, model.Metadata.Modes, "%s Modes", model.ID) } + continue } + require.NotNil(t, model.Metadata, "%s Metadata", model.ID) + assert.Equal(t, want, model.Metadata.Modes, "%s Modes", model.ID) + assert.Equal(t, core.CategoriesForModes(want), model.Metadata.Categories, "%s Categories", model.ID) } } - if len(headers) != 2 || headers[0] != "Bearer first" || headers[1] != "Bearer second" { - t.Fatalf("Authorization headers = %#v", headers) + + requests := capture.All() + require.Len(t, requests, 2) + for i, sent := range requests { + assert.Equal(t, "/v1/models", sent.Path, "request %d", i) + assert.Equal(t, "1000", sent.Query.Get("page_size"), "request %d", i) } + assert.Equal(t, "Bearer first", requests[0].Header.Get("Authorization")) + assert.Equal(t, "Bearer second", requests[1].Header.Get("Authorization")) } func TestInvalidCohereRequestsReturnClientErrors(t *testing.T) { @@ -637,26 +508,13 @@ func TestInvalidCohereRequestsReturnClientErrors(t *testing.T) { assertInvalidRequest(t, err) } -func TestProviderImplementsNativePassthrough(t *testing.T) { - provider := NewWithHTTPClient("key", "https://example.com", nil, llmclient.Hooks{}) - var _ core.PassthroughProvider = provider -} - func TestPassthroughForwardsNativeCohereRequest(t *testing.T) { const upstreamBody = `{"message":"rate limited"}` - var gotPath, gotAuth, gotHeader string - var gotBody []byte - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.RequestURI() - gotAuth = r.Header.Get("Authorization") - gotHeader = r.Header.Get("X-Client-Trace") - gotBody, _ = io.ReadAll(r.Body) + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("X-Upstream-Request-ID", "cohere-request-1") w.WriteHeader(http.StatusTooManyRequests) _, _ = io.WriteString(w, upstreamBody) - })) - defer server.Close() + }) provider := NewWithHTTPClient("test-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ @@ -665,45 +523,27 @@ func TestPassthroughForwardsNativeCohereRequest(t *testing.T) { Body: io.NopCloser(strings.NewReader(`{"model":"rerank-v3.5","query":"gateway"}`)), Headers: http.Header{"X-Client-Trace": {"trace-123"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/v2/rerank?priority=high" { - t.Fatalf("path = %q, want native endpoint and query", gotPath) - } - if gotAuth != "Bearer test-key" { - t.Fatalf("authorization = %q, want Cohere bearer token", gotAuth) - } - if gotHeader != "trace-123" { - t.Fatalf("X-Client-Trace = %q, want forwarded header", gotHeader) - } - if !strings.Contains(string(gotBody), `"rerank-v3.5"`) { - t.Fatalf("body = %q, want native request body", gotBody) - } - if resp.StatusCode != http.StatusTooManyRequests { - t.Fatalf("status = %d, want 429", resp.StatusCode) - } - if got := resp.Headers["X-Upstream-Request-Id"]; len(got) != 1 || got[0] != "cohere-request-1" { - t.Fatalf("upstream response header = %#v", got) - } + sent := capture.Last(t) + assert.Equal(t, "/v2/rerank", sent.Path) + assert.Equal(t, "high", sent.Query.Get("priority")) + assert.Equal(t, "Bearer test-key", sent.Header.Get("Authorization")) + assert.Equal(t, "trace-123", sent.Header.Get("X-Client-Trace")) + assert.Contains(t, string(sent.Body), `"rerank-v3.5"`) + + assert.Equal(t, http.StatusTooManyRequests, resp.StatusCode) + assert.Equal(t, []string{"cohere-request-1"}, resp.Headers["X-Upstream-Request-Id"]) + body, err := io.ReadAll(resp.Body) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if string(body) != upstreamBody { - t.Fatalf("body = %q, want upstream error body", body) - } + require.NoError(t, err) + assert.Equal(t, upstreamBody, string(body)) } func assertInvalidRequest(t *testing.T, err error) { t.Helper() - if err == nil { - t.Fatal("error = nil, want invalid request") - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok || gatewayErr.Type != core.ErrorTypeInvalidRequest { - t.Fatalf("error = %#v, want invalid request", err) - } + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, core.ErrorTypeInvalidRequest, gatewayErr.Type) } diff --git a/internal/providers/cohere/passthrough_semantics_test.go b/internal/providers/cohere/passthrough_semantics_test.go index f7fe36405..3f3acfa99 100644 --- a/internal/providers/cohere/passthrough_semantics_test.go +++ b/internal/providers/cohere/passthrough_semantics_test.go @@ -4,6 +4,8 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricherRecognizesCohereV2Inference(t *testing.T) { @@ -15,8 +17,7 @@ func TestPassthroughSemanticEnricherRecognizesCohereV2Inference(t *testing.T) { info := passthroughSemanticEnricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ Provider: "cohere", RawEndpoint: endpoint, NormalizedEndpoint: endpoint, }) - if info == nil || info.GenAIOperation != want { - t.Fatalf("GenAIOperation for %q = %+v, want %q", endpoint, info, want) - } + require.NotNil(t, info, "endpoint %q", endpoint) + assert.Equal(t, want, info.GenAIOperation, "endpoint %q", endpoint) } } diff --git a/internal/providers/deepseek/deepseek_test.go b/internal/providers/deepseek/deepseek_test.go index 5e2670d8b..28cfdf6e9 100644 --- a/internal/providers/deepseek/deepseek_test.go +++ b/internal/providers/deepseek/deepseek_test.go @@ -5,108 +5,60 @@ import ( "encoding/json" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-deepseek", - "created":1677652288, - "model":"deepseek-v4-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "deepseek-v4-pro", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +var _ core.PassthroughProvider = (*Provider)(nil) + +// DeepSeek is a thin wrapper over the shared chat-centric adapter, so the +// shared contract covers its surface. DeepSeek exposes no embeddings +// endpoint, so Embeddings must fail fast without an upstream call, and the +// provider must not advertise native batch, file, audio, or +// response-lifecycle support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "deepseek", + DefaultBaseURL: "https://api.deepseek.com", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "deepseek-v4-pro" { - t.Fatalf("resp.Model = %q, want deepseek-v4-pro", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer deepseek-key" { - t.Fatalf("authorization = %q, want Bearer deepseek-key", gotAuth) - } + + provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + assert.False(t, ok, "provider should not implement core.NativeResponseLifecycleProvider") } func TestChatCompletion_MapsReasoningToDeepSeekReasoningEffort(t *testing.T) { - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-deepseek", - "created":1677652288, - "model":"deepseek-v4-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "deepseek-v4-pro", Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: "medium"}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotBody["reasoning"] != nil { - t.Fatalf("request body should not include nested reasoning, got %#v", gotBody["reasoning"]) - } - if gotBody["reasoning_effort"] != "high" { - t.Fatalf("reasoning_effort = %#v, want high", gotBody["reasoning_effort"]) - } + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning") + assert.Equal(t, "high", sent["reasoning_effort"]) } func TestChatCompletion_PadsMissingReasoningContentForAssistantToolCalls(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-deepseek", - "created":1, - "model":"deepseek-v4-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) var req core.ChatRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"deepseek-v4-pro", "messages":[ {"role":"user","content":"check"}, @@ -115,27 +67,21 @@ func TestChatCompletion_PadsMissingReasoningContentForAssistantToolCalls(t *test {"role":"assistant","content":null,"reasoning_content":"client reasoning","tool_calls":[{"id":"call_2","type":"function","function":{"name":"lookup","arguments":"{}"}}]} ], "tools":[{"type":"function","function":{"name":"lookup","parameters":{"type":"object"}}}] - }`), &req); err != nil { - t.Fatalf("json.Unmarshal() error = %v", err) - } + }`), &req) + require.NoError(t, err) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + _, err = provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) - messages, _ := gotBody["messages"].([]any) + messages, ok := capture.Last(t).JSON(t)["messages"].([]any) + require.True(t, ok) + require.Len(t, messages, 4) missingReasoning, _ := messages[1].(map[string]any) - if missingReasoning["reasoning_content"] != " " { - t.Fatalf("synthesized reasoning_content = %#v, want one space", missingReasoning["reasoning_content"]) - } + assert.Equal(t, " ", missingReasoning["reasoning_content"]) preservedReasoning, _ := messages[3].(map[string]any) - if preservedReasoning["reasoning_content"] != "client reasoning" { - t.Fatalf("preserved reasoning_content = %#v, want client value", preservedReasoning["reasoning_content"]) - } - if req.Messages[1].ExtraFields.Lookup("reasoning_content") != nil { - t.Fatal("ChatCompletion() mutated the caller's request") - } + assert.Equal(t, "client reasoning", preservedReasoning["reasoning_content"]) + assert.Nil(t, req.Messages[1].ExtraFields.Lookup("reasoning_content"), "caller's request must not be mutated") } func TestAdaptChatRequest_DoesNotPadWithoutTools(t *testing.T) { @@ -147,170 +93,94 @@ func TestAdaptChatRequest_DoesNotPadWithoutTools(t *testing.T) { } adapted, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if adapted != req { - t.Fatal("adaptChatRequest() copied an unchanged request") - } - if adapted.Messages[0].ExtraFields.Lookup("reasoning_content") != nil { - t.Fatal("reasoning_content should not be added when the request has no tools") - } + require.NoError(t, err) + assert.Same(t, req, adapted) + assert.Nil(t, adapted.Messages[0].ExtraFields.Lookup("reasoning_content")) } -func TestResponses_TranslatesToChatCompletions(t *testing.T) { - var gotPath string - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-deepseek", - "created":1677652288, - "model":"deepseek-v4-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"translated"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5} - }`)) - })) - defer server.Close() +func TestResponses_MapsTokensAndReasoningEffort(t *testing.T) { + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-deepseek", + "created":1677652288, + "model":"deepseek-v4-pro", + "choices":[{"index":0,"message":{"role":"assistant","content":"translated"},"finish_reason":"stop"}], + "usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5} + }`) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) maxOutputTokens := 64 - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "deepseek-v4-pro", Input: "Reply with exactly ok", MaxOutputTokens: &maxOutputTokens, Reasoning: &core.Reasoning{Effort: "xhigh"}, }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotBody["max_output_tokens"] != nil { - t.Fatalf("request body should not include max_output_tokens, got %#v", gotBody["max_output_tokens"]) - } - if gotBody["max_tokens"] != float64(64) { - t.Fatalf("max_tokens = %#v, want 64", gotBody["max_tokens"]) - } - if gotBody["reasoning_effort"] != "max" { - t.Fatalf("reasoning_effort = %#v, want max", gotBody["reasoning_effort"]) - } - messages, ok := gotBody["messages"].([]any) - if !ok || len(messages) != 1 { - t.Fatalf("messages = %#v, want one chat message", gotBody["messages"]) - } - message, _ := messages[0].(map[string]any) - if message["role"] != "user" || message["content"] != "Reply with exactly ok" { - t.Fatalf("message = %#v, want converted user message", message) - } - if resp.Object != "response" || resp.Status != "completed" { - t.Fatalf("response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) - } - if len(resp.Output) != 1 || len(resp.Output[0].Content) != 1 || resp.Output[0].Content[0].Text != "translated" { - t.Fatalf("unexpected responses output: %+v", resp.Output) - } - if resp.Usage == nil || resp.Usage.TotalTokens != 5 { - t.Fatalf("usage = %+v, want total_tokens=5", resp.Usage) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "/chat/completions", req.Path) + sent := req.JSON(t) + assert.NotContains(t, sent, "max_output_tokens") + assert.Equal(t, float64(64), sent["max_tokens"]) + assert.Equal(t, "max", sent["reasoning_effort"]) + assert.Equal(t, []any{map[string]any{"role": "user", "content": "Reply with exactly ok"}}, sent["messages"]) + + assert.Equal(t, "response", resp.Object) + assert.Equal(t, "completed", resp.Status) + require.Len(t, resp.Output, 1) + require.Len(t, resp.Output[0].Content, 1) + assert.Equal(t, "translated", resp.Output[0].Content[0].Text) + require.NotNil(t, resp.Usage) + assert.Equal(t, 5, resp.Usage.TotalTokens) } func TestResponses_ReplaysReasoningContentForToolCall(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-deepseek", - "created":1, - "model":"deepseek-v4-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2} - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) var req core.ResponsesRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"deepseek-v4-pro", "input":[ {"type":"reasoning","summary":[],"content":[{"type":"reasoning_text","text":"Need the weather."}]}, {"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"}, {"type":"function_call_output","call_id":"call_1","output":"sunny"} ] - }`), &req); err != nil { - t.Fatalf("json.Unmarshal() error = %v", err) - } + }`), &req) + require.NoError(t, err) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.Responses(context.Background(), &req); err != nil { - t.Fatalf("Responses() error = %v", err) - } + _, err = provider.Responses(context.Background(), &req) + require.NoError(t, err) - messages, _ := gotBody["messages"].([]any) - if len(messages) != 2 { - t.Fatalf("messages = %#v, want assistant call and tool result", gotBody["messages"]) - } + messages, ok := capture.Last(t).JSON(t)["messages"].([]any) + require.True(t, ok) + require.Len(t, messages, 2) assistant, _ := messages[0].(map[string]any) - if assistant["role"] != "assistant" || assistant["reasoning_content"] != "Need the weather." { - t.Fatalf("assistant = %#v", assistant) - } - if calls, _ := assistant["tool_calls"].([]any); len(calls) != 1 { - t.Fatalf("assistant tool_calls = %#v", assistant["tool_calls"]) - } + assert.Equal(t, "assistant", assistant["role"]) + assert.Equal(t, "Need the weather.", assistant["reasoning_content"]) + assert.Len(t, assistant["tool_calls"], 1) } func TestStreamResponses_TranslatesToChatCompletions(t *testing.T) { - var gotPath string - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-deepseek\",\"object\":\"chat.completion.chunk\",\"created\":1677652288,\"model\":\"deepseek-v4-pro\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"},\"finish_reason\":null}]}\n\n")) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "data: {\"id\":\"chatcmpl-deepseek\",\"object\":\"chat.completion.chunk\",\"created\":1677652288,\"model\":\"deepseek-v4-pro\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"},\"finish_reason\":null}]}\n\ndata: [DONE]\n\n") provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "deepseek-v4-pro", Input: "hi", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotBody["stream"] != true { - t.Fatalf("stream = %#v, want true", gotBody["stream"]) - } + require.NoError(t, err) raw := string(body) - if !strings.Contains(raw, "response.output_text.delta") || !strings.Contains(raw, "data: [DONE]") { - t.Fatalf("converted stream missing responses events or done marker: %s", raw) - } + assert.Contains(t, raw, "response.output_text.delta") + assert.Contains(t, raw, "data: [DONE]") + + req := capture.Last(t) + assert.Equal(t, "/chat/completions", req.Path) + assert.Equal(t, true, req.JSON(t)["stream"]) } func TestNormalizeReasoningEffort(t *testing.T) { @@ -324,209 +194,81 @@ func TestNormalizeReasoningEffort(t *testing.T) { } for input, expected := range tests { t.Run(input, func(t *testing.T) { - if got := normalizeReasoningEffort(input); got != expected { - t.Fatalf("normalizeReasoningEffort(%q) = %q, want %q", input, got, expected) - } + assert.Equal(t, expected, normalizeReasoningEffort(input)) }) } } -func TestProvider_DoesNotExposeOptionalNativeInterfaces(t *testing.T) { - provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("deepseek provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("deepseek provider should not implement native file provider") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("deepseek provider should not implement native response lifecycle provider") - } -} - func TestPassthrough_ForwardsRequestWithBearerAuth(t *testing.T) { - var gotPath, gotAuth, gotMethod string - var gotBody []byte - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - gotMethod = r.Method - gotBody, _ = io.ReadAll(r.Body) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"object":"fim_completion","choices":[{"text":"world"}]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"fim_completion","choices":[{"text":"world"}]}`) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - - body := strings.NewReader(`{"model":"deepseek-v4-pro","prompt":"hello "}`) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/beta/completions", - Body: io.NopCloser(body), + Body: io.NopCloser(strings.NewReader(`{"model":"deepseek-v4-pro","prompt":"hello "}`)), + Headers: http.Header{"Content-Type": []string{"application/json"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/beta/completions" { - t.Fatalf("path = %q, want /beta/completions", gotPath) - } - if gotAuth != "Bearer deepseek-key" { - t.Fatalf("authorization = %q, want Bearer deepseek-key", gotAuth) - } - if gotMethod != http.MethodPost { - t.Fatalf("method = %q, want POST", gotMethod) - } - if !strings.Contains(string(gotBody), "deepseek-v4-pro") { - t.Fatalf("body = %q, want body containing deepseek-v4-pro", gotBody) - } - if resp.StatusCode != http.StatusOK { - t.Fatalf("status = %d, want 200", resp.StatusCode) - } + assert.Equal(t, http.StatusOK, resp.StatusCode) + + req := capture.Last(t) + assert.Equal(t, http.MethodPost, req.Method) + assert.Equal(t, "/beta/completions", req.Path) + assert.Equal(t, "Bearer deepseek-key", req.Header.Get("Authorization")) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Contains(t, string(req.Body), "deepseek-v4-pro") } func TestPassthrough_NilRequest_ReturnsError(t *testing.T) { provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) _, err := provider.Passthrough(context.Background(), nil) - if err == nil { - t.Fatal("expected error for nil passthrough request, got nil") - } + require.Error(t, err) } -func TestPassthrough_ForwardsRequestHeaders(t *testing.T) { - var gotContentType string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotContentType = r.Header.Get("Content-Type") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{}`)) - })) - defer server.Close() +func TestPassthrough_PreservesNon2xxStatusAndBody(t *testing.T) { + const upstreamBody = `{"error":{"message":"rate_limit_exceeded","type":"rate_limit_error"}}` + server, _ := providertest.JSONServer(t, http.StatusTooManyRequests, upstreamBody) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "/beta/completions", Body: io.NopCloser(strings.NewReader(`{}`)), - Headers: http.Header{"Content-Type": []string{"application/json"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotContentType != "application/json" { - t.Fatalf("Content-Type = %q, want application/json", gotContentType) - } -} - -func TestPassthrough_PreservesNon2xxStatus(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusTooManyRequests) - _, _ = w.Write([]byte(`{"error":{"code":"rate_limit_exceeded"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ - Method: http.MethodPost, - Endpoint: "/beta/completions", - Body: io.NopCloser(strings.NewReader(`{}`)), - }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - defer resp.Body.Close() - if resp.StatusCode != http.StatusTooManyRequests { - t.Fatalf("status = %d, want 429", resp.StatusCode) - } + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusTooManyRequests, resp.StatusCode) + assert.Equal(t, upstreamBody, string(body)) } func TestPassthrough_ForwardsQueryString(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.RequestURI() - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{}`) provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodGet, Endpoint: "/beta/completions?stream=true", Body: io.NopCloser(strings.NewReader(``)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/beta/completions?stream=true" { - t.Fatalf("path = %q, want /beta/completions?stream=true", gotPath) - } -} -func TestPassthrough_PreservesResponseBody(t *testing.T) { - const upstreamBody = `{"error":{"message":"rate_limit_exceeded","type":"rate_limit_error"}}` - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusTooManyRequests) - _, _ = w.Write([]byte(upstreamBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("deepseek-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ - Method: http.MethodPost, - Endpoint: "/beta/completions", - Body: io.NopCloser(strings.NewReader(`{}`)), - }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - defer resp.Body.Close() - body, err := io.ReadAll(resp.Body) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if string(body) != upstreamBody { - t.Fatalf("response body = %q, want %q", string(body), upstreamBody) - } -} - -func TestProvider_ImplementsPassthroughProvider(t *testing.T) { - provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) - var _ core.PassthroughProvider = provider + req := capture.Last(t) + assert.Equal(t, "/beta/completions", req.Path) + assert.Equal(t, "true", req.Query.Get("stream")) } func TestResponses_NilRequest_ReturnsError(t *testing.T) { provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) - _, err := provider.Responses(context.Background(), nil) - if err == nil { - t.Fatal("expected error for nil Responses request, got nil") - } -} - -func TestStreamResponses_NilRequest_ReturnsError(t *testing.T) { - provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) - _, err := provider.StreamResponses(context.Background(), nil) - if err == nil { - t.Fatal("expected error for nil StreamResponses request, got nil") - } -} -func TestEmbeddings_ReturnsUnsupported(t *testing.T) { - provider := NewWithHTTPClient("deepseek-key", "", nil, llmclient.Hooks{}) + _, err := provider.Responses(context.Background(), nil) + require.Error(t, err) - _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{Model: "embedding-model", Input: "hi"}) - if err == nil { - t.Fatal("expected unsupported embeddings error, got nil") - } + _, err = provider.StreamResponses(context.Background(), nil) + require.Error(t, err) } diff --git a/internal/providers/deepseek/passthrough_semantics_test.go b/internal/providers/deepseek/passthrough_semantics_test.go index 44b6a3a1d..e954a0f26 100644 --- a/internal/providers/deepseek/passthrough_semantics_test.go +++ b/internal/providers/deepseek/passthrough_semantics_test.go @@ -4,20 +4,20 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricher_ProviderType(t *testing.T) { e := passthroughSemanticEnricher - if got := e.ProviderType(); got != "deepseek" { - t.Fatalf("ProviderType() = %q, want deepseek", got) - } + got := e.ProviderType() + require.Equal(t, "deepseek", got) } func TestPassthroughSemanticEnricher_NilInfo_ReturnsNil(t *testing.T) { e := passthroughSemanticEnricher - if got := e.Enrich(nil, nil, nil); got != nil { - t.Fatalf("Enrich(nil) = %v, want nil", got) - } + got := e.Enrich(nil, nil, nil) + require.Nil(t, got) } func TestPassthroughSemanticEnricher_Enrich(t *testing.T) { @@ -63,15 +63,11 @@ func TestPassthroughSemanticEnricher_Enrich(t *testing.T) { NormalizedEndpoint: tc.normalizedEndpoint, } got := e.Enrich(nil, nil, info) - if got == nil { - t.Fatal("Enrich() returned nil, want enriched info") - } - if tc.wantSemanticOp != "" && got.SemanticOperation != tc.wantSemanticOp { - t.Errorf("SemanticOperation = %q, want %q", got.SemanticOperation, tc.wantSemanticOp) - } - if got.AuditPath != tc.wantAuditPath { - t.Errorf("AuditPath = %q, want %q", got.AuditPath, tc.wantAuditPath) + require.NotNil(t, got) + if tc.wantSemanticOp != "" { + assert.Equal(t, tc.wantSemanticOp, got.SemanticOperation) } + assert.Equal(t, tc.wantAuditPath, got.AuditPath) }) } } diff --git a/internal/providers/elevenlabs/audio_test.go b/internal/providers/elevenlabs/audio_test.go index 687260edc..b109d0e74 100644 --- a/internal/providers/elevenlabs/audio_test.go +++ b/internal/providers/elevenlabs/audio_test.go @@ -7,62 +7,77 @@ import ( "mime" "mime/multipart" "net/http" - "net/http/httptest" "strings" "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -func TestCreateSpeech_UsesVoiceIDInPathAndDefaultsToMP3(t *testing.T) { - var gotPath, gotQuery, gotAuth string - var gotBody speechRequest - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotQuery = r.URL.RawQuery - gotAuth = r.Header.Get("xi-api-key") - body, _ := io.ReadAll(r.Body) - _ = json.Unmarshal(body, &gotBody) - w.Header().Set("Content-Type", "audio/mpeg") +// audioServer answers every request with a one-byte audio payload. +func audioServer(t *testing.T, contentType string) (string, *http.Client, *providertest.Capture) { + t.Helper() + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { + if contentType != "" { + w.Header().Set("Content-Type", contentType) + } _, _ = w.Write([]byte{0x49, 0x44, 0x33}) - })) - defer server.Close() + }) + return server.URL, server.Client(), capture +} - provider := NewWithHTTPClient("elk_test", server.URL, server.Client(), llmclient.Hooks{}) +// multipartFields decodes the recorded multipart body into its form fields, +// returning the file part's filename separately. +func multipartFields(t *testing.T, req providertest.Recorded) (fields map[string]string, filename string) { + t.Helper() + _, params, err := mime.ParseMediaType(req.Header.Get("Content-Type")) + require.NoError(t, err) + + fields = map[string]string{} + reader := multipart.NewReader(bytes.NewReader(req.Body), params["boundary"]) + for { + part, err := reader.NextPart() + if err == io.EOF { + return fields, filename + } + require.NoError(t, err) + data, err := io.ReadAll(part) + require.NoError(t, err) + fields[part.FormName()] = string(data) + if part.FormName() == "file" { + filename = part.FileName() + } + } +} + +func TestCreateSpeech_UsesVoiceIDInPathAndDefaultsToMP3(t *testing.T) { + url, client, capture := audioServer(t, "audio/mpeg") + + provider := NewWithHTTPClient("elk_test", url, client, llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "eleven_multilingual_v2", Input: "hello there", Voice: "21m00Tcm4TlvDq8ikWAM", }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } + require.NoError(t, err) - if gotPath != "/v1/text-to-speech/21m00Tcm4TlvDq8ikWAM" { - t.Fatalf("path = %q, want voice_id in path", gotPath) - } - if gotQuery != "output_format=mp3_44100_128" { - t.Fatalf("query = %q, want mp3_44100_128 output format", gotQuery) - } - if gotAuth != "elk_test" { - t.Fatalf("xi-api-key = %q, want elk_test", gotAuth) - } - if gotBody.Text != "hello there" || gotBody.ModelID != "eleven_multilingual_v2" { - t.Fatalf("request body = %+v", gotBody) - } - if gotBody.VoiceSetting != nil { - t.Fatalf("voice_settings = %+v, want nil when speed unset", gotBody.VoiceSetting) - } - if resp.ContentType != "audio/mpeg" { - t.Fatalf("content type = %q, want audio/mpeg", resp.ContentType) - } - if !bytes.Equal(resp.Data, []byte{0x49, 0x44, 0x33}) { - t.Fatalf("audio data = %v", resp.Data) - } + req := capture.Last(t) + assert.Equal(t, "/v1/text-to-speech/21m00Tcm4TlvDq8ikWAM", req.Path) + assert.Equal(t, "output_format=mp3_44100_128", req.Query.Encode()) + assert.Equal(t, "elk_test", req.Header.Get("xi-api-key")) + + var gotBody speechRequest + require.NoError(t, json.Unmarshal(req.Body, &gotBody)) + assert.Equal(t, "hello there", gotBody.Text) + assert.Equal(t, "eleven_multilingual_v2", gotBody.ModelID) + assert.Nil(t, gotBody.VoiceSetting) + assert.Equal(t, "audio/mpeg", resp.ContentType) + assert.Equal(t, []byte{0x49, 0x44, 0x33}, resp.Data) } func TestCreateSpeech_MapsResponseFormats(t *testing.T) { @@ -79,27 +94,16 @@ func TestCreateSpeech_MapsResponseFormats(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotQuery string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotQuery = r.URL.RawQuery - _, _ = w.Write([]byte{0x01}) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) + url, client, capture := audioServer(t, "") + + provider := NewWithHTTPClient("key", url, client, llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "eleven_multilingual_v2", Input: "hi", Voice: "voice-id", ResponseFormat: tt.responseFormat, }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if gotQuery != tt.wantQuery { - t.Fatalf("query = %q, want %q", gotQuery, tt.wantQuery) - } - if resp.ContentType != tt.wantContent { - t.Fatalf("content type = %q, want %q", resp.ContentType, tt.wantContent) - } + require.NoError(t, err) + assert.Equal(t, tt.wantQuery, capture.Last(t).Query.Encode()) + assert.Equal(t, tt.wantContent, resp.ContentType) }) } } @@ -116,23 +120,18 @@ func TestCreateSpeech_ClampsSpeedToSupportedRange(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotBody speechRequest - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - body, _ := io.ReadAll(r.Body) - _ = json.Unmarshal(body, &gotBody) - _, _ = w.Write([]byte{0x01}) - })) - defer server.Close() - - provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ + url, client, capture := audioServer(t, "") + + provider := NewWithHTTPClient("key", url, client, llmclient.Hooks{}) + _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "eleven_multilingual_v2", Input: "hi", Voice: "voice-id", Speed: tt.speed, - }); err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if gotBody.VoiceSetting == nil || gotBody.VoiceSetting.Speed != tt.wantSpeed { - t.Fatalf("voice_settings = %+v, want speed %v", gotBody.VoiceSetting, tt.wantSpeed) - } + }) + require.NoError(t, err) + + var gotBody speechRequest + require.NoError(t, json.Unmarshal(capture.Last(t).Body, &gotBody)) + require.NotNil(t, gotBody.VoiceSetting) + assert.Equal(t, tt.wantSpeed, gotBody.VoiceSetting.Speed) }) } } @@ -154,35 +153,24 @@ func TestCreateSpeech_ValidatesRequest(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, err := provider.CreateSpeech(context.Background(), tt.req) - if err == nil || !strings.Contains(err.Error(), tt.want) { - t.Fatalf("CreateSpeech() error = %v, want substring %q", err, tt.want) - } + require.Error(t, err) + assert.Contains(t, err.Error(), tt.want) }) } } func TestCreateSpeech_ReturnsUpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte(`{"detail":{"status":"invalid_api_key","message":"bad key"}}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusUnauthorized, `{"detail":{"status":"invalid_api_key","message":"bad key"}}`) provider := NewWithHTTPClient("bad-key", server.URL, server.Client(), llmclient.Hooks{}) _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "m", Input: "hi", Voice: "v", }) - gatewayErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if gatewayErr.StatusCode != http.StatusUnauthorized || gatewayErr.Type != core.ErrorTypeAuthentication { - t.Fatalf("gateway error = %+v, want 401 authentication", gatewayErr) - } - if gatewayErr.Message != "bad key" { - t.Fatalf("message = %q, want the unwrapped detail.message, not the raw JSON body", gatewayErr.Message) - } + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusUnauthorized, gatewayErr.StatusCode) + assert.Equal(t, core.ErrorTypeAuthentication, gatewayErr.Type) + assert.Equal(t, "bad key", gatewayErr.Message) } func TestRefineElevenLabsError_UnwrapsDetailShapes(t *testing.T) { @@ -200,58 +188,15 @@ func TestRefineElevenLabsError_UnwrapsDetailShapes(t *testing.T) { t.Run(tt.name, func(t *testing.T) { original := core.ParseProviderError("elevenlabs", http.StatusBadRequest, []byte(tt.body), nil) refined := refineElevenLabsError(original) - gatewayErr, ok := refined.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", refined) - } - if gatewayErr.Message != tt.want { - t.Fatalf("message = %q, want %q", gatewayErr.Message, tt.want) - } + var gatewayErr *core.GatewayError + require.ErrorAs(t, refined, &gatewayErr) + assert.Equal(t, tt.want, gatewayErr.Message) }) } } func TestCreateTranscription_SendsMultipartAndReturnsJSON(t *testing.T) { - var gotPath, gotAuth string - var gotModelID, gotLanguage, gotGranularity, gotFilename, gotFileContent string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("xi-api-key") - - _, params, err := mime.ParseMediaType(r.Header.Get("Content-Type")) - if err != nil { - http.Error(w, "bad content type", http.StatusBadRequest) - return - } - reader := multipart.NewReader(r.Body, params["boundary"]) - for { - part, err := reader.NextPart() - if err == io.EOF { - break - } - if err != nil { - http.Error(w, "multipart error", http.StatusBadRequest) - return - } - data, _ := io.ReadAll(part) - switch part.FormName() { - case "model_id": - gotModelID = string(data) - case "language_code": - gotLanguage = string(data) - case "timestamps_granularity": - gotGranularity = string(data) - case "file": - gotFilename = part.FileName() - gotFileContent = string(data) - } - } - - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"language_code":"en","text":"hello world","words":[{"text":"hello","type":"word","start":0,"end":0.5},{"text":" ","type":"spacing","start":0.5,"end":0.6},{"text":"world","type":"word","start":0.6,"end":1.1}]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"language_code":"en","text":"hello world","words":[{"text":"hello","type":"word","start":0,"end":0.5},{"text":" ","type":"spacing","start":0.5,"end":0.6},{"text":"world","type":"word","start":0.6,"end":1.1}]}`) provider := NewWithHTTPClient("elk_test", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ @@ -260,80 +205,44 @@ func TestCreateTranscription_SendsMultipartAndReturnsJSON(t *testing.T) { File: []byte("fake-audio-bytes"), Language: "en", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "/v1/speech-to-text", req.Path) + assert.Equal(t, "elk_test", req.Header.Get("xi-api-key")) + + fields, filename := multipartFields(t, req) + assert.Equal(t, "scribe_v1", fields["model_id"]) + assert.Equal(t, "en", fields["language_code"]) + assert.Equal(t, "none", fields["timestamps_granularity"]) + assert.Equal(t, "clip.mp3", filename) + assert.Equal(t, "fake-audio-bytes", fields["file"]) + assert.Equal(t, "application/json", resp.ContentType) - if gotPath != "/v1/speech-to-text" || gotAuth != "elk_test" { - t.Fatalf("path/auth = %q/%q", gotPath, gotAuth) - } - if gotModelID != "scribe_v1" || gotLanguage != "en" { - t.Fatalf("model_id/language_code = %q/%q", gotModelID, gotLanguage) - } - if gotGranularity != "none" { - t.Fatalf("timestamps_granularity = %q, want none", gotGranularity) - } - if gotFilename != "clip.mp3" || gotFileContent != "fake-audio-bytes" { - t.Fatalf("file = %q/%q", gotFilename, gotFileContent) - } - if resp.ContentType != "application/json" { - t.Fatalf("content type = %q, want application/json", resp.ContentType) - } var decoded struct { Text string `json:"text"` } - if err := json.Unmarshal(resp.Data, &decoded); err != nil || decoded.Text != "hello world" { - t.Fatalf("response body = %s, err = %v", resp.Data, err) - } + require.NoError(t, json.Unmarshal(resp.Data, &decoded)) + assert.Equal(t, "hello world", decoded.Text) } func TestCreateTranscription_WordGranularityFromRequest(t *testing.T) { - var gotGranularity string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, params, err := mime.ParseMediaType(r.Header.Get("Content-Type")) - if err != nil { - http.Error(w, "bad content type", http.StatusBadRequest) - return - } - reader := multipart.NewReader(r.Body, params["boundary"]) - for { - part, err := reader.NextPart() - if err == io.EOF { - break - } - if err != nil { - http.Error(w, "multipart error", http.StatusBadRequest) - return - } - if part.FormName() == "timestamps_granularity" { - data, _ := io.ReadAll(part) - gotGranularity = string(data) - } - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"text":"hi"}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"text":"hi"}`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ + _, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "scribe_v1", File: []byte("audio"), TimestampGranularities: []string{"word"}, - }); err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } - if gotGranularity != "word" { - t.Fatalf("timestamps_granularity = %q, want word", gotGranularity) - } + }) + require.NoError(t, err) + + fields, _ := multipartFields(t, capture.Last(t)) + assert.Equal(t, "word", fields["timestamps_granularity"]) } func TestCreateTranscription_VerboseJSONIncludesWordsAndDuration(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"language_code":"en","text":"hi there","words":[{"text":"hi","type":"word","start":0,"end":0.3},{"text":"there","type":"word","start":0.4,"end":0.9}]}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `{"language_code":"en","text":"hi there","words":[{"text":"hi","type":"word","start":0,"end":0.3},{"text":"there","type":"word","start":0.4,"end":0.9}]}`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ @@ -341,9 +250,7 @@ func TestCreateTranscription_VerboseJSONIncludesWordsAndDuration(t *testing.T) { File: []byte("audio"), ResponseFormat: "verbose_json", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } + require.NoError(t, err) var decoded struct { Language string `json:"language"` @@ -355,23 +262,16 @@ func TestCreateTranscription_VerboseJSONIncludesWordsAndDuration(t *testing.T) { End float64 `json:"end"` } `json:"words"` } - if err := json.Unmarshal(resp.Data, &decoded); err != nil { - t.Fatalf("unmarshal error = %v, body = %s", err, resp.Data) - } - if decoded.Language != "en" || decoded.Text != "hi there" || decoded.Duration != 0.9 { - t.Fatalf("verbose response = %+v", decoded) - } - if len(decoded.Words) != 2 || decoded.Words[1].Word != "there" { - t.Fatalf("words = %+v", decoded.Words) - } + require.NoError(t, json.Unmarshal(resp.Data, &decoded)) + assert.Equal(t, "en", decoded.Language) + assert.Equal(t, "hi there", decoded.Text) + assert.Equal(t, 0.9, decoded.Duration) + require.Len(t, decoded.Words, 2) + assert.Equal(t, "there", decoded.Words[1].Word) } func TestCreateTranscription_TextFormatReturnsPlainText(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"text":"plain text result"}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `{"text":"plain text result"}`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ @@ -379,15 +279,9 @@ func TestCreateTranscription_TextFormatReturnsPlainText(t *testing.T) { File: []byte("audio"), ResponseFormat: "text", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } - if string(resp.Data) != "plain text result" { - t.Fatalf("data = %q, want plain text result", resp.Data) - } - if !strings.HasPrefix(resp.ContentType, "text/plain") { - t.Fatalf("content type = %q, want text/plain", resp.ContentType) - } + require.NoError(t, err) + assert.Equal(t, "plain text result", string(resp.Data)) + assert.True(t, strings.HasPrefix(resp.ContentType, "text/plain"), "content type = %q, want text/plain", resp.ContentType) } func TestCreateTranscription_ValidatesRequest(t *testing.T) { @@ -406,9 +300,8 @@ func TestCreateTranscription_ValidatesRequest(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, err := provider.CreateTranscription(context.Background(), tt.req) - if err == nil || !strings.Contains(err.Error(), tt.want) { - t.Fatalf("CreateTranscription() error = %v, want substring %q", err, tt.want) - } + require.Error(t, err) + assert.Contains(t, err.Error(), tt.want) }) } } diff --git a/internal/providers/elevenlabs/elevenlabs_test.go b/internal/providers/elevenlabs/elevenlabs_test.go index 0729f2a7c..b6d85bb8e 100644 --- a/internal/providers/elevenlabs/elevenlabs_test.go +++ b/internal/providers/elevenlabs/elevenlabs_test.go @@ -3,157 +3,91 @@ package elevenlabs import ( "context" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestNew_ConstructsRegisteredProvider(t *testing.T) { - provider, ok := New(providers.ProviderConfig{ - APIKey: "elk_test", - BaseURL: "https://elevenlabs.example", - }, providers.ProviderOptions{}).(*Provider) - if !ok || provider.client == nil { - t.Fatalf("New() = %T, want initialized *Provider", provider) - } - if Registration.Discovery.DefaultBaseURL != defaultBaseURL { - t.Fatalf("registration base URL = %q, want %q", Registration.Discovery.DefaultBaseURL, defaultBaseURL) - } -} - -func TestProvider_ImplementsExpectedInterfaces(t *testing.T) { - provider := NewWithHTTPClient("key", "", nil, llmclient.Hooks{}) - if _, ok := any(provider).(core.Provider); !ok { - t.Fatal("elevenlabs provider should implement core.Provider") - } - if _, ok := any(provider).(core.AudioProvider); !ok { - t.Fatal("elevenlabs provider should implement core.AudioProvider") - } - if _, ok := any(provider).(core.PassthroughProvider); !ok { - t.Fatal("elevenlabs provider should implement core.PassthroughProvider") - } -} - func TestSetBaseURL_ChangesRequestTarget(t *testing.T) { - var gotMethod, gotPath, gotAuth string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotMethod = r.Method - gotPath = r.URL.Path - gotAuth = r.Header.Get("xi-api-key") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`[]`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `[]`) provider := NewWithHTTPClient("elk_test", "https://unused.example", server.Client(), llmclient.Hooks{}) provider.SetBaseURL(server.URL) - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotMethod != http.MethodGet || gotPath != "/v1/models" { - t.Fatalf("method/path = %q/%q, want GET /v1/models", gotMethod, gotPath) - } - if gotAuth != "elk_test" { - t.Fatalf("xi-api-key = %q, want elk_test", gotAuth) - } + _, err := provider.ListModels(context.Background()) + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, http.MethodGet, req.Method) + assert.Equal(t, "/v1/models", req.Path) + assert.Equal(t, "elk_test", req.Header.Get("xi-api-key")) } func TestListModels_FiltersToTextToSpeechAndAddsScribe(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`[ - {"model_id":"eleven_multilingual_v2","name":"Eleven Multilingual v2","can_do_text_to_speech":true,"languages":[{"language_id":"en"},{"language_id":"es"}]}, - {"model_id":"eleven_english_sts_v2","name":"Eleven English STS v2","can_do_text_to_speech":false} - ]`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `[ + {"model_id":"eleven_multilingual_v2","name":"Eleven Multilingual v2","can_do_text_to_speech":true,"languages":[{"language_id":"en"},{"language_id":"es"}]}, + {"model_id":"eleven_english_sts_v2","name":"Eleven English STS v2","can_do_text_to_speech":false} + ]`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } + require.NoError(t, err) byID := make(map[string]core.Model, len(resp.Data)) for _, model := range resp.Data { byID[model.ID] = model } - if _, ok := byID["eleven_english_sts_v2"]; ok { - t.Fatal("ListModels() should exclude models that cannot do text-to-speech") - } + assert.NotContains(t, byID, "eleven_english_sts_v2") + assert.Contains(t, byID, "scribe_v1") + tts, ok := byID["eleven_multilingual_v2"] - if !ok { - t.Fatal("ListModels() should include text-to-speech models") - } - if tts.Metadata == nil || len(tts.Metadata.Modes) != 1 || tts.Metadata.Modes[0] != "audio_speech" { - t.Fatalf("tts model metadata = %+v, want audio_speech mode", tts.Metadata) - } - if !tts.Metadata.Capabilities["multilingual"] { - t.Fatalf("tts model capabilities = %+v, want multilingual", tts.Metadata.Capabilities) - } - if _, ok := byID["scribe_v1"]; !ok { - t.Fatal("ListModels() should include the static scribe_v1 model") - } + require.True(t, ok) + require.NotNil(t, tts.Metadata) + assert.Equal(t, []string{"audio_speech"}, tts.Metadata.Modes) + assert.True(t, tts.Metadata.Capabilities["multilingual"], "tts model capabilities = %+v, want multilingual", tts.Metadata.Capabilities) + scribe, ok := byID["scribe_v2"] - if !ok { - t.Fatal("ListModels() should include the static scribe_v2 model") - } - if scribe.Metadata == nil || len(scribe.Metadata.Modes) != 1 || scribe.Metadata.Modes[0] != "audio_transcription" { - t.Fatalf("scribe model metadata = %+v, want audio_transcription mode", scribe.Metadata) - } + require.True(t, ok) + require.NotNil(t, scribe.Metadata) + assert.Equal(t, []string{"audio_transcription"}, scribe.Metadata.Modes) } func TestListModels_FallsBackToStaticModelsOnFirstFetchFailure(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - http.Error(w, "unavailable", http.StatusServiceUnavailable) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusServiceUnavailable, `unavailable`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v, want static fallback on first-ever fetch failure", err) - } - if len(resp.Data) != len(staticTranscriptionModels) { - t.Fatalf("ListModels() data = %+v, want only the static transcription models", resp.Data) - } - if _, ok := func() (core.Model, bool) { - for _, m := range resp.Data { - if m.ID == "scribe_v2" { - return m, true - } - } - return core.Model{}, false - }(); !ok { - t.Fatal("ListModels() fallback should include scribe_v2") + require.NoError(t, err) + require.Len(t, resp.Data, len(staticTranscriptionModels), "want only the static transcription models") + + ids := make([]string, 0, len(resp.Data)) + for _, m := range resp.Data { + ids = append(ids, m.ID) } + assert.Contains(t, ids, "scribe_v2") } func TestListModels_PropagatesErrorOnceCatalogHasSucceededOnce(t *testing.T) { fail := false - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + server, _ := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { if fail { http.Error(w, "unavailable", http.StatusServiceUnavailable) return } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`[{"model_id":"eleven_multilingual_v2","name":"Eleven Multilingual v2","can_do_text_to_speech":true}]`)) - })) - defer server.Close() + }) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("first ListModels() error = %v, want success", err) - } + _, err := provider.ListModels(context.Background()) + require.NoError(t, err) fail = true - if _, err := provider.ListModels(context.Background()); err == nil { - t.Fatal("ListModels() error = nil, want propagated catalog error once a fetch has already succeeded, so the registry's stale-inventory carry-forward keeps the larger prior list instead of this call shrinking it") - } + _, err = provider.ListModels(context.Background()) + require.Error(t, err) } func TestUnsupportedCapabilities_ReturnInvalidRequestErrors(t *testing.T) { @@ -182,23 +116,14 @@ func TestUnsupportedCapabilities_ReturnInvalidRequestErrors(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { err := tt.call() - if err == nil || !strings.Contains(err.Error(), tt.want) { - t.Fatalf("error = %v, want substring %q", err, tt.want) - } + require.Error(t, err) + assert.Contains(t, err.Error(), tt.want) }) } } func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { - var gotPath, gotAuth string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("xi-api-key") - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusAccepted) - _, _ = w.Write([]byte(`{"accepted":true}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusAccepted, `{"accepted":true}`) provider := NewWithHTTPClient("elk_test", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ @@ -206,18 +131,11 @@ func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { Endpoint: "voices", Headers: http.Header{}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/voices" { - t.Fatalf("path = %q, want /voices", gotPath) - } - if gotAuth != "elk_test" { - t.Fatalf("xi-api-key = %q, want elk_test", gotAuth) - } - if resp.StatusCode != http.StatusAccepted { - t.Fatalf("status = %d, want 202", resp.StatusCode) - } + req := capture.Last(t) + assert.Equal(t, "/voices", req.Path) + assert.Equal(t, "elk_test", req.Header.Get("xi-api-key")) + assert.Equal(t, http.StatusAccepted, resp.StatusCode) } diff --git a/internal/providers/fireworks/fireworks_test.go b/internal/providers/fireworks/fireworks_test.go index 6251c71c2..d3653701e 100644 --- a/internal/providers/fireworks/fireworks_test.go +++ b/internal/providers/fireworks/fireworks_test.go @@ -1,132 +1,26 @@ package fireworks import ( - "context" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-fireworks", - "created":1677652288, - "model":"accounts/fireworks/models/gpt-oss-120b", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("fw-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "accounts/fireworks/models/gpt-oss-120b", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +// Fireworks is a thin wrapper over the shared chat-centric adapter and +// forwards embeddings upstream, so the shared contract covers its surface. +// It must not advertise native batch, file, or audio support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "fireworks", + DefaultBaseURL: "https://api.fireworks.ai/inference/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, + Embeddings: true, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "accounts/fireworks/models/gpt-oss-120b" { - t.Fatalf("resp.Model = %q, want accounts/fireworks/models/gpt-oss-120b", resp.Model) - } - if resp.Usage.TotalTokens != 4 { - t.Fatalf("resp.Usage = %+v, want total_tokens=4", resp.Usage) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer fw-key" { - t.Fatalf("authorization = %q, want Bearer fw-key", gotAuth) - } -} - -func TestEmbeddings_DelegatesToCompatibleProvider(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "model":"nomic-ai/nomic-embed-text-v1.5", - "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], - "usage":{"prompt_tokens":3,"total_tokens":3} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("fw-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: "nomic-ai/nomic-embed-text-v1.5", - Input: "hello", - }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Model != "nomic-ai/nomic-embed-text-v1.5" { - t.Fatalf("resp.Model = %q, want nomic-ai/nomic-embed-text-v1.5", resp.Model) - } - if gotPath != "/embeddings" { - t.Fatalf("path = %q, want /embeddings", gotPath) - } - if gotAuth != "Bearer fw-key" { - t.Fatalf("authorization = %q, want Bearer fw-key", gotAuth) - } -} - -func TestListModels_UsesModelsEndpoint(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[{"id":"accounts/fireworks/models/gpt-oss-120b","object":"model","created":1677652288,"owned_by":"fireworks"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("fw-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "accounts/fireworks/models/gpt-oss-120b" { - t.Fatalf("resp.Data = %+v, want one model accounts/fireworks/models/gpt-oss-120b", resp.Data) - } - if gotPath != "/models" { - t.Fatalf("path = %q, want /models", gotPath) - } -} - -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("fw-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("fireworks provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("fireworks provider should not implement native file provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("fireworks provider should not implement audio provider") - } + providertest.AssertNoNativeSurfaces(t, NewWithHTTPClient("fw-key", "", nil, llmclient.Hooks{})) } diff --git a/internal/providers/fireworks/reasoning_test.go b/internal/providers/fireworks/reasoning_test.go index 045376337..40f6ea884 100644 --- a/internal/providers/fireworks/reasoning_test.go +++ b/internal/providers/fireworks/reasoning_test.go @@ -2,13 +2,14 @@ package fireworks import ( "context" - "encoding/json" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestChatCompletion_MapsReasoningToReasoningEffort(t *testing.T) { @@ -22,15 +23,7 @@ func TestChatCompletion_MapsReasoningToReasoningEffort(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var raw map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&raw); err != nil { - t.Errorf("decode request: %v", err) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"c1","object":"chat.completion","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) provider := NewWithHTTPClient("test-api-key", server.URL, nil, llmclient.Hooks{}) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ @@ -38,20 +31,15 @@ func TestChatCompletion_MapsReasoningToReasoningEffort(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: tt.effort}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := raw["reasoning"]; ok { - t.Errorf("request body includes nested reasoning: %v", raw["reasoning"]) - } - got, ok := raw["reasoning_effort"] + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning") if tt.wantEffort == "" { - if ok { - t.Errorf("reasoning_effort = %v, want absent", got) - } - } else if got != tt.wantEffort { - t.Errorf("reasoning_effort = %v, want %q", got, tt.wantEffort) + assert.NotContains(t, sent, "reasoning_effort") + return } + assert.Equal(t, tt.wantEffort, sent["reasoning_effort"]) }) } } diff --git a/internal/providers/googlecommon/auth_test.go b/internal/providers/googlecommon/auth_test.go index fa3036606..4944ffc1b 100644 --- a/internal/providers/googlecommon/auth_test.go +++ b/internal/providers/googlecommon/auth_test.go @@ -16,6 +16,9 @@ import ( "strings" "testing" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "golang.org/x/oauth2" ) @@ -24,12 +27,8 @@ func TestServiceAccountJSONDecodesURLSafeBase64(t *testing.T) { encoded := base64.RawURLEncoding.EncodeToString(want) got, err := serviceAccountJSON(Config{ServiceAccountJSONBase64: encoded}) - if err != nil { - t.Fatalf("serviceAccountJSON() error = %v", err) - } - if string(got) != string(want) { - t.Fatalf("decoded bytes = %q, want %q", string(got), string(want)) - } + require.NoError(t, err) + assert.Equal(t, want, got) } func TestServiceAccountJSONDecodesPaddedURLSafeBase64(t *testing.T) { @@ -37,22 +36,14 @@ func TestServiceAccountJSONDecodesPaddedURLSafeBase64(t *testing.T) { encoded := base64.URLEncoding.EncodeToString(want) got, err := serviceAccountJSON(Config{ServiceAccountJSONBase64: encoded}) - if err != nil { - t.Fatalf("serviceAccountJSON() error = %v", err) - } - if string(got) != string(want) { - t.Fatalf("decoded bytes = %q, want %q", string(got), string(want)) - } + require.NoError(t, err) + assert.Equal(t, want, got) } func TestServiceAccountJSONReportsOriginalBase64DecodeError(t *testing.T) { _, err := serviceAccountJSON(Config{ServiceAccountJSONBase64: "not valid base64!"}) - if err == nil { - t.Fatal("expected invalid base64 error") - } - if !strings.Contains(err.Error(), "standard base64 decode failed") { - t.Fatalf("error = %v, want standard decode context", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "standard base64 decode failed") } func TestFindCredentialsAndHTTPClientAuthSelection(t *testing.T) { @@ -68,9 +59,9 @@ func TestFindCredentialsAndHTTPClientAuthSelection(t *testing.T) { cfg: func(t *testing.T, tokenURL string) Config { credentials := serviceAccountCredentials(t, tokenURL) encoded := base64.StdEncoding.EncodeToString([]byte(credentials)) - if _, err := serviceAccountJSON(Config{ServiceAccountJSONBase64: encoded}); err != nil { - t.Fatalf("serviceAccountJSON() error = %v", err) - } + _, err := serviceAccountJSON(Config{ServiceAccountJSONBase64: encoded}) + require.NoError(t, err) + return Config{ServiceAccountJSONBase64: encoded} }, wantToken: "service-account-token", @@ -91,9 +82,9 @@ func TestFindCredentialsAndHTTPClientAuthSelection(t *testing.T) { t.Run(tt.name, func(t *testing.T) { var gotForm url.Values tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := r.ParseForm(); err != nil { - t.Fatalf("ParseForm() error = %v", err) - } + err := r.ParseForm() + require.NoError(t, err) + gotForm = r.PostForm token := tt.wantToken if gotForm.Get("grant_type") == "urn:ietf:params:oauth:grant-type:jwt-bearer" { @@ -111,34 +102,23 @@ func TestFindCredentialsAndHTTPClientAuthSelection(t *testing.T) { defer tokenServer.Close() creds, err := FindCredentials(context.Background(), tt.cfg(t, tokenServer.URL)) - if err != nil { - t.Fatalf("FindCredentials() error = %v", err) - } - source := creds.TokenSource + require.NoError(t, err) - upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if got := r.Header.Get("Authorization"); got != "Bearer "+tt.wantToken { - t.Fatalf("Authorization = %q, want Bearer %s", got, tt.wantToken) - } + upstream, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) - })) - defer upstream.Close() + }) - client := HTTPClient(upstream.Client(), source, "") + client := HTTPClient(upstream.Client(), creds.TokenSource, "") resp, err := client.Get(upstream.URL) - if err != nil { - t.Fatalf("HTTPClient request error = %v", err) - } + require.NoError(t, err) _ = resp.Body.Close() + assert.Equal(t, "Bearer "+tt.wantToken, capture.Last(t).Header.Get("Authorization")) if tt.wantScope != "" { - scope := tokenRequestScope(t, gotForm) - if scope != tt.wantScope { - t.Fatalf("scope = %q, want %q", scope, tt.wantScope) - } + assert.Equal(t, tt.wantScope, tokenRequestScope(t, gotForm)) } - if tt.wantADCFile && gotForm.Get("refresh_token") != "adc-refresh-token" { - t.Fatalf("refresh_token = %q, want adc-refresh-token", gotForm.Get("refresh_token")) + if tt.wantADCFile { + assert.Equal(t, "adc-refresh-token", gotForm.Get("refresh_token")) } }) } @@ -163,108 +143,76 @@ func adcCredentialsFileWithQuotaProject(t *testing.T, tokenURL, quotaProject str contents["quota_project_id"] = quotaProject } encoded, err := json.Marshal(contents) - if err != nil { - t.Fatalf("failed to marshal ADC credentials: %v", err) - } - if err := os.WriteFile(path, encoded, 0o600); err != nil { - t.Fatalf("failed to write ADC credentials: %v", err) - } + require.NoError(t, err) + err = os.WriteFile(path, encoded, 0o600) + require.NoError(t, err) + return path } -func TestFindCredentialsReadsADCQuotaProject(t *testing.T) { - tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": "adc-token", - "token_type": "Bearer", - "expires_in": 3600, - }) - })) - defer tokenServer.Close() +// staticTokenServer answers every token request with accessToken. +func staticTokenServer(t *testing.T, accessToken string) *httptest.Server { + t.Helper() + server, _ := providertest.JSONServer(t, http.StatusOK, `{"access_token":"`+accessToken+`","token_type":"Bearer","expires_in":3600}`) + return server +} +func TestFindCredentialsReadsADCQuotaProject(t *testing.T) { + tokenServer := staticTokenServer(t, "adc-token") t.Setenv("GOOGLE_APPLICATION_CREDENTIALS", adcCredentialsFileWithQuotaProject(t, tokenServer.URL, "billing-target")) creds, err := FindCredentials(context.Background(), Config{}) - if err != nil { - t.Fatalf("FindCredentials() error = %v", err) - } - if creds.QuotaProjectID != "billing-target" { - t.Fatalf("QuotaProjectID = %q, want billing-target", creds.QuotaProjectID) - } + require.NoError(t, err) + assert.Equal(t, "billing-target", creds.QuotaProjectID) } func TestFindCredentialsReadsServiceAccountProject(t *testing.T) { - tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": "sa-token", - "token_type": "Bearer", - "expires_in": 3600, - }) - })) - defer tokenServer.Close() + tokenServer := staticTokenServer(t, "sa-token") saJSON := serviceAccountCredentialsWithProject(t, tokenServer.URL, "sa-home-project") creds, err := FindCredentials(context.Background(), Config{ServiceAccountJSON: saJSON}) - if err != nil { - t.Fatalf("FindCredentials() error = %v", err) - } - if creds.QuotaProjectID != "sa-home-project" { - t.Fatalf("QuotaProjectID = %q, want sa-home-project", creds.QuotaProjectID) - } + require.NoError(t, err) + assert.Equal(t, "sa-home-project", creds.QuotaProjectID) } -func TestHTTPClientSetsQuotaProjectHeader(t *testing.T) { - var gotHeader string - upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotHeader = r.Header.Get(QuotaProjectHeader) - w.WriteHeader(http.StatusNoContent) - })) - defer upstream.Close() - - source := oauth2.StaticTokenSource(&oauth2.Token{AccessToken: "test-token", TokenType: "Bearer"}) - client := HTTPClient(upstream.Client(), source, "billing-target") - resp, err := client.Get(upstream.URL) - if err != nil { - t.Fatalf("HTTPClient request error = %v", err) - } - _ = resp.Body.Close() - if gotHeader != "billing-target" { - t.Fatalf("%s header = %q, want billing-target", QuotaProjectHeader, gotHeader) +func TestHTTPClientQuotaProjectHeader(t *testing.T) { + tests := []struct { + name string + quotaProject string + }{ + {name: "set when configured", quotaProject: "billing-target"}, + {name: "omitted when empty", quotaProject: ""}, } -} + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNoContent) + }) -func TestHTTPClientOmitsQuotaProjectHeaderWhenEmpty(t *testing.T) { - var hasHeader bool - upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, hasHeader = r.Header[QuotaProjectHeader] - w.WriteHeader(http.StatusNoContent) - })) - defer upstream.Close() - - source := oauth2.StaticTokenSource(&oauth2.Token{AccessToken: "test-token", TokenType: "Bearer"}) - client := HTTPClient(upstream.Client(), source, "") - resp, err := client.Get(upstream.URL) - if err != nil { - t.Fatalf("HTTPClient request error = %v", err) - } - _ = resp.Body.Close() - if hasHeader { - t.Fatalf("upstream saw %s header; expected it to be omitted", QuotaProjectHeader) + source := oauth2.StaticTokenSource(&oauth2.Token{AccessToken: "test-token", TokenType: "Bearer"}) + client := HTTPClient(upstream.Client(), source, tt.quotaProject) + resp, err := client.Get(upstream.URL) + require.NoError(t, err) + _ = resp.Body.Close() + + header := capture.Last(t).Header + if tt.quotaProject == "" { + assert.NotContains(t, header, QuotaProjectHeader) + return + } + assert.Equal(t, tt.quotaProject, header.Get(QuotaProjectHeader)) + }) } } func serviceAccountCredentialsWithProject(t *testing.T, tokenURL, projectID string) string { t.Helper() key, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - t.Fatalf("failed to generate test RSA key: %v", err) - } + require.NoError(t, err) + keyBytes, err := x509.MarshalPKCS8PrivateKey(key) - if err != nil { - t.Fatalf("failed to marshal test RSA key: %v", err) - } + require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{ Type: "PRIVATE KEY", Bytes: keyBytes, @@ -275,41 +223,19 @@ func serviceAccountCredentialsWithProject(t *testing.T, tokenURL, projectID stri "private_key_id": "test-key-id", "private_key": string(keyPEM), "token_uri": tokenURL, - "project_id": projectID, } - encoded, err := json.Marshal(contents) - if err != nil { - t.Fatalf("failed to marshal service account credentials: %v", err) + if projectID != "" { + contents["project_id"] = projectID } + encoded, err := json.Marshal(contents) + require.NoError(t, err) + return string(encoded) } func serviceAccountCredentials(t *testing.T, tokenURL string) string { t.Helper() - key, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - t.Fatalf("failed to generate test RSA key: %v", err) - } - keyBytes, err := x509.MarshalPKCS8PrivateKey(key) - if err != nil { - t.Fatalf("failed to marshal test RSA key: %v", err) - } - keyPEM := pem.EncodeToMemory(&pem.Block{ - Type: "PRIVATE KEY", - Bytes: keyBytes, - }) - contents := map[string]string{ - "type": "service_account", - "client_email": "service@example.com", - "private_key_id": "test-key-id", - "private_key": string(keyPEM), - "token_uri": tokenURL, - } - encoded, err := json.Marshal(contents) - if err != nil { - t.Fatalf("failed to marshal service account credentials: %v", err) - } - return string(encoded) + return serviceAccountCredentialsWithProject(t, tokenURL, "") } func tokenRequestScope(t *testing.T, form url.Values) string { @@ -319,18 +245,16 @@ func tokenRequestScope(t *testing.T, form url.Values) string { } assertion := form.Get("assertion") parts := strings.Split(assertion, ".") - if len(parts) < 2 { - t.Fatalf("JWT assertion = %q, want header.payload.signature", assertion) - } + require.GreaterOrEqual(t, len(parts), 2, "JWT assertion = %q, want header.payload.signature", assertion) + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) - if err != nil { - t.Fatalf("failed to decode JWT payload: %v", err) - } + require.NoError(t, err) + var claims struct { Scope string `json:"scope"` } - if err := json.Unmarshal(payload, &claims); err != nil { - t.Fatalf("failed to decode JWT claims: %v", err) - } + err = json.Unmarshal(payload, &claims) + require.NoError(t, err) + return claims.Scope } diff --git a/internal/providers/groq/audio_test.go b/internal/providers/groq/audio_test.go index 408f97b25..8fdab4602 100644 --- a/internal/providers/groq/audio_test.go +++ b/internal/providers/groq/audio_test.go @@ -2,13 +2,15 @@ package groq import ( "context" + "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" - "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestCreateTranscription_StripsVendorMember(t *testing.T) { @@ -51,14 +53,10 @@ func TestCreateTranscription_StripsVendorMember(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { for _, endpoint := range []string{"/audio/transcriptions", "/audio/translations"} { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != endpoint { - t.Errorf("path = %q, want %q", r.URL.Path, endpoint) - } - _, _ = w.Write([]byte(tt.upstream)) - })) - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, tt.upstream) + }) + provider := newTestProvider(server.URL) req := &core.AudioTranscriptionRequest{ Model: "whisper-large-v3-turbo", @@ -75,42 +73,26 @@ func TestCreateTranscription_StripsVendorMember(t *testing.T) { } else { resp, err = provider.CreateTranslation(context.Background(), req) } - server.Close() - if err != nil { - t.Fatalf("%s error = %v", endpoint, err) - } - if got := string(resp.Data); got != tt.want { - t.Errorf("%s body = %s, want %s", endpoint, got, tt.want) - } + require.NoError(t, err) + assert.Equal(t, endpoint, capture.Last(t).Path) + assert.Equal(t, tt.want, string(resp.Data), "%s body", endpoint) } }) } } func TestCreateTranscription_PropagatesUpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":{"message":"bad audio","type":"invalid_request_error"}}`)) - })) - defer server.Close() - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusBadRequest, `{"error":{"message":"bad audio","type":"invalid_request_error"}}`) + provider := newTestProvider(server.URL) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "whisper-large-v3-turbo", File: []byte("audio"), Filename: "a.mp3", }) - if err == nil { - t.Fatalf("CreateTranscription() error = nil, want the upstream error (resp = %+v)", resp) - } - if resp != nil { - t.Errorf("response = %+v, want nil", resp) - } - if !strings.Contains(err.Error(), "bad audio") { - t.Errorf("error = %v, want it to carry the upstream message", err) - } + require.Error(t, err) + assert.Nil(t, resp) + assert.Contains(t, err.Error(), "bad audio") } func TestWithoutJSONMember_LeavesMalformedBodiesAlone(t *testing.T) { @@ -126,9 +108,8 @@ func TestWithoutJSONMember_LeavesMalformedBodiesAlone(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := string(withoutJSONMember([]byte(tt.body), vendorTranscriptionMember)); got != tt.body { - t.Errorf("withoutJSONMember() = %q, want %q", got, tt.body) - } + got := string(withoutJSONMember([]byte(tt.body), vendorTranscriptionMember)) + assert.Equal(t, tt.body, got) }) } } diff --git a/internal/providers/groq/groq_test.go b/internal/providers/groq/groq_test.go index ba911762c..ddae59e70 100644 --- a/internal/providers/groq/groq_test.go +++ b/internal/providers/groq/groq_test.go @@ -5,31 +5,66 @@ import ( "encoding/json" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestNew(t *testing.T) { - // Use NewWithHTTPClient to get concrete type for internal testing - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) +const testAPIKey = "test-api-key" - if provider.compat == nil { - t.Error("compat provider should not be nil") - } +// newTestProvider points a provider at the given upstream URL. +func newTestProvider(baseURL string) *Provider { + p := NewWithHTTPClient(testAPIKey, nil, llmclient.Hooks{}) + p.SetBaseURL(baseURL) + return p } -func TestNew_ReturnsProvider(t *testing.T) { - provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "groq", + DefaultBaseURL: "https://api.groq.com/openai/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + p := NewWithHTTPClient(apiKey, client, hooks) + p.SetBaseURL(baseURL) + return p + }, + Embeddings: true, + }) +} - if provider == nil { - t.Error("provider should not be nil") +const chatCompletionJSON = `{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama-3.3-70b-versatile", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30 } -} +}` + +const chatChunkSSE = `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} + +data: [DONE] +` func TestChatCompletion(t *testing.T) { tests := []struct { @@ -40,50 +75,17 @@ func TestChatCompletion(t *testing.T) { checkResponse func(*testing.T, *core.ChatResponse) }{ { - name: "successful request", - statusCode: http.StatusOK, - responseBody: `{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama-3.3-70b-versatile", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello! How can I help you today?" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30 - } - }`, - expectedError: false, + name: "successful request", + statusCode: http.StatusOK, + responseBody: chatCompletionJSON, checkResponse: func(t *testing.T, resp *core.ChatResponse) { - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } - if resp.Model != "llama-3.3-70b-versatile" { - t.Errorf("Model = %q, want %q", resp.Model, "llama-3.3-70b-versatile") - } - if len(resp.Choices) != 1 { - t.Fatalf("len(Choices) = %d, want 1", len(resp.Choices)) - } - if resp.Choices[0].Message.Content != "Hello! How can I help you today?" { - t.Errorf("Message content = %q, want %q", resp.Choices[0].Message.Content, "Hello! How can I help you today?") - } - if resp.Usage.PromptTokens != 10 { - t.Errorf("PromptTokens = %d, want 10", resp.Usage.PromptTokens) - } - if resp.Usage.CompletionTokens != 20 { - t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens) - } - if resp.Usage.TotalTokens != 30 { - t.Errorf("TotalTokens = %d, want 30", resp.Usage.TotalTokens) - } + assert.Equal(t, "chatcmpl-123", resp.ID) + assert.Equal(t, "llama-3.3-70b-versatile", resp.Model) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Choices[0].Message.Content) + assert.Equal(t, 10, resp.Usage.PromptTokens) + assert.Equal(t, 20, resp.Usage.CompletionTokens) + assert.Equal(t, 30, resp.Usage.TotalTokens) }, }, { @@ -108,55 +110,25 @@ func TestChatCompletion(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - // Verify request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } + server, capture := providertest.JSONServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(server.URL) - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "llama-3.3-70b-versatile", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - resp, err := provider.ChatCompletion(context.Background(), req) + sent := capture.Last(t) + assert.Equal(t, "application/json", sent.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, sent.Header.Get("Authorization")) + assert.Equal(t, "llama-3.3-70b-versatile", sent.JSON(t)["model"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + require.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } @@ -169,15 +141,9 @@ func TestStreamChatCompletion(t *testing.T) { expectedError bool }{ { - name: "successful streaming request", - statusCode: http.StatusOK, - responseBody: `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} - -data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} - -data: [DONE] -`, - expectedError: false, + name: "successful streaming request", + statusCode: http.StatusOK, + responseBody: chatChunkSSE, }, { name: "API error", @@ -189,68 +155,34 @@ data: [DONE] for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } - + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + _, _ = io.WriteString(w, tt.responseBody) + }) + provider := newTestProvider(server.URL) - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + body, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - req := &core.ChatRequest{ - Model: "llama-3.3-70b-versatile", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - body, err := provider.StreamChatCompletion(context.Background(), req) + sent := capture.Last(t) + assert.Equal(t, "application/json", sent.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, sent.Header.Get("Authorization")) + stream, _ := sent.JSON(t)["stream"].(bool) + assert.True(t, stream) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if body == nil { - t.Fatal("body should not be nil") - } - defer func() { _ = body.Close() }() - - // Read and verify the streaming response - respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } - if string(respBody) != tt.responseBody { - t.Errorf("response body = %q, want %q", string(respBody), tt.responseBody) - } + require.Error(t, err) + return } + require.NoError(t, err) + require.NotNil(t, body) + defer func() { _ = body.Close() }() + + respBody, err := io.ReadAll(body) + require.NoError(t, err) + assert.Equal(t, tt.responseBody, string(respBody)) }) } } @@ -269,34 +201,15 @@ func TestListModels(t *testing.T) { responseBody: `{ "object": "list", "data": [ - { - "id": "llama-3.3-70b-versatile", - "object": "model", - "created": 1687882411, - "owned_by": "groq" - }, - { - "id": "mixtral-8x7b-32768", - "object": "model", - "created": 1687882410, - "owned_by": "groq" - } + {"id": "llama-3.3-70b-versatile", "object": "model", "created": 1687882411, "owned_by": "groq"}, + {"id": "mixtral-8x7b-32768", "object": "model", "created": 1687882410, "owned_by": "groq"} ] }`, - expectedError: false, checkResponse: func(t *testing.T, resp *core.ModelsResponse) { - if resp.Object != "list" { - t.Errorf("Object = %q, want %q", resp.Object, "list") - } - if len(resp.Data) != 2 { - t.Fatalf("len(Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].ID != "llama-3.3-70b-versatile" { - t.Errorf("Data[0].ID = %q, want %q", resp.Data[0].ID, "llama-3.3-70b-versatile") - } - if resp.Data[0].OwnedBy != "groq" { - t.Errorf("Data[0].OwnedBy = %q, want %q", resp.Data[0].OwnedBy, "groq") - } + assert.Equal(t, "list", resp.Object) + require.Len(t, resp.Data, 2) + assert.Equal(t, "llama-3.3-70b-versatile", resp.Data[0].ID) + assert.Equal(t, "groq", resp.Data[0].OwnedBy) }, }, { @@ -309,282 +222,95 @@ func TestListModels(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request method and path - if r.Method != http.MethodGet { - t.Errorf("Method = %q, want %q", r.Method, http.MethodGet) - } - if r.URL.Path != "/models" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/models") - } - - // Verify authorization header - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(server.URL) resp, err := provider.ListModels(context.Background()) + sent := capture.Last(t) + assert.Equal(t, http.MethodGet, sent.Method) + assert.Equal(t, "/models", sent.Path) + assert.Equal(t, "Bearer "+testAPIKey, sent.Header.Get("Authorization")) + if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + require.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } -func TestChatCompletionWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Simulate a slow response - <-r.Context().Done() - w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() +// blockUntilCancelled answers only once the caller gives up, so a cancelled +// context must surface as an error rather than a response. +func blockUntilCancelled(w http.ResponseWriter, r *http.Request) { + <-r.Context().Done() + w.WriteHeader(http.StatusRequestTimeout) +} - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) +func TestChatCompletionWithContext(t *testing.T) { + server, _ := providertest.Server(t, blockUntilCancelled) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately + cancel() - req := &core.ChatRequest{ - Model: "llama-3.3-70b-versatile", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - _, err := provider.ChatCompletion(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + _, err := provider.ChatCompletion(ctx, &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) + assert.Error(t, err) } func TestResponses(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request path for chat completions (Groq converts Responses to chat) - if r.URL.Path != "/chat/completions" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/chat/completions") - } + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama-3.3-70b-versatile", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello! How can I help you today?" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30 - } - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ResponsesRequest{ + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "llama-3.3-70b-versatile", Input: "Hello", - } - - resp, err := provider.Responses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } - if resp.Object != "response" { - t.Errorf("Object = %q, want %q", resp.Object, "response") - } - if resp.Model != "llama-3.3-70b-versatile" { - t.Errorf("Model = %q, want %q", resp.Model, "llama-3.3-70b-versatile") - } - if resp.Status != "completed" { - t.Errorf("Status = %q, want %q", resp.Status, "completed") - } - if len(resp.Output) != 1 { - t.Fatalf("len(Output) = %d, want 1", len(resp.Output)) - } - if len(resp.Output[0].Content) != 1 { - t.Fatalf("len(Output[0].Content) = %d, want 1", len(resp.Output[0].Content)) - } - if resp.Output[0].Content[0].Text != "Hello! How can I help you today?" { - t.Errorf("Output text = %q, want %q", resp.Output[0].Content[0].Text, "Hello! How can I help you today?") - } - if resp.Usage == nil { - t.Fatal("Usage should not be nil") - } - if resp.Usage.InputTokens != 10 { - t.Errorf("InputTokens = %d, want 10", resp.Usage.InputTokens) - } - if resp.Usage.OutputTokens != 20 { - t.Errorf("OutputTokens = %d, want 20", resp.Usage.OutputTokens) - } + }) + require.NoError(t, err) + assert.Equal(t, "/chat/completions", capture.Last(t).Path) + assert.Equal(t, "chatcmpl-123", resp.ID) + assert.Equal(t, "response", resp.Object) + assert.Equal(t, "llama-3.3-70b-versatile", resp.Model) + assert.Equal(t, "completed", resp.Status) + require.Len(t, resp.Output, 1) + require.Len(t, resp.Output[0].Content, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Output[0].Content[0].Text) + require.NotNil(t, resp.Usage) + assert.Equal(t, 10, resp.Usage.InputTokens) + assert.Equal(t, 20, resp.Usage.OutputTokens) } func TestResponsesWithArrayInput(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request body is converted to chat format - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - - var req map[string]any - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - - // Verify messages array exists (converted from input) - messages, ok := req["messages"].([]any) - if !ok { - t.Fatal("messages should be an array") - } - // Should have system message + 2 input messages - if len(messages) != 3 { - t.Errorf("len(messages) = %d, want 3", len(messages)) - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama-3.3-70b-versatile", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello!" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15 - } - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ResponsesRequest{ + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "llama-3.3-70b-versatile", Input: []any{ - map[string]any{ - "role": "user", - "content": "Hello", - }, - map[string]any{ - "role": "assistant", - "content": "Hi there!", - }, + map[string]any{"role": "user", "content": "Hello"}, + map[string]any{"role": "assistant", "content": "Hi there!"}, }, Instructions: "Be helpful", - } - - resp, err := provider.Responses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + }) + require.NoError(t, err) + assert.Equal(t, "chatcmpl-123", resp.ID) - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } + // Instructions become a system message ahead of the two input messages. + messages, ok := capture.Last(t).JSON(t)["messages"].([]any) + require.True(t, ok) + assert.Len(t, messages, 3) } func TestResponses_PreservesOpaqueFieldsThroughChatAdapter(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/chat/completions" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/chat/completions") - } - - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal chat request: %v", err) - } - if req.ExtraFields.Lookup("response_format") == nil { - t.Fatal("response_format missing after responses-to-chat conversion") - } - if len(req.Messages) != 1 { - t.Fatalf("len(Messages) = %d, want 1", len(req.Messages)) - } - if req.Messages[0].ExtraFields.Lookup("x_message_hint") == nil { - t.Fatal("message extras missing after conversion") - } - parts, ok := req.Messages[0].Content.([]core.ContentPart) - if !ok { - t.Fatalf("Messages[0].Content type = %T, want []core.ContentPart", req.Messages[0].Content) - } - if parts[0].ExtraFields.Lookup("cache_control") == nil { - t.Fatal("content part extras missing after conversion") - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-opaque", - "object": "chat.completion", - "created": 1677652288, - "model": "llama-3.3-70b-versatile", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "ok" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2 - } - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) - req := &core.ResponsesRequest{ + _, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "llama-3.3-70b-versatile", Input: []core.ResponsesInputElement{ { @@ -606,99 +332,62 @@ func TestResponses_PreservesOpaqueFieldsThroughChatAdapter(t *testing.T) { ExtraFields: core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ "response_format": json.RawMessage(`{"type":"json_schema"}`), }), - } + }) + require.NoError(t, err) - if _, err := provider.Responses(context.Background(), req); err != nil { - t.Fatalf("unexpected error: %v", err) - } -} + sent := capture.Last(t) + assert.Equal(t, "/chat/completions", sent.Path) -func TestStreamResponses(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } + var req core.ChatRequest + require.NoError(t, json.Unmarshal(sent.Body, &req)) + assert.NotNil(t, req.ExtraFields.Lookup("response_format")) + require.Len(t, req.Messages, 1) + assert.NotNil(t, req.Messages[0].ExtraFields.Lookup("x_message_hint")) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} - -data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} - -data: [DONE] -`)) - })) - defer server.Close() + parts, ok := req.Messages[0].Content.([]core.ContentPart) + require.True(t, ok, "Messages[0].Content type = %T, want []core.ContentPart", req.Messages[0].Content) + assert.NotNil(t, parts[0].ExtraFields.Lookup("cache_control")) +} - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) +func TestStreamResponses(t *testing.T) { + server, capture := providertest.SSEServer(t, chatChunkSSE) + provider := newTestProvider(server.URL) - req := &core.ResponsesRequest{ + body, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "llama-3.3-70b-versatile", Input: "Hello", - } - - body, err := provider.StreamResponses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if body == nil { - t.Fatal("body should not be nil") - } + }) + require.NoError(t, err) + require.NotNil(t, body) defer func() { _ = body.Close() }() + stream, _ := capture.Last(t).JSON(t)["stream"].(bool) + assert.True(t, stream) + respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } + require.NoError(t, err) responseStr := string(respBody) - if !strings.Contains(responseStr, "response.created") { - t.Error("response should contain response.created event") - } - if !strings.Contains(responseStr, "response.output_text.delta") { - t.Error("response should contain response.output_text.delta event") - } - if !strings.Contains(responseStr, "[DONE]") { - t.Error("response should end with [DONE]") - } + assert.Contains(t, responseStr, "response.created") + assert.Contains(t, responseStr, "response.output_text.delta") + assert.Contains(t, responseStr, "[DONE]") } func TestResponsesWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Simulate a slow response - <-r.Context().Done() - w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.Server(t, blockUntilCancelled) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately + cancel() - req := &core.ResponsesRequest{ + _, err := provider.Responses(ctx, &core.ResponsesRequest{ Model: "llama-3.3-70b-versatile", Input: "Hello", - } - - _, err := provider.Responses(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + }) + assert.Error(t, err) } func TestGroqResponsesStreamConverter(t *testing.T) { - // Test the stream converter with mock chat completion stream mockStream := `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":null}]} @@ -709,54 +398,32 @@ data: [DONE] reader := io.NopCloser(strings.NewReader(mockStream)) converter := providers.NewOpenAIResponsesStreamConverter(reader, "llama-3.3-70b-versatile", "groq") - // Read all data from converter data, err := io.ReadAll(converter) - if err != nil { - t.Fatalf("failed to read from converter: %v", err) - } + require.NoError(t, err) result := string(data) - - // Check that the stream contains expected events - if !strings.Contains(result, "response.created") { - t.Error("stream should contain response.created event") - } - if !strings.Contains(result, "response.output_text.delta") { - t.Error("stream should contain response.output_text.delta event") - } - if !strings.Contains(result, "Hello") { - t.Error("stream should contain 'Hello' content") - } - if !strings.Contains(result, " world") { - t.Error("stream should contain ' world' content") - } - if !strings.Contains(result, "response.completed") { - t.Error("stream should contain response.completed event") - } - if !strings.Contains(result, "[DONE]") { - t.Error("stream should contain [DONE] marker") - } + assert.Contains(t, result, "response.created") + assert.Contains(t, result, "response.output_text.delta") + assert.Contains(t, result, "Hello") + assert.Contains(t, result, " world") + assert.Contains(t, result, "response.completed") + assert.Contains(t, result, "[DONE]") } func TestGroqResponsesStreamConverter_Close(t *testing.T) { reader := io.NopCloser(strings.NewReader("data: [DONE]\n")) converter := providers.NewOpenAIResponsesStreamConverter(reader, "test-model", "groq") - err := converter.Close() - if err != nil { - t.Errorf("Close() returned error: %v", err) - } + require.NoError(t, converter.Close()) - // Subsequent reads should return EOF + // Subsequent reads should return EOF. buf := make([]byte, 100) n, err := converter.Read(buf) - if n != 0 || err != io.EOF { - t.Errorf("Read after Close: n=%d, err=%v, want n=0, err=EOF", n, err) - } + assert.Equal(t, 0, n) + assert.Equal(t, io.EOF, err) } func TestGroqResponsesStreamConverter_EmptyDelta(t *testing.T) { - // Test that empty deltas are not emitted mockStream := `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{},"finish_reason":null}]} data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":""},"finish_reason":null}]} @@ -770,154 +437,83 @@ data: [DONE] converter := providers.NewOpenAIResponsesStreamConverter(reader, "llama-3.3-70b-versatile", "groq") data, err := io.ReadAll(converter) - if err != nil { - t.Fatalf("failed to read from converter: %v", err) - } + require.NoError(t, err) + // Empty deltas are not emitted: only the "Hello" delta event remains. result := string(data) - - // Count delta event lines - should only have one with "Hello" - // Each event has "event: response.output_text.delta\n" line - deltaCount := strings.Count(result, "event: response.output_text.delta") - if deltaCount != 1 { - t.Errorf("expected 1 delta event line, got %d", deltaCount) - } - - // Verify the Hello content is present - if !strings.Contains(result, `"delta":"Hello"`) { - t.Error("expected delta with Hello content") - } -} - -func TestNewWithHTTPClient(t *testing.T) { - provider := NewWithHTTPClient("test-api-key", &http.Client{}, llmclient.Hooks{}) - - if provider.compat == nil { - t.Error("compat provider should not be nil") - } -} - -func TestSetBaseURL(t *testing.T) { - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - customURL := "https://custom.groq.api.com/v1" - - provider.SetBaseURL(customURL) - - // We can't directly check the baseURL as it's encapsulated in llmclient - // but we can verify the provider still works by making a test request - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider.SetBaseURL(server.URL) - _, err := provider.ListModels(context.Background()) - if err != nil { - t.Errorf("SetBaseURL should allow using custom URL: %v", err) - } + assert.Equal(t, 1, strings.Count(result, "event: response.output_text.delta")) + assert.Contains(t, result, `"delta":"Hello"`) } func TestCreateSpeech(t *testing.T) { audio := []byte("fake-mp3-bytes") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/audio/speech" { - t.Errorf("path = %q, want /audio/speech", r.URL.Path) - } - if auth := r.Header.Get("Authorization"); !strings.HasPrefix(auth, "Bearer ") { - t.Error("Authorization header should start with 'Bearer '") - } - var req core.AudioSpeechRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - t.Fatalf("decode request: %v", err) - } - if req.Model != "playai-tts" || req.Voice != "Fritz-PlayAI" { - t.Errorf("forwarded request = %+v, want model/voice preserved", req) - } + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "audio/mpeg") _, _ = w.Write(audio) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }) + provider := newTestProvider(server.URL) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "playai-tts", Input: "Hello from Groq.", Voice: "Fritz-PlayAI", }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if resp.ContentType != "audio/mpeg" { - t.Errorf("ContentType = %q, want audio/mpeg", resp.ContentType) - } - if string(resp.Data) != string(audio) { - t.Errorf("Data = %q, want %q", resp.Data, audio) - } + require.NoError(t, err) + + sent := capture.Last(t) + assert.Equal(t, "/audio/speech", sent.Path) + assert.Equal(t, "Bearer "+testAPIKey, sent.Header.Get("Authorization")) + body := sent.JSON(t) + assert.Equal(t, "playai-tts", body["model"]) + assert.Equal(t, "Fritz-PlayAI", body["voice"]) + + assert.Equal(t, "audio/mpeg", resp.ContentType) + assert.Equal(t, string(audio), string(resp.Data)) } func TestCreateTranscription(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/audio/transcriptions" { - t.Errorf("path = %q, want /audio/transcriptions", r.URL.Path) - } + var gotModel, gotFile string + server, capture := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { if err := r.ParseMultipartForm(1 << 20); err != nil { - t.Fatalf("parse multipart: %v", err) - } - if got := r.FormValue("model"); got != "whisper-large-v3" { - t.Errorf("model field = %q, want whisper-large-v3", got) + http.Error(w, err.Error(), http.StatusBadRequest) + return } + gotModel = r.FormValue("model") file, _, err := r.FormFile("file") if err != nil { - t.Fatalf("file part: %v", err) + http.Error(w, err.Error(), http.StatusBadRequest) + return } defer func() { _ = file.Close() }() data, _ := io.ReadAll(file) - if string(data) != "fake-wav-bytes" { - t.Errorf("file content = %q, want fake-wav-bytes", data) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"text":"hello from groq"}`)) - })) - defer server.Close() + gotFile = string(data) - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"text":"hello from groq"}`) + }) + provider := newTestProvider(server.URL) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "whisper-large-v3", Filename: "speech.wav", File: []byte("fake-wav-bytes"), }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } - if !strings.Contains(string(resp.Data), "hello from groq") { - t.Errorf("Data = %s, want transcription text", resp.Data) - } + require.NoError(t, err) + assert.Equal(t, "/audio/transcriptions", capture.Last(t).Path) + assert.Equal(t, "whisper-large-v3", gotModel) + assert.Equal(t, "fake-wav-bytes", gotFile) + assert.Contains(t, string(resp.Data), "hello from groq") } func TestCreateTranscription_UpstreamError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"error":{"message":"invalid audio"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusBadRequest, `{"error":{"message":"invalid audio"}}`) + provider := newTestProvider(server.URL) _, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "whisper-large-v3", Filename: "speech.wav", File: []byte("bad"), }) - if err == nil { - t.Fatal("CreateTranscription() error = nil, want upstream error") - } - if !strings.Contains(err.Error(), "invalid audio") { - t.Errorf("error = %v, want upstream message propagated", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid audio") } diff --git a/internal/providers/groq/reasoning_response_test.go b/internal/providers/groq/reasoning_response_test.go index f2393b9e6..c57607f1f 100644 --- a/internal/providers/groq/reasoning_response_test.go +++ b/internal/providers/groq/reasoning_response_test.go @@ -6,6 +6,8 @@ import ( "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" ) @@ -39,22 +41,18 @@ func TestNormalizeChatResponse(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { var msg core.ResponseMessage - if err := json.Unmarshal([]byte(tt.message), &msg); err != nil { - t.Fatalf("unmarshal message: %v", err) - } + err := json.Unmarshal([]byte(tt.message), &msg) + require.NoError(t, err) + resp := &core.ChatResponse{Choices: []core.Choice{{Message: msg}}} normalizeChatResponse(resp) encoded, err := json.Marshal(resp.Choices[0].Message) - if err != nil { - t.Fatalf("marshal message: %v", err) - } - if !strings.Contains(string(encoded), tt.want) { - t.Errorf("message = %s, want it to contain %s", encoded, tt.want) - } - if tt.absent != "" && strings.Contains(string(encoded), tt.absent) { - t.Errorf("message = %s, want it to drop %s", encoded, tt.absent) + require.NoError(t, err) + assert.Contains(t, string(encoded), tt.want) + if tt.absent != "" { + assert.NotContains(t, string(encoded), tt.absent) } }) } @@ -115,18 +113,13 @@ func TestNormalizeChatStream(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got, err := io.ReadAll(normalizeChatStream(io.NopCloser(strings.NewReader(tt.in)))) - if err != nil { - t.Fatalf("read stream: %v", err) - } - if string(got) != tt.want { - t.Errorf("stream = %q, want %q", got, tt.want) - } + require.NoError(t, err) + assert.Equal(t, tt.want, string(got)) }) } } func TestNormalizeChatStreamNilIsSafe(t *testing.T) { - if got := normalizeChatStream(nil); got != nil { - t.Errorf("normalizeChatStream(nil) = %v, want nil", got) - } + got := normalizeChatStream(nil) + assert.Nil(t, got) } diff --git a/internal/providers/groq/reasoning_test.go b/internal/providers/groq/reasoning_test.go index a70362ad2..ea1fea3bf 100644 --- a/internal/providers/groq/reasoning_test.go +++ b/internal/providers/groq/reasoning_test.go @@ -4,13 +4,16 @@ import ( "context" "encoding/json" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" - "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +const minimalChatCompletionJSON = `{"id":"c1","object":"chat.completion","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}` + func TestChatCompletion_MapsReasoningPerModelFamily(t *testing.T) { tests := []struct { name string @@ -32,37 +35,23 @@ func TestChatCompletion_MapsReasoningPerModelFamily(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var raw map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&raw); err != nil { - t.Errorf("decode request: %v", err) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"c1","object":"chat.completion","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, minimalChatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: tt.model, Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: tt.effort}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := raw["reasoning"]; ok { - t.Errorf("request body includes nested reasoning: %v", raw["reasoning"]) - } - got, ok := raw["reasoning_effort"] + require.NoError(t, err) + + raw := capture.Last(t).JSON(t) + assert.NotContains(t, raw, "reasoning", "request body includes nested reasoning") if tt.wantEffort == "" { - if ok { - t.Errorf("reasoning_effort = %v, want absent", got) - } - } else if got != tt.wantEffort { - t.Errorf("reasoning_effort = %v, want %q", got, tt.wantEffort) + assert.NotContains(t, raw, "reasoning_effort") + return } + assert.Equal(t, tt.wantEffort, raw["reasoning_effort"]) }) } } @@ -83,41 +72,26 @@ func TestChatCompletion_DefaultsReasoningFormatPerModelFamily(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var raw map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&raw); err != nil { - t.Errorf("decode request: %v", err) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"c1","object":"chat.completion","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, minimalChatCompletionJSON) + provider := newTestProvider(server.URL) req := &core.ChatRequest{Model: tt.model, Messages: []core.Message{{Role: "user", Content: "hi"}}} if tt.caller != "" { extra, err := core.MergeUnknownJSONFields(req.ExtraFields, map[string]json.RawMessage{ "reasoning_format": json.RawMessage(`"` + tt.caller + `"`), }) - if err != nil { - t.Fatalf("MergeUnknownJSONFields() error = %v", err) - } + require.NoError(t, err) req.ExtraFields = extra } - if _, err := provider.ChatCompletion(context.Background(), req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got, ok := raw["reasoning_format"] + _, err := provider.ChatCompletion(context.Background(), req) + require.NoError(t, err) + + raw := capture.Last(t).JSON(t) if tt.wantFormat == nil { - if ok { - t.Errorf("reasoning_format = %v, want absent", got) - } + assert.NotContains(t, raw, "reasoning_format") return } - if got != tt.wantFormat { - t.Errorf("reasoning_format = %v, want %v", got, tt.wantFormat) - } + assert.Equal(t, tt.wantFormat, raw["reasoning_format"]) }) } } diff --git a/internal/providers/health/tracker_test.go b/internal/providers/health/tracker_test.go index 5579af7a3..2b2ec71fe 100644 --- a/internal/providers/health/tracker_test.go +++ b/internal/providers/health/tracker_test.go @@ -10,6 +10,8 @@ import ( "unicode/utf8" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func newTestTracker(start time.Time) (*Tracker, *time.Time) { @@ -210,9 +212,7 @@ func TestTrackerHooksFeedRecord(t *testing.T) { hooks.OnRequestEnd(t.Context(), llmclient.ResponseInfo{Provider: "openai", Model: "gpt-4o", StatusCode: 200}) snapshot := tracker.Snapshot() - if snapshot["openai"].Requests != 1 { - t.Fatalf("Snapshot()[openai].Requests = %d, want 1", snapshot["openai"].Requests) - } + require.Equal(t, 1, snapshot["openai"].Requests) } func TestTrackerEmptyResponsesFlagModel(t *testing.T) { @@ -240,16 +240,15 @@ func TestTrackerEmptyResponsesFlagModel(t *testing.T) { } snapshot := tracker.Snapshot()["openai"] - if len(snapshot.Models) != 1 { - t.Fatalf("models = %+v, want one row", snapshot.Models) - } + require.Len(t, snapshot.Models, 1) + row := snapshot.Models[0] - if row.Requests != tt.wantRequests || row.Errors != tt.wantErrors || row.Flagged != tt.wantFlagged { - t.Fatalf("row = %+v, want requests=%d errors=%d flagged=%v", row, tt.wantRequests, tt.wantErrors, tt.wantFlagged) - } - if row.LastError == nil || row.LastError.StatusCode != 200 || !strings.Contains(row.LastError.Message, "no_choices") { - t.Fatalf("last error = %+v, want a 200 no_choices error", row.LastError) - } + assert.Equal(t, tt.wantRequests, row.Requests) + assert.Equal(t, tt.wantErrors, row.Errors) + assert.Equal(t, tt.wantFlagged, row.Flagged) + require.NotNil(t, row.LastError) + assert.Equal(t, 200, row.LastError.StatusCode) + assert.Contains(t, row.LastError.Message, "no_choices") }) } } @@ -270,9 +269,9 @@ func TestTrackerEmptyResponseInterleavedWithOtherRequests(t *testing.T) { tracker.Record(success) // D row := tracker.Snapshot()["openai"].Models[0] - if row.Requests != 4 || row.Errors != 3 || !row.Flagged { - t.Fatalf("row = %+v, want requests=4 errors=3 flagged", row) - } + assert.Equal(t, 4, row.Requests) + assert.Equal(t, 3, row.Errors) + assert.True(t, row.Flagged) } func TestTrackerEvictsStalestModel(t *testing.T) { @@ -284,15 +283,9 @@ func TestTrackerEvictsStalestModel(t *testing.T) { tracker.Record(llmclient.ResponseInfo{Provider: "openai", Model: "one-too-many", StatusCode: 200}) models := tracker.providers["openai"].models - if len(models) != maxTrackedModels { - t.Fatalf("tracked models = %d, want %d", len(models), maxTrackedModels) - } - if _, ok := models["model-000"]; ok { - t.Fatalf("expected stalest model model-000 to be evicted") - } - if _, ok := models["one-too-many"]; !ok { - t.Fatalf("expected newest model to be tracked") - } + require.Len(t, models, maxTrackedModels) + assert.NotContains(t, models, "model-000") + assert.Contains(t, models, "one-too-many") } func TestTrackerCapsEventsPerModel(t *testing.T) { @@ -300,9 +293,7 @@ func TestTrackerCapsEventsPerModel(t *testing.T) { for range maxEventsPerModel + 50 { tracker.Record(llmclient.ResponseInfo{Provider: "openai", Model: "gpt-4o", StatusCode: 200}) } - if got := len(tracker.providers["openai"].models["gpt-4o"].events); got != maxEventsPerModel { - t.Fatalf("events kept = %d, want %d", got, maxEventsPerModel) - } + require.Len(t, tracker.providers["openai"].models["gpt-4o"].events, maxEventsPerModel) } func TestTrackerTruncatesLongErrorMessages(t *testing.T) { @@ -314,12 +305,8 @@ func TestTrackerTruncatesLongErrorMessages(t *testing.T) { Error: errors.New(strings.Repeat("x", maxErrorMessageLen+100)), }) message := tracker.Snapshot()["openai"].Models[0].LastError.Message - if len(message) > maxErrorMessageLen+len("…") { - t.Fatalf("error message length = %d, want <= %d", len(message), maxErrorMessageLen+len("…")) - } - if !strings.HasSuffix(message, "…") { - t.Fatalf("expected truncated message to end with ellipsis") - } + assert.LessOrEqual(t, len(message), maxErrorMessageLen+len("…")) + assert.True(t, strings.HasSuffix(message, "…"), "message = %q, want truncation marker", message) } func TestTrackerSnapshotCapsModelRowsTroubledFirst(t *testing.T) { @@ -332,16 +319,12 @@ func TestTrackerSnapshotCapsModelRowsTroubledFirst(t *testing.T) { } snapshot := tracker.Snapshot()["openai"] - if len(snapshot.Models) != maxSnapshotModels { - t.Fatalf("model rows = %d, want %d", len(snapshot.Models), maxSnapshotModels) - } - if snapshot.Models[0].Model != "broken" || !snapshot.Models[0].Flagged { - t.Fatalf("expected flagged model first, got %+v", snapshot.Models[0]) - } + require.Len(t, snapshot.Models, maxSnapshotModels) + assert.Equal(t, "broken", snapshot.Models[0].Model) + assert.True(t, snapshot.Models[0].Flagged, "expected flagged model first") + // Provider totals still cover every tracked model, not just listed rows. - if snapshot.Requests != maxSnapshotModels+5+4 { - t.Fatalf("provider requests = %d, want %d", snapshot.Requests, maxSnapshotModels+5+4) - } + assert.Equal(t, maxSnapshotModels+5+4, snapshot.Requests) } func TestTrackerProviderLastErrorSurvivesModelCap(t *testing.T) { @@ -369,21 +352,12 @@ func TestTrackerProviderLastErrorSurvivesModelCap(t *testing.T) { }) snapshot := tracker.Snapshot()["router"] - listed := false for _, row := range snapshot.Models { - if row.Model == "quiet-model" { - listed = true - } - } - if listed { - t.Fatalf("expected quiet-model to be dropped by the snapshot cap") - } - if snapshot.LastError == nil || snapshot.LastError.Message != "newest failure" { - t.Fatalf("provider LastError = %+v, want newest failure", snapshot.LastError) - } - if snapshot.LastErrorModel != "quiet-model" { - t.Fatalf("LastErrorModel = %q, want quiet-model", snapshot.LastErrorModel) + assert.NotEqual(t, "quiet-model", row.Model, "quiet-model should be dropped from the capped listing") } + require.NotNil(t, snapshot.LastError) + assert.Equal(t, "newest failure", snapshot.LastError.Message) + assert.Equal(t, "quiet-model", snapshot.LastErrorModel) } func TestErrorMessageTruncationIsRuneSafe(t *testing.T) { @@ -396,12 +370,8 @@ func TestErrorMessageTruncationIsRuneSafe(t *testing.T) { Error: errors.New(strings.Repeat("é", maxErrorMessageLen)), }) message := tracker.Snapshot()["openai"].Models[0].LastError.Message - if !strings.HasSuffix(message, "…") { - t.Fatalf("expected truncated message, got %q", message) - } - if !utf8.ValidString(message) { - t.Fatalf("truncated message is not valid UTF-8: %q", message) - } + assert.True(t, strings.HasSuffix(message, "…"), "expected truncated message, got %q", message) + assert.True(t, utf8.ValidString(message), "truncated message is not valid UTF-8: %q", message) } func TestProviderHealthFlaggedModels(t *testing.T) { @@ -410,45 +380,34 @@ func TestProviderHealthFlaggedModels(t *testing.T) { {Model: "b"}, {Model: "c", Flagged: true}, }} - got := snapshot.FlaggedModels() - if len(got) != 2 || got[0] != "a" || got[1] != "c" { - t.Fatalf("FlaggedModels() = %v, want [a c]", got) - } + assert.Equal(t, []string{"a", "c"}, snapshot.FlaggedModels()) } func assertSnapshotsEqual(t *testing.T, got, want map[string]ProviderHealth) { t.Helper() - if len(got) != len(want) { - t.Fatalf("snapshot providers = %d (%v), want %d", len(got), got, len(want)) - } + require.Len(t, got, len(want), "snapshot providers = %v", got) + for name, wantProvider := range want { gotProvider, ok := got[name] - if !ok { - t.Fatalf("missing provider %q in snapshot", name) - } - if gotProvider.CircuitState != wantProvider.CircuitState || - gotProvider.WindowSeconds != wantProvider.WindowSeconds || - gotProvider.Requests != wantProvider.Requests || - gotProvider.Errors != wantProvider.Errors { - t.Fatalf("provider %q = %+v, want %+v", name, gotProvider, wantProvider) - } - if len(gotProvider.Models) != len(wantProvider.Models) { - t.Fatalf("provider %q models = %+v, want %+v", name, gotProvider.Models, wantProvider.Models) - } + require.True(t, ok, "missing provider %q in snapshot", name) + assert.Equal(t, wantProvider.CircuitState, gotProvider.CircuitState, "provider %q", name) + assert.Equal(t, wantProvider.WindowSeconds, gotProvider.WindowSeconds, "provider %q", name) + assert.Equal(t, wantProvider.Requests, gotProvider.Requests, "provider %q", name) + assert.Equal(t, wantProvider.Errors, gotProvider.Errors, "provider %q", name) + require.Len(t, gotProvider.Models, len(wantProvider.Models), "provider %q models = %+v", name, gotProvider.Models) + for i, wantModel := range wantProvider.Models { gotModel := gotProvider.Models[i] - if gotModel.Model != wantModel.Model || - gotModel.Requests != wantModel.Requests || - gotModel.Errors != wantModel.Errors || - gotModel.Flagged != wantModel.Flagged { - t.Fatalf("provider %q model[%d] = %+v, want %+v", name, i, gotModel, wantModel) - } - if (gotModel.LastError == nil) != (wantModel.LastError == nil) { - t.Fatalf("provider %q model[%d] last_error = %+v, want %+v", name, i, gotModel.LastError, wantModel.LastError) - } - if wantModel.LastError != nil && *gotModel.LastError != *wantModel.LastError { - t.Fatalf("provider %q model[%d] last_error = %+v, want %+v", name, i, *gotModel.LastError, *wantModel.LastError) + assert.Equal(t, wantModel.Model, gotModel.Model, "provider %q model[%d]", name, i) + assert.Equal(t, wantModel.Requests, gotModel.Requests, "provider %q model[%d]", name, i) + assert.Equal(t, wantModel.Errors, gotModel.Errors, "provider %q model[%d]", name, i) + assert.Equal(t, wantModel.Flagged, gotModel.Flagged, "provider %q model[%d]", name, i) + if wantModel.LastError == nil { + assert.Nil(t, gotModel.LastError, "provider %q model[%d] last_error", name, i) + continue } + require.NotNil(t, gotModel.LastError, "provider %q model[%d] last_error", name, i) + assert.Equal(t, *wantModel.LastError, *gotModel.LastError, "provider %q model[%d] last_error", name, i) } } } diff --git a/internal/providers/hetzner/hetzner_test.go b/internal/providers/hetzner/hetzner_test.go index aa927e18f..6687e7b55 100644 --- a/internal/providers/hetzner/hetzner_test.go +++ b/internal/providers/hetzner/hetzner_test.go @@ -1,325 +1,26 @@ package hetzner import ( - "context" - "encoding/json" - "errors" - "io" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -// TestNew_ReturnsProvider asserts that New returns a non-nil *Provider whose embedded -// ChatCompatible is non-nil. Matches kimicode's surface. -func TestNew_ReturnsProvider(t *testing.T) { - provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) - - if provider == nil { - t.Fatal("provider should not be nil") - } - - concrete, ok := provider.(*Provider) - if !ok { - t.Fatalf("New() returned %T, want *hetzner.Provider", provider) - } - if concrete.ChatCompatible == nil { - t.Error("embedded ChatCompatible should not be nil") - } -} - -// TestNewWithHTTPClient_ReturnsProvider asserts the explicit HTTP-client constructor -// returns a valid Provider with a non-nil ChatCompatible. -func TestNewWithHTTPClient_ReturnsProvider(t *testing.T) { - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) - - if provider == nil { - t.Fatal("provider should not be nil") - } - if provider.ChatCompatible == nil { - t.Error("embedded ChatCompatible should not be nil") - } -} - -// TestNewWithHTTPClient_NilHTTPClientDoesNotPanic asserts that passing nil for the -// HTTP client falls back to http.DefaultClient without panicking. -func TestNewWithHTTPClient_NilHTTPClientDoesNotPanic(t *testing.T) { - defer func() { - if r := recover(); r != nil { - t.Fatalf("NewWithHTTPClient(nil, ...) panicked: %v", r) - } - }() - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("provider should not be nil") - } -} - -// TestNewWithHTTPClient_ZeroHooksDoesNotPanic asserts that the hooks argument can be -// an empty struct (no hooks registered) without panicking. -func TestNewWithHTTPClient_ZeroHooksDoesNotPanic(t *testing.T) { - defer func() { - if r := recover(); r != nil { - t.Fatalf("NewWithHTTPClient(..., llmclient.Hooks{}) panicked: %v", r) - } - }() - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) - if provider == nil { - t.Fatal("provider should not be nil") - } -} - -// TestRegistration_TypeAndDiscovery asserts the Registration struct exposes the -// expected type, New function, and default base URL. -func TestRegistration_TypeAndDiscovery(t *testing.T) { - if Registration.Type != "hetzner" { - t.Errorf("Registration.Type = %q, want %q", Registration.Type, "hetzner") - } - if Registration.New == nil { - t.Error("Registration.New should not be nil") - } - if Registration.Discovery.DefaultBaseURL == "" { - t.Error("Registration.Discovery.DefaultBaseURL should not be empty") - } - want := "https://inference.hetzner.com/api/v1" - if Registration.Discovery.DefaultBaseURL != want { - t.Errorf("Registration.Discovery.DefaultBaseURL = %q, want %q", Registration.Discovery.DefaultBaseURL, want) - } -} - -// TestProvider_ImplementsCoreProvider is a compile-time check that *Provider -// satisfies the core.Provider interface used by the factory. -func TestProvider_ImplementsCoreProvider(t *testing.T) { - var _ core.Provider = (*Provider)(nil) -} - -// TestChatCompletion_UsesBearerAuthAndForwardsModel asserts that ChatCompletion -// posts to /chat/completions with the Bearer header and forwards the requested -// model unchanged. -func TestChatCompletion_UsesBearerAuthAndForwardsModel(t *testing.T) { - var gotPath string - var gotAuth string - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-hetzner", - "created":1677652288, - "model":"Qwen/Qwen3.6-35B-A3B-FP8", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("hetzner-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "Qwen/Qwen3.6-35B-A3B-FP8", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer hetzner-key" { - t.Fatalf("authorization = %q, want Bearer hetzner-key", gotAuth) - } - if gotBody["model"] != "Qwen/Qwen3.6-35B-A3B-FP8" { - t.Fatalf("request model = %#v, want Qwen/Qwen3.6-35B-A3B-FP8", gotBody["model"]) - } - if resp.Model != "Qwen/Qwen3.6-35B-A3B-FP8" { - t.Fatalf("response model = %q, want Qwen/Qwen3.6-35B-A3B-FP8", resp.Model) - } - if len(resp.Choices) != 1 || resp.Choices[0].Message.Content != "hello" { - t.Fatalf("unexpected response: %+v", resp) - } -} - -// TestStreamChatCompletion_UsesSSE asserts that streaming requests go to -// /chat/completions with the Bearer header, set stream=true, and return SSE data -// the adapter normalizes. -func TestStreamChatCompletion_UsesSSE(t *testing.T) { - var gotPath string - var gotAuth string - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-hetzner\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hi\"}}]}\n\ndata: [DONE]\n\n") - })) - defer server.Close() - - provider := NewWithHTTPClient("hetzner-key", server.URL, server.Client(), llmclient.Hooks{}) - stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ - Model: "Qwen/Qwen3.6-35B-A3B-FP8", - Messages: []core.Message{{Role: "user", Content: "hi"}}, - }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } - defer stream.Close() - body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer hetzner-key" { - t.Fatalf("authorization = %q, want Bearer hetzner-key", gotAuth) - } - if gotBody["model"] != "Qwen/Qwen3.6-35B-A3B-FP8" || gotBody["stream"] != true { - t.Fatalf("stream request body = %#v", gotBody) - } - if !strings.Contains(string(body), "data: [DONE]") { - t.Fatalf("stream body = %q, want SSE terminator", body) - } -} - -// TestListModels_ForwardsToModelsEndpoint asserts that ListModels calls -// /v1/models and returns the parsed model list unchanged. -func TestListModels_ForwardsToModelsEndpoint(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"Qwen/Qwen3.6-35B-A3B-FP8","object":"model","owned_by":"alibaba"}]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("hetzner-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotPath != "/models" { - t.Fatalf("path = %q, want /models", gotPath) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "Qwen/Qwen3.6-35B-A3B-FP8" { - t.Fatalf("models = %+v, want one hetzner model", resp.Data) - } -} - -// TestEmbeddings_ReturnsUnsupportedError asserts that Embeddings returns a typed -// "not supported" error without calling upstream — Hetzner documents no embeddings -// endpoint, so the provider overrides the embedded adapter to fail fast. The -// httptest server asserts zero requests: a regression that forwards embeddings -// upstream fails this test deterministically instead of hitting the network. -// The typed contract is asserted via errors.As against *core.GatewayError so the -// test would fail on a plain error with the same text. -func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { - var requests int - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - requests++ - http.Error(w, "should not be called", http.StatusInternalServerError) - })) - defer server.Close() - - provider := NewWithHTTPClient("hetzner-key", server.URL, server.Client(), llmclient.Hooks{}) - _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{Model: "any"}) - if err == nil { - t.Fatal("Embeddings() error = nil, want typed unsupported error") - } - - var gwErr *core.GatewayError - if !errors.As(err, &gwErr) { - t.Fatalf("Embeddings() error type = %T, want *core.GatewayError", err) - } - if gwErr.Type != core.ErrorTypeInvalidRequest { - t.Errorf("error Type = %v, want %v", gwErr.Type, core.ErrorTypeInvalidRequest) - } - if gwErr.StatusCode != http.StatusBadRequest { - t.Errorf("error StatusCode = %v, want %v", gwErr.StatusCode, http.StatusBadRequest) - } - if !strings.Contains(err.Error(), "hetzner does not support embeddings") { - t.Errorf("Embeddings() error = %v, want message containing \"hetzner does not support embeddings\"", err) - } - if requests != 0 { - t.Fatalf("upstream received %d requests, want 0 (embeddings must not be forwarded)", requests) - } -} - -// TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces mirrors the kilo -// guard: hetzner wraps *ChatCompatible which does not satisfy the optional native -// interfaces. If Hetzner ever gains native batch/file/audio support, the test -// fails and the implementation must add explicit method overrides to remove -// capabilities it cannot honour upstream. -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("hetzner-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("hetzner provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("hetzner provider should not implement native file provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("hetzner provider should not implement audio provider") - } -} - -// TestResponses_TranslatesToChatCompletions asserts that a Responses API request is -// translated to a chat-completions call (the doc claims /v1/responses is served via -// chat translation; this test keeps that claim honest). -func TestResponses_TranslatesToChatCompletions(t *testing.T) { - var gotPath string - var gotBody struct { - Model string `json:"model"` - } - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-hetzner", - "created":1677652288, - "model":"Qwen/Qwen3.6-35B-A3B-FP8", - "choices":[{"index":0,"message":{"role":"assistant","content":"translated"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("hetzner-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ - Model: "Qwen/Qwen3.6-35B-A3B-FP8", - Input: "hi", +// Hetzner is a thin wrapper over the shared chat-centric adapter, so the +// shared contract covers its surface. Hetzner documents no embeddings +// endpoint, so Embeddings must fail fast without an upstream call, and the +// provider must not advertise native batch, file, or audio support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "hetzner", + DefaultBaseURL: "https://inference.hetzner.com/api/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) + }, }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotBody.Model != "Qwen/Qwen3.6-35B-A3B-FP8" { - t.Fatalf("request model = %q, want Qwen/Qwen3.6-35B-A3B-FP8", gotBody.Model) - } - if resp.Object != "response" || resp.Status != "completed" { - t.Fatalf("response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) - } + providertest.AssertNoNativeSurfaces(t, NewWithHTTPClient("hetzner-key", "", nil, llmclient.Hooks{})) } diff --git a/internal/providers/kilo/kilo_test.go b/internal/providers/kilo/kilo_test.go index d7afd5acb..898ad30f7 100644 --- a/internal/providers/kilo/kilo_test.go +++ b/internal/providers/kilo/kilo_test.go @@ -2,39 +2,41 @@ package kilo import ( "context" - "encoding/json" "io" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_UsesBearerAuthAndPreservesModelAndTools(t *testing.T) { - var gotPath string - var gotAuth string - var gotBody map[string]any +// Kilo is a thin wrapper over the shared chat-centric adapter, so the shared +// contract covers its surface. Kilo exposes no embeddings endpoint, so +// Embeddings must fail fast without an upstream call, and the provider must +// not advertise native batch, file, or audio support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "kilo", + DefaultBaseURL: "https://api.kilo.ai/api/gateway", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) + }, + }) + providertest.AssertNoNativeSurfaces(t, NewWithHTTPClient("kilo-key", "", nil, llmclient.Hooks{})) +} - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-kilo", - "created":1677652288, - "model":"anthropic/claude-sonnet-4.5", - "choices":[{"index":0,"message":{"role":"assistant","content":null,"tool_calls":[{"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{}"}}]},"finish_reason":"tool_calls"}], - "usage":{"prompt_tokens":10,"completion_tokens":2,"total_tokens":12} - }`)) - })) - defer server.Close() +func TestChatCompletion_PreservesToolsAndToolChoice(t *testing.T) { + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-kilo", + "created":1677652288, + "model":"anthropic/claude-sonnet-4.5", + "choices":[{"index":0,"message":{"role":"assistant","content":null,"tool_calls":[{"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{}"}}]},"finish_reason":"tool_calls"}], + "usage":{"prompt_tokens":10,"completion_tokens":2,"total_tokens":12} + }`) provider := NewWithHTTPClient("kilo-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ @@ -49,46 +51,18 @@ func TestChatCompletion_UsesBearerAuthAndPreservesModelAndTools(t *testing.T) { }}, ToolChoice: "auto", }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer kilo-key" { - t.Fatalf("authorization = %q, want Bearer kilo-key", gotAuth) - } - if gotBody["model"] != "anthropic/claude-sonnet-4.5" { - t.Fatalf("request model = %#v, want slash-delimited model unchanged", gotBody["model"]) - } - if gotBody["tool_choice"] != "auto" { - t.Fatalf("tool_choice = %#v, want auto", gotBody["tool_choice"]) - } - tools, ok := gotBody["tools"].([]any) - if !ok || len(tools) != 1 { - t.Fatalf("tools = %#v, want one tool", gotBody["tools"]) - } - if resp.Model != "anthropic/claude-sonnet-4.5" || len(resp.Choices) != 1 || len(resp.Choices[0].Message.ToolCalls) != 1 { - t.Fatalf("unexpected response: %+v", resp) - } -} + require.NoError(t, err) -func TestStreamChatCompletion_UsesSSEAndPreservesStreamOptions(t *testing.T) { - var gotPath string - var gotAuth string - var gotBody map[string]any + sent := capture.Last(t).JSON(t) + assert.Equal(t, "anthropic/claude-sonnet-4.5", sent["model"]) + assert.Equal(t, "auto", sent["tool_choice"]) + assert.Len(t, sent["tools"], 1) + require.Len(t, resp.Choices, 1) + assert.Len(t, resp.Choices[0].Message.ToolCalls, 1) +} - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-kilo\",\"object\":\"chat.completion.chunk\",\"choices\":[]}\n\ndata: [DONE]\n\n") - })) - defer server.Close() +func TestStreamChatCompletion_PreservesStreamOptions(t *testing.T) { + server, capture := providertest.SSEServer(t, "data: {\"id\":\"chatcmpl-kilo\",\"object\":\"chat.completion.chunk\",\"choices\":[]}\n\ndata: [DONE]\n\n") provider := NewWithHTTPClient("kilo-key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ @@ -96,69 +70,14 @@ func TestStreamChatCompletion_UsesSSEAndPreservesStreamOptions(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hi"}}, StreamOptions: &core.StreamOptions{IncludeUsage: true}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() - body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotPath != "/chat/completions" || gotAuth != "Bearer kilo-key" { - t.Fatalf("request path/auth = %q/%q", gotPath, gotAuth) - } - if gotBody["model"] != "openai/gpt-5.5" || gotBody["stream"] != true { - t.Fatalf("stream request body = %#v", gotBody) - } - streamOptions, ok := gotBody["stream_options"].(map[string]any) - if !ok || streamOptions["include_usage"] != true { - t.Fatalf("stream_options = %#v, want include_usage=true", gotBody["stream_options"]) - } - if !strings.Contains(string(body), "data: [DONE]") { - t.Fatalf("stream body = %q, want SSE terminator", body) - } -} - -func TestListModels_PreservesProviderQualifiedIDs(t *testing.T) { - var gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"google/gemini-3.1-pro","object":"model","owned_by":"google"}]}`)) - })) - defer server.Close() - provider := NewWithHTTPClient("kilo-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotPath != "/models" { - t.Fatalf("path = %q, want /models", gotPath) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "google/gemini-3.1-pro" { - t.Fatalf("models = %+v, want provider-qualified Kilo model", resp.Data) - } -} - -func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { - provider := NewWithHTTPClient("kilo-key", "", nil, llmclient.Hooks{}) - _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{Model: "any"}) - if err == nil || !strings.Contains(err.Error(), "kilo does not support embeddings") { - t.Fatalf("Embeddings() error = %v, want unsupported error", err) - } -} - -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("kilo-key", "", nil, llmclient.Hooks{}) + body, err := io.ReadAll(stream) + require.NoError(t, err) + assert.Contains(t, string(body), "data: [DONE]") - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("kilo provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("kilo provider should not implement native file provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("kilo provider should not implement audio provider") - } + sent := capture.Last(t).JSON(t) + assert.Equal(t, true, sent["stream"]) + assert.Equal(t, map[string]any{"include_usage": true}, sent["stream_options"]) } diff --git a/internal/providers/kilo/passthrough_semantics_test.go b/internal/providers/kilo/passthrough_semantics_test.go index ed5b30774..d38c48619 100644 --- a/internal/providers/kilo/passthrough_semantics_test.go +++ b/internal/providers/kilo/passthrough_semantics_test.go @@ -4,21 +4,18 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricher(t *testing.T) { - if got := passthroughSemanticEnricher.ProviderType(); got != "kilo" { - t.Fatalf("ProviderType() = %q, want kilo", got) - } + assert.Equal(t, "kilo", passthroughSemanticEnricher.ProviderType()) got := passthroughSemanticEnricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ RawEndpoint: "v1/chat/completions", NormalizedEndpoint: "chat/completions", }) - if got == nil { - t.Fatal("Enrich() returned nil") - } - if got.SemanticOperation != "kilo.chat_completions" || got.AuditPath != "/v1/chat/completions" { - t.Fatalf("enriched info = %+v", got) - } + require.NotNil(t, got) + assert.Equal(t, "kilo.chat_completions", got.SemanticOperation) + assert.Equal(t, "/v1/chat/completions", got.AuditPath) } diff --git a/internal/providers/kimicode/kimicode_test.go b/internal/providers/kimicode/kimicode_test.go index 8724db9b0..245f70a02 100644 --- a/internal/providers/kimicode/kimicode_test.go +++ b/internal/providers/kimicode/kimicode_test.go @@ -1,105 +1,25 @@ package kimicode import ( - "context" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -// kimibgeM3EmbedModel is the model ID used by these tests for the Kimi Code /embeddings endpoint. -// -// NOTE: "bge_m3_embed" is not part of Kimi Code's documented public model catalogue. It is -// retained here because the Kimi Code provider package currently forwards embedding model IDs -// unchanged through the OpenAI-compatible adapter. Because the model is undocumented upstream, -// its name, behaviour, or availability may change without notice; if Kimi Code rotates the ID -// these tests (and the provider's embedding round-trip) will need to be updated. -const kimibgeM3EmbedModel = "bge_m3_embed" - -func TestNew_ReturnsProvider(t *testing.T) { - provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) - - if provider == nil { - t.Fatal("provider should not be nil") - } - - concrete, ok := provider.(*Provider) - if !ok { - t.Fatalf("New() returned %T, want *kimicode.Provider", provider) - } - if concrete.ChatCompatible == nil { - t.Error("embedded ChatCompatible should not be nil") - } -} - -func TestNewWithHTTPClient_ReturnsProvider(t *testing.T) { - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) - - if provider == nil { - t.Fatal("provider should not be nil") - } - if provider.ChatCompatible == nil { - t.Error("embedded ChatCompatible should not be nil") - } -} - -func TestRegistration_TypeIsKimicode(t *testing.T) { - if Registration.Type != "kimicode" { - t.Errorf("Registration.Type = %q, want %q", Registration.Type, "kimicode") - } - if Registration.New == nil { - t.Error("Registration.New should not be nil") - } - if Registration.Discovery.DefaultBaseURL == "" { - t.Error("Registration.Discovery.DefaultBaseURL should not be empty") - } -} - -func TestEmbeddings_RoundTrip(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object": "list", - "model": "bge_m3_embed", - "data": [ - {"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0} - ], - "usage": {"prompt_tokens": 3, "total_tokens": 3} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("kimi-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: kimibgeM3EmbedModel, - Input: "hello", +// Kimi Code is a thin wrapper over the shared chat-centric adapter and +// forwards embeddings upstream unchanged, so the shared contract covers its +// surface. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "kimicode", + DefaultBaseURL: "https://api.kimi.com/coding/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) + }, + Embeddings: true, }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp == nil { - t.Fatal("Embeddings() response should not be nil") - } - if resp.Model != kimibgeM3EmbedModel { - t.Errorf("resp.Model = %q, want %q", resp.Model, kimibgeM3EmbedModel) - } - if len(resp.Data) != 1 { - t.Fatalf("len(resp.Data) = %d, want 1", len(resp.Data)) - } - if gotPath != "/embeddings" { - t.Errorf("path = %q, want /embeddings", gotPath) - } - if gotAuth != "Bearer kimi-key" { - t.Errorf("authorization = %q, want %q", gotAuth, "Bearer kimi-key") - } } diff --git a/internal/providers/llamacpp/llamacpp_test.go b/internal/providers/llamacpp/llamacpp_test.go index abd2d2c85..94e7f4721 100644 --- a/internal/providers/llamacpp/llamacpp_test.go +++ b/internal/providers/llamacpp/llamacpp_test.go @@ -4,15 +4,19 @@ import ( "context" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +var _ core.PassthroughProvider = (*Provider)(nil) + func TestChatCompletion_UsesOptionalBearerAuthAndChatEndpoint(t *testing.T) { tests := []struct { name string @@ -25,60 +29,36 @@ func TestChatCompletion_UsesOptionalBearerAuthAndChatEndpoint(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-llamacpp", - "created":1677652288, - "model":"gemma-3-4b-it", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-llamacpp", + "created":1677652288, + "model":"gemma-3-4b-it", + "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] + }`) provider := NewWithHTTPClient(tt.apiKey, server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "gemma-3-4b-it", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: "gemma-3-4b-it", + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "gemma-3-4b-it" { - t.Fatalf("resp.Model = %q, want gemma-3-4b-it", resp.Model) - } - if gotPath != "/v1/chat/completions" { - t.Fatalf("path = %q, want /v1/chat/completions", gotPath) - } - if gotAuth != tt.wantAuth { - t.Fatalf("authorization = %q, want %q", gotAuth, tt.wantAuth) - } + require.NoError(t, err) + assert.Equal(t, "gemma-3-4b-it", resp.Model) + + req := capture.Last(t) + assert.Equal(t, "/v1/chat/completions", req.Path) + assert.Equal(t, tt.wantAuth, req.Header.Get("Authorization")) }) } } func TestEmbeddings_DelegatesToCompatibleProvider(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "model":"nomic-embed-text-v1.5", - "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], - "usage":{"prompt_tokens":3,"total_tokens":3} - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "object":"list", + "model":"nomic-embed-text-v1.5", + "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], + "usage":{"prompt_tokens":3,"total_tokens":3} + }`) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) @@ -86,35 +66,16 @@ func TestEmbeddings_DelegatesToCompatibleProvider(t *testing.T) { Model: "nomic-embed-text-v1.5", Input: "hello", }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Model != "nomic-embed-text-v1.5" { - t.Fatalf("resp.Model = %q, want nomic-embed-text-v1.5", resp.Model) - } - if gotPath != "/v1/embeddings" { - t.Fatalf("path = %q, want /v1/embeddings", gotPath) - } + require.NoError(t, err) + assert.Equal(t, "nomic-embed-text-v1.5", resp.Model) + assert.Equal(t, "/v1/embeddings", capture.Last(t).Path) } -func TestProvider_ExposesPassthroughButNotOptionalNativeInterfaces(t *testing.T) { +func TestProvider_DoesNotExposeOptionalNativeInterfaces(t *testing.T) { provider := NewWithHTTPClient("", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.PassthroughProvider); !ok { - t.Fatal("llamacpp provider should implement passthrough provider") - } - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("llamacpp provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("llamacpp provider should not implement native file provider") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("llamacpp provider should not implement native response lifecycle provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("llamacpp provider should not implement audio provider") - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + require.False(t, ok) } func TestPassthrough_RoutesNativeEndpointsToServerRoot(t *testing.T) { @@ -136,18 +97,7 @@ func TestPassthrough_RoutesNativeEndpointsToServerRoot(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - var gotQuery string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotQuery = r.URL.RawQuery - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"ok":true}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"ok":true}`) provider := NewWithHTTPClient("llamacpp-key", server.URL+"/v1", server.Client(), llmclient.Hooks{}) @@ -157,33 +107,19 @@ func TestPassthrough_RoutesNativeEndpointsToServerRoot(t *testing.T) { Body: io.NopCloser(strings.NewReader("{}")), Headers: http.Header{"Content-Type": []string{"application/json"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } - if gotQuery != tt.wantQuery { - t.Fatalf("query = %q, want %q", gotQuery, tt.wantQuery) - } - if gotAuth != "Bearer llamacpp-key" { - t.Fatalf("authorization = %q, want Bearer llamacpp-key", gotAuth) - } + req := capture.Last(t) + assert.Equal(t, tt.wantPath, req.Path) + assert.Equal(t, tt.wantQuery, req.Query.Encode()) + assert.Equal(t, "Bearer llamacpp-key", req.Header.Get("Authorization")) }) } } func TestPassthrough_NativeEndpointRotatesConfiguredKeys(t *testing.T) { - var gotAuths []string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotAuths = append(gotAuths, r.Header.Get("Authorization")) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"ok":true}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"ok":true}`) provider := New(providers.ProviderConfig{ APIKey: "key-1", @@ -199,40 +135,29 @@ func TestPassthrough_NativeEndpointRotatesConfiguredKeys(t *testing.T) { Endpoint: "health", Headers: http.Header{}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) resp.Body.Close() } - if len(gotAuths) != 2 || gotAuths[0] == gotAuths[1] { - t.Fatalf("native passthrough should rotate keys per request, got %v", gotAuths) - } - for _, auth := range gotAuths { - if auth != "Bearer key-1" && auth != "Bearer key-2" { - t.Fatalf("unexpected authorization %q", auth) - } + var gotAuths []string + for _, req := range capture.All() { + gotAuths = append(gotAuths, req.Header.Get("Authorization")) } + assert.ElementsMatch(t, []string{"Bearer key-1", "Bearer key-2"}, gotAuths) } func TestNew_DefaultOptionsAuthenticatesAndRelaysNativeErrors(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Header.Get("Authorization") != "Bearer llamacpp-key" { - t.Errorf("authorization = %q, want Bearer llamacpp-key", r.Header.Get("Authorization")) - } - switch r.URL.Path { - case "/health": + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/health": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"status":"ok"}`)) - case "/slots": + }, + "/slots": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusNotImplemented) _, _ = w.Write([]byte(`{"error":{"message":"slots endpoint is disabled"}}`)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - } - })) - defer server.Close() + }, + }) provider := New(providers.ProviderConfig{ APIKey: "llamacpp-key", @@ -244,14 +169,12 @@ func TestNew_DefaultOptionsAuthenticatesAndRelaysNativeErrors(t *testing.T) { Endpoint: "health", Headers: http.Header{}, }) - if err != nil { - t.Fatalf("Passthrough(health) error = %v", err) - } + require.NoError(t, err) + body, _ := io.ReadAll(resp.Body) resp.Body.Close() - if resp.StatusCode != http.StatusOK || !strings.Contains(string(body), `"ok"`) { - t.Fatalf("health status = %d body = %s, want 200 with ok", resp.StatusCode, body) - } + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Contains(t, string(body), `"ok"`) // Provider-native errors relay status and body verbatim instead of being // converted into gateway errors. @@ -260,15 +183,14 @@ func TestNew_DefaultOptionsAuthenticatesAndRelaysNativeErrors(t *testing.T) { Endpoint: "slots", Headers: http.Header{}, }) - if err != nil { - t.Fatalf("Passthrough(slots) error = %v", err) - } + require.NoError(t, err) + body, _ = io.ReadAll(resp.Body) resp.Body.Close() - if resp.StatusCode != http.StatusNotImplemented { - t.Fatalf("slots status = %d, want 501", resp.StatusCode) - } - if !strings.Contains(string(body), "slots endpoint is disabled") { - t.Fatalf("slots body = %s, want relayed upstream error", body) + assert.Equal(t, http.StatusNotImplemented, resp.StatusCode) + assert.Contains(t, string(body), "slots endpoint is disabled") + + for _, req := range capture.All() { + assert.Equal(t, "Bearer llamacpp-key", req.Header.Get("Authorization"), "path %s", req.Path) } } diff --git a/internal/providers/llamacpp/models_test.go b/internal/providers/llamacpp/models_test.go index 9289f7acc..c842af4ba 100644 --- a/internal/providers/llamacpp/models_test.go +++ b/internal/providers/llamacpp/models_test.go @@ -3,7 +3,6 @@ package llamacpp import ( "context" "net/http" - "net/http/httptest" "testing" "time" @@ -11,8 +10,31 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +// jsonRoute answers with status and body as JSON. +func jsonRoute(status int, body string) http.HandlerFunc { + return func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _, _ = w.Write([]byte(body)) + } +} + +// countPath returns how many recorded requests hit path. +func countPath(capture *providertest.Capture, path string) int { + n := 0 + for _, req := range capture.All() { + if req.Path == path { + n++ + } + } + return n +} + // legacyListing is the /v1/models payload of builds whose meta object predates // n_ctx, leaving only the GGUF's trained context. const legacyListing = `{ @@ -112,143 +134,79 @@ func TestListModels_SurfacesServerReportedMetadata(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - propsFetched := false - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - switch r.URL.Path { - case "/v1/models": - _, _ = w.Write([]byte(tt.listing)) - case "/props": - propsFetched = true - w.WriteHeader(tt.propsStatus) - _, _ = w.Write([]byte(tt.props)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - w.WriteHeader(http.StatusNotFound) - } - })) - defer server.Close() + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/v1/models": jsonRoute(http.StatusOK, tt.listing), + "/props": jsonRoute(tt.propsStatus, tt.props), + }) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 { - t.Fatalf("len(resp.Data) = %d, want 1", len(resp.Data)) - } - if propsFetched != tt.wantPropsFetched { - t.Fatalf("props fetched = %v, want %v", propsFetched, tt.wantPropsFetched) - } + require.NoError(t, err) + require.Len(t, resp.Data, 1) + assert.Equal(t, tt.wantPropsFetched, countPath(capture, "/props") == 1) model := resp.Data[0] wantID := tt.wantModelID if wantID == "" { wantID = "Meta-Llama-3.1-8B-Instruct" } - if model.ID != wantID { - t.Fatalf("model.ID = %q, want %q", model.ID, wantID) - } - if model.Metadata == nil { - t.Fatalf("model.Metadata = nil, want context window %d", tt.wantContextWindow) - } - if model.Metadata.ContextWindow == nil || *model.Metadata.ContextWindow != tt.wantContextWindow { - t.Fatalf("context window = %v, want %d", model.Metadata.ContextWindow, tt.wantContextWindow) - } - if len(model.Metadata.Capabilities) != len(tt.wantCapabilities) { - t.Fatalf("capabilities = %v, want %v", model.Metadata.Capabilities, tt.wantCapabilities) - } + assert.Equal(t, wantID, model.ID) + require.NotNil(t, model.Metadata) + require.NotNil(t, model.Metadata.ContextWindow) + assert.Equal(t, tt.wantContextWindow, *model.Metadata.ContextWindow) + assert.Len(t, model.Metadata.Capabilities, len(tt.wantCapabilities)) for name, want := range tt.wantCapabilities { - if model.Metadata.Capabilities[name] != want { - t.Fatalf("capability %q = %v, want %v", name, model.Metadata.Capabilities[name], want) - } + assert.Equal(t, want, model.Metadata.Capabilities[name], "capability %q", name) } // Modes must stay empty so the registry's ID heuristic can still // classify local embedding and reranking GGUFs. - if len(model.Metadata.Modes) != 0 { - t.Fatalf("modes = %v, want none", model.Metadata.Modes) - } + assert.Empty(t, model.Metadata.Modes) }) } } func TestListModels_RouterModeKeepsPerModelContextAndSkipsProps(t *testing.T) { - propsFetched := false - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - switch r.URL.Path { - case "/v1/models": - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[ - {"id":"gemma-3-4b","object":"model","meta":{"n_ctx":8192,"n_ctx_train":131072}}, - {"id":"qwen3-8b","object":"model","meta":{"n_ctx":32768,"n_ctx_train":262144}} - ] - }`)) - case "/props": - propsFetched = true - _, _ = w.Write([]byte(`{"default_generation_settings":{"n_ctx":512}}`)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - } - })) - defer server.Close() + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/v1/models": jsonRoute(http.StatusOK, `{ + "object":"list", + "data":[ + {"id":"gemma-3-4b","object":"model","meta":{"n_ctx":8192,"n_ctx_train":131072}}, + {"id":"qwen3-8b","object":"model","meta":{"n_ctx":32768,"n_ctx_train":262144}} + ] + }`), + "/props": jsonRoute(http.StatusOK, `{"default_generation_settings":{"n_ctx":512}}`), + }) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if propsFetched { - t.Fatal("props was fetched for a multi-model listing; it describes a single loaded model") - } + require.NoError(t, err) + assert.Zero(t, countPath(capture, "/props"), "router mode must not probe /props") want := map[string]int{"gemma-3-4b": 8192, "qwen3-8b": 32768} for _, model := range resp.Data { - if model.Metadata == nil || model.Metadata.ContextWindow == nil { - t.Fatalf("model %q lost its context window", model.ID) - } - if got := *model.Metadata.ContextWindow; got != want[model.ID] { - t.Fatalf("model %q context window = %d, want %d", model.ID, got, want[model.ID]) - } + require.NotNil(t, model.Metadata) + require.NotNil(t, model.Metadata.ContextWindow, "model %q lost its context window", model.ID) + assert.Equal(t, want[model.ID], *model.Metadata.ContextWindow, "model %q context window", model.ID) } } func TestListModels_LeavesMetadataUnsetWhenServerReportsNothing(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - switch r.URL.Path { - case "/v1/models": - // LM Studio and other plain OpenAI-compatible servers omit "meta". - _, _ = w.Write([]byte(`{"data":[{"id":"local-model"}]}`)) - case "/props": - _, _ = w.Write([]byte(`{"default_generation_settings":{"n_ctx":0}}`)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - } - })) - defer server.Close() + server, _ := providertest.RouteServer(t, map[string]http.HandlerFunc{ + // LM Studio and other plain OpenAI-compatible servers omit "meta". + "/v1/models": jsonRoute(http.StatusOK, `{"data":[{"id":"local-model"}]}`), + "/props": jsonRoute(http.StatusOK, `{"default_generation_settings":{"n_ctx":0}}`), + }) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 { - t.Fatalf("len(resp.Data) = %d, want 1", len(resp.Data)) - } - if resp.Object != "list" { - t.Fatalf("resp.Object = %q, want list", resp.Object) - } - if resp.Data[0].Object != "model" { - t.Fatalf("model.Object = %q, want model", resp.Data[0].Object) - } - if resp.Data[0].Metadata != nil { - t.Fatalf("model.Metadata = %+v, want nil so lower metadata layers still apply", resp.Data[0].Metadata) - } + require.NoError(t, err) + require.Len(t, resp.Data, 1) + assert.Equal(t, "list", resp.Object) + assert.Equal(t, "model", resp.Data[0].Object) + assert.Nil(t, resp.Data[0].Metadata) } // TestListModels_FailingPropsLeavesNativeRoutesUsable pins the isolation of the @@ -257,24 +215,11 @@ func TestListModels_LeavesMetadataUnsetWhenServerReportsNothing(t *testing.T) { // rootClient here cost four attempts per listing and locked /health out // entirely after six discovery cycles. func TestListModels_FailingPropsLeavesNativeRoutesUsable(t *testing.T) { - var propsAttempts, healthUpstream int - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - switch r.URL.Path { - case "/v1/models": - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"m","object":"model","meta":{"n_ctx_train":8192}}]}`)) - case "/props": - propsAttempts++ - w.WriteHeader(http.StatusServiceUnavailable) // retryable status - _, _ = w.Write([]byte(`{"error":"unavailable"}`)) - case "/health": - healthUpstream++ - _, _ = w.Write([]byte(`{"status":"ok"}`)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - } - })) - defer server.Close() + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/v1/models": jsonRoute(http.StatusOK, `{"object":"list","data":[{"id":"m","object":"model","meta":{"n_ctx_train":8192}}]}`), + "/props": jsonRoute(http.StatusServiceUnavailable, `{"error":"unavailable"}`), // retryable status + "/health": jsonRoute(http.StatusOK, `{"status":"ok"}`), + }) retry := config.DefaultRetryConfig() retry.InitialBackoff = time.Millisecond // the attempt count is what matters @@ -288,26 +233,18 @@ func TestListModels_FailingPropsLeavesNativeRoutesUsable(t *testing.T) { const listings = 6 for i := range listings { resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() #%d error = %v", i, err) - } + require.NoError(t, err) + // The listing still succeeds on meta.n_ctx_train despite /props failing. - if resp.Data[0].Metadata == nil || *resp.Data[0].Metadata.ContextWindow != 8192 { - t.Fatalf("listing #%d lost its fallback context window", i) - } + require.NotNil(t, resp.Data[0].Metadata) + assert.Equal(t, 8192, *resp.Data[0].Metadata.ContextWindow, "listing #%d lost its fallback context window", i) } - if propsAttempts != listings { - t.Fatalf("props attempts = %d, want %d (one per listing, no retries)", propsAttempts, listings) - } - - if _, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + assert.Equal(t, listings, countPath(capture, "/props")) + _, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodGet, Endpoint: "health", Headers: http.Header{}, - }); err != nil { - t.Fatalf("native /health rejected after failing /props calls: %v", err) - } - if healthUpstream != 1 { - t.Fatalf("health upstream hits = %d, want 1", healthUpstream) - } + }) + require.NoError(t, err) + assert.Equal(t, 1, countPath(capture, "/health")) } diff --git a/internal/providers/llmd/llmd_test.go b/internal/providers/llmd/llmd_test.go index eb0f8f901..e1621b7c6 100644 --- a/internal/providers/llmd/llmd_test.go +++ b/internal/providers/llmd/llmd_test.go @@ -2,7 +2,6 @@ package llmd import ( "context" - "errors" "io" "net/http" "net/http/httptest" @@ -12,18 +11,15 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestRegistrationRequiresEndpointAndAllowsKeyless(t *testing.T) { - if Registration.Type != "llmd" { - t.Fatalf("Registration.Type = %q, want llmd", Registration.Type) - } - if !Registration.Discovery.RequireBaseURL { - t.Fatal("llmd base URL must be required") - } - if !Registration.Discovery.AllowAPIKeyless { - t.Fatal("llmd must allow a keyless Gateway") - } + require.Equal(t, "llmd", Registration.Type) + require.True(t, Registration.Discovery.RequireBaseURL) + require.True(t, Registration.Discovery.AllowAPIKeyless) } func TestChatCompletionInjectsTrustedLLMDHeaders(t *testing.T) { @@ -69,48 +65,35 @@ func TestChatCompletionInjectsTrustedLLMDHeaders(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var got http.Header - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - got = r.Header.Clone() - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-llmd", - "created":1677652288, - "model":"Qwen/Qwen2.5-0.5B-Instruct", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-llmd", + "created":1677652288, + "model":"Qwen/Qwen2.5-0.5B-Instruct", + "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}] + }`) provider := NewWithHTTPClient(tt.apiKey, server.URL, tt.controls, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(tt.ctx, &core.ChatRequest{ Model: "Qwen/Qwen2.5-0.5B-Instruct", Messages: []core.Message{{Role: "user", Content: "hello"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.ID != "chatcmpl-llmd" || resp.Model != "Qwen/Qwen2.5-0.5B-Instruct" { - t.Errorf("response identity = (%q, %q), want normalized ID and model", resp.ID, resp.Model) - } - if len(resp.Choices) != 1 || resp.Choices[0].Message.Role != "assistant" || resp.Choices[0].Message.Content != "ok" { - t.Errorf("response choices = %#v, want one assistant choice containing ok", resp.Choices) - } + require.NoError(t, err) + assert.Equal(t, "chatcmpl-llmd", resp.ID) + assert.Equal(t, "Qwen/Qwen2.5-0.5B-Instruct", resp.Model) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "assistant", resp.Choices[0].Message.Role) + assert.Equal(t, "ok", resp.Choices[0].Message.Content) + + got := capture.Last(t).Header for key, want := range tt.wantHeaders { - assertHeader(t, got, key, want) + assert.Equal(t, want, got.Get(key), "header %s", key) } }) } } func TestPassthroughReplacesClientSuppliedControlHeaders(t *testing.T) { - var got []http.Header - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - got = append(got, r.Header.Clone()) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"tokens":[1,2,3]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"tokens":[1,2,3]}`) provider := NewWithHTTPClient("router-token", server.URL+"/v1", ControlConfig{ InferenceObjective: "trusted-objective", @@ -136,21 +119,21 @@ func TestPassthroughReplacesClientSuppliedControlHeaders(t *testing.T) { "X-Gateway-Destination-Endpoint-Served": {"10.0.0.1:8000"}, }, }) - if err != nil { - t.Fatalf("Passthrough(%q) error = %v", endpoint, err) - } + require.NoError(t, err) + _ = resp.Body.Close() } - if len(got) != 2 { - t.Fatalf("upstream requests = %d, want root and /v1 requests", len(got)) - } - for i, headers := range got { - assertHeader(t, headers, "Authorization", "Bearer router-token") - assertHeader(t, headers, canonicalObjectiveHeader, "trusted-objective") - assertHeader(t, headers, legacyObjectiveHeader, "trusted-objective") - assertHeader(t, headers, canonicalFairnessHeader, "/trusted/tenant") - assertHeader(t, headers, legacyFairnessHeader, "/trusted/tenant") + got := capture.All() + require.Len(t, got, 2) + + for i, req := range got { + headers := req.Header + assert.Equal(t, "Bearer router-token", headers.Get("Authorization"), "request %d", i) + assert.Equal(t, "trusted-objective", headers.Get(canonicalObjectiveHeader), "request %d", i) + assert.Equal(t, "trusted-objective", headers.Get(legacyObjectiveHeader), "request %d", i) + assert.Equal(t, "/trusted/tenant", headers.Get(canonicalFairnessHeader), "request %d", i) + assert.Equal(t, "/trusted/tenant", headers.Get(legacyFairnessHeader), "request %d", i) for _, key := range []string{ "X-Api-Key", "Api-Key", @@ -159,21 +142,18 @@ func TestPassthroughReplacesClientSuppliedControlHeaders(t *testing.T) { "X-Gateway-Model-Name-Rewrite", "X-Gateway-Destination-Endpoint-Served", } { - if value := headers.Get(key); value != "" { - t.Errorf("request %d: %s = %q, want stripped", i, key, value) - } + assert.Empty(t, headers.Get(key), "request %d: %s should be stripped", i, key) } } } func TestPassthroughPreservesDroppedReasonOnRawErrorResponses(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + server, _ := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") w.Header().Set(droppedReasonHeader, "rejected-saturated") w.WriteHeader(http.StatusTooManyRequests) _, _ = w.Write([]byte(`{"error":{"message":"request dropped"}}`)) - })) - defer server.Close() + }) provider := NewWithHTTPClient("", server.URL+"/v1", ControlConfig{}, server.Client(), llmclient.Hooks{}) for _, endpoint := range []string{"tokenize", "chat/completions"} { @@ -182,25 +162,16 @@ func TestPassthroughPreservesDroppedReasonOnRawErrorResponses(t *testing.T) { Endpoint: endpoint, Body: io.NopCloser(strings.NewReader(`{}`)), }) - if err != nil { - t.Fatalf("Passthrough(%q) error = %v", endpoint, err) - } + require.NoError(t, err) + _ = resp.Body.Close() - if resp.StatusCode != http.StatusTooManyRequests { - t.Errorf("Passthrough(%q) status = %d, want %d", endpoint, resp.StatusCode, http.StatusTooManyRequests) - } - assertHeader(t, http.Header(resp.Headers), droppedReasonHeader, "rejected-saturated") + assert.Equal(t, http.StatusTooManyRequests, resp.StatusCode, "Passthrough(%q)", endpoint) + assert.Equal(t, "rejected-saturated", http.Header(resp.Headers).Get(droppedReasonHeader), "Passthrough(%q)", endpoint) } } func TestPassthroughSelectsV1AndRouterRootPaths(t *testing.T) { - var gotPaths []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPaths = append(gotPaths, r.URL.Path) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{}`) provider := NewWithHTTPClient("", server.URL+"/v1", ControlConfig{}, server.Client(), llmclient.Hooks{}) for _, endpoint := range []string{"completions", "messages", "inference/v1/generate", "tokenize"} { @@ -209,35 +180,27 @@ func TestPassthroughSelectsV1AndRouterRootPaths(t *testing.T) { Endpoint: endpoint, Body: io.NopCloser(strings.NewReader(`{}`)), }) - if err != nil { - t.Fatalf("Passthrough(%q) error = %v", endpoint, err) - } + require.NoError(t, err) + _ = resp.Body.Close() } - want := []string{"/v1/completions", "/v1/messages", "/inference/v1/generate", "/tokenize"} - if len(gotPaths) != len(want) { - t.Fatalf("paths = %v, want %v", gotPaths, want) - } - for i := range want { - if gotPaths[i] != want[i] { - t.Errorf("path[%d] = %q, want %q", i, gotPaths[i], want[i]) - } + var gotPaths []string + for _, req := range capture.All() { + gotPaths = append(gotPaths, req.Path) } + assert.Equal(t, []string{"/v1/completions", "/v1/messages", "/inference/v1/generate", "/tokenize"}, gotPaths) } func TestCompatibleAndRootClientsShareKeyRotation(t *testing.T) { - var gotAuth []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotAuth = append(gotAuth, r.Header.Get("Authorization")) + server, capture := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") if r.URL.Path == "/v1/chat/completions" { _, _ = w.Write([]byte(`{"id":"chatcmpl-1","model":"test","choices":[]}`)) return } _, _ = w.Write([]byte(`{}`)) - })) - defer server.Close() + }) provider := newProvider("", server.URL+"/v1", ControlConfig{}, providers.ProviderOptions{ Keys: providers.NewKeyring("router-a", "router-b"), @@ -247,26 +210,20 @@ func TestCompatibleAndRootClientsShareKeyRotation(t *testing.T) { Endpoint: "tokenize", Body: io.NopCloser(strings.NewReader(`{}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) + _ = passthrough.Body.Close() - if _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + _, err = provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "test", Messages: []core.Message{{Role: "user", Content: "hello"}}, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + }) + require.NoError(t, err) - want := []string{"Bearer router-a", "Bearer router-b"} - if len(gotAuth) != len(want) { - t.Fatalf("Authorization headers = %v, want %v", gotAuth, want) - } - for i := range want { - if gotAuth[i] != want[i] { - t.Errorf("Authorization[%d] = %q, want %q", i, gotAuth[i], want[i]) - } + var gotAuth []string + for _, req := range capture.All() { + gotAuth = append(gotAuth, req.Header.Get("Authorization")) } + assert.Equal(t, []string{"Bearer router-a", "Bearer router-b"}, gotAuth) } func TestFairnessHeaderCanBeDisabled(t *testing.T) { @@ -275,13 +232,8 @@ func TestFairnessHeaderCanBeDisabled(t *testing.T) { provider := &Provider{controls: ControlConfig{FairnessFromUserPath: false}} provider.setHeaders(request, "") - - if value := request.Header.Get(canonicalFairnessHeader); value != "" { - t.Fatalf("canonical fairness header = %q, want empty", value) - } - if value := request.Header.Get(legacyFairnessHeader); value != "" { - t.Fatalf("legacy fairness header = %q, want empty", value) - } + assert.Empty(t, request.Header.Get(canonicalFairnessHeader)) + assert.Empty(t, request.Header.Get(legacyFairnessHeader)) } func TestFairnessHeaderIgnoresClientAssertedSnapshotPath(t *testing.T) { @@ -294,13 +246,8 @@ func TestFairnessHeaderIgnoresClientAssertedSnapshotPath(t *testing.T) { provider := &Provider{controls: ControlConfig{FairnessFromUserPath: true}} provider.setHeaders(request, "") - - if value := request.Header.Get(canonicalFairnessHeader); value != "" { - t.Fatalf("canonical fairness header = %q, want empty", value) - } - if value := request.Header.Get(legacyFairnessHeader); value != "" { - t.Fatalf("legacy fairness header = %q, want empty", value) - } + assert.Empty(t, request.Header.Get(canonicalFairnessHeader)) + assert.Empty(t, request.Header.Get(legacyFairnessHeader)) } func TestExposeDroppedReasonReturnsOnlySafeLLMDHeader(t *testing.T) { @@ -311,36 +258,19 @@ func TestExposeDroppedReasonReturnsOnlySafeLLMDHeader(t *testing.T) { } wrapped := exposeDroppedReason(upstream) - if !errors.Is(wrapped, upstream) { - t.Fatal("wrapped error must preserve the upstream error chain") - } + require.ErrorIs(t, wrapped, upstream) + headerErr, ok := wrapped.(interface{ ResponseHeaders() http.Header }) - if !ok { - t.Fatalf("wrapped error type %T does not expose response headers", wrapped) - } + require.True(t, ok, "wrapped error type %T does not expose response headers", wrapped) + headers := headerErr.ResponseHeaders() - assertHeader(t, headers, droppedReasonHeader, "rejected-saturated") - if value := headers.Get("Set-Cookie"); value != "" { - t.Fatalf("Set-Cookie = %q, want filtered", value) - } + assert.Equal(t, "rejected-saturated", headers.Get(droppedReasonHeader)) + assert.Empty(t, headers.Get("Set-Cookie")) } func TestProviderDoesNotAdvertiseUnsupportedNativeSurfaces(t *testing.T) { provider := NewWithHTTPClient("", "http://llmd.invalid/v1", ControlConfig{}, nil, llmclient.Hooks{}) - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("llmd provider must not advertise the separate llm-d Batch Gateway") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("llmd provider must not advertise native files") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("llmd provider must not advertise Responses lifecycle operations") - } -} - -func assertHeader(t *testing.T, headers http.Header, key, want string) { - t.Helper() - if got := headers.Get(key); got != want { - t.Errorf("%s = %q, want %q", key, got, want) - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + require.False(t, ok) } diff --git a/internal/providers/meta/meta_test.go b/internal/providers/meta/meta_test.go index 3d171b3bd..0f701e199 100644 --- a/internal/providers/meta/meta_test.go +++ b/internal/providers/meta/meta_test.go @@ -1,101 +1,26 @@ package meta import ( - "context" - "encoding/json" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - _ = json.NewDecoder(r.Body).Decode(&gotBody) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-meta", - "created":1677652288, - "model":"muse-spark-1.1", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("meta-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "muse-spark-1.1", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +// Meta is a thin wrapper over the shared chat-centric adapter and forwards +// embeddings upstream, so the shared contract covers its surface. It must +// not advertise native batch, file, or audio support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "meta", + DefaultBaseURL: "https://api.meta.ai/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, + Embeddings: true, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "muse-spark-1.1" { - t.Fatalf("resp.Model = %q, want muse-spark-1.1", resp.Model) - } - if resp.Usage.TotalTokens != 4 { - t.Fatalf("resp.Usage = %+v, want total_tokens=4", resp.Usage) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer meta-key" { - t.Fatalf("authorization = %q, want Bearer meta-key", gotAuth) - } - if gotBody["model"] != "muse-spark-1.1" { - t.Fatalf("request model = %v, want muse-spark-1.1", gotBody["model"]) - } -} - -func TestListModels_UsesModelsEndpoint(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[{"id":"muse-spark-1.1","object":"model","created":1677652288,"owned_by":"meta"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("meta-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "muse-spark-1.1" { - t.Fatalf("resp.Data = %+v, want one model muse-spark-1.1", resp.Data) - } - if gotPath != "/models" { - t.Fatalf("path = %q, want /models", gotPath) - } -} - -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("meta-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("meta provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("meta provider should not implement native file provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("meta provider should not implement audio provider") - } + providertest.AssertNoNativeSurfaces(t, NewWithHTTPClient("meta-key", "", nil, llmclient.Hooks{})) } diff --git a/internal/providers/minimax/audio_test.go b/internal/providers/minimax/audio_test.go index 01858e024..c84cd7561 100644 --- a/internal/providers/minimax/audio_test.go +++ b/internal/providers/minimax/audio_test.go @@ -1,50 +1,22 @@ package minimax import ( - "bytes" "context" - "errors" - "io" "net/http" - "net/http/httptest" "strconv" - "strings" "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -func TestProvider_ImplementsAudioProvider(t *testing.T) { - provider := NewWithHTTPClient("key", "", nil, llmclient.Hooks{}) - if _, ok := any(provider).(core.AudioProvider); !ok { - t.Fatal("minimax provider should implement core.AudioProvider") - } -} - func TestCreateSpeech_UsesNativeEndpointAndDecodesHex(t *testing.T) { - var gotPath string - var gotAuth string - var gotRequest speechRequest - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - body, err := io.ReadAll(r.Body) - if err != nil { - http.Error(w, "read error", http.StatusInternalServerError) - return - } - if err := json.Unmarshal(body, &gotRequest); err != nil { - http.Error(w, "invalid json", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"data":{"audio":"000102ff","status":2},"base_resp":{"status_code":0,"status_msg":"success"}}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":{"audio":"000102ff","status":2},"base_resp":{"status_code":0,"status_msg":"success"}}`) provider := NewWithHTTPClient("minimax-key", server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ @@ -54,48 +26,27 @@ func TestCreateSpeech_UsesNativeEndpointAndDecodesHex(t *testing.T) { ResponseFormat: "wav", Speed: 1.5, }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } + require.NoError(t, err) - if gotPath != "/v1/t2a_v2" { - t.Fatalf("path = %q, want /v1/t2a_v2", gotPath) - } - if gotAuth != "Bearer minimax-key" { - t.Fatalf("authorization = %q, want Bearer minimax-key", gotAuth) - } - if gotRequest.Model != "speech-2.8-hd" || gotRequest.Text != "hello" { - t.Fatalf("request model/text = %q/%q", gotRequest.Model, gotRequest.Text) - } - if gotRequest.Stream { - t.Fatal("stream = true, want false") - } - if gotRequest.OutputFormat != "hex" { - t.Fatalf("output_format = %q, want hex", gotRequest.OutputFormat) - } - if gotRequest.VoiceSetting.VoiceID != "English_expressive_narrator" || gotRequest.VoiceSetting.Speed != 1.5 { - t.Fatalf("voice_setting = %+v", gotRequest.VoiceSetting) - } - if gotRequest.AudioSetting.Format != "wav" { - t.Fatalf("audio_setting.format = %q, want wav", gotRequest.AudioSetting.Format) - } - if resp.ContentType != "audio/wav" { - t.Fatalf("content type = %q, want audio/wav", resp.ContentType) - } - if !bytes.Equal(resp.Data, []byte{0x00, 0x01, 0x02, 0xff}) { - t.Fatalf("audio data = %v", resp.Data) - } + req := capture.Last(t) + assert.Equal(t, "/v1/t2a_v2", req.Path) + assert.Equal(t, "Bearer minimax-key", req.Header.Get("Authorization")) + + var gotRequest speechRequest + require.NoError(t, json.Unmarshal(req.Body, &gotRequest)) + assert.Equal(t, "speech-2.8-hd", gotRequest.Model) + assert.Equal(t, "hello", gotRequest.Text) + assert.False(t, gotRequest.Stream) + assert.Equal(t, "hex", gotRequest.OutputFormat) + assert.Equal(t, "English_expressive_narrator", gotRequest.VoiceSetting.VoiceID) + assert.Equal(t, 1.5, gotRequest.VoiceSetting.Speed) + assert.Equal(t, "wav", gotRequest.AudioSetting.Format) + assert.Equal(t, "audio/wav", resp.ContentType) + assert.Equal(t, []byte{0x00, 0x01, 0x02, 0xff}, resp.Data) } func TestCreateSpeech_DefaultsToMP3AndNormalSpeed(t *testing.T) { - var gotRequest speechRequest - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - body, _ := io.ReadAll(r.Body) - _ = json.Unmarshal(body, &gotRequest) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"data":{"audio":"ff","status":2},"base_resp":{"status_code":0}}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":{"audio":"ff","status":2},"base_resp":{"status_code":0}}`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ @@ -103,15 +54,13 @@ func TestCreateSpeech_DefaultsToMP3AndNormalSpeed(t *testing.T) { Input: "hello", Voice: "voice-id", }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if gotRequest.AudioSetting.Format != "mp3" || gotRequest.VoiceSetting.Speed != 1 { - t.Fatalf("defaults = format %q, speed %v", gotRequest.AudioSetting.Format, gotRequest.VoiceSetting.Speed) - } - if resp.ContentType != "audio/mpeg" { - t.Fatalf("content type = %q, want audio/mpeg", resp.ContentType) - } + require.NoError(t, err) + + var gotRequest speechRequest + require.NoError(t, json.Unmarshal(capture.Last(t).Body, &gotRequest)) + assert.Equal(t, "mp3", gotRequest.AudioSetting.Format) + assert.Equal(t, float64(1), gotRequest.VoiceSetting.Speed) + assert.Equal(t, "audio/mpeg", resp.ContentType) } func TestCreateSpeech_ValidatesNativeConstraints(t *testing.T) { @@ -133,9 +82,8 @@ func TestCreateSpeech_ValidatesNativeConstraints(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, err := provider.CreateSpeech(context.Background(), tt.req) - if err == nil || !strings.Contains(err.Error(), tt.want) { - t.Fatalf("CreateSpeech() error = %v, want substring %q", err, tt.want) - } + require.Error(t, err) + assert.Contains(t, err.Error(), tt.want) }) } } @@ -164,51 +112,32 @@ func TestCreateSpeech_MapsNativeStatusCodes(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - body, _ := json.Marshal(map[string]any{ - "data": nil, - "base_resp": map[string]any{"status_code": tt.nativeStatus, "status_msg": tt.statusMsg}, - }) - _, _ = w.Write(body) - })) - defer server.Close() + body, err := json.Marshal(map[string]any{ + "data": nil, + "base_resp": map[string]any{"status_code": tt.nativeStatus, "status_msg": tt.statusMsg}, + }) + require.NoError(t, err) + server, _ := providertest.JSONServer(t, http.StatusOK, string(body)) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) - _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ + _, err = provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "speech-2.8-hd", Input: "hello", Voice: "voice-id", }) var gatewayErr *core.GatewayError - if !errors.As(err, &gatewayErr) { - t.Fatalf("CreateSpeech() error = %v, want *core.GatewayError", err) - } - if gatewayErr.StatusCode != tt.wantHTTPStatus { - t.Fatalf("status = %d, want %d", gatewayErr.StatusCode, tt.wantHTTPStatus) - } - if gatewayErr.Type != tt.wantType { - t.Fatalf("type = %q, want %q", gatewayErr.Type, tt.wantType) - } - if !strings.Contains(gatewayErr.Message, tt.statusMsg) { - t.Fatalf("message = %q, want substring %q", gatewayErr.Message, tt.statusMsg) - } - if !strings.Contains(gatewayErr.Message, strconv.Itoa(tt.nativeStatus)) { - t.Fatalf("message = %q, want native status %d", gatewayErr.Message, tt.nativeStatus) - } - if gatewayErr.Provider != "minimax" { - t.Fatalf("provider = %q, want minimax", gatewayErr.Provider) - } + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, tt.wantHTTPStatus, gatewayErr.StatusCode) + assert.Equal(t, tt.wantType, gatewayErr.Type) + assert.Contains(t, gatewayErr.Message, tt.statusMsg) + assert.Contains(t, gatewayErr.Message, strconv.Itoa(tt.nativeStatus)) + assert.Equal(t, "minimax", gatewayErr.Provider) }) } } func TestCreateSpeech_RejectsMalformedAudio(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"data":{"audio":"not-hex","status":2},"base_resp":{"status_code":0}}`)) - })) - defer server.Close() + server, _ := providertest.JSONServer(t, http.StatusOK, `{"data":{"audio":"not-hex","status":2},"base_resp":{"status_code":0}}`) provider := NewWithHTTPClient("key", server.URL, server.Client(), llmclient.Hooks{}) _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ @@ -216,15 +145,13 @@ func TestCreateSpeech_RejectsMalformedAudio(t *testing.T) { Input: "hello", Voice: "voice-id", }) - if err == nil || !strings.Contains(err.Error(), "not valid hexadecimal") { - t.Fatalf("CreateSpeech() error = %v, want malformed audio error", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "not valid hexadecimal") } func TestCreateTranscription_IsUnsupported(t *testing.T) { provider := NewWithHTTPClient("key", "", nil, llmclient.Hooks{}) _, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{}) - if err == nil || !strings.Contains(err.Error(), "does not support speech-to-text") { - t.Fatalf("CreateTranscription() error = %v", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "does not support speech-to-text") } diff --git a/internal/providers/minimax/minimax_test.go b/internal/providers/minimax/minimax_test.go index d98aa79d5..cf4242a94 100644 --- a/internal/providers/minimax/minimax_test.go +++ b/internal/providers/minimax/minimax_test.go @@ -2,75 +2,30 @@ package minimax import ( "context" - "io" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-minimax", - "created":1677652288, - "model":"MiniMax-M3", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("minimax-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "MiniMax-M3", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "minimax", + DefaultBaseURL: "https://api.minimax.io/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, + Embeddings: true, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "MiniMax-M3" { - t.Fatalf("resp.Model = %q, want MiniMax-M3", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer minimax-key" { - t.Fatalf("authorization = %q, want Bearer minimax-key", gotAuth) - } } func TestChatCompletion_ClampsZeroTemperature(t *testing.T) { - var gotBody []byte - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var err error - gotBody, err = io.ReadAll(r.Body) - if err != nil { - http.Error(w, "read error", http.StatusInternalServerError) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-minimax", - "created":1677652288, - "model":"MiniMax-M3", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) provider := NewWithHTTPClient("minimax-key", server.URL, server.Client(), llmclient.Hooks{}) temp := 0.0 @@ -79,77 +34,44 @@ func TestChatCompletion_ClampsZeroTemperature(t *testing.T) { Messages: []core.Message{{Role: "user", Content: "hi"}}, Temperature: &temp, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - - // Verify the body sent to the server does not contain temperature=0 - bodyStr := string(gotBody) - if strings.Contains(bodyStr, `"temperature":0`) { - t.Fatalf("request body should not contain temperature=0, got: %s", bodyStr) - } - if !strings.Contains(bodyStr, `"temperature":1`) { - t.Fatalf("request body should contain temperature=1, got: %s", bodyStr) - } -} - -func TestProvider_DefaultBaseURL(t *testing.T) { - provider := NewWithHTTPClient("key", "", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("expected non-nil provider") - } + require.NoError(t, err) + assert.Equal(t, float64(1), capture.Last(t).JSON(t)["temperature"]) } -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { +// MiniMax serves speech natively (see audio.go), so only the batch and file +// surfaces must stay hidden. +func TestProvider_DoesNotExposeBatchOrFileInterfaces(t *testing.T) { provider := NewWithHTTPClient("minimax-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("minimax provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("minimax provider should not implement native file provider") - } + _, ok := any(provider).(core.NativeBatchProvider) + require.False(t, ok) + _, ok = any(provider).(core.NativeFileProvider) + require.False(t, ok) } func TestClampTemperature_NilRequest(t *testing.T) { - result := clampTemperature(nil) - if result != nil { - t.Fatal("expected nil for nil input") - } + require.Nil(t, clampTemperature(nil)) } func TestClampTemperature_NilTemperature(t *testing.T) { req := &core.ChatRequest{Model: "MiniMax-M3"} - result := clampTemperature(req) - if result.Temperature != nil { - t.Fatal("expected nil temperature to remain nil") - } + require.Nil(t, clampTemperature(req).Temperature) } func TestClampTemperature_ZeroTemperature(t *testing.T) { temp := 0.0 req := &core.ChatRequest{Model: "MiniMax-M3", Temperature: &temp} result := clampTemperature(req) - if result.Temperature == nil { - t.Fatal("expected non-nil temperature after clamping") - } - if *result.Temperature != defaultTemperature { - t.Fatalf("expected temperature=%v after clamping zero, got %v", defaultTemperature, *result.Temperature) - } + require.NotNil(t, result.Temperature) + assert.Equal(t, defaultTemperature, *result.Temperature) + // Original request should not be mutated - if *req.Temperature != 0.0 { - t.Fatal("original request should not be mutated") - } + assert.Equal(t, 0.0, *req.Temperature) } func TestClampTemperature_PositiveTemperature(t *testing.T) { temp := 0.7 req := &core.ChatRequest{Model: "MiniMax-M3", Temperature: &temp} result := clampTemperature(req) - if result != req { - t.Fatal("expected same pointer for valid temperature") - } - if *result.Temperature != 0.7 { - t.Fatalf("expected temperature=0.7, got %v", *result.Temperature) - } + require.Same(t, req, result) + assert.Equal(t, 0.7, *result.Temperature) } diff --git a/internal/providers/ollama/ollama_test.go b/internal/providers/ollama/ollama_test.go index f75e913d7..665539ae5 100644 --- a/internal/providers/ollama/ollama_test.go +++ b/internal/providers/ollama/ollama_test.go @@ -3,51 +3,69 @@ package ollama import ( "context" "encoding/json" - "errors" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestNew(t *testing.T) { - apiKey := "test-api-key" - // Use NewWithHTTPClient to get concrete type for internal testing - provider := NewWithHTTPClient(apiKey, nil, llmclient.Hooks{}) - - if got := provider.keys.Primary(); got != apiKey { - t.Errorf("primary key = %q, want %q", got, apiKey) - } - if provider.compat == nil { - t.Error("compat should not be nil") - } - if provider.nativeClient == nil { - t.Error("nativeClient should not be nil") +const chatCompletionJSON = `{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama3.2", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30 } -} +}` -func TestNew_ReturnsProvider(t *testing.T) { - provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) +const chatChunkSSE = `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} - if provider == nil { - t.Error("provider should not be nil") - } -} +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} + +data: [DONE] +` -func TestNew_WithoutAPIKey(t *testing.T) { - // Ollama doesn't require an API key +// newTestProvider builds a keyless provider pointed at baseURL. +func newTestProvider(baseURL string) *Provider { provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) + provider.SetBaseURL(baseURL) + return provider +} - if got := provider.keys.Primary(); got != "" { - t.Errorf("primary key = %q, want empty", got) +func TestNew(t *testing.T) { + tests := []struct { + name string + apiKey string + }{ + {name: "with api key", apiKey: "test-api-key"}, + // Ollama doesn't require an API key + {name: "without api key"}, } - if provider.compat == nil { - t.Error("compat should not be nil") + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider := NewWithHTTPClient(tt.apiKey, nil, llmclient.Hooks{}) + assert.Equal(t, tt.apiKey, provider.keys.Primary()) + assert.NotNil(t, provider.compat) + assert.NotNil(t, provider.nativeClient) + }) } } @@ -60,50 +78,17 @@ func TestChatCompletion(t *testing.T) { checkResponse func(*testing.T, *core.ChatResponse) }{ { - name: "successful request", - statusCode: http.StatusOK, - responseBody: `{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama3.2", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello! How can I help you today?" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30 - } - }`, - expectedError: false, + name: "successful request", + statusCode: http.StatusOK, + responseBody: chatCompletionJSON, checkResponse: func(t *testing.T, resp *core.ChatResponse) { - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } - if resp.Model != "llama3.2" { - t.Errorf("Model = %q, want %q", resp.Model, "llama3.2") - } - if len(resp.Choices) != 1 { - t.Fatalf("len(Choices) = %d, want 1", len(resp.Choices)) - } - if resp.Choices[0].Message.Content != "Hello! How can I help you today?" { - t.Errorf("Message content = %q, want %q", resp.Choices[0].Message.Content, "Hello! How can I help you today?") - } - if resp.Usage.PromptTokens != 10 { - t.Errorf("PromptTokens = %d, want 10", resp.Usage.PromptTokens) - } - if resp.Usage.CompletionTokens != 20 { - t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens) - } - if resp.Usage.TotalTokens != 30 { - t.Errorf("TotalTokens = %d, want 30", resp.Usage.TotalTokens) - } + assert.Equal(t, "chatcmpl-123", resp.ID) + assert.Equal(t, "llama3.2", resp.Model) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Choices[0].Message.Content) + assert.Equal(t, 10, resp.Usage.PromptTokens) + assert.Equal(t, 20, resp.Usage.CompletionTokens) + assert.Equal(t, 30, resp.Usage.TotalTokens) }, }, { @@ -116,123 +101,50 @@ func TestChatCompletion(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - - // Verify request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } + server, capture := providertest.JSONServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(server.URL) - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: "llama3.2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "llama3.2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - resp, err := provider.ChatCompletion(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, "llama3.2", req.JSON(t)["model"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + require.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } -func TestChatCompletion_WithAPIKey(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify authorization header is set when API key is provided - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - if authHeader != "Bearer test-api-key" { - t.Errorf("Authorization = %q, want %q", authHeader, "Bearer test-api-key") - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama3.2", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "Hi"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "llama3.2", - Messages: []core.Message{{Role: "user", Content: "Hello"}}, - } - - _, err := provider.ChatCompletion(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } -} - -func TestChatCompletion_WithoutAPIKey(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify no authorization header when API key is not provided - authHeader := r.Header.Get("Authorization") - if authHeader != "" { - t.Errorf("Authorization header should be empty, got %q", authHeader) - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama3.2", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "Hi"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "llama3.2", - Messages: []core.Message{{Role: "user", Content: "Hello"}}, +func TestChatCompletion_AuthorizationHeader(t *testing.T) { + tests := []struct { + name string + apiKey string + wantAuth string + }{ + {name: "with api key", apiKey: "test-api-key", wantAuth: "Bearer test-api-key"}, + {name: "without api key"}, } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := NewWithHTTPClient(tt.apiKey, nil, llmclient.Hooks{}) + provider.SetBaseURL(server.URL) - _, err := provider.ChatCompletion(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) + _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: "llama3.2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) + require.NoError(t, err) + assert.Equal(t, tt.wantAuth, capture.Last(t).Header.Get("Authorization")) + }) } } @@ -244,15 +156,9 @@ func TestStreamChatCompletion(t *testing.T) { expectedError bool }{ { - name: "successful streaming request", - statusCode: http.StatusOK, - responseBody: `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} - -data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} - -data: [DONE] -`, - expectedError: false, + name: "successful streaming request", + statusCode: http.StatusOK, + responseBody: chatChunkSSE, }, { name: "server error", @@ -264,64 +170,32 @@ data: [DONE] for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } - + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(tt.statusCode) _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }) + provider := newTestProvider(server.URL) - req := &core.ChatRequest{ - Model: "llama3.2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } + body, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ + Model: "llama3.2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - body, err := provider.StreamChatCompletion(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, true, req.JSON(t)["stream"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if body == nil { - t.Fatal("body should not be nil") - } - defer func() { _ = body.Close() }() - - // Read and verify the streaming response - respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } - if string(respBody) != tt.responseBody { - t.Errorf("response body = %q, want %q", string(respBody), tt.responseBody) - } + require.Error(t, err) + return } + require.NoError(t, err) + require.NotNil(t, body) + defer func() { _ = body.Close() }() + + respBody, err := io.ReadAll(body) + require.NoError(t, err) + assert.Equal(t, tt.responseBody, string(respBody)) }) } } @@ -354,20 +228,11 @@ func TestListModels(t *testing.T) { } ] }`, - expectedError: false, checkResponse: func(t *testing.T, resp *core.ModelsResponse) { - if resp.Object != "list" { - t.Errorf("Object = %q, want %q", resp.Object, "list") - } - if len(resp.Data) != 2 { - t.Fatalf("len(Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].ID != "llama3.2" { - t.Errorf("Data[0].ID = %q, want %q", resp.Data[0].ID, "llama3.2") - } - if resp.Data[0].OwnedBy != "library" { - t.Errorf("Data[0].OwnedBy = %q, want %q", resp.Data[0].OwnedBy, "library") - } + assert.Equal(t, "list", resp.Object) + require.Len(t, resp.Data, 2) + assert.Equal(t, "llama3.2", resp.Data[0].ID) + assert.Equal(t, "library", resp.Data[0].OwnedBy) }, }, { @@ -380,43 +245,30 @@ func TestListModels(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/models": func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(tt.statusCode) + _, _ = w.Write([]byte(tt.responseBody)) + }, // Best-effort capability probe issued per listed model. - if r.URL.Path == "/api/show" && r.Method == http.MethodPost { - w.WriteHeader(http.StatusOK) + "/api/show": func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte(`{"capabilities":["completion"]}`)) - return - } - // Verify request method and path - if r.Method != http.MethodGet { - t.Errorf("Method = %q, want %q", r.Method, http.MethodGet) - } - if r.URL.Path != "/models" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/models") - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }, + }) + provider := newTestProvider(server.URL) resp, err := provider.ListModels(context.Background()) + listing := capture.All()[0] + assert.Equal(t, http.MethodGet, listing.Method) + assert.Equal(t, "/models", listing.Path) + if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + require.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } @@ -426,14 +278,19 @@ func TestListModels(t *testing.T) { // repeat listings don't re-probe, and leave models unstamped when the probe // fails so the ID heuristic can still apply. func TestListModels_StampsShowCapabilities(t *testing.T) { - showCalls := map[string]int{} - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/show" { + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/models": func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(`{"object":"list","data":[ + {"id":"llama3.2","object":"model","owned_by":"library"}, + {"id":"nomic-embed-text","object":"model","owned_by":"library"}, + {"id":"mystery-model","object":"model","owned_by":"library"} + ]}`)) + }, + "/api/show": func(w http.ResponseWriter, r *http.Request) { var req struct { Model string `json:"model"` } _ = json.NewDecoder(r.Body).Decode(&req) - showCalls[req.Model]++ switch req.Model { case "nomic-embed-text": _, _ = w.Write([]byte(`{"capabilities":["embedding"]}`)) @@ -442,313 +299,141 @@ func TestListModels_StampsShowCapabilities(t *testing.T) { default: w.WriteHeader(http.StatusInternalServerError) } - return - } - _, _ = w.Write([]byte(`{"object":"list","data":[ - {"id":"llama3.2","object":"model","owned_by":"library"}, - {"id":"nomic-embed-text","object":"model","owned_by":"library"}, - {"id":"mystery-model","object":"model","owned_by":"library"} - ]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }, + }) + provider := newTestProvider(server.URL) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) + byID := map[string]core.Model{} for _, m := range resp.Data { byID[m.ID] = m } embed := byID["nomic-embed-text"] - if embed.Metadata == nil || len(embed.Metadata.Modes) != 1 || embed.Metadata.Modes[0] != "embedding" { - t.Errorf("nomic-embed-text metadata = %+v, want embedding modes", embed.Metadata) - } - if embed.Metadata == nil || len(embed.Metadata.Categories) != 1 || embed.Metadata.Categories[0] != core.CategoryEmbedding { - t.Errorf("nomic-embed-text categories = %+v, want [embedding]", embed.Metadata) - } + require.NotNil(t, embed.Metadata) + assert.Equal(t, []string{"embedding"}, embed.Metadata.Modes) + assert.Equal(t, []core.ModelCategory{core.CategoryEmbedding}, embed.Metadata.Categories) + chat := byID["llama3.2"] - if chat.Metadata == nil || len(chat.Metadata.Modes) != 1 || chat.Metadata.Modes[0] != "chat" { - t.Errorf("llama3.2 metadata = %+v, want chat modes (tools capability skipped)", chat.Metadata) - } - if byID["mystery-model"].Metadata != nil { - t.Errorf("mystery-model metadata = %+v, want nil after failed probe", byID["mystery-model"].Metadata) - } + require.NotNil(t, chat.Metadata) + assert.Equal(t, []string{"chat"}, chat.Metadata.Modes) + assert.Nil(t, byID["mystery-model"].Metadata) // Second listing: successes served from cache, the failure re-probed. - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("unexpected error on second listing: %v", err) - } - if showCalls["nomic-embed-text"] != 1 || showCalls["llama3.2"] != 1 { - t.Errorf("show calls = %v, want cached results for successful probes", showCalls) - } - if showCalls["mystery-model"] != 2 { - t.Errorf("mystery-model show calls = %d, want re-probe after failure", showCalls["mystery-model"]) + _, err = provider.ListModels(context.Background()) + require.NoError(t, err) + + showCalls := map[string]int{} + for _, req := range capture.All() { + if req.Path == "/api/show" { + showCalls[req.JSON(t)["model"].(string)]++ + } } + assert.Equal(t, map[string]int{"nomic-embed-text": 1, "llama3.2": 1, "mystery-model": 2}, showCalls, "successful probes must be cached") } func TestChatCompletionWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + server, _ := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { // Simulate a slow response <-r.Context().Done() w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) cancel() // Cancel immediately - req := &core.ChatRequest{ - Model: "llama3.2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - _, err := provider.ChatCompletion(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + _, err := provider.ChatCompletion(ctx, &core.ChatRequest{ + Model: "llama3.2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) + require.Error(t, err) } func TestResponses(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request path for chat completions (Ollama converts Responses to chat) - if r.URL.Path != "/chat/completions" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/chat/completions") - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama3.2", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello! How can I help you today?" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30 - } - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ResponsesRequest{ + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "llama3.2", Input: "Hello", - } - - resp, err := provider.Responses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } - if resp.Object != "response" { - t.Errorf("Object = %q, want %q", resp.Object, "response") - } - if resp.Model != "llama3.2" { - t.Errorf("Model = %q, want %q", resp.Model, "llama3.2") - } - if resp.Status != "completed" { - t.Errorf("Status = %q, want %q", resp.Status, "completed") - } - if len(resp.Output) != 1 { - t.Fatalf("len(Output) = %d, want 1", len(resp.Output)) - } - if len(resp.Output[0].Content) != 1 { - t.Fatalf("len(Output[0].Content) = %d, want 1", len(resp.Output[0].Content)) - } - if resp.Output[0].Content[0].Text != "Hello! How can I help you today?" { - t.Errorf("Output text = %q, want %q", resp.Output[0].Content[0].Text, "Hello! How can I help you today?") - } - if resp.Usage == nil { - t.Fatal("Usage should not be nil") - } - if resp.Usage.InputTokens != 10 { - t.Errorf("InputTokens = %d, want 10", resp.Usage.InputTokens) - } - if resp.Usage.OutputTokens != 20 { - t.Errorf("OutputTokens = %d, want 20", resp.Usage.OutputTokens) - } + }) + require.NoError(t, err) + // Ollama converts Responses to chat completions. + assert.Equal(t, "/chat/completions", capture.Last(t).Path) + assert.Equal(t, "chatcmpl-123", resp.ID) + assert.Equal(t, "response", resp.Object) + assert.Equal(t, "llama3.2", resp.Model) + assert.Equal(t, "completed", resp.Status) + require.Len(t, resp.Output, 1) + require.Len(t, resp.Output[0].Content, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Output[0].Content[0].Text) + require.NotNil(t, resp.Usage) + assert.Equal(t, 10, resp.Usage.InputTokens) + assert.Equal(t, 20, resp.Usage.OutputTokens) } func TestResponsesWithArrayInput(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request body is converted to chat format - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - - var req map[string]any - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - - // Verify messages array exists (converted from input) - messages, ok := req["messages"].([]any) - if !ok { - t.Fatal("messages should be an array") - } - // Should have system message + 2 input messages - if len(messages) != 3 { - t.Errorf("len(messages) = %d, want 3", len(messages)) - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "llama3.2", - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": "Hello!" - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15 - } - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, chatCompletionJSON) + provider := newTestProvider(server.URL) - req := &core.ResponsesRequest{ + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "llama3.2", Input: []any{ - map[string]any{ - "role": "user", - "content": "Hello", - }, - map[string]any{ - "role": "assistant", - "content": "Hi there!", - }, + map[string]any{"role": "user", "content": "Hello"}, + map[string]any{"role": "assistant", "content": "Hi there!"}, }, Instructions: "Be helpful", - } - - resp, err := provider.Responses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + }) + require.NoError(t, err) + assert.Equal(t, "chatcmpl-123", resp.ID) - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } + // Input is converted to messages: system message + 2 input messages. + messages, ok := capture.Last(t).JSON(t)["messages"].([]any) + require.True(t, ok) + assert.Len(t, messages, 3) } func TestStreamResponses(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + server, capture := providertest.SSEServer(t, chatChunkSSE) + provider := newTestProvider(server.URL) -data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama3.2","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} - -data: [DONE] -`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ResponsesRequest{ + body, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "llama3.2", Input: "Hello", - } - - body, err := provider.StreamResponses(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if body == nil { - t.Fatal("body should not be nil") - } + }) + require.NoError(t, err) + require.NotNil(t, body) defer func() { _ = body.Close() }() respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } + require.NoError(t, err) + assert.Equal(t, true, capture.Last(t).JSON(t)["stream"]) responseStr := string(respBody) - if !strings.Contains(responseStr, "response.created") { - t.Error("response should contain response.created event") - } - if !strings.Contains(responseStr, "response.output_text.delta") { - t.Error("response should contain response.output_text.delta event") - } - if !strings.Contains(responseStr, "[DONE]") { - t.Error("response should end with [DONE]") - } + assert.Contains(t, responseStr, "response.created") + assert.Contains(t, responseStr, "response.output_text.delta") + assert.Contains(t, responseStr, "[DONE]") } func TestResponsesWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + server, _ := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { // Simulate a slow response <-r.Context().Done() w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) cancel() // Cancel immediately - req := &core.ResponsesRequest{ + _, err := provider.Responses(ctx, &core.ResponsesRequest{ Model: "llama3.2", Input: "Hello", - } - - _, err := provider.Responses(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + }) + require.Error(t, err) } func TestOllamaResponsesStreamConverter(t *testing.T) { @@ -763,192 +448,78 @@ data: [DONE] reader := io.NopCloser(strings.NewReader(mockStream)) converter := providers.NewOpenAIResponsesStreamConverter(reader, "llama3.2", "ollama") - // Read all data from converter data, err := io.ReadAll(converter) - if err != nil { - t.Fatalf("failed to read from converter: %v", err) - } + require.NoError(t, err) result := string(data) - - // Check that the stream contains expected events - if !strings.Contains(result, "response.created") { - t.Error("stream should contain response.created event") - } - if !strings.Contains(result, "response.output_text.delta") { - t.Error("stream should contain response.output_text.delta event") - } - if !strings.Contains(result, "Hello") { - t.Error("stream should contain 'Hello' content") - } - if !strings.Contains(result, " world") { - t.Error("stream should contain ' world' content") - } - if !strings.Contains(result, "response.completed") { - t.Error("stream should contain response.completed event") - } - if !strings.Contains(result, "[DONE]") { - t.Error("stream should contain [DONE] marker") - } -} - -func TestNewWithHTTPClient(t *testing.T) { - customClient := &http.Client{} - apiKey := "test-api-key" - - provider := NewWithHTTPClient(apiKey, customClient, llmclient.Hooks{}) - - if got := provider.keys.Primary(); got != apiKey { - t.Errorf("primary key = %q, want %q", got, apiKey) - } - if provider.compat == nil { - t.Error("compat should not be nil") - } - if provider.nativeClient == nil { - t.Error("nativeClient should not be nil") - } -} - -func TestSetBaseURL(t *testing.T) { - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - customURL := "http://custom.ollama.server:11434/v1" - - provider.SetBaseURL(customURL) - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() - - provider.SetBaseURL(server.URL) - _, err := provider.ListModels(context.Background()) - if err != nil { - t.Errorf("SetBaseURL should allow using custom URL: %v", err) - } + assert.Contains(t, result, "response.created") + assert.Contains(t, result, "response.output_text.delta") + assert.Contains(t, result, "Hello") + assert.Contains(t, result, " world") + assert.Contains(t, result, "response.completed") + assert.Contains(t, result, "[DONE]") } func TestSetBaseURL_TrailingSlash(t *testing.T) { - var nativePath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - nativePath = r.URL.Path - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"model":"nomic-embed-text","embeddings":[[0.1,0.2]]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/v1/") + server, capture := providertest.JSONServer(t, http.StatusOK, `{"model":"nomic-embed-text","embeddings":[[0.1,0.2]]}`) + provider := newTestProvider(server.URL + "/v1/") _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "nomic-embed-text", Input: "hello", }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if nativePath != "/api/embed" { - t.Errorf("native client path = %q, want /api/embed (trailing slash not normalized)", nativePath) - } + require.NoError(t, err) + assert.Equal(t, "/api/embed", capture.Last(t).Path) } func TestEmbeddings(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/api/embed" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/api/embed") - } - if r.Method != http.MethodPost { - t.Errorf("Method = %q, want %q", r.Method, http.MethodPost) - } - - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req ollamaEmbedRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if req.Model != "nomic-embed-text" { - t.Errorf("Model = %q, want %q", req.Model, "nomic-embed-text") - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "model": "nomic-embed-text", - "embeddings": [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], - "prompt_eval_count": 8 - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/v1") + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "model": "nomic-embed-text", + "embeddings": [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + "prompt_eval_count": 8 + }`) + provider := newTestProvider(server.URL + "/v1") resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "nomic-embed-text", Input: []string{"hello", "world"}, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "/api/embed", req.Path) + assert.Equal(t, http.MethodPost, req.Method) + var sent ollamaEmbedRequest + require.NoError(t, json.Unmarshal(req.Body, &sent)) + assert.Equal(t, "nomic-embed-text", sent.Model) + + assert.Equal(t, "list", resp.Object) + assert.Equal(t, "nomic-embed-text", resp.Model) + require.Len(t, resp.Data, 2) + assert.Equal(t, "embedding", resp.Data[0].Object) - if resp.Object != "list" { - t.Errorf("Object = %q, want %q", resp.Object, "list") - } - if resp.Model != "nomic-embed-text" { - t.Errorf("Model = %q, want %q", resp.Model, "nomic-embed-text") - } - if len(resp.Data) != 2 { - t.Fatalf("len(Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].Object != "embedding" { - t.Errorf("Data[0].Object = %q, want %q", resp.Data[0].Object, "embedding") - } var floats []float64 - if err := json.Unmarshal(resp.Data[0].Embedding, &floats); err != nil { - t.Fatalf("failed to unmarshal embedding: %v", err) - } - if len(floats) != 3 { - t.Errorf("len(embedding floats) = %d, want 3", len(floats)) - } - if resp.Data[1].Index != 1 { - t.Errorf("Data[1].Index = %d, want 1", resp.Data[1].Index) - } - if resp.Usage.PromptTokens != 8 { - t.Errorf("PromptTokens = %d, want 8", resp.Usage.PromptTokens) - } - if resp.Usage.TotalTokens != 8 { - t.Errorf("TotalTokens = %d, want 8", resp.Usage.TotalTokens) - } + require.NoError(t, json.Unmarshal(resp.Data[0].Embedding, &floats)) + assert.Len(t, floats, 3) + assert.Equal(t, 1, resp.Data[1].Index) + assert.Equal(t, 8, resp.Usage.PromptTokens) + assert.Equal(t, 8, resp.Usage.TotalTokens) } func TestEmbeddings_ModelFallback(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "model": "", - "embeddings": [[0.1]], - "prompt_eval_count": 1 - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/v1") + server, _ := providertest.JSONServer(t, http.StatusOK, `{ + "model": "", + "embeddings": [[0.1]], + "prompt_eval_count": 1 + }`) + provider := newTestProvider(server.URL + "/v1") resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "nomic-embed-text", Input: "hello", }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if resp.Model != "nomic-embed-text" { - t.Errorf("Model = %q, want %q (should fall back to request model)", resp.Model, "nomic-embed-text") - } + require.NoError(t, err) + assert.Equal(t, "nomic-embed-text", resp.Model) } // TestEmbeddings_NoVectorsErrors guards the common misconfiguration where an @@ -957,36 +528,20 @@ func TestEmbeddings_ModelFallback(t *testing.T) { // error body, which unmarshals into zero embeddings. The adapter must surface // an error instead of returning an empty, OpenAI-shaped list. func TestEmbeddings_NoVectorsErrors(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"error":"Unexpected endpoint or method. (POST /api/embed)"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/v1") + server, _ := providertest.JSONServer(t, http.StatusOK, `{"error":"Unexpected endpoint or method. (POST /api/embed)"}`) + provider := newTestProvider(server.URL + "/v1") resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "text-embedding-nomic-embed-text-v1.5", Input: "hello world", }) - if err == nil { - t.Fatal("expected provider error for empty embeddings, got nil") - } - if resp != nil { - t.Fatalf("expected nil response on error, got %d data entries", len(resp.Data)) - } + require.Error(t, err) + require.Nil(t, resp) var gatewayErr *core.GatewayError - if !errors.As(err, &gatewayErr) { - t.Fatalf("expected *core.GatewayError, got %T", err) - } - if gatewayErr.HTTPStatusCode() != http.StatusBadGateway { - t.Fatalf("status = %d, want %d", gatewayErr.HTTPStatusCode(), http.StatusBadGateway) - } - if !strings.Contains(gatewayErr.Message, `"openai" or "vllm" provider`) { - t.Fatalf("unexpected error message: %q", gatewayErr.Message) - } + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusBadGateway, gatewayErr.HTTPStatusCode()) + assert.Contains(t, gatewayErr.Message, `"openai" or "vllm" provider`) } // TestEmbeddings_EmptyInputNoError ensures an empty input batch (an empty @@ -994,37 +549,27 @@ func TestEmbeddings_NoVectorsErrors(t *testing.T) { // LM-Studio-as-ollama misconfiguration: zero vectors for an empty batch is a // legitimate result, not a provider error. func TestEmbeddings_EmptyInputNoError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"model":"nomic-embed-text","embeddings":[],"prompt_eval_count":0}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL + "/v1") + server, _ := providertest.JSONServer(t, http.StatusOK, `{"model":"nomic-embed-text","embeddings":[],"prompt_eval_count":0}`) + provider := newTestProvider(server.URL + "/v1") for _, empty := range []any{[]any{}, []string{}} { resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "nomic-embed-text", Input: empty, }) - if err != nil { - t.Fatalf("unexpected error for empty batch %#v: %v", empty, err) - } - if resp == nil || len(resp.Data) != 0 { - t.Fatalf("input %#v: expected empty data response, got %+v", empty, resp) - } + require.NoError(t, err) + require.NotNil(t, resp) + assert.Empty(t, resp.Data) } // Scalar/nil inputs are NOT empty batches: zero vectors must stay on the // loud-error path so a misconfigured OpenAI-compatible endpoint returning a // 200 error body for "" / null isn't silently swallowed as an empty list. for _, scalar := range []any{"", nil} { - if _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ + _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "nomic-embed-text", Input: scalar, - }); err == nil { - t.Fatalf("input %#v: expected provider error for zero vectors, got nil", scalar) - } + }) + assert.Error(t, err, "input %#v", scalar) } } diff --git a/internal/providers/opencodego/opencodego_test.go b/internal/providers/opencodego/opencodego_test.go index 2778f9bac..2b6f193f4 100644 --- a/internal/providers/opencodego/opencodego_test.go +++ b/internal/providers/opencodego/opencodego_test.go @@ -4,11 +4,13 @@ import ( "context" "io" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // newTestProvider builds a provider whose OpenAI-compatible and Anthropic @@ -17,81 +19,41 @@ func newTestProvider(serverURL string, client *http.Client) *Provider { return NewWithHTTPClient("sk-opencode", serverURL, client, llmclient.Hooks{}) } -func TestChatCompletion_OpenAIStyleModel_UsesChatCompletions(t *testing.T) { - var gotPath, gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-opencode", - "created":1677652288, - "model":"glm-5.1", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - resp, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "glm-5.1", - Messages: []core.Message{{Role: "user", Content: "hi"}}, +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "opencode_go", + DefaultBaseURL: "https://opencode.ai/zen/go/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) + }, + Embeddings: false, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "glm-5.1" { - t.Fatalf("resp.Model = %q, want glm-5.1", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer sk-opencode" { - t.Fatalf("authorization = %q, want Bearer sk-opencode", gotAuth) - } + providertest.AssertNoNativeSurfaces(t, newTestProvider("", nil)) } func TestChatCompletion_AnthropicStyleModel_UsesMessages(t *testing.T) { - var gotPath, gotAuth, gotAPIKey, gotVersion string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - gotAPIKey = r.Header.Get("x-api-key") - gotVersion = r.Header.Get("anthropic-version") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_opencode", - "model":"qwen3.7-max", - "content":[{"type":"text","text":"hello"}], - "stop_reason":"end_turn", - "usage":{"input_tokens":5,"output_tokens":2} - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"msg_opencode", + "model":"qwen3.7-max", + "content":[{"type":"text","text":"hello"}], + "stop_reason":"end_turn", + "usage":{"input_tokens":5,"output_tokens":2} + }`) resp, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3.7-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotPath != "/messages" { - t.Fatalf("path = %q, want /messages", gotPath) - } - if gotAPIKey != "sk-opencode" { - t.Fatalf("x-api-key = %q, want sk-opencode", gotAPIKey) - } - if gotAuth != "" { - t.Fatalf("authorization = %q, want empty (messages uses x-api-key)", gotAuth) - } - if gotVersion == "" { - t.Fatal("anthropic-version header missing on /messages request") - } - if len(resp.Choices) != 1 || resp.Choices[0].Message.Content != "hello" { - t.Fatalf("unexpected response: %+v", resp.Choices) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "/messages", req.Path) + assert.Equal(t, "sk-opencode", req.Header.Get("x-api-key")) + assert.Empty(t, req.Header.Get("Authorization")) + assert.NotEmpty(t, req.Header.Get("anthropic-version")) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "hello", resp.Choices[0].Message.Content) } func TestStreamChatCompletion_RoutesByModel(t *testing.T) { @@ -106,97 +68,27 @@ func TestStreamChatCompletion_RoutesByModel(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "data: [DONE]\n\n") stream, err := newTestProvider(server.URL, server.Client()).StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: tt.model, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) + _, _ = io.Copy(io.Discard, stream) _ = stream.Close() - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } + assert.Equal(t, tt.wantPath, capture.Last(t).Path) }) } } func TestMessagesModels_EnvOverride(t *testing.T) { t.Setenv(messagesModelsEnvVar, "foo-model, bar-model ") - p := NewWithHTTPClient("sk-opencode", "", nil, llmclient.Hooks{}) + p := newTestProvider("", nil) - if !p.usesMessages("foo-model") || !p.usesMessages("bar-model") { - t.Fatal("override models should route to /messages") - } - if p.usesMessages("qwen3.7-max") { - t.Fatal("default model should not apply when override is set") - } - if !p.usesMessages("opencode_go/foo-model") { - t.Fatal("provider-qualified model should match after prefix strip") - } -} - -func TestListModels_NormalizesResponse(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "data":[ - {"id":"kimi-k2.7-code","object":"model","created":1781462836,"owned_by":"opencode"}, - {"id":"glm-5.1","object":"model","created":1781462836,"owned_by":"opencode"} - ] - }`)) - })) - defer server.Close() - - resp, err := newTestProvider(server.URL, server.Client()).ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if gotPath != "/models" { - t.Fatalf("path = %q, want /models", gotPath) - } - if len(resp.Data) != 2 || resp.Data[0].ID != "kimi-k2.7-code" { - t.Fatalf("unexpected models response: %+v", resp.Data) - } -} - -func TestEmbeddings_Unsupported(t *testing.T) { - _, err := newTestProvider("", nil).Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: "glm-5.1", - Input: "hello", - }) - if err == nil { - t.Fatal("Embeddings() error = nil, want invalid_request_error") - } - gwErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if gwErr.HTTPStatusCode() != http.StatusBadRequest { - t.Fatalf("status = %d, want %d", gwErr.HTTPStatusCode(), http.StatusBadRequest) - } -} - -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := newTestProvider("", nil) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("opencode_go provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("opencode_go provider should not implement native file provider") - } + assert.True(t, p.usesMessages("foo-model")) + assert.True(t, p.usesMessages("bar-model")) + assert.False(t, p.usesMessages("qwen3.7-max")) + assert.True(t, p.usesMessages("opencode_go/foo-model")) } diff --git a/internal/providers/opencodego/reasoning_test.go b/internal/providers/opencodego/reasoning_test.go index b6ea0da6f..5fed65e28 100644 --- a/internal/providers/opencodego/reasoning_test.go +++ b/internal/providers/opencodego/reasoning_test.go @@ -8,51 +8,35 @@ import ( "testing" "github.com/goccy/go-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -// captureBody serves a minimal chat completion and records the outgoing body. -func captureBody(t *testing.T, body *map[string]any) *httptest.Server { +// chatServer serves a minimal chat completion and records the outgoing body. +func chatServer(t *testing.T) (*httptest.Server, *providertest.Capture) { t.Helper() - return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - raw, err := io.ReadAll(r.Body) - if err != nil { - t.Errorf("read request body: %v", err) - return - } - if err := json.Unmarshal(raw, body); err != nil { - t.Errorf("unmarshal request body: %v", err) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-opencode", - "created":1677652288, - "model":"ox-alpha-free", - "choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}] - }`)) - })) + return providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-opencode", + "created":1677652288, + "model":"ox-alpha-free", + "choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}] + }`) } func TestChatCompletion_InjectsDefaultReasoningEffortWhenAbsent(t *testing.T) { - var got map[string]any - server := captureBody(t, &got) - defer server.Close() + server, capture := chatServer(t) _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got["reasoning_effort"] != "low" { - t.Fatalf("reasoning_effort = %v, want low", got["reasoning_effort"]) - } - if _, ok := got["reasoning"]; ok { - t.Fatal("nested reasoning object should not be forwarded") - } + require.NoError(t, err) + got := capture.Last(t).JSON(t) + assert.Equal(t, "low", got["reasoning_effort"]) + assert.NotContains(t, got, "reasoning") } func TestChatCompletion_MapsExplicitReasoningEffort(t *testing.T) { @@ -72,47 +56,33 @@ func TestChatCompletion_MapsExplicitReasoningEffort(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var got map[string]any - server := captureBody(t, &got) - defer server.Close() + server, capture := chatServer(t) _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: tt.effort}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got["reasoning_effort"] != tt.want { - t.Fatalf("reasoning_effort = %v, want %v", got["reasoning_effort"], tt.want) - } + require.NoError(t, err) + assert.Equal(t, tt.want, capture.Last(t).JSON(t)["reasoning_effort"]) }) } } func TestChatCompletion_KeepsClientSuppliedFlatReasoningEffort(t *testing.T) { - var got map[string]any - server := captureBody(t, &got) - defer server.Close() + server, capture := chatServer(t) extra, err := core.MergeUnknownJSONFields(core.UnknownJSONFields{}, map[string]json.RawMessage{ "reasoning_effort": json.RawMessage(`"max"`), }) - if err != nil { - t.Fatalf("MergeUnknownJSONFields() error = %v", err) - } - - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ + require.NoError(t, err) + _, err = newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, ExtraFields: extra, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got["reasoning_effort"] != "max" { - t.Fatalf("reasoning_effort = %v, want max (client value preserved)", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.Equal(t, "max", capture.Last(t).JSON(t)["reasoning_effort"]) } // TestChatCompletion_NestedReasoningWinsOverFlatField pins the precedence a @@ -122,46 +92,32 @@ func TestChatCompletion_KeepsClientSuppliedFlatReasoningEffort(t *testing.T) { // only authoritative when no canonical reasoning is present, where it stops the // default from being injected over it. func TestChatCompletion_NestedReasoningWinsOverFlatField(t *testing.T) { - var got map[string]any - server := captureBody(t, &got) - defer server.Close() + server, capture := chatServer(t) extra, err := core.MergeUnknownJSONFields(core.UnknownJSONFields{}, map[string]json.RawMessage{ "reasoning_effort": json.RawMessage(`"max"`), }) - if err != nil { - t.Fatalf("MergeUnknownJSONFields() error = %v", err) - } - - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ + require.NoError(t, err) + _, err = newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: "medium"}, ExtraFields: extra, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got["reasoning_effort"] != "low" { - t.Fatalf("reasoning_effort = %v, want low (canonical reasoning.effort wins)", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.Equal(t, "low", capture.Last(t).JSON(t)["reasoning_effort"]) } func TestChatCompletion_DefaultReasoningEffortEnvOverride(t *testing.T) { t.Setenv(defaultReasoningEffortEnvVar, "high") - var got map[string]any - server := captureBody(t, &got) - defer server.Close() - - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ + server, capture := chatServer(t) + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got["reasoning_effort"] != "high" { - t.Fatalf("reasoning_effort = %v, want high", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.Equal(t, "high", capture.Last(t).JSON(t)["reasoning_effort"]) } func TestChatCompletion_DefaultReasoningEffortDisabled(t *testing.T) { @@ -169,120 +125,69 @@ func TestChatCompletion_DefaultReasoningEffortDisabled(t *testing.T) { t.Run(value, func(t *testing.T) { t.Setenv(defaultReasoningEffortEnvVar, value) - var got map[string]any - server := captureBody(t, &got) - defer server.Close() - - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ + server, capture := chatServer(t) + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := got["reasoning_effort"]; ok { - t.Fatalf("reasoning_effort = %v, want absent when injection is disabled", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.NotContains(t, capture.Last(t).JSON(t), "reasoning_effort", "want absent when injection is disabled") }) } } func TestStreamChatCompletion_InjectsDefaultReasoningEffort(t *testing.T) { - var got map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - raw, err := io.ReadAll(r.Body) - if err != nil { - t.Errorf("read request body: %v", err) - return - } - if err := json.Unmarshal(raw, &got); err != nil { - t.Errorf("unmarshal request body: %v", err) - return - } - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "data: [DONE]\n\n") body, err := newTestProvider(server.URL, server.Client()).StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: "ox-alpha-free", Messages: []core.Message{{Role: "user", Content: "hi"}}, Stream: true, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) + _, _ = io.Copy(io.Discard, body) _ = body.Close() - if got["reasoning_effort"] != "low" { - t.Fatalf("reasoning_effort = %v, want low", got["reasoning_effort"]) - } + assert.Equal(t, "low", capture.Last(t).JSON(t)["reasoning_effort"]) } // TestResponses_InjectsDefaultReasoningEffort covers the /v1/responses path: // it is translated to a chat completion by ResponsesViaChat, so the adaptation // must reach the upstream body there too. func TestResponses_InjectsDefaultReasoningEffort(t *testing.T) { - var got map[string]any - server := captureBody(t, &got) - defer server.Close() - - if _, err := newTestProvider(server.URL, server.Client()).Responses(context.Background(), &core.ResponsesRequest{ + server, capture := chatServer(t) + _, err := newTestProvider(server.URL, server.Client()).Responses(context.Background(), &core.ResponsesRequest{ Model: "ox-alpha-free", Input: "hi", - }); err != nil { - t.Fatalf("Responses() error = %v", err) - } - if got["reasoning_effort"] != "low" { - t.Fatalf("reasoning_effort = %v, want low", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.Equal(t, "low", capture.Last(t).JSON(t)["reasoning_effort"]) } // TestChatCompletion_MessagesModelUnaffected pins that the injection lives on // the /chat/completions path only: the Anthropic-native dialect has no // reasoning_effort field. func TestChatCompletion_MessagesModelUnaffected(t *testing.T) { - var got map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - raw, err := io.ReadAll(r.Body) - if err != nil { - t.Errorf("read request body: %v", err) - return - } - if err := json.Unmarshal(raw, &got); err != nil { - t.Errorf("unmarshal request body: %v", err) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_opencode", - "model":"qwen3.7-max", - "content":[{"type":"text","text":"hello"}], - "stop_reason":"end_turn", - "usage":{"input_tokens":5,"output_tokens":2} - }`)) - })) - defer server.Close() - - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "id":"msg_opencode", + "model":"qwen3.7-max", + "content":[{"type":"text","text":"hello"}], + "stop_reason":"end_turn", + "usage":{"input_tokens":5,"output_tokens":2} + }`) + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), &core.ChatRequest{ Model: "qwen3.7-max", Messages: []core.Message{{Role: "user", Content: "hi"}}, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := got["reasoning_effort"]; ok { - t.Fatalf("reasoning_effort = %v, want absent on the /messages dialect", got["reasoning_effort"]) - } + }) + require.NoError(t, err) + assert.NotContains(t, capture.Last(t).JSON(t), "reasoning_effort", "want absent on the /messages dialect") } func TestAdaptChatRequest_NilRequest(t *testing.T) { req, err := adaptChatRequest(defaultReasoningEffort)(nil) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if req != nil { - t.Fatalf("req = %#v, want nil", req) - } + require.NoError(t, err) + require.Nil(t, req) } func TestLoadDefaultReasoningEffort(t *testing.T) { @@ -300,9 +205,7 @@ func TestLoadDefaultReasoningEffort(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Setenv(defaultReasoningEffortEnvVar, tt.env) - if got := loadDefaultReasoningEffort(); got != tt.want { - t.Fatalf("loadDefaultReasoningEffort() = %q, want %q", got, tt.want) - } + assert.Equal(t, tt.want, loadDefaultReasoningEffort()) }) } } diff --git a/internal/providers/opencodego/session_test.go b/internal/providers/opencodego/session_test.go index 0b9b69306..ddf2b81c2 100644 --- a/internal/providers/opencodego/session_test.go +++ b/internal/providers/opencodego/session_test.go @@ -7,36 +7,21 @@ import ( "net/http/httptest" "os" "strings" - "sync" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" "github.com/enterpilot/gomodel/internal/version" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -// headerServer records the request headers of the last call and answers with -// a minimal chat completion (or SSE stream) so both endpoint dialects succeed. -func headerServer(t *testing.T) (*httptest.Server, func() http.Header) { +// headerServer records every request and answers with a minimal chat +// completion (or SSE stream) so both endpoint dialects succeed. +func headerServer(t *testing.T) (*httptest.Server, *providertest.Capture) { t.Helper() - server, last, _ := headerServerWithPath(t) - return server, last -} - -// headerServerWithPath is headerServer with a third accessor for the request -// path of the last call. -func headerServerWithPath(t *testing.T) (*httptest.Server, func() http.Header, func() string) { - t.Helper() - var ( - mu sync.Mutex - last http.Header - lastSeen string - ) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - mu.Lock() - last = r.Header.Clone() - lastSeen = r.URL.Path - mu.Unlock() + return providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { if r.Header.Get("Accept") == "text/event-stream" { w.Header().Set("Content-Type", "text/event-stream") _, _ = w.Write([]byte("data: [DONE]\n\n")) @@ -48,17 +33,7 @@ func headerServerWithPath(t *testing.T) (*httptest.Server, func() http.Header, f return } _, _ = w.Write([]byte(`{"id":"chatcmpl-1","created":1,"model":"glm-5.1","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}]}`)) - })) - t.Cleanup(server.Close) - return server, func() http.Header { - mu.Lock() - defer mu.Unlock() - return last - }, func() string { - mu.Lock() - defer mu.Unlock() - return lastSeen - } + }) } func snapshotContext(ctx context.Context, headers map[string][]string) context.Context { @@ -71,100 +46,77 @@ func chatRequest(model string) *core.ChatRequest { } func TestRequestHeaders_DetectedSessionForwarded(t *testing.T) { - server, last := headerServer(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "session-123") for _, model := range []string{"glm-5.1", "qwen3.7-max"} { t.Run(model, func(t *testing.T) { - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest(model)); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got := last() - if v := got.Get(sessionHeader); v != "session-123" { - t.Fatalf("%s = %q, want session-123", sessionHeader, v) - } - if v := got.Get(clientHeader); v != defaultClient { - t.Fatalf("%s = %q, want %s", clientHeader, v, defaultClient) - } - if v := got.Get("User-Agent"); v != "gomodel/"+version.Version { - t.Fatalf("User-Agent = %q, want gomodel/%s", v, version.Version) - } + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest(model)) + require.NoError(t, err) + + got := capture.Last(t).Header + assert.Equal(t, "session-123", got.Get(sessionHeader)) + assert.Equal(t, defaultClient, got.Get(clientHeader)) + assert.Equal(t, "gomodel/"+version.Version, got.Get("User-Agent")) }) } } func TestRequestHeaders_StreamCarriesSession(t *testing.T) { - server, last := headerServer(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "session-stream") for _, model := range []string{"glm-5.1", "qwen3.7-max"} { t.Run(model, func(t *testing.T) { body, err := newTestProvider(server.URL, server.Client()).StreamChatCompletion(ctx, chatRequest(model)) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) + _, _ = io.ReadAll(body) _ = body.Close() - if v := last().Get(sessionHeader); v != "session-stream" { - t.Fatalf("%s = %q, want session-stream", sessionHeader, v) - } + assert.Equal(t, "session-stream", capture.Last(t).Header.Get(sessionHeader)) }) } } func TestRequestHeaders_InboundOpenCodeHeadersWin(t *testing.T) { - server, last := headerServer(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "scoped-detected") ctx = snapshotContext(ctx, map[string][]string{ "X-Opencode-Session": {"ses_client"}, "X-Opencode-Client": {"pi"}, }) + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest("glm-5.1")) + require.NoError(t, err) - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest("glm-5.1")); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got := last() - if v := got.Get(sessionHeader); v != "ses_client" { - t.Fatalf("%s = %q, want ses_client", sessionHeader, v) - } - if v := got.Get(clientHeader); v != "pi" { - t.Fatalf("%s = %q, want pi", clientHeader, v) - } + got := capture.Last(t).Header + assert.Equal(t, "ses_client", got.Get(sessionHeader)) + assert.Equal(t, "pi", got.Get(clientHeader)) } func TestRequestHeaders_NoSessionSendsNoSessionHeader(t *testing.T) { - server, last := headerServer(t) + server, capture := headerServer(t) + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), chatRequest("glm-5.1")) + require.NoError(t, err) - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(context.Background(), chatRequest("glm-5.1")); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got := last() - if _, ok := got[http.CanonicalHeaderKey(sessionHeader)]; ok { - t.Fatalf("%s should be absent without a session, got %q", sessionHeader, got.Get(sessionHeader)) - } - if v := got.Get(clientHeader); v != defaultClient { - t.Fatalf("%s = %q, want %s", clientHeader, v, defaultClient) - } + got := capture.Last(t).Header + assert.NotContains(t, got, http.CanonicalHeaderKey(sessionHeader), "%s should be absent without a session", sessionHeader) + assert.Equal(t, defaultClient, got.Get(clientHeader)) } func TestRequestHeaders_Disabled(t *testing.T) { t.Setenv(sessionHeaderEnvVar, "false") - server, last := headerServer(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "session-123") ctx = snapshotContext(ctx, map[string][]string{"X-Opencode-Session": {"ses_client"}}) for _, model := range []string{"glm-5.1", "qwen3.7-max"} { t.Run(model, func(t *testing.T) { - if _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest(model)); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got := last() - if _, ok := got[http.CanonicalHeaderKey(sessionHeader)]; ok { - t.Fatalf("%s should be absent when disabled, got %q", sessionHeader, got.Get(sessionHeader)) - } - if v := got.Get(clientHeader); v != defaultClient { - t.Fatalf("%s = %q, want %s (identification is not gated)", clientHeader, v, defaultClient) - } + _, err := newTestProvider(server.URL, server.Client()).ChatCompletion(ctx, chatRequest(model)) + require.NoError(t, err) + + got := capture.Last(t).Header + assert.NotContains(t, got, http.CanonicalHeaderKey(sessionHeader), "%s should be absent when disabled", sessionHeader) + assert.Equal(t, defaultClient, got.Get(clientHeader), "identification is not gated") }) } } @@ -174,7 +126,7 @@ func TestNew_FactoryConstructorWiresHeadersOnBothPaths(t *testing.T) { // header or reroute the /messages model through /chat/completions. t.Setenv(sessionHeaderEnvVar, "") t.Setenv(messagesModelsEnvVar, "qwen3.7-max") - server, last, lastPath := headerServerWithPath(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "session-factory") provider := New(providers.ProviderConfig{APIKey: "sk-opencode", BaseURL: server.URL}, providers.ProviderOptions{}) @@ -187,22 +139,18 @@ func TestNew_FactoryConstructorWiresHeadersOnBothPaths(t *testing.T) { } for _, tt := range tests { t.Run(tt.model, func(t *testing.T) { - if _, err := provider.ChatCompletion(ctx, chatRequest(tt.model)); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - got, path := last(), lastPath() - if path != tt.wantPath { - t.Fatalf("path = %q, want %s", path, tt.wantPath) - } - if v := got.Get(sessionHeader); v != "session-factory" { - t.Fatalf("%s = %q, want session-factory", sessionHeader, v) - } + _, err := provider.ChatCompletion(ctx, chatRequest(tt.model)) + require.NoError(t, err) + + got := capture.Last(t) + assert.Equal(t, tt.wantPath, got.Path) + assert.Equal(t, "session-factory", got.Header.Get(sessionHeader)) }) } } func TestPassthrough_FillsMissingIdentificationHeaders(t *testing.T) { - server, last := headerServer(t) + server, capture := headerServer(t) ctx := core.WithSessionID(context.Background(), "session-pass") resp, err := newTestProvider(server.URL, server.Client()).Passthrough(ctx, &core.PassthroughRequest{ @@ -211,28 +159,21 @@ func TestPassthrough_FillsMissingIdentificationHeaders(t *testing.T) { Body: io.NopCloser(strings.NewReader(`{"model":"glm-5.1","messages":[{"role":"user","content":"hi"}]}`)), Headers: http.Header{"Content-Type": {"application/json"}, "X-Opencode-Client": {"curl-script"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) + _ = resp.Body.Close() - got := last() - if v := got.Get(sessionHeader); v != "session-pass" { - t.Fatalf("%s = %q, want session-pass", sessionHeader, v) - } - if v := got.Get(clientHeader); v != "curl-script" { - t.Fatalf("%s = %q, want caller value curl-script", clientHeader, v) - } - if v := got.Get("User-Agent"); v != "gomodel/"+version.Version { - t.Fatalf("User-Agent = %q, want gomodel/%s", v, version.Version) - } + got := capture.Last(t).Header + assert.Equal(t, "session-pass", got.Get(sessionHeader)) + assert.Equal(t, "curl-script", got.Get(clientHeader)) + assert.Equal(t, "gomodel/"+version.Version, got.Get("User-Agent")) } func TestWithDefaultHeaders(t *testing.T) { defaults := http.Header{"X-Opencode-Session": {"gw"}, "User-Agent": {"gomodel/dev"}} merged := withDefaultHeaders(nil, defaults) - if merged.Get("X-Opencode-Session") != "gw" || merged.Get("User-Agent") != "gomodel/dev" { - t.Fatalf("nil headers should take every default, got %v", merged) - } + assert.Equal(t, "gw", merged.Get("X-Opencode-Session")) + assert.Equal(t, "gomodel/dev", merged.Get("User-Agent"), "nil headers should take every default") + // A non-canonical caller key still counts as present and is not duplicated. caller := http.Header{} caller[strings.ToLower("X-Opencode-Session")] = []string{"mine"} @@ -243,27 +184,17 @@ func TestWithDefaultHeaders(t *testing.T) { sessionValues = append(sessionValues, values...) } } - if len(sessionValues) != 1 || sessionValues[0] != "mine" { - t.Fatalf("caller value should win without duplication, got %v", merged) - } - if merged.Get("User-Agent") != "gomodel/dev" { - t.Fatalf("missing default should be added, got %v", merged) - } - if len(caller) != 1 { - t.Fatalf("caller headers must not be mutated, got %v", caller) - } + assert.Equal(t, []string{"mine"}, sessionValues, "caller value should win without duplication") + assert.Equal(t, "gomodel/dev", merged.Get("User-Agent"), "missing default should be added") + assert.Len(t, caller, 1) } func TestInboundHeader_RejectsLineBreaks(t *testing.T) { ctx := snapshotContext(context.Background(), map[string][]string{ "X-Opencode-Session": {"bad\r\nvalue", " good "}, }) - if got := inboundHeader(ctx, sessionHeader); got != "good" { - t.Fatalf("inboundHeader() = %q, want good", got) - } - if got := inboundHeader(context.Background(), sessionHeader); got != "" { - t.Fatalf("inboundHeader() without snapshot = %q, want empty", got) - } + assert.Equal(t, "good", inboundHeader(ctx, sessionHeader)) + assert.Empty(t, inboundHeader(context.Background(), sessionHeader)) } func TestLoadSessionHeaderEnabled(t *testing.T) { @@ -288,9 +219,7 @@ func TestLoadSessionHeaderEnabled(t *testing.T) { } else { t.Setenv(sessionHeaderEnvVar, tt.value) } - if got := loadSessionHeaderEnabled(); got != tt.want { - t.Fatalf("loadSessionHeaderEnabled() = %v, want %v", got, tt.want) - } + assert.Equal(t, tt.want, loadSessionHeaderEnabled()) }) } } diff --git a/internal/providers/openrouter/openrouter_test.go b/internal/providers/openrouter/openrouter_test.go index 1d791456f..02d73d555 100644 --- a/internal/providers/openrouter/openrouter_test.go +++ b/internal/providers/openrouter/openrouter_test.go @@ -2,332 +2,211 @@ package openrouter import ( "context" - "errors" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +// newTestProvider builds a provider pointed at baseURL. +func newTestProvider(baseURL string, client *http.Client) *Provider { + provider := NewWithHTTPClient("test-api-key", client, llmclient.Hooks{}) + provider.SetBaseURL(baseURL) + return provider +} + +// modelsByID indexes a listing by model ID. +func modelsByID(resp *core.ModelsResponse) map[string]core.Model { + byID := map[string]core.Model{} + for _, m := range resp.Data { + byID[m.ID] = m + } + return byID +} + +// pricingOf returns the model's pricing, or nil when it carries none. +func pricingOf(m core.Model) any { + if m.Metadata == nil || m.Metadata.Pricing == nil { + return nil + } + return m.Metadata.Pricing +} + // ListModels must keep OpenRouter's architecture modalities and context // length, mapping output modalities onto gateway modes so the catalog's long // tail is categorized without remote-registry entries. func TestListModels_StampsArchitectureModalities(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/models" { - t.Errorf("Path = %q, want /models", r.URL.Path) - } - // The endpoint defaults to text-output models; without this parameter - // embedding models would never enter the catalog. - if got := r.URL.Query().Get("output_modalities"); got != "all" { - t.Errorf("output_modalities = %q, want all", got) - } - _, _ = w.Write([]byte(`{"data":[ - {"id":"openai/gpt-4o-mini","created":1721260800,"context_length":128000, - "architecture":{"input_modalities":["text","image"],"output_modalities":["text"]}}, - {"id":"google/gemini-3-pro-image","created":1721260800, - "architecture":{"input_modalities":["text"],"output_modalities":["image"]}}, - {"id":"voyageai/voyage-4-lite","created":1721260800, - "architecture":{"input_modalities":["text"],"output_modalities":["embeddings"]}}, - {"id":"fish-audio/s1","created":1721260800, - "architecture":{"input_modalities":["text"],"output_modalities":["speech"]}}, - {"id":"mistralai/voxtral-mini-3b-2507","created":1721260800, - "architecture":{"input_modalities":["audio"],"output_modalities":["transcription"]}}, - {"id":"cohere/rerank-only","created":1721260800, - "architecture":{"input_modalities":["text"],"output_modalities":["rerank"]}}, - {"id":"acme/video-only","created":1721260800, - "architecture":{"input_modalities":["text"],"output_modalities":["video"]}}, - {"id":"mystery/no-architecture","created":1721260800} - ]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, `{"data":[ + {"id":"openai/gpt-4o-mini","created":1721260800,"context_length":128000, + "architecture":{"input_modalities":["text","image"],"output_modalities":["text"]}}, + {"id":"google/gemini-3-pro-image","created":1721260800, + "architecture":{"input_modalities":["text"],"output_modalities":["image"]}}, + {"id":"voyageai/voyage-4-lite","created":1721260800, + "architecture":{"input_modalities":["text"],"output_modalities":["embeddings"]}}, + {"id":"fish-audio/s1","created":1721260800, + "architecture":{"input_modalities":["text"],"output_modalities":["speech"]}}, + {"id":"mistralai/voxtral-mini-3b-2507","created":1721260800, + "architecture":{"input_modalities":["audio"],"output_modalities":["transcription"]}}, + {"id":"cohere/rerank-only","created":1721260800, + "architecture":{"input_modalities":["text"],"output_modalities":["rerank"]}}, + {"id":"acme/video-only","created":1721260800, + "architecture":{"input_modalities":["text"],"output_modalities":["video"]}}, + {"id":"mystery/no-architecture","created":1721260800} + ]}`) + provider := newTestProvider(server.URL, server.Client()) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(resp.Data) != 6 { - t.Fatalf("len(Data) = %d, want 6 (rerank-only and video-only skipped): %+v", len(resp.Data), resp.Data) - } - byID := map[string]core.Model{} - for _, m := range resp.Data { - byID[m.ID] = m - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "/models", req.Path) + // The endpoint defaults to text-output models; without this parameter + // embedding models would never enter the catalog. + assert.Equal(t, "all", req.Query.Get("output_modalities")) + + require.Len(t, resp.Data, 6) + byID := modelsByID(resp) chat := byID["openai/gpt-4o-mini"] - if chat.Metadata == nil || len(chat.Metadata.Modes) != 1 || chat.Metadata.Modes[0] != "chat" { - t.Errorf("gpt-4o-mini metadata = %+v, want chat modes", chat.Metadata) - } - if chat.Metadata == nil || chat.Metadata.ContextWindow == nil || *chat.Metadata.ContextWindow != 128000 { - t.Errorf("gpt-4o-mini context window = %+v, want 128000", chat.Metadata) - } + require.NotNil(t, chat.Metadata) + assert.Equal(t, []string{"chat"}, chat.Metadata.Modes) + require.NotNil(t, chat.Metadata.ContextWindow) + assert.Equal(t, 128000, *chat.Metadata.ContextWindow) + image := byID["google/gemini-3-pro-image"] - if image.Metadata == nil || len(image.Metadata.Modes) != 1 || image.Metadata.Modes[0] != "image_generation" { - t.Errorf("image model metadata = %+v, want image_generation modes", image.Metadata) - } - if image.Metadata == nil || len(image.Metadata.Categories) != 1 || image.Metadata.Categories[0] != core.CategoryImage { - t.Errorf("image model categories = %+v, want [image]", image.Metadata) - } + require.NotNil(t, image.Metadata) + assert.Equal(t, []string{"image_generation"}, image.Metadata.Modes) + assert.Equal(t, []core.ModelCategory{core.CategoryImage}, image.Metadata.Categories) + embed := byID["voyageai/voyage-4-lite"] - if embed.Metadata == nil || len(embed.Metadata.Modes) != 1 || embed.Metadata.Modes[0] != "embedding" { - t.Errorf("voyage-4-lite metadata = %+v, want embedding modes", embed.Metadata) - } - if embed.Metadata == nil || len(embed.Metadata.Categories) != 1 || embed.Metadata.Categories[0] != core.CategoryEmbedding { - t.Errorf("voyage-4-lite categories = %+v, want [embedding]", embed.Metadata) - } + require.NotNil(t, embed.Metadata) + assert.Equal(t, []string{"embedding"}, embed.Metadata.Modes) + assert.Equal(t, []core.ModelCategory{core.CategoryEmbedding}, embed.Metadata.Categories) + speech := byID["fish-audio/s1"] - if speech.Metadata == nil || len(speech.Metadata.Modes) != 1 || speech.Metadata.Modes[0] != "audio_speech" { - t.Errorf("speech model metadata = %+v, want audio_speech modes", speech.Metadata) - } - if speech.Metadata == nil || len(speech.Metadata.Categories) != 1 || speech.Metadata.Categories[0] != core.CategoryAudio { - t.Errorf("speech model categories = %+v, want [audio]", speech.Metadata) - } + require.NotNil(t, speech.Metadata) + assert.Equal(t, []string{"audio_speech"}, speech.Metadata.Modes) + assert.Equal(t, []core.ModelCategory{core.CategoryAudio}, speech.Metadata.Categories) + stt := byID["mistralai/voxtral-mini-3b-2507"] - if stt.Metadata == nil || len(stt.Metadata.Modes) != 1 || stt.Metadata.Modes[0] != "audio_transcription" { - t.Errorf("transcription model metadata = %+v, want audio_transcription modes", stt.Metadata) - } - if _, ok := byID["cohere/rerank-only"]; ok { - t.Error("rerank-only model must be skipped: no gateway surface reaches it on OpenRouter") - } - if _, ok := byID["acme/video-only"]; ok { - t.Error("video-only model must be skipped: no gateway surface reaches it on OpenRouter") - } + require.NotNil(t, stt.Metadata) + assert.Equal(t, []string{"audio_transcription"}, stt.Metadata.Modes) + + assert.NotContains(t, byID, "cohere/rerank-only") + assert.NotContains(t, byID, "acme/video-only") + noArch, ok := byID["mystery/no-architecture"] - if !ok { - t.Fatal("no-architecture model must be retained: missing signal is not proof of unservability") - } - if noArch.Metadata != nil { - t.Errorf("no-architecture metadata = %+v, want nil", noArch.Metadata) - } + require.True(t, ok) + assert.Nil(t, noArch.Metadata) } // A failed upstream listing must propagate as an error, not an empty catalog. func TestListModels_UpstreamErrorPropagates(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusInternalServerError) - _, _ = w.Write([]byte(`{"error":{"message":"upstream exploded"}}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusInternalServerError, `{"error":{"message":"upstream exploded"}}`) + provider := newTestProvider(server.URL, server.Client()) - resp, err := provider.ListModels(context.Background()) - if err == nil { - t.Fatalf("expected error, got response: %+v", resp) - } - if _, ok := errors.AsType[*core.GatewayError](err); !ok { - t.Fatalf("error type = %T, want *core.GatewayError: %v", err, err) - } + _, err := provider.ListModels(context.Background()) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) } // Audio flows through the embedded OpenAI-compatible implementation; // OpenRouter-specific request mutation (attribution headers) must still apply // on that path so audio traffic is attributed like every other call. func TestAudio_UsesOpenAISurfaceWithAttributionHeaders(t *testing.T) { - type seen struct { - path string - referer string - title string - } - requests := make(chan seen, 2) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - requests <- seen{ - path: r.URL.Path, - referer: r.Header.Get("HTTP-Referer"), - title: r.Header.Get("X-OpenRouter-Title"), - } - switch r.URL.Path { - case "/audio/speech": + server, capture := providertest.RouteServer(t, map[string]http.HandlerFunc{ + "/audio/speech": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "audio/mpeg") _, _ = w.Write([]byte("mp3-bytes")) - case "/audio/transcriptions": + }, + "/audio/transcriptions": func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"text":"hello"}`)) - default: - t.Errorf("unexpected path %q", r.URL.Path) - w.WriteHeader(http.StatusNotFound) - } - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + }, + }) + provider := newTestProvider(server.URL, server.Client()) speech, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "fish-audio/s1", Input: "hello world", Voice: "alloy", }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if speech.ContentType != "audio/mpeg" { - t.Errorf("speech ContentType = %q, want audio/mpeg", speech.ContentType) - } + require.NoError(t, err) + assert.Equal(t, "audio/mpeg", speech.ContentType) transcription, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "mistralai/voxtral-mini-3b-2507", File: []byte("wav-bytes"), Filename: "clip.wav", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } - if !strings.Contains(string(transcription.Data), "hello") { - t.Errorf("transcription Data = %q, want to contain hello", transcription.Data) - } + require.NoError(t, err) + assert.Contains(t, string(transcription.Data), "hello") - for _, want := range []string{"/audio/speech", "/audio/transcriptions"} { - got := <-requests - if got.path != want { - t.Errorf("path = %q, want %q", got.path, want) - } - if got.referer != defaultSiteURL { - t.Errorf("HTTP-Referer on %s = %q, want %q", want, got.referer, defaultSiteURL) - } - if got.title != defaultAppName { - t.Errorf("X-OpenRouter-Title on %s = %q, want %q", want, got.title, defaultAppName) - } + requests := capture.All() + require.Len(t, requests, 2) + for i, want := range []string{"/audio/speech", "/audio/transcriptions"} { + assert.Equal(t, want, requests[i].Path) + assert.Equal(t, defaultSiteURL, requests[i].Header.Get("HTTP-Referer"), "HTTP-Referer on %s", want) + assert.Equal(t, defaultAppName, requests[i].Header.Get("X-OpenRouter-Title"), "X-OpenRouter-Title on %s", want) } } func TestChatCompletion_AddsDefaultAttributionHeaders(t *testing.T) { - var gotReferer string - var gotTitle string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotReferer = r.Header.Get("HTTP-Referer") - gotTitle = r.Header.Get("X-OpenRouter-Title") - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-123", - "object":"chat.completion", - "created":1677652288, - "model":"openai/gpt-4o-mini", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL, server.Client()) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "openai/gpt-4o-mini", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: "openai/gpt-4o-mini", + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotAuth != "Bearer test-api-key" { - t.Fatalf("authorization = %q, want Bearer test-api-key", gotAuth) - } - if gotReferer != defaultSiteURL { - t.Fatalf("HTTP-Referer = %q, want %q", gotReferer, defaultSiteURL) - } - if gotTitle != defaultAppName { - t.Fatalf("X-OpenRouter-Title = %q, want %q", gotTitle, defaultAppName) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "Bearer test-api-key", req.Header.Get("Authorization")) + assert.Equal(t, defaultSiteURL, req.Header.Get("HTTP-Referer")) + assert.Equal(t, defaultAppName, req.Header.Get("X-OpenRouter-Title")) } func TestChatCompletion_ForwardsGoModelSessionID(t *testing.T) { - gotSessionID := make(chan string, 1) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotSessionID <- r.Header.Get("X-Session-Id") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-123","object":"chat.completion","created":1677652288, - "model":"openai/gpt-4o-mini", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL, server.Client()) + ctx := core.WithSessionID(context.Background(), "conversation-42") _, err := provider.ChatCompletion(ctx, &core.ChatRequest{ Model: "openai/gpt-4o-mini", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if got := <-gotSessionID; got != "conversation-42" { - t.Fatalf("X-Session-Id = %q, want conversation-42", got) - } + require.NoError(t, err) + assert.Equal(t, "conversation-42", capture.Last(t).Header.Get("X-Session-Id")) } func TestChatCompletion_UsesEnvOverridesForAttributionHeaders(t *testing.T) { t.Setenv("OPENROUTER_SITE_URL", "https://example.com") t.Setenv("OPENROUTER_APP_NAME", "Example App") - var gotReferer string - var gotTitle string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotReferer = r.Header.Get("HTTP-Referer") - gotTitle = r.Header.Get("X-OpenRouter-Title") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-123", - "object":"chat.completion", - "created":1677652288, - "model":"openai/gpt-4o-mini", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL, server.Client()) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "openai/gpt-4o-mini", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: "openai/gpt-4o-mini", + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if gotReferer != "https://example.com" { - t.Fatalf("HTTP-Referer = %q, want https://example.com", gotReferer) - } - if gotTitle != "Example App" { - t.Fatalf("X-OpenRouter-Title = %q, want Example App", gotTitle) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, "https://example.com", req.Header.Get("HTTP-Referer")) + assert.Equal(t, "Example App", req.Header.Get("X-OpenRouter-Title")) } func TestPassthrough_PreservesUserProvidedAttributionHeaders(t *testing.T) { - var gotReferer string - var gotTitle string - var gotLegacyTitle string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotReferer = r.Header.Get("HTTP-Referer") - gotTitle = r.Header.Get("X-OpenRouter-Title") - gotLegacyTitle = r.Header.Get("X-Title") - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusTooManyRequests) - _, _ = w.Write([]byte(`{"error":"rate limited"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusTooManyRequests, `{"error":"rate limited"}`) + provider := newTestProvider(server.URL, server.Client()) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, @@ -339,20 +218,13 @@ func TestPassthrough_PreservesUserProvidedAttributionHeaders(t *testing.T) { "X-Title": {"Caller App"}, }, }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) defer func() { _ = resp.Body.Close() }() - if gotReferer != "https://caller.example" { - t.Fatalf("HTTP-Referer = %q, want https://caller.example", gotReferer) - } - if gotLegacyTitle != "Caller App" { - t.Fatalf("X-Title = %q, want Caller App", gotLegacyTitle) - } - if gotTitle != "" { - t.Fatalf("X-OpenRouter-Title = %q, want empty when caller provided X-Title", gotTitle) - } + req := capture.Last(t) + assert.Equal(t, "https://caller.example", req.Header.Get("HTTP-Referer")) + assert.Equal(t, "Caller App", req.Header.Get("X-Title")) + assert.Empty(t, req.Header.Get("X-OpenRouter-Title")) } // OpenRouter prices its own catalog, including the ":free" variants it @@ -360,112 +232,78 @@ func TestPassthrough_PreservesUserProvidedAttributionHeaders(t *testing.T) { // enrichment treats what a provider reports as the override, and price-based // model filtering and cost-based load balancing both read them. func TestListModels_StampsPricing(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{"data":[ - {"id":"openai/gpt-4o-mini","context_length":128000, - "architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"0.00000015","completion":"0.0000006"}}, - {"id":"deepseek/deepseek-r1:free", - "architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"0","completion":"0"}}, - {"id":"openrouter/auto", - "architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"-1","completion":"-1"}}, - {"id":"acme/unpriced","architecture":{"output_modalities":["text"]}} - ]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"data":[ + {"id":"openai/gpt-4o-mini","context_length":128000, + "architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"0.00000015","completion":"0.0000006"}}, + {"id":"deepseek/deepseek-r1:free", + "architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"0","completion":"0"}}, + {"id":"openrouter/auto", + "architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"-1","completion":"-1"}}, + {"id":"acme/unpriced","architecture":{"output_modalities":["text"]}} + ]}`) + provider := newTestProvider(server.URL, server.Client()) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - byID := map[string]core.Model{} - for _, m := range resp.Data { - byID[m.ID] = m - } + require.NoError(t, err) + byID := modelsByID(resp) paid := byID["openai/gpt-4o-mini"].Metadata - if paid == nil || paid.Pricing == nil { - t.Fatalf("gpt-4o-mini pricing = %+v, want per-Mtok rates", paid) - } - if paid.Pricing.Currency != "USD" { - t.Errorf("Currency = %q, want USD", paid.Pricing.Currency) - } + require.NotNil(t, paid) + require.NotNil(t, paid.Pricing) + assert.Equal(t, "USD", paid.Pricing.Currency) + // Per-token rates are scaled to per million tokens. - if paid.Pricing.InputPerMtok == nil || *paid.Pricing.InputPerMtok != 0.15 { - t.Errorf("InputPerMtok = %v, want 0.15", paid.Pricing.InputPerMtok) - } - if paid.Pricing.OutputPerMtok == nil || *paid.Pricing.OutputPerMtok != 0.6 { - t.Errorf("OutputPerMtok = %v, want 0.6", paid.Pricing.OutputPerMtok) - } + require.NotNil(t, paid.Pricing.InputPerMtok) + assert.Equal(t, 0.15, *paid.Pricing.InputPerMtok) + require.NotNil(t, paid.Pricing.OutputPerMtok) + assert.Equal(t, 0.6, *paid.Pricing.OutputPerMtok) free := byID["deepseek/deepseek-r1:free"].Metadata - if free == nil || free.Pricing == nil || free.Pricing.InputPerMtok == nil { - t.Fatalf("free model pricing = %+v, want an explicit zero rate", free) - } - if *free.Pricing.InputPerMtok != 0 || *free.Pricing.OutputPerMtok != 0 { - t.Errorf("free model rates = %v/%v, want 0/0", *free.Pricing.InputPerMtok, *free.Pricing.OutputPerMtok) - } + require.NotNil(t, free) + require.NotNil(t, free.Pricing) + require.NotNil(t, free.Pricing.InputPerMtok) + assert.Equal(t, float64(0), *free.Pricing.InputPerMtok) + require.NotNil(t, free.Pricing.OutputPerMtok) + assert.Equal(t, float64(0), *free.Pricing.OutputPerMtok) // "-1" means OpenRouter cannot state the rate up front; reporting it as a // negative price would make an auto-routed model look cheaper than free. - if auto := byID["openrouter/auto"].Metadata; auto != nil && auto.Pricing != nil { - t.Errorf("auto-router pricing = %+v, want none", auto.Pricing) - } - if unpriced := byID["acme/unpriced"].Metadata; unpriced != nil && unpriced.Pricing != nil { - t.Errorf("unpriced model pricing = %+v, want none", unpriced.Pricing) - } + assert.Nil(t, pricingOf(byID["openrouter/auto"]), "auto-router must carry no pricing") + assert.Nil(t, pricingOf(byID["acme/unpriced"]), "unpriced model must carry no pricing") } // strconv.ParseFloat accepts "NaN" and "Inf", and scaling a huge per-token rate // to per-Mtok can overflow. Either would corrupt every downstream price // comparison and cost calculation, so such rates report no price at all. func TestListModels_RejectsNonFinitePricing(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{"data":[ - {"id":"acme/nan","architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"NaN","completion":"NaN"}}, - {"id":"acme/inf","architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"Inf","completion":"Inf"}}, - {"id":"acme/overflow","architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"1e308","completion":"1e308"}}, - {"id":"acme/partial","architecture":{"output_modalities":["text"]}, - "pricing":{"prompt":"0.000001","completion":"NaN"}} - ]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.JSONServer(t, http.StatusOK, `{"data":[ + {"id":"acme/nan","architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"NaN","completion":"NaN"}}, + {"id":"acme/inf","architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"Inf","completion":"Inf"}}, + {"id":"acme/overflow","architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"1e308","completion":"1e308"}}, + {"id":"acme/partial","architecture":{"output_modalities":["text"]}, + "pricing":{"prompt":"0.000001","completion":"NaN"}} + ]}`) + provider := newTestProvider(server.URL, server.Client()) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - byID := map[string]core.Model{} - for _, m := range resp.Data { - byID[m.ID] = m - } + require.NoError(t, err) + byID := modelsByID(resp) for _, id := range []string{"acme/nan", "acme/inf", "acme/overflow"} { - if meta := byID[id].Metadata; meta != nil && meta.Pricing != nil { - t.Errorf("%s pricing = %+v, want none", id, meta.Pricing) - } + assert.Nil(t, pricingOf(byID[id]), "%s must carry no pricing", id) } // One unusable rate must not discard the other, usable one. partial := byID["acme/partial"].Metadata - if partial == nil || partial.Pricing == nil || partial.Pricing.InputPerMtok == nil { - t.Fatalf("acme/partial pricing = %+v, want the parseable input rate", partial) - } - if *partial.Pricing.InputPerMtok != 1 { - t.Errorf("InputPerMtok = %v, want 1", *partial.Pricing.InputPerMtok) - } - if partial.Pricing.OutputPerMtok != nil { - t.Errorf("OutputPerMtok = %v, want none", *partial.Pricing.OutputPerMtok) - } + require.NotNil(t, partial) + require.NotNil(t, partial.Pricing) + require.NotNil(t, partial.Pricing.InputPerMtok) + assert.Equal(t, float64(1), *partial.Pricing.InputPerMtok) + assert.Nil(t, partial.Pricing.OutputPerMtok) } diff --git a/internal/providers/openrouter/passthrough_semantics_test.go b/internal/providers/openrouter/passthrough_semantics_test.go index 1be4571c6..1f0faf392 100644 --- a/internal/providers/openrouter/passthrough_semantics_test.go +++ b/internal/providers/openrouter/passthrough_semantics_test.go @@ -4,20 +4,20 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricherUsesOpenRouterType(t *testing.T) { enricher := Registration.PassthroughSemanticEnricher - if enricher == nil { - t.Fatal("registration passthrough enricher is nil") - } - if got := enricher.ProviderType(); got != "openrouter" { - t.Fatalf("ProviderType() = %q, want openrouter", got) - } + require.NotNil(t, enricher) + got := enricher.ProviderType() + require.Equal(t, "openrouter", got) + info := enricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ Provider: "openrouter", NormalizedEndpoint: "chat/completions", }) - if info == nil || info.GenAIOperation != "chat" || info.SemanticOperation != "openrouter.chat_completions" || info.AuditPath != "/v1/chat/completions" { - t.Fatalf("enriched info = %+v, want OpenRouter chat semantics", info) - } + require.NotNil(t, info) + require.Equal(t, "chat", info.GenAIOperation) + require.Equal(t, "openrouter.chat_completions", info.SemanticOperation) + require.Equal(t, "/v1/chat/completions", info.AuditPath) } diff --git a/internal/providers/oracle/oracle_test.go b/internal/providers/oracle/oracle_test.go index 92d57d311..2d4fd5005 100644 --- a/internal/providers/oracle/oracle_test.go +++ b/internal/providers/oracle/oracle_test.go @@ -3,65 +3,39 @@ package oracle import ( "context" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestListModels_ReturnsUpstreamInventory(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/models" { - http.NotFound(w, r) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"openai.gpt-oss-120b","object":"model","owned_by":"oracle"}]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[{"id":"openai.gpt-oss-120b","object":"model","owned_by":"oracle"}]}`) provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}) provider.SetBaseURL(server.URL) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "openai.gpt-oss-120b" { - t.Fatalf("unexpected models response: %+v", resp.Data) - } + require.NoError(t, err) + assert.Equal(t, "/models", capture.Last(t).Path) + require.Len(t, resp.Data, 1) + assert.Equal(t, "openai.gpt-oss-120b", resp.Data[0].ID) } func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}) _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{Model: "text-embedding-3-small"}) - if err == nil { - t.Fatal("expected error, got nil") - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if gatewayErr.Type != core.ErrorTypeInvalidRequest { - t.Fatalf("gatewayErr.Type = %q, want %q", gatewayErr.Type, core.ErrorTypeInvalidRequest) - } - if gatewayErr.Message != "oracle does not support embeddings" { - t.Fatalf("gatewayErr.Message = %q, want oracle does not support embeddings", gatewayErr.Message) - } + providertest.AssertUnsupported(t, err) + assert.Contains(t, err.Error(), "oracle does not support embeddings") } func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("oracle provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("oracle provider should not implement native file provider") - } - if _, ok := any(provider).(core.PassthroughProvider); ok { - t.Fatal("oracle provider should not implement passthrough provider") - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.PassthroughProvider) + assert.False(t, ok, "provider should not implement core.PassthroughProvider") } diff --git a/internal/providers/sglang/sglang_test.go b/internal/providers/sglang/sglang_test.go index acc0be83c..ca62186c6 100644 --- a/internal/providers/sglang/sglang_test.go +++ b/internal/providers/sglang/sglang_test.go @@ -5,13 +5,15 @@ import ( "encoding/json" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestChatCompletionUsesOptionalBearerAuthAndV1Endpoint(t *testing.T) { @@ -24,74 +26,42 @@ func TestChatCompletionUsesOptionalBearerAuthAndV1Endpoint(t *testing.T) { {name: "without API key"}, } { t.Run(tt.name, func(t *testing.T) { - var gotPath, gotAuth string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-sglang", - "created":1677652288, - "model":"HuggingFaceTB/SmolLM2-135M-Instruct", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) provider := NewWithHTTPClient(tt.apiKey, server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "HuggingFaceTB/SmolLM2-135M-Instruct", + Model: providertest.Model, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "HuggingFaceTB/SmolLM2-135M-Instruct" { - t.Fatalf("resp.Model = %q", resp.Model) - } - if gotPath != "/v1/chat/completions" { - t.Fatalf("path = %q, want /v1/chat/completions", gotPath) - } - if gotAuth != tt.wantAuth { - t.Fatalf("authorization = %q, want %q", gotAuth, tt.wantAuth) - } + require.NoError(t, err) + assert.Equal(t, providertest.Model, resp.Model) + + req := capture.Last(t) + assert.Equal(t, "/v1/chat/completions", req.Path) + assert.Equal(t, tt.wantAuth, req.Header.Get("Authorization")) }) } } func TestChatCompletionPreservesSGLangExtensionFields(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - t.Errorf("decode request: %v", err) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"chatcmpl-sglang","model":"test","choices":[]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"chatcmpl-sglang","model":"test","choices":[]}`) var req core.ChatRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"test", "messages":[{"role":"user","content":"hi"}], "chat_template_kwargs":{"enable_thinking":false}, "separate_reasoning":true - }`), &req); err != nil { - t.Fatalf("decode ChatRequest: %v", err) - } + }`), &req) + require.NoError(t, err) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + _, err = provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) - kwargs, ok := gotBody["chat_template_kwargs"].(map[string]any) - if !ok || kwargs["enable_thinking"] != false { - t.Fatalf("chat_template_kwargs = %#v", gotBody["chat_template_kwargs"]) - } - if gotBody["separate_reasoning"] != true { - t.Fatalf("separate_reasoning = %#v", gotBody["separate_reasoning"]) - } + sent := capture.Last(t).JSON(t) + assert.Equal(t, map[string]any{"enable_thinking": false}, sent["chat_template_kwargs"]) + assert.Equal(t, true, sent["separate_reasoning"]) } func TestOpenAICompatibleEndpoints(t *testing.T) { @@ -132,21 +102,11 @@ func TestOpenAICompatibleEndpoints(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(tt.response)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, tt.response) provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) - if err := tt.call(provider); err != nil { - t.Fatalf("call error = %v", err) - } - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } + require.NoError(t, tt.call(provider)) + assert.Equal(t, tt.wantPath, capture.Last(t).Path) }) } } @@ -175,42 +135,24 @@ func TestOpenAICompatibleStreamingEndpoints(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() + server, capture := providertest.SSEServer(t, "data: [DONE]\n\n") provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) body, err := tt.call(provider) - if err != nil { - t.Fatalf("streaming call error = %v", err) - } + require.NoError(t, err) defer body.Close() - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } + assert.Equal(t, tt.wantPath, capture.Last(t).Path) }) } } +// SGLang serves passthrough and the OpenAI-compatible surface, but nothing +// that would need native batch, file, audio, or response-lifecycle support. func TestProviderExposesOnlyVerifiedOptionalInterfaces(t *testing.T) { provider := NewWithHTTPClient("", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.PassthroughProvider); !ok { - t.Fatal("sglang provider should implement passthrough provider") - } - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("sglang provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("sglang provider should not implement native file provider") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("sglang provider should not implement native response lifecycle provider") - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + assert.False(t, ok, "provider should not implement core.NativeResponseLifecycleProvider") } func TestPassthroughRoutesNativeAndOpenAIEndpoints(t *testing.T) { @@ -231,14 +173,7 @@ func TestPassthroughRoutesNativeAndOpenAIEndpoints(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath, gotAuth string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{}`) provider := NewWithHTTPClient("sglang-key", server.URL+"/v1", server.Client(), llmclient.Hooks{}) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ @@ -247,29 +182,18 @@ func TestPassthroughRoutesNativeAndOpenAIEndpoints(t *testing.T) { Body: io.NopCloser(strings.NewReader("{}")), Headers: http.Header{"Content-Type": []string{"application/json"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != tt.wantPath { - t.Fatalf("path = %q, want %q", gotPath, tt.wantPath) - } - if gotAuth != "Bearer sglang-key" { - t.Fatalf("authorization = %q, want Bearer sglang-key", gotAuth) - } + req := capture.Last(t) + assert.Equal(t, tt.wantPath, req.Path) + assert.Equal(t, "Bearer sglang-key", req.Header.Get("Authorization")) }) } } func TestNewSharesKeyRotationWithNativePassthrough(t *testing.T) { - var gotAuth []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotAuth = append(gotAuth, r.Header.Get("Authorization")) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) keys := providers.NewKeyring("key-one", "key-two") provider := New(providers.ProviderConfig{ @@ -277,48 +201,39 @@ func TestNewSharesKeyRotationWithNativePassthrough(t *testing.T) { APIKey: "key-one", BaseURL: server.URL + "/v1", }, providers.ProviderOptions{Keys: keys}).(*Provider) + _, err := provider.ListModels(context.Background()) + require.NoError(t, err) - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodGet, Endpoint: "health", }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if len(gotAuth) != 2 || gotAuth[0] != "Bearer key-one" || gotAuth[1] != "Bearer key-two" { - t.Fatalf("authorization headers = %v", gotAuth) - } + requests := capture.All() + require.Len(t, requests, 2) + assert.Equal(t, "Bearer key-one", requests[0].Header.Get("Authorization")) + assert.Equal(t, "Bearer key-two", requests[1].Header.Get("Authorization")) } func TestSetBaseURLUpdatesOpenAIAndNativeClients(t *testing.T) { - var gotPaths []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPaths = append(gotPaths, r.URL.Path) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"object":"list","data":[]}`) provider := NewWithHTTPClient("", "http://127.0.0.1:1/v1", server.Client(), llmclient.Hooks{}) provider.SetBaseURL(server.URL + "/v1") - if _, err := provider.ListModels(context.Background()); err != nil { - t.Fatalf("ListModels() error = %v", err) - } + _, err := provider.ListModels(context.Background()) + require.NoError(t, err) + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodGet, Endpoint: "health", }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if len(gotPaths) != 2 || gotPaths[0] != "/v1/models" || gotPaths[1] != "/health" { - t.Fatalf("paths = %v, want [/v1/models /health]", gotPaths) - } + requests := capture.All() + require.Len(t, requests, 2) + assert.Equal(t, "/v1/models", requests[0].Path) + assert.Equal(t, "/health", requests[1].Path) } diff --git a/internal/providers/vertex/vertex_test.go b/internal/providers/vertex/vertex_test.go index 785e3b907..98b9decb4 100644 --- a/internal/providers/vertex/vertex_test.go +++ b/internal/providers/vertex/vertex_test.go @@ -11,71 +11,68 @@ import ( "encoding/pem" "math" "net/http" - "net/http/httptest" + "net/url" "os" "path/filepath" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" "github.com/enterpilot/gomodel/internal/providers/googlecommon" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "golang.org/x/oauth2" ) +const ( + nativeBasePath = "/v1/projects/prod-ai/locations/us-central1/publishers/google" + generateContentJSON = `{ + "responseId": "vertex-auth", + "candidates": [{ + "content": {"role": "model", "parts": [{"text": "ok"}]}, + "finishReason": "STOP" + }] + }` +) + +// tokenServer answers every OAuth token exchange with accessToken; the +// recorder keeps the form body so the test can inspect the grant type. +func tokenServer(t *testing.T, accessToken string) (string, *providertest.Capture) { + t.Helper() + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": accessToken, + "token_type": "Bearer", + "expires_in": 3600, + }) + }) + return server.URL, capture +} + func TestProviderDoesNotExposeFilesOrBatches(t *testing.T) { provider := newProvider(testConfig(), providers.ProviderOptions{}, authedTestClient(http.DefaultClient)) - - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("Vertex provider must not expose native files") - } - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("Vertex provider must not expose native batches") - } + _, ok := any(provider).(core.NativeFileProvider) + assert.False(t, ok, "provider should not implement core.NativeFileProvider") + _, ok = any(provider).(core.NativeBatchProvider) + assert.False(t, ok, "provider should not implement core.NativeBatchProvider") } func TestEmbeddingsUsesNativePrediction(t *testing.T) { - var operation string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1/projects/prod-ai/locations/us-central1/publishers/google/models/text-embedding-005:predict" { - t.Errorf("Path = %q, want Vertex native predict endpoint", r.URL.Path) - } - if got := r.Header.Get("Authorization"); got != "Bearer vertex-token" { - t.Errorf("Authorization = %q, want Bearer vertex-token", got) - } - if got := r.Header.Get("x-goog-api-key"); got != "" { - t.Errorf("x-goog-api-key = %q, want empty for Vertex OAuth", got) - } - - var payload vertexEmbeddingPredictRequest - if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { - t.Fatalf("failed to decode request: %v", err) - } - if len(payload.Instances) != 2 { - t.Fatalf("instances = %+v, want 2", payload.Instances) - } - if payload.Instances[0].Content != "hello" || payload.Instances[1].Content != "world" { - t.Fatalf("instances = %+v, want hello/world", payload.Instances) - } - if got := payload.Parameters["outputDimensionality"]; got != float64(3) { - t.Fatalf("outputDimensionality = %#v, want 3", got) - } - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "predictions": [ - {"embeddings": {"values": [0.1, 0.2, 0.3], "statistics": {"token_count": 4}}}, - {"embeddings": {"values": [0.4, 0.5, 0.6], "statistics": {"token_count": 5}}} - ] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "predictions": [ + {"embeddings": {"values": [0.1, 0.2, 0.3], "statistics": {"token_count": 4}}}, + {"embeddings": {"values": [0.4, 0.5, 0.6], "statistics": {"token_count": 5}}} + ] + }`) + var operation string dimensions := 3 cfg := testConfig() - cfg.BaseURL = server.URL + "/v1/projects/prod-ai/locations/us-central1/publishers/google" + cfg.BaseURL = server.URL + nativeBasePath provider := newProvider(cfg, providers.ProviderOptions{Hooks: llmclient.Hooks{ OnRequestStart: func(ctx context.Context, info llmclient.RequestInfo) context.Context { operation = info.Operation @@ -88,24 +85,25 @@ func TestEmbeddingsUsesNativePrediction(t *testing.T) { Input: []string{"hello", "world"}, Dimensions: &dimensions, }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Provider != "vertex" { - t.Fatalf("Provider = %q, want vertex", resp.Provider) - } - if len(resp.Data) != 2 { - t.Fatalf("data = %+v, want 2 embeddings", resp.Data) - } - if got := string(resp.Data[0].Embedding); got != `[0.1,0.2,0.3]` { - t.Fatalf("embedding = %s, want [0.1,0.2,0.3]", got) - } - if resp.Usage.PromptTokens != 9 || resp.Usage.TotalTokens != 9 { - t.Fatalf("usage = %+v, want 9 prompt/total tokens", resp.Usage) - } - if operation != llmclient.OperationEmbeddings { - t.Fatalf("operation = %q, want embeddings", operation) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, nativeBasePath+"/models/text-embedding-005:predict", req.Path) + assert.Equal(t, "Bearer vertex-token", req.Header.Get("Authorization")) + assert.Empty(t, req.Header.Get("x-goog-api-key")) + var payload vertexEmbeddingPredictRequest + require.NoError(t, json.Unmarshal(req.Body, &payload)) + require.Len(t, payload.Instances, 2) + assert.Equal(t, "hello", payload.Instances[0].Content) + assert.Equal(t, "world", payload.Instances[1].Content) + assert.Equal(t, float64(3), payload.Parameters["outputDimensionality"]) + + assert.Equal(t, "vertex", resp.Provider) + require.Len(t, resp.Data, 2) + assert.Equal(t, `[0.1,0.2,0.3]`, string(resp.Data[0].Embedding)) + assert.Equal(t, 9, resp.Usage.PromptTokens) + assert.Equal(t, 9, resp.Usage.TotalTokens) + assert.Equal(t, llmclient.OperationEmbeddings, operation) } func TestEmbeddingsRejectsEmptyStringInBatch(t *testing.T) { @@ -115,12 +113,8 @@ func TestEmbeddingsRejectsEmptyStringInBatch(t *testing.T) { Model: "google/text-embedding-005", Input: []string{"hello", "", "world"}, }) - if err == nil { - t.Fatal("expected empty batch embedding input to be rejected") - } - if !strings.Contains(err.Error(), "embedding input must not be empty") { - t.Fatalf("error = %v, want empty input error", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), "embedding input must not be empty") } func TestOpenAIEmbeddingResponseSupportsBase64Encoding(t *testing.T) { @@ -135,62 +129,43 @@ func TestOpenAIEmbeddingResponseSupportsBase64Encoding(t *testing.T) { }, }}, }) - if err != nil { - t.Fatalf("openAIEmbeddingResponse() error = %v", err) - } - if len(resp.Data) != 1 { - t.Fatalf("data = %+v, want one embedding", resp.Data) - } + require.NoError(t, err) + require.Len(t, resp.Data, 1) var encoded string - if err := json.Unmarshal(resp.Data[0].Embedding, &encoded); err != nil { - t.Fatalf("embedding is not JSON string: %v", err) - } + require.NoError(t, json.Unmarshal(resp.Data[0].Embedding, &encoded)) decoded, err := base64.StdEncoding.DecodeString(encoded) - if err != nil { - t.Fatalf("embedding is not valid base64: %v", err) - } - if len(decoded) != 8 { - t.Fatalf("decoded length = %d, want 8", len(decoded)) - } + require.NoError(t, err) + require.Len(t, decoded, 8) + values := []float32{ math.Float32frombits(binary.LittleEndian.Uint32(decoded[0:4])), math.Float32frombits(binary.LittleEndian.Uint32(decoded[4:8])), } - if values[0] != 0.5 || values[1] != -1.25 { - t.Fatalf("decoded values = %v, want [0.5 -1.25]", values) - } - if resp.Usage.PromptTokens != 3 || resp.Usage.TotalTokens != 3 { - t.Fatalf("usage = %+v, want 3 prompt/total tokens", resp.Usage) - } + assert.Equal(t, []float32{0.5, -1.25}, values) + assert.Equal(t, 3, resp.Usage.PromptTokens) + assert.Equal(t, 3, resp.Usage.TotalTokens) } func TestNewAcceptsBaseURLWithoutProjectLocation(t *testing.T) { provider := newProvider(providers.ProviderConfig{ Type: "vertex", AuthType: "gcp_adc", - BaseURL: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/publishers/google", + BaseURL: "https://proxy.example.com" + nativeBasePath, }, providers.ProviderOptions{}, authedTestClient(http.DefaultClient)) - - if err := provider.ready(); err != nil { - t.Fatalf("ready() error = %v, want nil for custom Vertex base URL", err) - } + require.NoError(t, provider.ready()) } func TestNewRejectsUnsupportedAuthType(t *testing.T) { provider := newProvider(providers.ProviderConfig{ Type: "vertex", AuthType: "api_key", - BaseURL: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/publishers/google", + BaseURL: "https://proxy.example.com" + nativeBasePath, }, providers.ProviderOptions{}, authedTestClient(http.DefaultClient)) err := provider.ready() - if err == nil { - t.Fatal("expected unsupported auth type error") - } - if !strings.Contains(err.Error(), `unsupported vertex AI auth type "api_key"`) { - t.Fatalf("error = %v, want unsupported auth type", err) - } + require.Error(t, err) + assert.Contains(t, err.Error(), `unsupported vertex AI auth type "api_key"`) } func TestNewAuthFormsInjectBearerToken(t *testing.T) { @@ -242,64 +217,32 @@ func TestNewAuthFormsInjectBearerToken(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotGrantType string - tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := r.ParseForm(); err != nil { - t.Fatalf("ParseForm() error = %v", err) - } - gotGrantType = r.PostForm.Get("grant_type") - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": tt.token, - "token_type": "Bearer", - "expires_in": 3600, - }) - })) - defer tokenServer.Close() - - upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1/projects/prod-ai/locations/us-central1/publishers/google/models/gemini-2.5-flash:generateContent" { - t.Errorf("Path = %q, want Vertex native generateContent endpoint", r.URL.Path) - } - if got := r.Header.Get("Authorization"); got != "Bearer "+tt.token { - t.Errorf("Authorization = %q, want Bearer %s", got, tt.token) - } - if got := r.Header.Get("x-goog-api-key"); got != "" { - t.Errorf("x-goog-api-key = %q, want empty for Vertex OAuth", got) - } - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "responseId": "vertex-auth", - "candidates": [{ - "content": {"role": "model", "parts": [{"text": "ok"}]}, - "finishReason": "STOP" - }] - }`)) - })) - defer upstream.Close() + tokenURL, tokenCapture := tokenServer(t, tt.token) + upstream, upstreamCapture := providertest.JSONServer(t, http.StatusOK, generateContentJSON) cfg := testConfig() cfg.AuthType = tt.authType cfg.APIMode = "native" - cfg.BaseURL = upstream.URL + "/v1/projects/prod-ai/locations/us-central1/publishers/google" - tt.configure(t, &cfg, tokenServer.URL) + cfg.BaseURL = upstream.URL + nativeBasePath + tt.configure(t, &cfg, tokenURL) provider := New(cfg, providers.ProviderOptions{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "google/gemini-2.5-flash", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, + Model: "google/gemini-2.5-flash", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp == nil || resp.Provider != "vertex" { - t.Fatalf("response = %+v, want vertex response", resp) - } - if gotGrantType != tt.grantType { - t.Fatalf("grant_type = %q, want %q", gotGrantType, tt.grantType) - } + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, "vertex", resp.Provider) + + form, err := url.ParseQuery(string(tokenCapture.Last(t).Body)) + require.NoError(t, err) + assert.Equal(t, tt.grantType, form.Get("grant_type")) + + req := upstreamCapture.Last(t) + assert.Equal(t, nativeBasePath+"/models/gemini-2.5-flash:generateContent", req.Path) + assert.Equal(t, "Bearer "+tt.token, req.Header.Get("Authorization")) + assert.Empty(t, req.Header.Get("x-goog-api-key")) }) } } @@ -318,7 +261,7 @@ func TestVertexBaseURLs(t *testing.T) { VertexLocation: "us-central1", }, wantCompat: "https://aiplatform.googleapis.com/v1/projects/prod-ai/locations/us-central1/endpoints/openapi", - wantNative: "https://aiplatform.googleapis.com/v1/projects/prod-ai/locations/us-central1/publishers/google", + wantNative: "https://aiplatform.googleapis.com" + nativeBasePath, }, { name: "custom OpenAI-compatible vertex URL derives native sibling", @@ -326,27 +269,23 @@ func TestVertexBaseURLs(t *testing.T) { BaseURL: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/endpoints/openapi/", }, wantCompat: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/endpoints/openapi", - wantNative: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/publishers/google", + wantNative: "https://proxy.example.com" + nativeBasePath, }, { name: "custom native vertex URL derives OpenAI-compatible sibling", cfg: providers.ProviderConfig{ - BaseURL: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/publishers/google/", + BaseURL: "https://proxy.example.com" + nativeBasePath + "/", }, wantCompat: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/endpoints/openapi", - wantNative: "https://proxy.example.com/v1/projects/prod-ai/locations/us-central1/publishers/google", + wantNative: "https://proxy.example.com" + nativeBasePath, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { gotCompat, gotNative := googlecommon.VertexBaseURLs(tt.cfg.BaseURL, tt.cfg.VertexProject, tt.cfg.VertexLocation) - if gotCompat != tt.wantCompat { - t.Fatalf("OpenAI-compatible base = %q, want %q", gotCompat, tt.wantCompat) - } - if gotNative != tt.wantNative { - t.Fatalf("native base = %q, want %q", gotNative, tt.wantNative) - } + assert.Equal(t, tt.wantCompat, gotCompat) + assert.Equal(t, tt.wantNative, gotNative) }) } } @@ -386,12 +325,8 @@ func vertexADCCredentialsFileWithQuotaProject(t *testing.T, tokenURL, quotaProje contents["quota_project_id"] = quotaProject } encoded, err := json.Marshal(contents) - if err != nil { - t.Fatalf("failed to marshal ADC credentials: %v", err) - } - if err := os.WriteFile(path, encoded, 0o600); err != nil { - t.Fatalf("failed to write ADC credentials: %v", err) - } + require.NoError(t, err) + require.NoError(t, os.WriteFile(path, encoded, 0o600)) return path } @@ -421,39 +356,21 @@ func TestNewSetsQuotaProjectHeaderOnVertexRequests(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": "token", - "token_type": "Bearer", - "expires_in": 3600, - }) - })) - defer tokenServer.Close() - - var gotProj string - upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotProj = r.Header.Get("X-Goog-User-Project") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"responseId":"r","candidates":[{"content":{"role":"model","parts":[{"text":"ok"}]},"finishReason":"STOP"}]}`)) - })) - defer upstream.Close() + tokenURL, _ := tokenServer(t, "token") + upstream, capture := providertest.JSONServer(t, http.StatusOK, generateContentJSON) cfg := testConfig() cfg.APIMode = "native" - cfg.BaseURL = upstream.URL + "/v1/projects/prod-ai/locations/us-central1/publishers/google" - tt.configure(t, &cfg, tokenServer.URL) + cfg.BaseURL = upstream.URL + nativeBasePath + tt.configure(t, &cfg, tokenURL) provider := New(cfg, providers.ProviderOptions{}) - if _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "google/gemini-2.5-flash", Messages: []core.Message{{Role: "user", Content: "hi"}}, - }); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotProj != tt.wantProj { - t.Fatalf("X-Goog-User-Project = %q, want %q", gotProj, tt.wantProj) - } + }) + require.NoError(t, err) + assert.Equal(t, tt.wantProj, capture.Last(t).Header.Get("X-Goog-User-Project")) }) } } @@ -461,22 +378,18 @@ func TestNewSetsQuotaProjectHeaderOnVertexRequests(t *testing.T) { func vertexServiceAccountCredentialsFile(t *testing.T, tokenURL string) string { t.Helper() path := filepath.Join(t.TempDir(), "service-account.json") - if err := os.WriteFile(path, []byte(vertexServiceAccountCredentials(t, tokenURL)), 0o600); err != nil { - t.Fatalf("failed to write service account credentials: %v", err) - } + require.NoError(t, os.WriteFile(path, []byte(vertexServiceAccountCredentials(t, tokenURL)), 0o600)) return path } func vertexServiceAccountCredentials(t *testing.T, tokenURL string) string { t.Helper() key, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - t.Fatalf("failed to generate test RSA key: %v", err) - } + require.NoError(t, err) + keyBytes, err := x509.MarshalPKCS8PrivateKey(key) - if err != nil { - t.Fatalf("failed to marshal test RSA key: %v", err) - } + require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{ Type: "PRIVATE KEY", Bytes: keyBytes, @@ -489,41 +402,28 @@ func vertexServiceAccountCredentials(t *testing.T, tokenURL string) string { "token_uri": tokenURL, } encoded, err := json.Marshal(contents) - if err != nil { - t.Fatalf("failed to marshal service account credentials: %v", err) - } + require.NoError(t, err) return string(encoded) } func TestCreateImageDelegatesToGeminiPredict(t *testing.T) { t.Setenv("USE_GOOGLE_GEMINI_NATIVE_API", "true") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1/projects/prod-ai/locations/us-central1/publishers/google/models/imagen-4.0-generate-001:predict" { - t.Errorf("Path = %q, want Vertex Imagen predict endpoint", r.URL.Path) - } - if got := r.Header.Get("Authorization"); got != "Bearer vertex-token" { - t.Errorf("Authorization = %q, want Bearer vertex-token", got) - } - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"predictions": [{"bytesBase64Encoded": "aW1n", "mimeType": "image/png"}]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"predictions": [{"bytesBase64Encoded": "aW1n", "mimeType": "image/png"}]}`) cfg := testConfig() - cfg.BaseURL = server.URL + "/v1/projects/prod-ai/locations/us-central1/publishers/google" + cfg.BaseURL = server.URL + nativeBasePath provider := newProvider(cfg, providers.ProviderOptions{}, authedTestClient(server.Client())) resp, err := provider.CreateImage(context.Background(), &core.ImageGenerationRequest{ Model: "google/imagen-4.0-generate-001", Prompt: "a mountain", }) - if err != nil { - t.Fatalf("CreateImage() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].B64JSON != "aW1n" { - t.Fatalf("data = %+v, want the predicted image", resp.Data) - } - if resp.Provider != "vertex" { - t.Fatalf("Provider = %q, want vertex", resp.Provider) - } + require.NoError(t, err) + + req := capture.Last(t) + assert.Equal(t, nativeBasePath+"/models/imagen-4.0-generate-001:predict", req.Path) + assert.Equal(t, "Bearer vertex-token", req.Header.Get("Authorization")) + require.Len(t, resp.Data, 1) + assert.Equal(t, "aW1n", resp.Data[0].B64JSON) + assert.Equal(t, "vertex", resp.Provider) } diff --git a/internal/providers/vllm/reasoning_test.go b/internal/providers/vllm/reasoning_test.go index aa277c3a7..4d3b87a7c 100644 --- a/internal/providers/vllm/reasoning_test.go +++ b/internal/providers/vllm/reasoning_test.go @@ -4,98 +4,70 @@ import ( "context" "encoding/json" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +// sentMessages returns the messages array of the last recorded request. +func sentMessages(t *testing.T, capture *providertest.Capture) []any { + t.Helper() + messages, ok := capture.Last(t).JSON(t)["messages"].([]any) + require.True(t, ok, "request body has no messages array") + return messages +} + func TestChatCompletion_RenamesLegacyReasoningContentOnAssistantMessages(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-vllm", - "created":1, - "model":"Qwen3.8-27B", - "choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) var req core.ChatRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"Qwen3.8-27B", "messages":[ {"role":"user","content":"what is the median life-expectancy of a cat"}, {"role":"assistant","content":"12-15 years","reasoning_content":"the user wants a quick factual answer"}, {"role":"user","content":"and a dog's?"} ] - }`), &req); err != nil { - t.Fatalf("json.Unmarshal() error = %v", err) - } + }`), &req) + require.NoError(t, err) provider := NewWithHTTPClient("", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + _, err = provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) - messages, _ := gotBody["messages"].([]any) + messages := sentMessages(t, capture) + require.Len(t, messages, 3) assistantMsg, _ := messages[1].(map[string]any) - if assistantMsg["reasoning"] != "the user wants a quick factual answer" { - t.Fatalf("reasoning = %#v, want the replayed reasoning_content value", assistantMsg["reasoning"]) - } - if _, present := assistantMsg["reasoning_content"]; present { - t.Fatal("reasoning_content should be renamed away, not duplicated alongside reasoning") - } - if req.Messages[1].ExtraFields.Lookup("reasoning") != nil { - t.Fatal("ChatCompletion() mutated the caller's request") - } + assert.Equal(t, "the user wants a quick factual answer", assistantMsg["reasoning"]) + assert.NotContains(t, assistantMsg, "reasoning_content") + assert.Nil(t, req.Messages[1].ExtraFields.Lookup("reasoning"), "caller's request must not be mutated") } func TestChatCompletion_DoesNotOverrideExistingReasoningField(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-vllm", - "created":1, - "model":"Qwen3.8-27B", - "choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) var req core.ChatRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"Qwen3.8-27B", "messages":[ {"role":"assistant","content":"12-15 years","reasoning":"current field","reasoning_content":"stale legacy value"} ] - }`), &req); err != nil { - t.Fatalf("json.Unmarshal() error = %v", err) - } + }`), &req) + require.NoError(t, err) provider := NewWithHTTPClient("", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + _, err = provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) - messages, _ := gotBody["messages"].([]any) + messages := sentMessages(t, capture) + require.Len(t, messages, 1) assistantMsg, _ := messages[0].(map[string]any) - if assistantMsg["reasoning"] != "current field" { - t.Fatalf("reasoning = %#v, want the untouched current-field value", assistantMsg["reasoning"]) - } + assert.Equal(t, "current field", assistantMsg["reasoning"]) } func TestAdaptChatRequest_NoOpWithoutLegacyReasoningContent(t *testing.T) { @@ -104,12 +76,8 @@ func TestAdaptChatRequest_NoOpWithoutLegacyReasoningContent(t *testing.T) { } adapted, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if adapted != req { - t.Fatal("adaptChatRequest() copied a request it didn't need to change") - } + require.NoError(t, err) + assert.Same(t, req, adapted) } func TestAdaptChatRequest_IgnoresNonAssistantMessages(t *testing.T) { @@ -121,22 +89,14 @@ func TestAdaptChatRequest_IgnoresNonAssistantMessages(t *testing.T) { }) adapted, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v", err) - } - if adapted != req { - t.Fatal("adaptChatRequest() should not adapt non-assistant messages") - } + require.NoError(t, err) + assert.Same(t, req, adapted) } func TestAdaptChatRequest_NilRequest(t *testing.T) { adapted, err := adaptChatRequest(nil) - if err != nil { - t.Fatalf("adaptChatRequest(nil) error = %v", err) - } - if adapted != nil { - t.Fatal("adaptChatRequest(nil) should return nil") - } + require.NoError(t, err) + assert.Nil(t, adapted) } // TestChatCompletion_AppliesAdaptChatRequestThroughStandardConstructor covers @@ -145,46 +105,26 @@ func TestAdaptChatRequest_NilRequest(t *testing.T) { // provider via NewWithHTTPClient, which would not catch AdaptChatRequest // being wired into one constructor but not the other. func TestChatCompletion_AppliesAdaptChatRequestThroughStandardConstructor(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-vllm", - "created":1, - "model":"Qwen3.8-27B", - "choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) var req core.ChatRequest - if err := json.Unmarshal([]byte(`{ + err := json.Unmarshal([]byte(`{ "model":"Qwen3.8-27B", "messages":[ {"role":"assistant","content":"12-15 years","reasoning_content":"prior turn reasoning"} ] - }`), &req); err != nil { - t.Fatalf("json.Unmarshal() error = %v", err) - } + }`), &req) + require.NoError(t, err) provider, ok := New(providers.ProviderConfig{BaseURL: server.URL}, providers.ProviderOptions{}).(*Provider) - if !ok { - t.Fatal("New() did not return a *Provider") - } + require.True(t, ok) + _, err = provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - - messages, _ := gotBody["messages"].([]any) + messages := sentMessages(t, capture) + require.Len(t, messages, 1) assistantMsg, _ := messages[0].(map[string]any) - if assistantMsg["reasoning"] != "prior turn reasoning" { - t.Fatalf("reasoning = %#v, want AdaptChatRequest applied through New()", assistantMsg["reasoning"]) - } + assert.Equal(t, "prior turn reasoning", assistantMsg["reasoning"]) } // TestAdaptChatRequest_SkipsMalformedReasoningContentWithoutError documents @@ -205,10 +145,6 @@ func TestAdaptChatRequest_SkipsMalformedReasoningContentWithoutError(t *testing. }) adapted, err := adaptChatRequest(req) - if err != nil { - t.Fatalf("adaptChatRequest() error = %v, want nil (malformed value is silently skipped)", err) - } - if adapted != req { - t.Fatal("adaptChatRequest() should not have adapted a message whose reasoning_content failed to decode") - } + require.NoError(t, err) + assert.Same(t, req, adapted) } diff --git a/internal/providers/vllm/vllm_test.go b/internal/providers/vllm/vllm_test.go index 19ea0a9e3..2ad12afcc 100644 --- a/internal/providers/vllm/vllm_test.go +++ b/internal/providers/vllm/vllm_test.go @@ -4,14 +4,18 @@ import ( "context" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +var _ core.PassthroughProvider = (*Provider)(nil) + func TestChatCompletion_UsesOptionalBearerAuthAndChatEndpoint(t *testing.T) { tests := []struct { name string @@ -24,188 +28,98 @@ func TestChatCompletion_UsesOptionalBearerAuthAndChatEndpoint(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-vllm", - "created":1677652288, - "model":"meta-llama/Llama-3.1-8B-Instruct", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) provider := NewWithHTTPClient(tt.apiKey, server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "meta-llama/Llama-3.1-8B-Instruct", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, - }, + Model: providertest.Model, + Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "meta-llama/Llama-3.1-8B-Instruct" { - t.Fatalf("resp.Model = %q, want meta-llama/Llama-3.1-8B-Instruct", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != tt.wantAuth { - t.Fatalf("authorization = %q, want %q", gotAuth, tt.wantAuth) - } + require.NoError(t, err) + assert.Equal(t, providertest.Model, resp.Model) + + req := capture.Last(t) + assert.Equal(t, "/chat/completions", req.Path) + assert.Equal(t, tt.wantAuth, req.Header.Get("Authorization")) }) } } func TestEmbeddings_DelegatesToCompatibleProvider(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "model":"BAAI/bge-small-en-v1.5", - "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], - "usage":{"prompt_tokens":3,"total_tokens":3} - }`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{ + "object":"list", + "model":"BAAI/bge-small-en-v1.5", + "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], + "usage":{"prompt_tokens":3,"total_tokens":3} + }`) provider := NewWithHTTPClient("", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "BAAI/bge-small-en-v1.5", Input: "hello", }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Model != "BAAI/bge-small-en-v1.5" { - t.Fatalf("resp.Model = %q, want BAAI/bge-small-en-v1.5", resp.Model) - } - if gotPath != "/embeddings" { - t.Fatalf("path = %q, want /embeddings", gotPath) - } + require.NoError(t, err) + assert.Equal(t, "BAAI/bge-small-en-v1.5", resp.Model) + assert.Equal(t, "/embeddings", capture.Last(t).Path) } -func TestProvider_ExposesPassthroughButNotOptionalNativeInterfaces(t *testing.T) { +// vLLM serves passthrough and the OpenAI-compatible surface, but nothing +// that would need native batch, file, audio, or response-lifecycle support. +func TestProvider_DoesNotExposeOptionalNativeInterfaces(t *testing.T) { provider := NewWithHTTPClient("", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.PassthroughProvider); !ok { - t.Fatal("vllm provider should implement passthrough provider") - } - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("vllm provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("vllm provider should not implement native file provider") - } - if _, ok := any(provider).(core.NativeResponseLifecycleProvider); ok { - t.Fatal("vllm provider should not implement native response lifecycle provider") - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.NativeResponseLifecycleProvider) + assert.False(t, ok, "provider should not implement core.NativeResponseLifecycleProvider") } func TestPassthrough_ForwardsProviderNativeEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"tokens":[1,2,3]}`)) - })) - defer server.Close() + server, capture := providertest.JSONServer(t, http.StatusOK, `{"tokens":[1,2,3]}`) provider := NewWithHTTPClient("vllm-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ - Method: http.MethodPost, - Endpoint: "tokenize", - Body: io.NopCloser(strings.NewReader("{}")), - Headers: http.Header{"Content-Type": []string{"application/json"}}, - }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } - defer resp.Body.Close() - - if gotPath != "/tokenize" { - t.Fatalf("path = %q, want /tokenize", gotPath) - } - if gotAuth != "Bearer vllm-key" { - t.Fatalf("authorization = %q, want Bearer vllm-key", gotAuth) - } -} - -func TestPassthrough_UsesRootForNativeEndpointsWhenBaseURLIncludesV1(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"tokens":[1,2,3]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "tokenize", Body: io.NopCloser(strings.NewReader("{}")), Headers: http.Header{"Content-Type": []string{"application/json"}}, }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/tokenize" { - t.Fatalf("path = %q, want /tokenize", gotPath) - } + req := capture.Last(t) + assert.Equal(t, "/tokenize", req.Path) + assert.Equal(t, "Bearer vllm-key", req.Header.Get("Authorization")) } -func TestPassthrough_UsesV1ForOpenAICompatibleEndpointsWhenBaseURLIncludesV1(t *testing.T) { - var gotPath string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-vllm", - "created":1677652288, - "model":"Qwen/Qwen2.5-0.5B-Instruct", - "choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) - - resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ - Method: http.MethodPost, - Endpoint: "chat/completions", - Body: io.NopCloser(strings.NewReader(`{ - "model":"Qwen/Qwen2.5-0.5B-Instruct", - "messages":[{"role":"user","content":"hi"}] - }`)), - Headers: http.Header{"Content-Type": []string{"application/json"}}, - }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) +func TestPassthrough_RoutesByEndpointWhenBaseURLIncludesV1(t *testing.T) { + tests := []struct { + name string + endpoint string + body string + wantPath string + }{ + {name: "native endpoint uses the server root", endpoint: "tokenize", body: "{}", wantPath: "/tokenize"}, + { + name: "OpenAI-compatible endpoint keeps /v1", + endpoint: "chat/completions", + body: `{"model":"Qwen/Qwen2.5-0.5B-Instruct","messages":[{"role":"user","content":"hi"}]}`, + wantPath: "/v1/chat/completions", + }, } - defer resp.Body.Close() - if gotPath != "/v1/chat/completions" { - t.Fatalf("path = %q, want /v1/chat/completions", gotPath) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server, capture := providertest.JSONServer(t, http.StatusOK, `{}`) + + provider := NewWithHTTPClient("", server.URL+"/v1", server.Client(), llmclient.Hooks{}) + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + Method: http.MethodPost, + Endpoint: tt.endpoint, + Body: io.NopCloser(strings.NewReader(tt.body)), + Headers: http.Header{"Content-Type": []string{"application/json"}}, + }) + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, tt.wantPath, capture.Last(t).Path) + }) } } diff --git a/internal/providers/xai/images_test.go b/internal/providers/xai/images_test.go index e42a0b488..f8f68b586 100644 --- a/internal/providers/xai/images_test.go +++ b/internal/providers/xai/images_test.go @@ -2,15 +2,13 @@ package xai import ( "context" - "encoding/json" - "io" "net/http" - "net/http/httptest" - "strings" "testing" "github.com/enterpilot/gomodel/internal/core" - "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // TestCreateImage verifies the xAI provider advertises image generation and @@ -36,9 +34,9 @@ func TestCreateImage(t *testing.T) { body: `{"created":1713833628,"data":[{"url":"https://imgen.x.ai/xai-imgen/img.jpg"},{"url":"https://imgen.x.ai/xai-imgen/img2.jpg"}]}`, wantBody: map[string]any{"model": "grok-imagine-image", "prompt": "A cat", "n": float64(2), "response_format": "url"}, check: func(t *testing.T, resp *core.ImageGenerationResponse) { - if resp.Created != 1713833628 || len(resp.Data) != 2 || resp.Data[0].URL != "https://imgen.x.ai/xai-imgen/img.jpg" { - t.Errorf("response = %+v", resp) - } + assert.Equal(t, int64(1713833628), resp.Created) + require.Len(t, resp.Data, 2) + assert.Equal(t, "https://imgen.x.ai/xai-imgen/img.jpg", resp.Data[0].URL) }, }, { @@ -47,12 +45,8 @@ func TestCreateImage(t *testing.T) { statusCode: http.StatusOK, body: `{}`, check: func(t *testing.T, resp *core.ImageGenerationResponse) { - if resp.Created == 0 { - t.Error("Created should default to now when upstream omits it") - } - if resp.Data == nil { - t.Error("Data should be an empty array, not null") - } + assert.NotEqual(t, int64(0), resp.Created) + assert.NotNil(t, resp.Data) }, }, { @@ -66,45 +60,27 @@ func TestCreateImage(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotMethod, gotPath string - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotMethod, gotPath = r.Method, r.URL.Path - if !strings.HasPrefix(r.Header.Get("Authorization"), "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - raw, _ := io.ReadAll(r.Body) - _ = json.Unmarshal(raw, &gotBody) - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.body)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, tt.statusCode, tt.body) + provider := newTestProvider(server.URL) var imager core.ImageProvider = provider resp, err := imager.CreateImage(context.Background(), tt.req) - if gotMethod != http.MethodPost || gotPath != "/images/generations" { - t.Errorf("request = %s %s, want POST /images/generations", gotMethod, gotPath) - } + req := capture.Last(t) + assert.Equal(t, http.MethodPost, req.Method) + assert.Equal(t, "/images/generations", req.Path) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + sent := req.JSON(t) for key, want := range tt.wantBody { - if got := gotBody[key]; got != want { - t.Errorf("forwarded[%q] = %v, want %v", key, got, want) - } + assert.Equal(t, want, sent[key], "forwarded field %q", key) } if tt.wantErr != "" { - if err == nil || !strings.Contains(err.Error(), tt.wantErr) { - t.Fatalf("error = %v, want message containing %q", err, tt.wantErr) - } + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantErr) return } - if err != nil { - t.Fatalf("CreateImage() error = %v", err) - } + require.NoError(t, err) tt.check(t, resp) }) } diff --git a/internal/providers/xai/realtime_test.go b/internal/providers/xai/realtime_test.go index 4278d4804..11ff856c5 100644 --- a/internal/providers/xai/realtime_test.go +++ b/internal/providers/xai/realtime_test.go @@ -8,6 +8,8 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestRealtimeTarget(t *testing.T) { @@ -15,49 +17,33 @@ func TestRealtimeTarget(t *testing.T) { p := New(providers.ProviderConfig{APIKey: apiKey}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "grok-voice-latest"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://api.x.ai/v1/realtime?") { - t.Errorf("url = %q, want xAI realtime endpoint", target.URL) - } - parsed, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("parse target url: %v", err) - } - if got := parsed.Query().Get("model"); got != "grok-voice-latest" { - t.Errorf("model query = %q, want %q", got, "grok-voice-latest") - } - if got := target.Headers.Get("Authorization"); got != "Bearer "+apiKey { - t.Errorf("Authorization = %q, want bearer with key", got) - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://api.x.ai/v1/realtime?"), "url = %q, want xAI realtime endpoint", target.URL) - if _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + parsed, err := url.Parse(target.URL) + require.NoError(t, err) + got := parsed.Query().Get("model") + assert.Equal(t, "grok-voice-latest", got) + got = target.Headers.Get("Authorization") + assert.Equal(t, "Bearer "+apiKey, got) + _, err = p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } func TestRealtimeTargetOmitsAuthWhenNoKey(t *testing.T) { p := New(providers.ProviderConfig{APIKey: ""}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if _, present := target.Headers["Authorization"]; present { - t.Error("Authorization header should be absent when no API key is configured") - } + require.NoError(t, err) + _, present := target.Headers["Authorization"] + assert.False(t, present) } func TestRealtimeTargetFollowsSetBaseURL(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) p.SetBaseURL("https://custom.x.example/v1") target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://custom.x.example/v1/realtime") { - t.Errorf("url = %q, want the SetBaseURL host", target.URL) - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://custom.x.example/v1/realtime"), "url = %q, want the SetBaseURL host", target.URL) } func TestRealtimeCallTarget(t *testing.T) { @@ -65,44 +51,27 @@ func TestRealtimeCallTarget(t *testing.T) { p := New(providers.ProviderConfig{APIKey: apiKey}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: "grok-voice-latest"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if target.URL != "https://api.x.ai/v1/realtime/calls" { - t.Errorf("url = %q, want the xAI realtime calls endpoint", target.URL) - } - if got := target.Headers.Get("Authorization"); got != "Bearer "+apiKey { - t.Errorf("Authorization = %q, want bearer with key", got) - } - - if _, err := p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + require.NoError(t, err) + assert.Equal(t, "https://api.x.ai/v1/realtime/calls", target.URL) + got := target.Headers.Get("Authorization") + assert.Equal(t, "Bearer "+apiKey, got) + _, err = p.RealtimeCallTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } func TestRealtimeClientSecretTarget(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeClientSecretTarget(context.Background(), &core.RealtimeRequest{Model: "grok-voice-latest"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if target.URL != "https://api.x.ai/v1/realtime/client_secrets" { - t.Errorf("url = %q, want the xAI client secrets endpoint", target.URL) - } + require.NoError(t, err) + assert.Equal(t, "https://api.x.ai/v1/realtime/client_secrets", target.URL) } func TestRealtimeTargetAttachesByCallID(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "grok-voice-latest", CallID: "rtc_9"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.Contains(target.URL, "call_id=rtc_9") { - t.Errorf("url = %q, want call_id attach query", target.URL) - } - if strings.Contains(target.URL, "model=") { - t.Errorf("url = %q, want no model query on sideband attach", target.URL) - } + require.NoError(t, err) + assert.Contains(t, target.URL, "call_id=rtc_9") + assert.NotContains(t, target.URL, "model=", "sideband attach must not carry a model query") } diff --git a/internal/providers/xai/xai_test.go b/internal/providers/xai/xai_test.go index 0b90c8f71..29c1ab887 100644 --- a/internal/providers/xai/xai_test.go +++ b/internal/providers/xai/xai_test.go @@ -2,37 +2,42 @@ package xai import ( "context" - "encoding/json" "io" "net/http" - "net/http/httptest" "strings" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" - "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func TestNew(t *testing.T) { - apiKey := "test-api-key" - // Use NewWithHTTPClient to get concrete type for internal testing - provider := NewWithHTTPClient(apiKey, nil, llmclient.Hooks{}) +const testAPIKey = "test-api-key" - if got := provider.keys.Primary(); got != apiKey { - t.Errorf("primary key = %q, want %q", got, apiKey) - } - if provider.compat == nil { - t.Error("compat should not be nil") - } +// newTestProvider builds a provider pointed at baseURL with the default client. +func newTestProvider(baseURL string) *Provider { + provider := NewWithHTTPClient(testAPIKey, nil, llmclient.Hooks{}) + provider.SetBaseURL(baseURL) + return provider } -func TestNew_ReturnsProvider(t *testing.T) { - provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) +// statusServer answers every request with status and body, so a table can +// drive success and error cases through the same recorder. +func statusServer(t *testing.T, status int, body string) (string, *providertest.Capture) { + t.Helper() + server, capture := providertest.Server(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(status) + _, _ = io.WriteString(w, body) + }) + return server.URL, capture +} - if provider == nil { - t.Error("provider should not be nil") - } +// blockUntilCancelled never answers, so a cancelled context is the only way out. +func blockUntilCancelled(w http.ResponseWriter, r *http.Request) { + <-r.Context().Done() + w.WriteHeader(http.StatusRequestTimeout) } // customHeaderRoundTripper is a RoundTripper that injects a custom header @@ -47,33 +52,11 @@ func (c *customHeaderRoundTripper) RoundTrip(req *http.Request) (*http.Response, return c.transport.RoundTrip(req) } -func TestNewWithHTTPClient(t *testing.T) { +func TestNewWithHTTPClient_UsesCustomClient(t *testing.T) { const customHeaderKey = "X-Custom-Test-Header" const customHeaderVal = "custom-test-value" - var receivedHeader string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Capture the custom header to verify custom client was used - receivedHeader = r.Header.Get(customHeaderKey) - - // Return a valid chat completion response - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "grok-2", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "Hello!"}, - "finish_reason": "stop" - }], - "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} - }`)) - })) - defer server.Close() - - // Create a custom HTTP client with a RoundTripper that injects a header + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) customClient := &http.Client{ Transport: &customHeaderRoundTripper{ transport: http.DefaultTransport, @@ -81,74 +64,22 @@ func TestNewWithHTTPClient(t *testing.T) { headerVal: customHeaderVal, }, } - - // Create provider with custom HTTP client - provider := NewWithHTTPClient("test-api-key", customClient, llmclient.Hooks{}) - - // Verify provider is non-nil - if provider == nil { - t.Fatal("provider should not be nil") - return - } - if provider.compat == nil { - t.Fatal("provider.compat should not be nil") - } - if got := provider.keys.Primary(); got != "test-api-key" { - t.Errorf("primary key = %q, want %q", got, "test-api-key") - } - - // Set base URL to our test server + provider := NewWithHTTPClient(testAPIKey, customClient, llmclient.Hooks{}) + assert.Equal(t, testAPIKey, provider.keys.Primary()) provider.SetBaseURL(server.URL) - // Make a request to verify custom client is wired correctly - req := &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - resp, err := provider.ChatCompletion(context.Background(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - // Verify response is valid (server was hit) - if resp == nil { - t.Fatal("response should not be nil") - } - if resp.ID != "chatcmpl-123" { - t.Errorf("response ID = %q, want %q", resp.ID, "chatcmpl-123") - } - - // Verify the custom header was injected by our custom RoundTripper - if receivedHeader != customHeaderVal { - t.Errorf("custom header = %q, want %q (custom HTTP client not wired correctly)", receivedHeader, customHeaderVal) - } + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) + require.NoError(t, err) + assert.Equal(t, "chatcmpl-test", resp.ID) + assert.Equal(t, customHeaderVal, capture.Last(t).Header.Get(customHeaderKey)) } func TestChatCompletion_ForwardsXGrokConvIDFromSnapshot(t *testing.T) { - var receivedConvID string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - receivedConvID = r.Header.Get(grokConvIDHeader) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "grok-2", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "Hello!"}, - "finish_reason": "stop" - }], - "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) ctx := core.WithRequestSnapshot(context.Background(), core.NewRequestSnapshot( http.MethodPost, @@ -163,50 +94,22 @@ func TestChatCompletion_ForwardsXGrokConvIDFromSnapshot(t *testing.T) { nil, )) _, err := provider.ChatCompletion(ctx, &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if receivedConvID != "client-conv-123" { - t.Fatalf("%s = %q, want client-conv-123", grokConvIDHeader, receivedConvID) - } + require.NoError(t, err) + assert.Equal(t, "client-conv-123", capture.Last(t).Header.Get(grokConvIDHeader)) } func TestChatCompletion_GeneratesStableXGrokConvIDWhenMissing(t *testing.T) { - var receivedConvIDs []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - receivedConvIDs = append(receivedConvIDs, r.Header.Get(grokConvIDHeader)) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "grok-2", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "Hello!"}, - "finish_reason": "stop" - }], - "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) initialMessages := []core.Message{ {Role: "system", Content: "Reply with the requested marker only."}, {Role: "user", Content: "Use this fixed reference text for the cache anchor."}, } - req1 := &core.ChatRequest{ - Model: "grok-2", - Messages: initialMessages, - } + req1 := &core.ChatRequest{Model: "grok-2", Messages: initialMessages} req2 := &core.ChatRequest{ Model: "grok-2", Messages: append(append([]core.Message{}, initialMessages...), @@ -214,38 +117,21 @@ func TestChatCompletion_GeneratesStableXGrokConvIDWhenMissing(t *testing.T) { core.Message{Role: "user", Content: "Now reply with marker two."}, ), } + _, err := provider.ChatCompletion(context.Background(), req1) + require.NoError(t, err) + _, err = provider.ChatCompletion(context.Background(), req2) + require.NoError(t, err) - if _, err := provider.ChatCompletion(context.Background(), req1); err != nil { - t.Fatalf("first ChatCompletion() error = %v", err) - } - if _, err := provider.ChatCompletion(context.Background(), req2); err != nil { - t.Fatalf("second ChatCompletion() error = %v", err) - } - if len(receivedConvIDs) != 2 { - t.Fatalf("received %d requests, want 2", len(receivedConvIDs)) - } - if receivedConvIDs[0] == "" { - t.Fatal("first generated x-grok-conv-id is empty") - } - if !strings.HasPrefix(receivedConvIDs[0], "gomodel-") { - t.Fatalf("generated x-grok-conv-id = %q, want gomodel-*", receivedConvIDs[0]) - } - if receivedConvIDs[1] != receivedConvIDs[0] { - t.Fatalf("generated x-grok-conv-id changed across appended conversation: %q then %q", receivedConvIDs[0], receivedConvIDs[1]) - } + requests := capture.All() + require.Len(t, requests, 2) + first := requests[0].Header.Get(grokConvIDHeader) + assert.True(t, strings.HasPrefix(first, "gomodel-"), "generated x-grok-conv-id = %q, want gomodel-*", first) + assert.Equal(t, first, requests[1].Header.Get(grokConvIDHeader)) } func TestStreamChatCompletion_ForwardsXGrokConvIDFromSnapshot(t *testing.T) { - var receivedConvID string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - receivedConvID = r.Header.Get(grokConvIDHeader) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", server.Client(), llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.SSEServer(t, "data: [DONE]\n\n") + provider := newTestProvider(server.URL) ctx := core.WithRequestSnapshot(context.Background(), core.NewRequestSnapshot( http.MethodPost, @@ -260,18 +146,12 @@ func TestStreamChatCompletion_ForwardsXGrokConvIDFromSnapshot(t *testing.T) { nil, )) body, err := provider.StreamChatCompletion(ctx, &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer func() { _ = body.Close() }() - if receivedConvID != "stream-conv-123" { - t.Fatalf("%s = %q, want stream-conv-123", grokConvIDHeader, receivedConvID) - } + assert.Equal(t, "stream-conv-123", capture.Last(t).Header.Get(grokConvIDHeader)) } func TestChatCompletion(t *testing.T) { @@ -304,29 +184,14 @@ func TestChatCompletion(t *testing.T) { "total_tokens": 30 } }`, - expectedError: false, checkResponse: func(t *testing.T, resp *core.ChatResponse) { - if resp.ID != "chatcmpl-123" { - t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") - } - if resp.Model != "grok-2" { - t.Errorf("Model = %q, want %q", resp.Model, "grok-2") - } - if len(resp.Choices) != 1 { - t.Fatalf("len(Choices) = %d, want 1", len(resp.Choices)) - } - if resp.Choices[0].Message.Content != "Hello! How can I help you today?" { - t.Errorf("Message content = %q, want %q", resp.Choices[0].Message.Content, "Hello! How can I help you today?") - } - if resp.Usage.PromptTokens != 10 { - t.Errorf("PromptTokens = %d, want 10", resp.Usage.PromptTokens) - } - if resp.Usage.CompletionTokens != 20 { - t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens) - } - if resp.Usage.TotalTokens != 30 { - t.Errorf("TotalTokens = %d, want 30", resp.Usage.TotalTokens) - } + assert.Equal(t, "chatcmpl-123", resp.ID) + assert.Equal(t, "grok-2", resp.Model) + require.Len(t, resp.Choices, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Choices[0].Message.Content) + assert.Equal(t, 10, resp.Usage.PromptTokens) + assert.Equal(t, 20, resp.Usage.CompletionTokens) + assert.Equal(t, 30, resp.Usage.TotalTokens) }, }, { @@ -351,55 +216,25 @@ func TestChatCompletion(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - // Verify request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } + baseURL, capture := statusServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(baseURL) - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - resp, err := provider.ChatCompletion(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + assert.Equal(t, "grok-2", req.JSON(t)["model"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + assert.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } @@ -416,24 +251,8 @@ func TestResponsesDropsMetadata(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var body map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - raw, err := io.ReadAll(r.Body) - if err != nil { - t.Errorf("read body: %v", err) - return - } - if err := json.Unmarshal(raw, &body); err != nil { - t.Errorf("unmarshal body: %v", err) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"resp_1","object":"response","status":"completed"}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, `{"id":"resp_1","object":"response","status":"completed"}`) + provider := newTestProvider(server.URL) req := &core.ResponsesRequest{Model: "grok-4.3", Metadata: tt.metadata} var err error @@ -446,14 +265,10 @@ func TestResponsesDropsMetadata(t *testing.T) { } else { _, err = provider.Responses(context.Background(), req) } - if err != nil { - t.Fatalf("responses: %v", err) - } - if _, ok := body["metadata"]; ok { - t.Errorf("outbound body carries metadata: %v", body["metadata"]) - } - if len(tt.metadata) > 0 && req.Metadata == nil { - t.Error("caller request was mutated; metadata must survive for the client echo") + require.NoError(t, err) + assert.NotContains(t, capture.Last(t).JSON(t), "metadata") + if len(tt.metadata) > 0 { + assert.NotNil(t, req.Metadata, "caller request was mutated; metadata must survive for the client echo") } }) } @@ -475,7 +290,6 @@ data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288 data: [DONE] `, - expectedError: false, }, { name: "API error", @@ -487,68 +301,30 @@ data: [DONE] for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } + baseURL, capture := statusServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(baseURL) - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ChatRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } + body, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) - body, err := provider.StreamChatCompletion(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + assert.Equal(t, true, req.JSON(t)["stream"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if body == nil { - t.Fatal("body should not be nil") - } - defer func() { _ = body.Close() }() - - // Read and verify the streaming response - respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } - if string(respBody) != tt.responseBody { - t.Errorf("response body = %q, want %q", string(respBody), tt.responseBody) - } + assert.Error(t, err) + return } + require.NoError(t, err) + require.NotNil(t, body) + defer func() { _ = body.Close() }() + + respBody, err := io.ReadAll(body) + require.NoError(t, err) + assert.Equal(t, tt.responseBody, string(respBody)) }) } } @@ -581,20 +357,11 @@ func TestListModels(t *testing.T) { } ] }`, - expectedError: false, checkResponse: func(t *testing.T, resp *core.ModelsResponse) { - if resp.Object != "list" { - t.Errorf("Object = %q, want %q", resp.Object, "list") - } - if len(resp.Data) != 2 { - t.Fatalf("len(Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].ID != "grok-2" { - t.Errorf("Data[0].ID = %q, want %q", resp.Data[0].ID, "grok-2") - } - if resp.Data[0].OwnedBy != "xai" { - t.Errorf("Data[0].OwnedBy = %q, want %q", resp.Data[0].OwnedBy, "xai") - } + assert.Equal(t, "list", resp.Object) + require.Len(t, resp.Data, 2) + assert.Equal(t, "grok-2", resp.Data[0].ID) + assert.Equal(t, "xai", resp.Data[0].OwnedBy) }, }, { @@ -607,72 +374,38 @@ func TestListModels(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request method and path - if r.Method != http.MethodGet { - t.Errorf("Method = %q, want %q", r.Method, http.MethodGet) - } - if r.URL.Path != "/models" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/models") - } - - // Verify authorization header - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + baseURL, capture := statusServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(baseURL) resp, err := provider.ListModels(context.Background()) + req := capture.Last(t) + assert.Equal(t, http.MethodGet, req.Method) + assert.Equal(t, "/models", req.Path) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + assert.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } func TestChatCompletionWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Simulate a slow response - <-r.Context().Done() - w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.Server(t, blockUntilCancelled) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) cancel() // Cancel immediately - req := &core.ChatRequest{ - Model: "grok-2", - Messages: []core.Message{ - {Role: "user", Content: "Hello"}, - }, - } - - _, err := provider.ChatCompletion(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + _, err := provider.ChatCompletion(ctx, &core.ChatRequest{ + Model: "grok-2", + Messages: []core.Message{{Role: "user", Content: "Hello"}}, + }) + assert.Error(t, err) } func TestResponses(t *testing.T) { @@ -708,41 +441,18 @@ func TestResponses(t *testing.T) { "total_tokens": 30 } }`, - expectedError: false, checkResponse: func(t *testing.T, resp *core.ResponsesResponse) { - if resp.ID != "resp_123" { - t.Errorf("ID = %q, want %q", resp.ID, "resp_123") - } - if resp.Object != "response" { - t.Errorf("Object = %q, want %q", resp.Object, "response") - } - if resp.Model != "grok-2" { - t.Errorf("Model = %q, want %q", resp.Model, "grok-2") - } - if resp.Status != "completed" { - t.Errorf("Status = %q, want %q", resp.Status, "completed") - } - if len(resp.Output) != 1 { - t.Fatalf("len(Output) = %d, want 1", len(resp.Output)) - } - if len(resp.Output[0].Content) != 1 { - t.Fatalf("len(Output[0].Content) = %d, want 1", len(resp.Output[0].Content)) - } - if resp.Output[0].Content[0].Text != "Hello! How can I help you today?" { - t.Errorf("Output text = %q, want %q", resp.Output[0].Content[0].Text, "Hello! How can I help you today?") - } - if resp.Usage == nil { - t.Fatal("Usage should not be nil") - } - if resp.Usage.InputTokens != 10 { - t.Errorf("InputTokens = %d, want 10", resp.Usage.InputTokens) - } - if resp.Usage.OutputTokens != 20 { - t.Errorf("OutputTokens = %d, want 20", resp.Usage.OutputTokens) - } - if resp.Usage.TotalTokens != 30 { - t.Errorf("TotalTokens = %d, want 30", resp.Usage.TotalTokens) - } + assert.Equal(t, "resp_123", resp.ID) + assert.Equal(t, "response", resp.Object) + assert.Equal(t, "grok-2", resp.Model) + assert.Equal(t, "completed", resp.Status) + require.Len(t, resp.Output, 1) + require.Len(t, resp.Output[0].Content, 1) + assert.Equal(t, "Hello! How can I help you today?", resp.Output[0].Content[0].Text) + require.NotNil(t, resp.Usage) + assert.Equal(t, 10, resp.Usage.InputTokens) + assert.Equal(t, 20, resp.Usage.OutputTokens) + assert.Equal(t, 30, resp.Usage.TotalTokens) }, }, { @@ -767,58 +477,26 @@ func TestResponses(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - // Verify request path - if r.URL.Path != "/responses" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/responses") - } - - // Verify request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ResponsesRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + baseURL, capture := statusServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(baseURL) - req := &core.ResponsesRequest{ + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: "grok-2", Input: "Hello", - } + }) - resp, err := provider.Responses(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "/responses", req.Path) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + assert.Equal(t, "grok-2", req.JSON(t)["model"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkResponse != nil { - tt.checkResponse(t, resp) - } + assert.Error(t, err) + return } + require.NoError(t, err) + tt.checkResponse(t, resp) }) } } @@ -846,29 +524,16 @@ data: {"type":"response.output_text.delta","delta":"!"} event: response.completed data: {"type":"response.completed","response":{"id":"resp_123","object":"response","status":"completed","model":"grok-2"}} `, - expectedError: false, checkStream: func(t *testing.T, body io.ReadCloser) { - if body == nil { - t.Fatal("body should not be nil") - } + require.NotNil(t, body) defer func() { _ = body.Close() }() - // Read and verify the streaming response respBody, err := io.ReadAll(body) - if err != nil { - t.Fatalf("failed to read response body: %v", err) - } - + require.NoError(t, err) responseStr := string(respBody) - if !strings.Contains(responseStr, "response.created") { - t.Error("response should contain response.created event") - } - if !strings.Contains(responseStr, "response.output_text.delta") { - t.Error("response should contain response.output_text.delta event") - } - if !strings.Contains(responseStr, "[DONE]") { - t.Error("response should end with [DONE]") - } + assert.Contains(t, responseStr, "response.created") + assert.Contains(t, responseStr, "response.output_text.delta") + assert.Contains(t, responseStr, "[DONE]") }, }, { @@ -887,161 +552,70 @@ data: {"type":"response.completed","response":{"id":"resp_123","object":"respons for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Verify request headers - if r.Header.Get("Content-Type") != "application/json" { - t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") - } - authHeader := r.Header.Get("Authorization") - if !strings.HasPrefix(authHeader, "Bearer ") { - t.Errorf("Authorization header should start with 'Bearer '") - } - - // Verify request path - if r.URL.Path != "/responses" { - t.Errorf("Path = %q, want %q", r.URL.Path, "/responses") - } - - // Verify stream is set in request body - body, err := io.ReadAll(r.Body) - if err != nil { - t.Fatalf("failed to read request body: %v", err) - } - var req core.ResponsesRequest - if err := json.Unmarshal(body, &req); err != nil { - t.Fatalf("failed to unmarshal request: %v", err) - } - if !req.Stream { - t.Error("Stream should be true in request") - } - - w.WriteHeader(tt.statusCode) - _, _ = w.Write([]byte(tt.responseBody)) - })) - defer server.Close() + baseURL, capture := statusServer(t, tt.statusCode, tt.responseBody) + provider := newTestProvider(baseURL) - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) - - req := &core.ResponsesRequest{ + body, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: "grok-2", Input: "Hello", - } + }) - body, err := provider.StreamResponses(context.Background(), req) + req := capture.Last(t) + assert.Equal(t, "/responses", req.Path) + assert.Equal(t, "application/json", req.Header.Get("Content-Type")) + assert.Equal(t, "Bearer "+testAPIKey, req.Header.Get("Authorization")) + assert.Equal(t, true, req.JSON(t)["stream"]) if tt.expectedError { - if err == nil { - t.Error("expected error, got nil") - } - } else { - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if tt.checkStream != nil { - tt.checkStream(t, body) - } + assert.Error(t, err) + return } + require.NoError(t, err) + tt.checkStream(t, body) }) } } func TestResponsesWithContext(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Simulate a slow response - <-r.Context().Done() - w.WriteHeader(http.StatusRequestTimeout) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, _ := providertest.Server(t, blockUntilCancelled) + provider := newTestProvider(server.URL) ctx, cancel := context.WithCancel(context.Background()) cancel() // Cancel immediately - req := &core.ResponsesRequest{ - Model: "grok-2", - Input: "Hello", - } - - _, err := provider.Responses(ctx, req) - if err == nil { - t.Error("expected error when context is cancelled, got nil") - } + _, err := provider.Responses(ctx, &core.ResponsesRequest{Model: "grok-2", Input: "Hello"}) + assert.Error(t, err) } func TestChatCompletion_MapsReasoningToXAIReasoningEffort(t *testing.T) { - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-xai", - "created":1677652288, - "model":"grok-4.5", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "grok-4.5", Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: "medium"}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotBody["reasoning"] != nil { - t.Fatalf("request body should not include nested reasoning, got %#v", gotBody["reasoning"]) - } - if gotBody["reasoning_effort"] != "medium" { - t.Fatalf("reasoning_effort = %#v, want medium", gotBody["reasoning_effort"]) - } + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning") + assert.Equal(t, "medium", sent["reasoning_effort"]) } func TestChatCompletion_OmitsReasoningEffortWhenReasoningAbsent(t *testing.T) { - var gotBody map[string]any - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-xai", - "created":1677652288, - "model":"grok-4.5", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "grok-4.5", Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := gotBody["reasoning_effort"]; ok { - t.Fatalf("reasoning_effort should be absent, got %#v", gotBody["reasoning_effort"]) - } - if _, ok := gotBody["reasoning"]; ok { - t.Fatalf("reasoning should be absent, got %#v", gotBody["reasoning"]) - } + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning_effort") + assert.NotContains(t, sent, "reasoning") } func TestNormalizeReasoningEffort(t *testing.T) { @@ -1073,9 +647,7 @@ func TestNormalizeReasoningEffort(t *testing.T) { {"grok-4.5", "custom-level", "custom-level"}, } for _, tt := range tests { - if got := normalizeReasoningEffort(tt.model, tt.effort); got != tt.want { - t.Errorf("normalizeReasoningEffort(%q, %q) = %q, want %q", tt.model, tt.effort, got, tt.want) - } + assert.Equal(t, tt.want, normalizeReasoningEffort(tt.model, tt.effort), "normalizeReasoningEffort(%q, %q)", tt.model, tt.effort) } } @@ -1098,41 +670,23 @@ func TestChatCompletion_DropsReasoningEffortForModelsThatRejectIt(t *testing.T) } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"c1","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: tt.model, Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: tt.effort}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := gotBody["reasoning"]; ok { - t.Errorf("request body includes nested reasoning: %#v", gotBody["reasoning"]) - } - got, ok := gotBody["reasoning_effort"] + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning") if tt.wantEffort == "" { - if ok { - t.Errorf("reasoning_effort = %#v, want absent", got) - } + assert.NotContains(t, sent, "reasoning_effort") return } - if got != tt.wantEffort { - t.Errorf("reasoning_effort = %#v, want %q", got, tt.wantEffort) - } + assert.Equal(t, tt.wantEffort, sent["reasoning_effort"]) }) } } @@ -1140,34 +694,19 @@ func TestChatCompletion_DropsReasoningEffortForModelsThatRejectIt(t *testing.T) func TestChatCompletion_DropsEmptyReasoningObject(t *testing.T) { for _, effort := range []string{"", " \t "} { t.Run("effort="+effort, func(t *testing.T) { - var gotBody map[string]any - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { - http.Error(w, "decode error", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"c1","model":"m","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("test-api-key", nil, llmclient.Hooks{}) - provider.SetBaseURL(server.URL) + server, capture := providertest.JSONServer(t, http.StatusOK, providertest.ChatCompletionJSON) + provider := newTestProvider(server.URL) _, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: "grok-4.5", Messages: []core.Message{{Role: "user", Content: "hi"}}, Reasoning: &core.Reasoning{Effort: effort}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if _, ok := gotBody["reasoning"]; ok { - t.Errorf("reasoning should be absent, got %#v", gotBody["reasoning"]) - } - if _, ok := gotBody["reasoning_effort"]; ok { - t.Errorf("reasoning_effort should be absent, got %#v", gotBody["reasoning_effort"]) - } + require.NoError(t, err) + + sent := capture.Last(t).JSON(t) + assert.NotContains(t, sent, "reasoning") + assert.NotContains(t, sent, "reasoning_effort") }) } } diff --git a/internal/providers/xiaomi/audio_test.go b/internal/providers/xiaomi/audio_test.go index b1dd3fd7b..475919c2c 100644 --- a/internal/providers/xiaomi/audio_test.go +++ b/internal/providers/xiaomi/audio_test.go @@ -4,7 +4,6 @@ import ( "context" "encoding/base64" "encoding/json" - "io" "net/http" "net/http/httptest" "strings" @@ -12,33 +11,32 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func newTTSServer(t *testing.T, audioBase64 string) (*httptest.Server, *[]byte) { +func newTTSServer(t *testing.T, audioBase64 string) (*httptest.Server, *providertest.Capture) { t.Helper() - var gotBody []byte - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var err error - gotBody, err = io.ReadAll(r.Body) - if err != nil { - http.Error(w, "read error", http.StatusInternalServerError) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-tts","created":1677652288,"model":"mimo-v2.5-tts", - "choices":[{"index":0,"message":{"role":"assistant","content":"","audio":{"id":"a1","data":"` + audioBase64 + `","format":"wav"}},"finish_reason":"stop"}], - "usage":{"prompt_tokens":10,"completion_tokens":50,"total_tokens":60} - }`)) - })) - return server, &gotBody + return providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-tts","created":1677652288,"model":"mimo-v2.5-tts", + "choices":[{"index":0,"message":{"role":"assistant","content":"","audio":{"id":"a1","data":"`+audioBase64+`","format":"wav"}},"finish_reason":"stop"}], + "usage":{"prompt_tokens":10,"completion_tokens":50,"total_tokens":60} + }`) +} + +func newASRServer(t *testing.T, text string) (*httptest.Server, *providertest.Capture) { + t.Helper() + return providertest.JSONServer(t, http.StatusOK, `{ + "id":"chatcmpl-asr","created":1677652288,"model":"mimo-v2.5-asr", + "choices":[{"index":0,"message":{"role":"assistant","content":"`+text+`"},"finish_reason":"stop"}], + "usage":{"prompt_tokens":20,"completion_tokens":3,"total_tokens":23} + }`) } func TestCreateSpeech_TranslatesToMiMoChatTTS(t *testing.T) { wavBytes := []byte("RIFF-fake-wav") - server, gotBody := newTTSServer(t, base64.StdEncoding.EncodeToString(wavBytes)) - defer server.Close() - + server, capture := newTTSServer(t, base64.StdEncoding.EncodeToString(wavBytes)) provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ @@ -47,15 +45,9 @@ func TestCreateSpeech_TranslatesToMiMoChatTTS(t *testing.T) { Voice: "Chloe", Instructions: "Bright bouncy tone", }) - if err != nil { - t.Fatalf("CreateSpeech() error = %v", err) - } - if resp.ContentType != "audio/wav" { - t.Fatalf("ContentType = %q, want audio/wav", resp.ContentType) - } - if string(resp.Data) != string(wavBytes) { - t.Fatalf("Data = %q, want decoded wav bytes", resp.Data) - } + require.NoError(t, err) + assert.Equal(t, "audio/wav", resp.ContentType) + assert.Equal(t, string(wavBytes), string(resp.Data)) var sent struct { Model string `json:"model"` @@ -65,49 +57,32 @@ func TestCreateSpeech_TranslatesToMiMoChatTTS(t *testing.T) { } `json:"messages"` Audio map[string]string `json:"audio"` } - if err := json.Unmarshal(*gotBody, &sent); err != nil { - t.Fatalf("failed to decode upstream body: %v", err) - } - if sent.Model != "mimo-v2.5-tts" { - t.Fatalf("model = %q, want mimo-v2.5-tts", sent.Model) - } - if len(sent.Messages) != 2 || sent.Messages[0].Role != "user" || sent.Messages[1].Role != "assistant" { - t.Fatalf("messages = %+v, want user instructions + assistant text", sent.Messages) - } - if sent.Messages[1].Content != "Hello world" { - t.Fatalf("assistant content = %q, want synthesis text", sent.Messages[1].Content) - } - if sent.Audio["format"] != "wav" || sent.Audio["voice"] != "Chloe" { - t.Fatalf("audio = %+v, want format=wav voice=Chloe", sent.Audio) - } + require.NoError(t, json.Unmarshal(capture.Last(t).Body, &sent)) + assert.Equal(t, "mimo-v2.5-tts", sent.Model) + require.Len(t, sent.Messages, 2) + assert.Equal(t, "user", sent.Messages[0].Role) + assert.Equal(t, "assistant", sent.Messages[1].Role) + assert.Equal(t, "Hello world", sent.Messages[1].Content) + assert.Equal(t, "wav", sent.Audio["format"]) + assert.Equal(t, "Chloe", sent.Audio["voice"]) } func TestCreateSpeech_MapsPCMAndRejectsUnsupportedFormats(t *testing.T) { - server, gotBody := newTTSServer(t, base64.StdEncoding.EncodeToString([]byte("pcm"))) - defer server.Close() - + server, capture := newTTSServer(t, base64.StdEncoding.EncodeToString([]byte("pcm"))) provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "mimo-v2.5-tts", Input: "hi", ResponseFormat: "pcm", }) - if err != nil { - t.Fatalf("CreateSpeech(pcm) error = %v", err) - } - if resp.ContentType != "audio/pcm" { - t.Fatalf("ContentType = %q, want audio/pcm", resp.ContentType) - } - if !strings.Contains(string(*gotBody), `"format":"pcm16"`) { - t.Fatalf("upstream body should request pcm16, got: %s", *gotBody) - } + require.NoError(t, err) + assert.Equal(t, "audio/pcm", resp.ContentType) + assert.Contains(t, string(capture.Last(t).Body), `"format":"pcm16"`) for _, format := range []string{"opus", "aac", "flac"} { _, err = provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "mimo-v2.5-tts", Input: "hi", ResponseFormat: format, }) - if err == nil { - t.Fatalf("CreateSpeech(%s) succeeded, want unsupported-format error", format) - } + assert.Error(t, err, "format %q", format) } } @@ -116,34 +91,24 @@ func TestCreateSpeech_MapsPCMAndRejectsUnsupportedFormats(t *testing.T) { // with wav, exactly like an omitted response_format. func TestCreateSpeech_TreatsMP3AsUnspecified(t *testing.T) { for _, format := range []string{"", "mp3", "MP3", " mp3 ", "wav"} { - server, gotBody := newTTSServer(t, base64.StdEncoding.EncodeToString([]byte("wav"))) - provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) + t.Run(format, func(t *testing.T) { + server, capture := newTTSServer(t, base64.StdEncoding.EncodeToString([]byte("wav"))) + provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) - resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ - Model: "mimo-v2.5-tts", Input: "hi", ResponseFormat: format, + resp, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ + Model: "mimo-v2.5-tts", Input: "hi", ResponseFormat: format, + }) + require.NoError(t, err) + assert.Equal(t, "audio/wav", resp.ContentType) + assert.Contains(t, string(capture.Last(t).Body), `"format":"wav"`) }) - if err != nil { - server.Close() - t.Fatalf("CreateSpeech(%q) error = %v", format, err) - } - if resp.ContentType != "audio/wav" { - server.Close() - t.Fatalf("CreateSpeech(%q) ContentType = %q, want audio/wav", format, resp.ContentType) - } - if !strings.Contains(string(*gotBody), `"format":"wav"`) { - server.Close() - t.Fatalf("CreateSpeech(%q) upstream body should request wav, got: %s", format, *gotBody) - } - server.Close() } } func TestCreateSpeech_RequiresInput(t *testing.T) { provider := NewWithHTTPClient("mimo-key", "", nil, llmclient.Hooks{}) _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{Model: "mimo-v2.5-tts"}) - if err == nil { - t.Fatal("CreateSpeech() succeeded, want input-required error") - } + require.Error(t, err) } func TestCreateSpeech_RejectsSpeedControl(t *testing.T) { @@ -151,38 +116,18 @@ func TestCreateSpeech_RejectsSpeedControl(t *testing.T) { _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "mimo-v2.5-tts", Input: "hi", Speed: 1.5, }) - if err == nil { - t.Fatal("CreateSpeech(speed=1.5) succeeded, want unsupported-speed error") - } + require.Error(t, err) server, _ := newTTSServer(t, base64.StdEncoding.EncodeToString([]byte("wav"))) - defer server.Close() provider = NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ + _, err = provider.CreateSpeech(context.Background(), &core.AudioSpeechRequest{ Model: "mimo-v2.5-tts", Input: "hi", Speed: 1, - }); err != nil { - t.Fatalf("CreateSpeech(speed=1) error = %v, want default speed accepted", err) - } + }) + require.NoError(t, err) } func TestCreateTranscription_TranslatesToMiMoChatASR(t *testing.T) { - var gotBody []byte - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var err error - gotBody, err = io.ReadAll(r.Body) - if err != nil { - http.Error(w, "read error", http.StatusInternalServerError) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-asr","created":1677652288,"model":"mimo-v2.5-asr", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello there"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":20,"completion_tokens":3,"total_tokens":23} - }`)) - })) - defer server.Close() - + server, capture := newASRServer(t, "hello there") provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) audio := []byte("RIFF-fake-wav") @@ -193,56 +138,34 @@ func TestCreateTranscription_TranslatesToMiMoChatASR(t *testing.T) { Language: "auto", Temperature: "0.2", }) - if err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } - if resp.ContentType != "application/json" { - t.Fatalf("ContentType = %q, want application/json", resp.ContentType) - } + require.NoError(t, err) + assert.Equal(t, "application/json", resp.ContentType) + var out struct { Text string `json:"text"` } - if err := json.Unmarshal(resp.Data, &out); err != nil || out.Text != "hello there" { - t.Fatalf("Data = %s, want {\"text\":\"hello there\"}", resp.Data) - } + require.NoError(t, json.Unmarshal(resp.Data, &out)) + assert.Equal(t, "hello there", out.Text) wantDataURI := "data:audio/wav;base64," + base64.StdEncoding.EncodeToString(audio) - body := string(gotBody) - if !strings.Contains(body, `"type":"input_audio"`) || !strings.Contains(body, wantDataURI) { - t.Fatalf("upstream body missing input_audio data URI, got: %s", body) - } - if strings.Contains(body, `"format"`) { - t.Fatalf("upstream body should not contain a format field, got: %s", body) - } - if !strings.Contains(body, `"asr_options":{"language":"auto"}`) { - t.Fatalf("upstream body missing asr_options, got: %s", body) - } - if !strings.Contains(body, `"temperature":0.2`) { - t.Fatalf("upstream body missing forwarded temperature, got: %s", body) - } + body := string(capture.Last(t).Body) + assert.Contains(t, body, `"type":"input_audio"`) + assert.Contains(t, body, wantDataURI) + assert.NotContains(t, body, `"format"`) + assert.Contains(t, body, `"asr_options":{"language":"auto"}`) + assert.Contains(t, body, `"temperature":0.2`) } func TestCreateTranscription_TextFormatAndValidation(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-asr","created":1677652288,"model":"mimo-v2.5-asr", - "choices":[{"index":0,"message":{"role":"assistant","content":"plain text"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - + server, _ := newASRServer(t, "plain text") provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.CreateTranscription(context.Background(), &core.AudioTranscriptionRequest{ Model: "mimo-v2.5-asr", Filename: "clip.wav", File: []byte("audio"), ResponseFormat: "text", }) - if err != nil { - t.Fatalf("CreateTranscription(text) error = %v", err) - } - if string(resp.Data) != "plain text" || !strings.HasPrefix(resp.ContentType, "text/plain") { - t.Fatalf("got %q (%s), want plain text body", resp.Data, resp.ContentType) - } + require.NoError(t, err) + assert.Equal(t, "plain text", string(resp.Data)) + assert.True(t, strings.HasPrefix(resp.ContentType, "text/plain"), "ContentType = %q, want text/plain", resp.ContentType) for _, unsupported := range []core.AudioTranscriptionRequest{ {Model: "mimo-v2.5-asr", File: []byte("audio"), ResponseFormat: "srt"}, @@ -253,9 +176,8 @@ func TestCreateTranscription_TextFormatAndValidation(t *testing.T) { {Model: "mimo-v2.5-asr"}, } { req := unsupported - if _, err := provider.CreateTranscription(context.Background(), &req); err == nil { - t.Fatalf("CreateTranscription(%+v) succeeded, want validation error", req) - } + _, err := provider.CreateTranscription(context.Background(), &req) + assert.Error(t, err, "request %+v", req) } } @@ -277,14 +199,7 @@ func TestCreateTranscription_FileReaderAndMIMEInference(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - var gotBody []byte - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotBody, _ = io.ReadAll(r.Body) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"id":"x","model":"mimo-v2.5-asr","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}]}`)) - })) - defer server.Close() - + server, capture := newASRServer(t, "ok") provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) audio := []byte("audio-bytes-" + tc.name) @@ -298,15 +213,11 @@ func TestCreateTranscription_FileReaderAndMIMEInference(t *testing.T) { } else { req.File = audio } - - if _, err := provider.CreateTranscription(context.Background(), req); err != nil { - t.Fatalf("CreateTranscription() error = %v", err) - } + _, err := provider.CreateTranscription(context.Background(), req) + require.NoError(t, err) wantData := tc.wantDataPrefix + base64.StdEncoding.EncodeToString(audio) - if !strings.Contains(string(gotBody), wantData) { - t.Fatalf("upstream body missing %q, got: %s", wantData, gotBody) - } + assert.Contains(t, string(capture.Last(t).Body), wantData) }) } } diff --git a/internal/providers/xiaomi/xiaomi_test.go b/internal/providers/xiaomi/xiaomi_test.go index 99fc9b1be..e244432f8 100644 --- a/internal/providers/xiaomi/xiaomi_test.go +++ b/internal/providers/xiaomi/xiaomi_test.go @@ -1,95 +1,33 @@ package xiaomi import ( - "context" - "errors" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/require" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-xiaomi", - "created":1677652288, - "model":"mimo-v2.5-pro", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], - "usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("mimo-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "mimo-v2.5-pro", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "xiaomi", + DefaultBaseURL: "https://api.xiaomimimo.com/v1", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, + Embeddings: false, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "mimo-v2.5-pro" { - t.Fatalf("resp.Model = %q, want mimo-v2.5-pro", resp.Model) - } - if resp.Usage.TotalTokens != 4 { - t.Fatalf("resp.Usage = %+v, want total_tokens=4", resp.Usage) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer mimo-key" { - t.Fatalf("authorization = %q, want Bearer mimo-key", gotAuth) - } -} - -func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { - provider := NewWithHTTPClient("mimo-key", "", nil, llmclient.Hooks{}) - - _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: "mimo-v2.5-pro", - Input: "hello", - }) - if err == nil { - t.Fatal("Embeddings() expected error, got nil") - } - var ge *core.GatewayError - if !errors.As(err, &ge) { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if ge.Type != core.ErrorTypeInvalidRequest { - t.Fatalf("error type = %q, want %q", ge.Type, core.ErrorTypeInvalidRequest) - } - if ge.HTTPStatusCode() != http.StatusBadRequest { - t.Fatalf("status = %d, want %d", ge.HTTPStatusCode(), http.StatusBadRequest) - } -} - -func TestProvider_DefaultBaseURL(t *testing.T) { - provider := NewWithHTTPClient("mimo-key", "", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("expected non-nil provider") - } } -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { +// Xiaomi serves audio through its chat endpoint (see audio.go), so only the +// batch and file surfaces must stay hidden. +func TestProvider_DoesNotExposeBatchOrFileInterfaces(t *testing.T) { provider := NewWithHTTPClient("mimo-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("xiaomi provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("xiaomi provider should not implement native file provider") - } + _, ok := any(provider).(core.NativeBatchProvider) + require.False(t, ok) + _, ok = any(provider).(core.NativeFileProvider) + require.False(t, ok) } diff --git a/internal/providers/zai/passthrough_semantics_test.go b/internal/providers/zai/passthrough_semantics_test.go index 9a68c69e3..1e9fd37c0 100644 --- a/internal/providers/zai/passthrough_semantics_test.go +++ b/internal/providers/zai/passthrough_semantics_test.go @@ -4,20 +4,20 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricherUsesZAIType(t *testing.T) { enricher := Registration.PassthroughSemanticEnricher - if enricher == nil { - t.Fatal("registration passthrough enricher is nil") - } - if got := enricher.ProviderType(); got != "zai" { - t.Fatalf("ProviderType() = %q, want zai", got) - } + require.NotNil(t, enricher) + got := enricher.ProviderType() + require.Equal(t, "zai", got) + info := enricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ Provider: "zai", NormalizedEndpoint: "embeddings", }) - if info == nil || info.GenAIOperation != "embeddings" || info.SemanticOperation != "zai.embeddings" || info.AuditPath != "/v1/embeddings" { - t.Fatalf("enriched info = %+v, want Z.ai embedding semantics", info) - } + require.NotNil(t, info) + require.Equal(t, "embeddings", info.GenAIOperation) + require.Equal(t, "zai.embeddings", info.SemanticOperation) + require.Equal(t, "/v1/embeddings", info.AuditPath) } diff --git a/internal/providers/zai/realtime_test.go b/internal/providers/zai/realtime_test.go index 0b86ccd40..8dbed49c1 100644 --- a/internal/providers/zai/realtime_test.go +++ b/internal/providers/zai/realtime_test.go @@ -8,6 +8,8 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestRealtimeTarget(t *testing.T) { @@ -15,22 +17,15 @@ func TestRealtimeTarget(t *testing.T) { p := New(providers.ProviderConfig{APIKey: apiKey}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "glm-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://api.z.ai/api/paas/v4/realtime?") { - t.Errorf("url = %q, want Z.ai realtime endpoint", target.URL) - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://api.z.ai/api/paas/v4/realtime?"), "url = %q, want Z.ai realtime endpoint", target.URL) + u, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("parse target url: %v", err) - } - if got := u.Query().Get("model"); got != "glm-realtime" { - t.Errorf("model query = %q, want %q", got, "glm-realtime") - } - if got := target.Headers.Get("Authorization"); got != "Bearer "+apiKey { - t.Errorf("Authorization = %q, want bearer with key", got) - } + require.NoError(t, err) + got := u.Query().Get("model") + assert.Equal(t, "glm-realtime", got) + got = target.Headers.Get("Authorization") + assert.Equal(t, "Bearer "+apiKey, got) } func TestRealtimeTargetFollowsSetBaseURL(t *testing.T) { @@ -38,12 +33,8 @@ func TestRealtimeTargetFollowsSetBaseURL(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) p.SetBaseURL("https://open.bigmodel.cn/api/paas/v4") target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "glm-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.HasPrefix(target.URL, "wss://open.bigmodel.cn/api/paas/v4/realtime?") { - t.Errorf("url = %q, want the configured region host", target.URL) - } + require.NoError(t, err) + assert.True(t, strings.HasPrefix(target.URL, "wss://open.bigmodel.cn/api/paas/v4/realtime?"), "url = %q, want the configured region host", target.URL) } func TestRealtimeTargetNormalizesCodingPlanBase(t *testing.T) { @@ -52,32 +43,23 @@ func TestRealtimeTargetNormalizesCodingPlanBase(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) p.SetBaseURL("https://api.z.ai/api/coding/paas/v4") target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "glm-realtime"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + require.NoError(t, err) + u, err := url.Parse(target.URL) - if err != nil { - t.Fatalf("parse target url: %v", err) - } - if u.Path != "/api/paas/v4/realtime" { - t.Errorf("path = %q, want /api/paas/v4/realtime", u.Path) - } + require.NoError(t, err) + assert.Equal(t, "/api/paas/v4/realtime", u.Path) } func TestRealtimeTargetOmitsAuthWhenNoKey(t *testing.T) { p := New(providers.ProviderConfig{APIKey: ""}, providers.ProviderOptions{}).(*Provider) target, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: "m"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if _, present := target.Headers["Authorization"]; present { - t.Error("Authorization header should be absent when no API key is configured") - } + require.NoError(t, err) + _, present := target.Headers["Authorization"] + assert.False(t, present) } func TestRealtimeTargetMissingModel(t *testing.T) { p := New(providers.ProviderConfig{APIKey: "k"}, providers.ProviderOptions{}).(*Provider) - if _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}); err == nil { - t.Fatal("expected error for missing model") - } + _, err := p.RealtimeTarget(context.Background(), &core.RealtimeRequest{Model: " "}) + require.Error(t, err) } diff --git a/internal/providers/zai/zai_test.go b/internal/providers/zai/zai_test.go index 5894da79b..e237161ab 100644 --- a/internal/providers/zai/zai_test.go +++ b/internal/providers/zai/zai_test.go @@ -1,98 +1,26 @@ package zai import ( - "context" "net/http" - "net/http/httptest" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" ) -func TestChatCompletion_UsesBearerAuthAndChatEndpoint(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"chatcmpl-zai", - "created":1677652288, - "model":"glm-5", - "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}] - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("zai-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ - Model: "glm-5", - Messages: []core.Message{ - {Role: "user", Content: "hi"}, +// Z.ai is a thin wrapper over the shared chat-centric adapter and forwards +// embeddings upstream, so the shared contract covers its surface. It must +// not advertise native batch, file, or audio support. +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: "zai", + DefaultBaseURL: "https://api.z.ai/api/paas/v4", + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return NewWithHTTPClient(apiKey, baseURL, client, hooks) }, + Embeddings: true, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if resp.Model != "glm-5" { - t.Fatalf("resp.Model = %q, want glm-5", resp.Model) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer zai-key" { - t.Fatalf("authorization = %q, want Bearer zai-key", gotAuth) - } -} - -func TestEmbeddings_DelegatesToCompatibleProvider(t *testing.T) { - var gotPath string - var gotAuth string - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotPath = r.URL.Path - gotAuth = r.Header.Get("Authorization") - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "object":"list", - "model":"embedding-model", - "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], - "usage":{"prompt_tokens":3,"total_tokens":3} - }`)) - })) - defer server.Close() - - provider := NewWithHTTPClient("zai-key", server.URL, server.Client(), llmclient.Hooks{}) - - resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ - Model: "embedding-model", - Input: "hello", - }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Model != "embedding-model" { - t.Fatalf("resp.Model = %q, want embedding-model", resp.Model) - } - if gotPath != "/embeddings" { - t.Fatalf("path = %q, want /embeddings", gotPath) - } - if gotAuth != "Bearer zai-key" { - t.Fatalf("authorization = %q, want Bearer zai-key", gotAuth) - } -} - -func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("zai-key", "", nil, llmclient.Hooks{}) - - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("zai provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("zai provider should not implement native file provider") - } + providertest.AssertNoNativeSurfaces(t, NewWithHTTPClient("zai-key", "", nil, llmclient.Hooks{})) }