From a1e4772670aa51792b6db0e5e5edef131f44e690 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Mon, 27 Jul 2026 06:03:51 +0530 Subject: [PATCH 1/4] test: increase eyrie client/adapters coverage from 22.9% to 84.6% MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add comprehensive test coverage for all major provider adapters: - gemini_test.go: 22 new tests covering legacy parser paths, buildBody branches (system, tools, safety, thinking, penalties, ContentParts), role mapping, tool choices, and streaming edge cases - bedrock_test.go: 6 new tests for EventStream reader (headers, invalid lengths, header overflow), default max tokens, and system prompt merge - anthropic_test.go: 13 new tests for response formats (json_schema, json_object, unknown types), metadata, and cached request builder - opencodego_test.go: 4 new tests for Name, Ping (success + fallback), and OACompatUnsupportedError - zai_test.go: 1 new test for Name with nil OpenAI router - protocol_router_test.go: 2 new tests for streamResultFromChat (full response with thinking/tool calls/usage, nil response) Key coverage improvements: - gemini.go: processStreamChunk 60%→90%, buildBody 68%→89%, streamLoop 96%→100% - bedrock.go: buildBody 65%→83%, ReadEvent 50%→79% - anthropic.go: buildAnthropicRequest 61%→88%, resolveMetadata 67%→100% - anthropic_cache.go: buildAnthropicCachedRequest 46%→100%, applyCacheBreakpoint 0%→100% - opencodego.go: Name 0%→100%, Ping 0%→100%, OACompatUnsupportedError 0%→100% - protocol_router.go: streamResultFromChat 68%→100% All tests pass across the entire eyrie codebase. Code is gofumpt-formatted. --- client/adapters/adapter_config_test.go | 268 ++++++ client/adapters/anthropic_test.go | 993 +++++++++++++++++++++++ client/adapters/azure_test.go | 206 +++++ client/adapters/bedrock_test.go | 567 +++++++++++++ client/adapters/deepseek_test.go | 191 +++++ client/adapters/dynamic_test.go | 158 ++++ client/adapters/gemini_test.go | 935 +++++++++++++++++++++ client/adapters/mimo_test.go | 187 +++++ client/adapters/openai_embedding_test.go | 165 ++++ client/adapters/openai_test.go | 476 +++++++++++ client/adapters/opencodego_test.go | 67 ++ client/adapters/poolside_ext_test.go | 82 ++ client/adapters/protocol_router_test.go | 65 ++ client/adapters/test_helpers_test.go | 14 + client/adapters/vertex_test.go | 237 ++++++ client/adapters/zai_test.go | 248 ++++++ conversation/engine_test.go | 8 +- engine/convert_test.go | 6 +- go.mod | 2 +- go.sum | 4 +- operationsgraph/operations_graph.go | 18 +- 21 files changed, 4877 insertions(+), 20 deletions(-) create mode 100644 client/adapters/adapter_config_test.go create mode 100644 client/adapters/anthropic_test.go create mode 100644 client/adapters/azure_test.go create mode 100644 client/adapters/bedrock_test.go create mode 100644 client/adapters/deepseek_test.go create mode 100644 client/adapters/dynamic_test.go create mode 100644 client/adapters/gemini_test.go create mode 100644 client/adapters/openai_embedding_test.go create mode 100644 client/adapters/openai_test.go create mode 100644 client/adapters/poolside_ext_test.go create mode 100644 client/adapters/vertex_test.go create mode 100644 client/adapters/zai_test.go diff --git a/client/adapters/adapter_config_test.go b/client/adapters/adapter_config_test.go new file mode 100644 index 0000000..a5cc8fb --- /dev/null +++ b/client/adapters/adapter_config_test.go @@ -0,0 +1,268 @@ +package adapters + +import ( + "log/slog" + "net/http" + "os" + "testing" + "time" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestAnthropicConfigSetters(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-test", "https://api.anthropic.com") + + // SetTimeout + c.SetTimeout(5 * time.Second) + if c.httpClient.Timeout != 5*time.Second { + t.Errorf("expected timeout 5s, got %v", c.httpClient.Timeout) + } + + // SetHTTPClient + custom := &http.Client{Timeout: 10 * time.Second} + c.SetHTTPClient(custom) + if c.httpClient != custom { + t.Error("expected httpClient to be replaced") + } + + // SetRetry + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 3}} + c.SetRetry(rc) + if c.retry.MaxRetries != 3 { + t.Errorf("expected MaxRetries=3, got %d", c.retry.MaxRetries) + } + + // SetLogger + logger := slog.New(slog.NewTextHandler(os.Stdout, nil)) + c.SetLogger(logger) + if c.logger != logger { + t.Error("expected logger to be set") + } + + // SetAPIKey + c.SetAPIKey("sk-test") + if c.apiKey != "sk-test" { + t.Errorf("expected apiKey=sk-test, got %q", c.apiKey) + } + + // SetBaseURL + c.SetBaseURL("https://test.example.com") + if c.baseURL != "https://test.example.com" { + t.Errorf("expected baseURL=https://test.example.com, got %q", c.baseURL) + } + + // SetDefaultModel + c.SetDefaultModel("claude-opus-4") + if c.defaultModel != "claude-opus-4" { + t.Errorf("expected defaultModel=claude-opus-4, got %q", c.defaultModel) + } + + // SetDefaultMaxTokens + c.SetDefaultMaxTokens(4096) + if c.defaultMaxTokens != 4096 { + t.Errorf("expected defaultMaxTokens=4096, got %d", c.defaultMaxTokens) + } + + // SetDefaultTemperature + c.SetDefaultTemperature(0.7) + if c.defaultTemperature == nil || *c.defaultTemperature != 0.7 { + t.Errorf("expected defaultTemperature=0.7, got %v", c.defaultTemperature) + } + + // SetGuardrails + g := core.NewGuardrails() + c.SetGuardrails(g) + if c.guardrails != g { + t.Error("expected guardrails to be set") + } + + // SetProviderName (no-op) + c.SetProviderName("custom") + // Should not panic or change anything + + // SetMimoAuth + c.SetMimoAuth() + if !c.useMimoAuth { + t.Error("expected useMimoAuth=true after SetMimoAuth") + } + + // SetProviderName (no-op for Anthropic) + c.SetProviderName("custom") +} + +func TestAnthropicConfigGetters(t *testing.T) { + t.Parallel() + g := core.NewGuardrails() + temp := 0.5 + c := &AnthropicClient{ + baseURL: "https://api.anthropic.com", + httpClient: &http.Client{Timeout: 30 * time.Second}, + retry: core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 5}}, + logger: slog.Default(), + guardrails: g, + defaultModel: "claude-sonnet-4", + defaultMaxTokens: 8192, + defaultTemperature: &temp, + version: "1.0.0", + } + + if c.BaseURL() != "https://api.anthropic.com" { + t.Errorf("BaseURL mismatch") + } + if c.HTTPClient().Timeout != 30*time.Second { + t.Errorf("HTTPClient mismatch") + } + if c.Retry().MaxRetries != 5 { + t.Errorf("Retry mismatch") + } + if c.Logger() == nil { + t.Error("Logger should not be nil") + } + if c.Guardrails() != g { + t.Error("Guardrails mismatch") + } + if c.DefaultModel() != "claude-sonnet-4" { + t.Errorf("DefaultModel mismatch") + } + if c.DefaultMaxTokens() != 8192 { + t.Errorf("DefaultMaxTokens mismatch") + } + if c.DefaultTemperature() == nil || *c.DefaultTemperature() != 0.5 { + t.Errorf("DefaultTemperature mismatch") + } + if c.Version() != "1.0.0" { + t.Errorf("Version mismatch") + } +} + +func TestOpenAIConfigSetters(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("sk-test", "https://api.openai.com", nil) + + c.SetTimeout(5 * time.Second) + if c.httpClient.Timeout != 5*time.Second { + t.Errorf("expected timeout 5s, got %v", c.httpClient.Timeout) + } + + custom := &http.Client{Timeout: 10 * time.Second} + c.SetHTTPClient(custom) + if c.httpClient != custom { + t.Error("expected httpClient to be replaced") + } + + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 3}} + c.SetRetry(rc) + if c.retry.MaxRetries != 3 { + t.Errorf("expected MaxRetries=3, got %d", c.retry.MaxRetries) + } + + logger := slog.New(slog.NewTextHandler(os.Stdout, nil)) + c.SetLogger(logger) + if c.logger != logger { + t.Error("expected logger to be set") + } + + c.SetAPIKey("sk-test") + if c.apiKey != "sk-test" { + t.Errorf("expected apiKey=sk-test, got %q", c.apiKey) + } + + c.SetBaseURL("https://test.example.com") + if c.baseURL != "https://test.example.com" { + t.Errorf("expected baseURL=https://test.example.com, got %q", c.baseURL) + } + + c.SetDefaultModel("gpt-5") + if c.defaultModel != "gpt-5" { + t.Errorf("expected defaultModel=gpt-5, got %q", c.defaultModel) + } + + c.SetDefaultMaxTokens(4096) + if c.defaultMaxTokens != 4096 { + t.Errorf("expected defaultMaxTokens=4096, got %d", c.defaultMaxTokens) + } + + c.SetDefaultTemperature(0.7) + if c.defaultTemperature == nil || *c.defaultTemperature != 0.7 { + t.Errorf("expected defaultTemperature=0.7, got %v", c.defaultTemperature) + } + + g := core.NewGuardrails() + c.SetGuardrails(g) + if c.guardrails != g { + t.Error("expected guardrails to be set") + } + + c.SetProviderName("custom") + if c.providerName != "custom" { + t.Errorf("expected providerName=custom, got %q", c.providerName) + } + + c.SetMimoAuth() + if !c.useMimoAuth { + t.Error("expected useMimoAuth=true after SetMimoAuth") + } +} + +func TestOpenAIConfigGetters(t *testing.T) { + t.Parallel() + compat := &OpenAICompatConfig{MaxTokensField: "max_tokens"} + c := &OpenAIClient{ + baseURL: "https://api.openai.com", + httpClient: &http.Client{Timeout: 30 * time.Second}, + retry: core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 5}}, + logger: slog.Default(), + compat: compat, + } + + if c.BaseURL() != "https://api.openai.com" { + t.Errorf("BaseURL mismatch") + } + if c.HTTPClient().Timeout != 30*time.Second { + t.Errorf("HTTPClient mismatch") + } + if c.Retry().MaxRetries != 5 { + t.Errorf("Retry mismatch") + } + if c.ProviderName() != "" { + t.Errorf("ProviderName mismatch") + } + if c.Compat() != compat { + t.Error("Compat mismatch") + } + if c.Logger() == nil { + t.Error("Logger should not be nil") + } + if c.Guardrails() != nil { + t.Error("Guardrails should be nil by default") + } + if c.DefaultModel() != "" { + t.Errorf("DefaultModel = %q", c.DefaultModel()) + } + if c.DefaultMaxTokens() != 0 { + t.Errorf("DefaultMaxTokens = %d", c.DefaultMaxTokens()) + } + if c.DefaultTemperature() != nil { + t.Error("DefaultTemperature should be nil") + } +} + +func TestBedrockConfigSetters(t *testing.T) { + t.Parallel() + c := &BedrockClient{} + + custom := &http.Client{Timeout: 10 * time.Second} + c.SetHTTPClient(custom) + if c.httpClient != custom { + t.Error("expected httpClient to be replaced with custom") + } + + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 3}} + c.SetRetry(rc) + if c.retry.MaxRetries != 3 { + t.Errorf("expected MaxRetries=3, got %d", c.retry.MaxRetries) + } +} diff --git a/client/adapters/anthropic_test.go b/client/adapters/anthropic_test.go new file mode 100644 index 0000000..ee8f518 --- /dev/null +++ b/client/adapters/anthropic_test.go @@ -0,0 +1,993 @@ +package adapters + +import ( + "context" + "encoding/json" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewAnthropicClient(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://custom.proxy.com") + if c == nil { + t.Fatal("NewAnthropicClient returned nil") + } + if c.apiKey != "sk-ant-key" { + t.Errorf("apiKey = %q", c.apiKey) + } + if c.baseURL != "https://custom.proxy.com" { + t.Errorf("baseURL = %q", c.baseURL) + } + if c.version != "2023-06-01" { + t.Errorf("version = %q", c.version) + } +} + +func TestNewAnthropicClient_EmptyBaseURL(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "") + if c.baseURL != "https://api.anthropic.com" { + t.Errorf("baseURL = %q, want https://api.anthropic.com", c.baseURL) + } +} + +func TestAnthropicClient_Name(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + if c.Name() != "anthropic" { + t.Errorf("Name() = %q", c.Name()) + } +} + +func TestAnthropicClient_Chat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "Hello from Anthropic!"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from Anthropic!" { + t.Errorf("content = %q", resp.Content) + } + if resp.FinishReason != "end_turn" { + t.Errorf("finishReason = %q", resp.FinishReason) + } +} + +func TestAnthropicClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestAnthropicClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusForbidden, map[string]any{ + "error": map[string]string{"message": "permission denied"}, + }), nil + }) + c := NewAnthropicClient("bad-key", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error for forbidden") + } +} + +func TestAnthropicClient_StreamChat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":10}}}\n\nevent: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello stream!\"}}\n\nevent: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewAnthropicClient("key", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "Hello stream!" { + t.Errorf("content = %q", content) + } +} + +func TestAnthropicClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestAnthropicClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + }) + c := NewAnthropicClient("key", "https://api.anthropic.com") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestAnthropicClient_Ping_AuthError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + }) + c := NewAnthropicClient("bad-key", "") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestAnthropicClient_CountTokens_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{"input_tokens": 42}), nil + }) + c := NewAnthropicClient("key", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + result, err := c.CountTokens(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Count me"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514"}) + if err != nil { + t.Fatalf("CountTokens: %v", err) + } + if result.InputTokens != 42 { + t.Errorf("InputTokens = %d", result.InputTokens) + } +} + +func TestAnthropicClient_CountTokens_EmptyModel(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + _, err := c.CountTokens(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestAnthropicClient_CountTokens_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{ + "error": map[string]string{"message": "invalid model"}, + }), nil + }) + c := NewAnthropicClient("key", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.CountTokens(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-invalid"}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestAnthropicClient_SetHTTPClientAndRetry(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + c2 := NewAnthropicClient("key2", "") + c.SetHTTPClient(c2.httpClient) + if c.httpClient != c2.httpClient { + t.Error("SetHTTPClient did not replace client") + } + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 5}} + c.SetRetry(rc) + if c.retry.MaxRetries != 5 { + t.Errorf("expected MaxRetries=5, got %d", c.retry.MaxRetries) + } +} + +func TestConvertToAnthropicTools(t *testing.T) { + t.Parallel() + result := ConvertToAnthropicTools([]core.EyrieTool{ + {Name: "get_weather", Description: "Get weather", Parameters: map[string]interface{}{"type": "object"}}, + }) + if len(result) != 1 { + t.Fatalf("len = %d", len(result)) + } + if result[0].Name != "get_weather" { + t.Errorf("Name = %q", result[0].Name) + } + if result[0].Description != "Get weather" { + t.Errorf("Description = %q", result[0].Description) + } +} + +func TestConvertToAnthropicTools_Empty(t *testing.T) { + t.Parallel() + result := ConvertToAnthropicTools(nil) + if result != nil { + t.Errorf("expected nil, got %v", result) + } +} + +func TestParseAnthropicResponse(t *testing.T) { + t.Parallel() + ar := anthropicResponse{ + ID: "msg_1", + Content: []struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + Thinking string `json:"thinking,omitempty"` + Signature string `json:"signature,omitempty"` + Data string `json:"data,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Input json.RawMessage `json:"input,omitempty"` + }{ + {Type: "text", Text: "Hello"}, + {Type: "thinking", Thinking: "deep thoughts"}, + {Type: "redacted_thinking", Data: "sensitive"}, + }, + StopReason: "end_turn", + Usage: struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + CacheCreationInputTokens int `json:"cache_creation_input_tokens"` + CacheReadInputTokens int `json:"cache_read_input_tokens"` + OutputTokensDetails struct { + ThinkingTokens int `json:"thinking_tokens"` + } `json:"output_tokens_details"` + }{ + InputTokens: 10, OutputTokens: 20, + CacheCreationInputTokens: 2, CacheReadInputTokens: 3, + }, + } + ar.Usage.OutputTokensDetails.ThinkingTokens = 5 + + resp := ParseAnthropicResponse(ar, "req_1", "org_1") + if resp.Content != "Hello" { + t.Errorf("Content = %q", resp.Content) + } + if resp.Thinking != "deep thoughts" { + t.Errorf("Thinking = %q", resp.Thinking) + } + if resp.FinishReason != "end_turn" { + t.Errorf("FinishReason = %q", resp.FinishReason) + } + if resp.RequestID != "req_1" { + t.Errorf("RequestID = %q", resp.RequestID) + } + if resp.OrganizationID != "org_1" { + t.Errorf("OrganizationID = %q", resp.OrganizationID) + } + if resp.Usage.PromptTokens != 10 || resp.Usage.CompletionTokens != 20 { + t.Errorf("usage tokens: %+v", resp.Usage) + } + if resp.Usage.ThinkingTokens != 5 { + t.Errorf("ThinkingTokens = %d", resp.Usage.ThinkingTokens) + } + if resp.Usage.CacheCreationTokens != 2 { + t.Errorf("CacheCreationTokens = %d", resp.Usage.CacheCreationTokens) + } + if resp.Usage.CacheReadTokens != 3 { + t.Errorf("CacheReadTokens = %d", resp.Usage.CacheReadTokens) + } +} + +func TestParseAnthropicResponse_ToolUse(t *testing.T) { + t.Parallel() + input, _ := json.Marshal(map[string]string{"city": "NYC"}) + ar := anthropicResponse{ + Content: []struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + Thinking string `json:"thinking,omitempty"` + Signature string `json:"signature,omitempty"` + Data string `json:"data,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Input json.RawMessage `json:"input,omitempty"` + }{ + {Type: "text", Text: "Let me check"}, + {Type: "tool_use", ID: "tc1", Name: "get_weather", Input: input}, + }, + StopReason: "tool_use", + } + resp := ParseAnthropicResponse(ar, "req_2", "") + if resp.Content != "Let me check" { + t.Errorf("Content = %q", resp.Content) + } + if len(resp.ToolCalls) != 1 { + t.Fatalf("ToolCalls = %d", len(resp.ToolCalls)) + } + if resp.ToolCalls[0].Name != "get_weather" { + t.Errorf("ToolCall.Name = %q", resp.ToolCalls[0].Name) + } + if resp.ToolCalls[0].ID != "tc1" { + t.Errorf("ToolCall.ID = %q", resp.ToolCalls[0].ID) + } +} + +func TestBuildAnthropicMessages_Basic(t *testing.T) { + t.Parallel() + msgs, system := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "system", Content: "Be helpful"}, + {Role: "user", Content: "Hello"}, + {Role: "assistant", Content: "Hi there"}, + }) + if system != "Be helpful" { + t.Errorf("system = %q", system) + } + if len(msgs) != 2 { + t.Fatalf("msgs = %d", len(msgs)) + } + if msgs[0]["role"] != "user" { + t.Errorf("msg[0] role = %v", msgs[0]["role"]) + } + if msgs[0]["content"] != "Hello" { + t.Errorf("msg[0] content = %v", msgs[0]["content"]) + } +} + +func TestBuildAnthropicMessages_ToolUse(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "assistant", Content: "Let me check", ToolUse: []core.ToolCall{{ID: "tu1", Name: "get_weather", Arguments: map[string]interface{}{"city": "NYC"}}}}, + }) + if len(msgs) != 1 { + t.Fatalf("msgs = %d", len(msgs)) + } + content := msgs[0]["content"].([]map[string]interface{}) + if len(content) != 2 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[1]["type"] != "tool_use" { + t.Errorf("block type = %v", content[1]["type"]) + } +} + +func TestBuildAnthropicMessages_ToolResults(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "user", ToolResults: []core.ToolResult{{ToolUseID: "tu1", Content: "72°F"}}}, + }) + if len(msgs) != 1 { + t.Fatalf("msgs = %d", len(msgs)) + } + content := msgs[0]["content"].([]map[string]interface{}) + if len(content) != 1 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[0]["type"] != "tool_result" { + t.Errorf("block type = %v", content[0]["type"]) + } + if content[0]["tool_use_id"] != "tu1" { + t.Errorf("tool_use_id = %v", content[0]["tool_use_id"]) + } +} + +func TestBuildAnthropicMessages_ToolResultWithError(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "user", ToolResults: []core.ToolResult{{ToolUseID: "tu1", Content: "error!", IsError: true}}}, + }) + content := msgs[0]["content"].([]map[string]interface{}) + if content[0]["is_error"] != true { + t.Errorf("expected is_error=true") + } +} + +func TestBuildAnthropicMessages_ContentParts(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "user", ContentParts: []core.ContentPart{ + {Type: "text", Text: "Describe this image"}, + {Type: "image_url", ImageURL: &core.ImageURLPart{URL: "data:image/png;base64,iVBORw0KGgo="}}, + }}, + }) + if len(msgs) != 1 { + t.Fatalf("msgs = %d", len(msgs)) + } + content := msgs[0]["content"].([]map[string]interface{}) + if len(content) != 2 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[1]["type"] != "image" { + t.Errorf("block type = %v", content[1]["type"]) + } +} + +func TestBuildAnthropicMessages_ContentParts_Audio(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "user", ContentParts: []core.ContentPart{ + {Type: "input_audio", InputAudio: &core.InputAudioPart{Data: "base64data", Format: "mp3"}}, + }}, + }) + content := msgs[0]["content"].([]map[string]interface{}) + if len(content) != 1 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[0]["type"] != "audio" { + t.Errorf("block type = %v", content[0]["type"]) + } +} + +func TestBuildAnthropicMessages_LegacyImages(t *testing.T) { + t.Parallel() + msgs, _ := BuildAnthropicMessages([]core.EyrieMessage{ + {Role: "user", Content: "Check", Images: []string{"data:image/jpeg;base64,/9j/4AAQ=="}}, + }) + content := msgs[0]["content"].([]map[string]interface{}) + if len(content) != 2 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[1]["type"] != "image" { + t.Errorf("block type = %v", content[1]["type"]) + } +} + +func TestResolveThinking(t *testing.T) { + t.Parallel() + tests := []struct { + mode string + budget int + want *AnthropicThinking + }{ + {"adaptive", 0, &AnthropicThinking{Type: "adaptive"}}, + {"disabled", 0, &AnthropicThinking{Type: "disabled"}}, + {"enabled", 1000, &AnthropicThinking{Type: "enabled", BudgetTokens: 1000}}, + {"", 500, &AnthropicThinking{Type: "enabled", BudgetTokens: 500}}, + {"", 0, nil}, + } + for _, tt := range tests { + got := ResolveThinking(core.ChatOptions{ThinkingMode: tt.mode, ThinkingBudgetTokens: tt.budget}) + if tt.want == nil { + if got != nil { + t.Errorf("ResolveThinking(%q,%d) = %v, want nil", tt.mode, tt.budget, got) + } + continue + } + if got == nil { + t.Errorf("ResolveThinking(%q,%d) = nil, want %+v", tt.mode, tt.budget, *tt.want) + continue + } + if got.Type != tt.want.Type || got.BudgetTokens != tt.want.BudgetTokens { + t.Errorf("ResolveThinking(%q,%d) = %+v, want %+v", tt.mode, tt.budget, *got, *tt.want) + } + } +} + +func TestResolveToolChoice(t *testing.T) { + t.Parallel() + if ResolveToolChoice(nil) != nil { + t.Error("expected nil for nil input") + } + tc := ResolveToolChoice(&core.ToolChoiceOption{Type: "tool", Name: "get_weather", DisableParallelToolUse: true}) + if tc == nil { + t.Fatal("expected non-nil") + } + if tc.Type != "tool" || tc.Name != "get_weather" || !tc.DisableParallelToolUse { + t.Errorf("ToolChoice = %+v", *tc) + } +} + +func TestResolveOutputConfig(t *testing.T) { + t.Parallel() + if ResolveOutputConfig(core.ChatOptions{}) != nil { + t.Error("expected nil for empty opts") + } + cfg := ResolveOutputConfig(core.ChatOptions{OutputEffort: "high"}) + if cfg == nil || cfg.Effort != "high" { + t.Errorf("OutputConfig = %+v", cfg) + } +} + +func TestAudioFormatToMediaType(t *testing.T) { + t.Parallel() + tests := []struct{ in, want string }{ + {"wav", "audio/wav"}, + {"mp3", "audio/mpeg"}, + {"flac", "audio/flac"}, + {"ogg", "audio/ogg"}, + {"aac", "audio/aac"}, + {"webm", "audio/webm"}, + {"audio/wav", "audio/wav"}, + {"unknown", "audio/unknown"}, + } + for _, tt := range tests { + got := AudioFormatToMediaType(tt.in) + if got != tt.want { + t.Errorf("AudioFormatToMediaType(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestThinkingForBudget(t *testing.T) { + t.Parallel() + if ThinkingForBudget(0) != nil { + t.Error("expected nil for budget 0") + } + tb := ThinkingForBudget(5000) + if tb == nil || tb.Type != "enabled" || tb.BudgetTokens != 5000 { + t.Errorf("ThinkingForBudget(5000) = %+v", *tb) + } +} + +func TestBuildAnthropicRequest_ResponseFormat(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + _, _, err := c.BuildAnthropicRequest(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_object"}, + }, false) + if err == nil { + t.Fatal("expected error for json_object without schema") + } +} + +func TestBuildAnthropicRequest_ResponseFormatJSONSchema(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("key", "") + req, _, err := c.BuildAnthropicRequest(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_schema", Schema: `{"type":"object"}`}, + }, false) + if err != nil { + t.Fatalf("BuildAnthropicRequest: %v", err) + } + if req == nil { + t.Fatal("expected non-nil request") + } +} + +func TestBuildAnthropicRequest_Headers(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + req, _, err := c.BuildAnthropicRequest(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", MaxTokens: 256, + }, false) + if err != nil { + t.Fatalf("BuildAnthropicRequest: %v", err) + } + if req.Header.Get("X-Api-Key") != "sk-ant-key" { + t.Errorf("X-Api-Key = %q", req.Header.Get("X-Api-Key")) + } + if req.Header.Get("Anthropic-Version") != "2023-06-01" { + t.Errorf("Anthropic-Version = %q", req.Header.Get("Anthropic-Version")) + } + if req.Header.Get("Content-Type") != "application/json" { + t.Errorf("Content-Type = %q", req.Header.Get("Content-Type")) + } +} + +func TestAnthropicClient_MimoAuthRetry(t *testing.T) { + t.Parallel() + calls := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + calls++ + if calls == 1 { + if req.Header.Get("api-key") != "tp-key" { + t.Errorf("expected api-key on first call, got %q", req.Header.Get("api-key")) + } + return jsonResponse(http.StatusUnauthorized, map[string]any{"error": "invalid"}), nil + } + if req.Header.Get("Authorization") != "Bearer tp-key" { + t.Errorf("expected Bearer on retry, got %q", req.Header.Get("Authorization")) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "retried"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 1, "output_tokens": 1}, + }), nil + }) + c := NewAnthropicClient("tp-key", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.SetMimoAuth() + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "retried" { + t.Errorf("content = %q", resp.Content) + } + if calls != 2 { + t.Errorf("calls = %d, want 2", calls) + } +} + +func TestAnthropicClient_Chat_DefaultMaxTokens(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["max_tokens"] != float64(4096) { + t.Errorf("max_tokens = %v, expected 4096", body["max_tokens"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514"}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestAnthropicClient_Chat_ResponseFormatJSONSchema(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["output_config"] == nil { + t.Error("expected output_config in request") + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_schema", Schema: `{"type":"object","properties":{"name":{"type":"string"}}}`}, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestAnthropicClient_Chat_ResponseFormatEmptySchema(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + })} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_schema", Schema: ""}, + }) + if err == nil { + t.Fatal("expected error for empty schema") + } +} + +func TestAnthropicClient_Chat_ResponseFormatInvalidSchema(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + })} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_schema", Schema: "not valid json"}, + }) + if err == nil { + t.Fatal("expected error for invalid schema") + } +} + +func TestAnthropicClient_Chat_ResponseFormatJSONObject(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + })} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_object"}, + }) + if err == nil { + t.Fatal("expected error for json_object without schema") + } +} + +func TestAnthropicClient_Chat_ResponseFormatUnknown(t *testing.T) { + t.Parallel() + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + })} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "unknown_type"}, + }) + if err == nil { + t.Fatal("expected error for unknown response format") + } +} + +func TestAnthropicClient_Chat_SystemPrompt(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + sys, ok := body["system"].(string) + if !ok { + t.Fatalf("system is not a string: %T", body["system"]) + } + if !strings.Contains(sys, "custom system") { + t.Errorf("system = %q, expected to contain 'custom system'", sys) + } + if !strings.Contains(sys, "from message") { + t.Errorf("system = %q, expected to contain 'from message'", sys) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{ + {Role: "system", Content: "from message"}, + {Role: "user", Content: "Hi"}, + }, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256, System: "custom system"}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestAnthropicClient_Chat_EnableCaching(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["max_tokens"] != float64(256) { + t.Errorf("max_tokens = %v", body["max_tokens"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + EnableCaching: true, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestAnthropicClient_Chat_MetadataUserID(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + meta, ok := body["metadata"].(map[string]interface{}) + if !ok { + t.Fatal("expected metadata in request") + } + if meta["user_id"] != "user-123" { + t.Errorf("user_id = %v", meta["user_id"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_ant_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewAnthropicClient("sk-ant-key", "https://api.anthropic.com") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + MetadataUserID: "user-123", + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestBuildAnthropicCachedRequest_Basic(t *testing.T) { + t.Parallel() + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{{Role: "user", Content: "Hello"}}, + "claude-3", 256, nil, false, nil, nil, nil, nil, nil, nil, + ) + if req["model"] != "claude-3" { + t.Errorf("model = %v", req["model"]) + } + if req["max_tokens"] != 256 { + t.Errorf("max_tokens = %v", req["max_tokens"]) + } + if req["stream"] != false { + t.Errorf("stream = %v", req["stream"]) + } +} + +func TestBuildAnthropicCachedRequest_WithSystem(t *testing.T) { + t.Parallel() + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{{Role: "system", Content: "You are helpful"}}, + "claude-3", 256, nil, false, nil, nil, nil, nil, nil, nil, + ) + sys, ok := req["system"].([]map[string]interface{}) + if !ok { + t.Fatal("expected system as array of maps") + } + if len(sys) != 1 { + t.Fatalf("expected 1 system block, got %d", len(sys)) + } + if sys[0]["cache_control"] == nil { + t.Error("expected cache_control on system") + } +} + +func TestBuildAnthropicCachedRequest_WithTools(t *testing.T) { + t.Parallel() + tools := []AnthropicTool{{Name: "test", Description: "a test", InputSchema: map[string]interface{}{"type": "object"}}} + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{{Role: "user", Content: "Hello"}}, + "claude-3", 256, nil, false, tools, nil, nil, nil, nil, nil, + ) + toolMaps, ok := req["tools"].([]map[string]interface{}) + if !ok { + t.Fatal("expected tools as array of maps") + } + if len(toolMaps) != 1 { + t.Fatalf("expected 1 tool, got %d", len(toolMaps)) + } + if toolMaps[0]["cache_control"] == nil { + t.Error("expected cache_control on last tool") + } +} + +func TestBuildAnthropicCachedRequest_WithAllOptions(t *testing.T) { + t.Parallel() + temp := 0.7 + thinking := &AnthropicThinking{Type: "enabled", BudgetTokens: 1024} + toolChoice := &AnthropicToolChoice{Type: "auto"} + topP := 0.9 + topK := 5 + stopSeqs := []string{"\n\n"} + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{{Role: "user", Content: "Hello"}}, + "claude-3", 256, &temp, false, nil, thinking, toolChoice, &topP, &topK, stopSeqs, + ) + if req["temperature"] != 0.7 { + t.Errorf("temperature = %v", req["temperature"]) + } + if req["thinking"] == nil { + t.Error("expected thinking") + } + if req["tool_choice"] == nil { + t.Error("expected tool_choice") + } + if req["top_p"] != 0.9 { + t.Errorf("top_p = %v", req["top_p"]) + } + if req["top_k"] != 5 { + t.Errorf("top_k = %v", req["top_k"]) + } + if req["stop_sequences"] == nil { + t.Error("expected stop_sequences") + } +} + +func TestBuildAnthropicCachedRequest_ApplyCacheBreakpoint(t *testing.T) { + t.Parallel() + // Test with string content + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{ + {Role: "user", Content: "First"}, + {Role: "user", Content: "Second"}, + }, + "claude-3", 256, nil, false, nil, nil, nil, nil, nil, nil, + ) + msgs := req["messages"].([]map[string]interface{}) + if len(msgs) < 2 { + t.Fatal("expected at least 2 messages") + } + // Second-to-last message should have cache_control + firstMsg := msgs[0] + content := firstMsg["content"] + contentArr, ok := content.([]map[string]interface{}) + if !ok { + t.Fatalf("expected content as array, got %T", content) + } + if contentArr[0]["cache_control"] == nil { + t.Error("expected cache_control on second-to-last message") + } +} + +func TestBuildAnthropicCachedRequest_ApplyCacheBreakpoint_Maps(t *testing.T) { + t.Parallel() + // Use ContentParts to trigger the []map[string]interface{} content path + req := BuildAnthropicCachedRequest( + []core.EyrieMessage{ + {Role: "user", Content: "First", ContentParts: []core.ContentPart{{Type: "text", Text: "First text"}}}, + {Role: "user", Content: "Second"}, + }, + "claude-3", 256, nil, false, nil, nil, nil, nil, nil, nil, + ) + msgs := req["messages"].([]map[string]interface{}) + if len(msgs) < 2 { + t.Fatal("expected at least 2 messages") + } + // Second-to-last message should have cache_control on its content + firstMsg := msgs[0] + content := firstMsg["content"] + contentArr, ok := content.([]map[string]interface{}) + if !ok { + t.Fatalf("expected content as array, got %T", content) + } + if contentArr[0]["cache_control"] == nil { + t.Error("expected cache_control on content array element") + } +} diff --git a/client/adapters/azure_test.go b/client/adapters/azure_test.go new file mode 100644 index 0000000..63ffdc8 --- /dev/null +++ b/client/adapters/azure_test.go @@ -0,0 +1,206 @@ +package adapters + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewAzureClient(t *testing.T) { + t.Parallel() + c := NewAzureClient("az-key", "https://my-azure.openai.azure.com", "2024-10-21") + if c == nil { + t.Fatal("NewAzureClient returned nil") + } + if c.apiKey != "az-key" { + t.Errorf("apiKey = %q, want az-key", c.apiKey) + } + if c.apiVersion != "2024-10-21" { + t.Errorf("apiVersion = %q, want 2024-10-21", c.apiVersion) + } + if !strings.HasSuffix(c.endpoint, "azure.com") { + t.Errorf("unexpected endpoint: %q", c.endpoint) + } +} + +func TestNewAzureClient_DefaultAPIVersion(t *testing.T) { + t.Parallel() + c := NewAzureClient("key", "https://example.openai.azure.com", "") + if c.apiVersion != "2024-10-21" { + t.Errorf("expected default api-version, got %q", c.apiVersion) + } +} + +func TestAzureClient_Name(t *testing.T) { + t.Parallel() + c := NewAzureClient("key", "https://example.openai.azure.com", "") + if c.Name() != "azure" { + t.Errorf("Name() = %q, want azure", c.Name()) + } +} + +func TestAzureClient_Chat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-azure-1", + "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "Hello Azure!"}, "finish_reason": "stop"}}, + "usage": map[string]int{"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + }), nil + }) + + c := NewAzureClient("az-key", "https://example.openai.azure.com", "") + c.httpClient = &http.Client{Transport: transport} + + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello Azure!" { + t.Errorf("content = %q, want Hello Azure!", resp.Content) + } + if resp.FinishReason != "stop" { + t.Errorf("finish_reason = %q, want stop", resp.FinishReason) + } + if resp.Usage == nil || resp.Usage.PromptTokens != 10 { + t.Errorf("unexpected usage: %+v", resp.Usage) + } +} + +func TestAzureClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewAzureClient("az-key", "https://example.openai.azure.com", "") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestAzureClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"message": "Invalid API key"}, + }), nil + }) + + c := NewAzureClient("bad-key", "https://example.openai.azure.com", "") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error for unauthorized") + } +} + +func TestAzureClient_StreamChat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + sse := "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ndata: {\"id\":\"2\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" Azure!\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(sse)), + }, nil + }) + + c := NewAzureClient("az-key", "https://example.openai.azure.com", "") + c.httpClient = &http.Client{Transport: transport} + + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "error" { + t.Fatalf("unexpected error: %s", event.Error) + } + if event.Type == "content" { + content += event.Content + } + } + if content != "Hello Azure!" { + t.Errorf("content = %q, want Hello Azure!", content) + } +} + +func TestAzureClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewAzureClient("az-key", "https://example.openai.azure.com", "") + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestAzureClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "data": []map[string]any{{"id": "gpt-4o", "status": "succeeded"}}, + }), nil + }) + + c := NewAzureClient("az-key", "https://example.openai.azure.com", "") + c.httpClient = &http.Client{Transport: transport} + + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestAzureClient_Ping_AuthError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"code": "401", "message": "Access denied"}, + }), nil + }) + + c := NewAzureClient("bad-key", "https://example.openai.azure.com", "") + c.httpClient = &http.Client{Transport: transport} + + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestAzureClient_APIVersionAndEndpoint(t *testing.T) { + t.Parallel() + c := NewAzureClient("key", "https://example.openai.azure.com", "2025-01-01") + if c.APIVersion() != "2025-01-01" { + t.Errorf("APIVersion = %q", c.APIVersion()) + } + if c.Endpoint() != "https://example.openai.azure.com" { + t.Errorf("Endpoint = %q", c.Endpoint()) + } +} + +func TestAzureClient_SetHTTPClientAndSetRetry(t *testing.T) { + t.Parallel() + c := NewAzureClient("key", "https://example.openai.azure.com", "") + c2 := NewAzureClient("key", "https://example.openai.azure.com", "") + + c.SetHTTPClient(c2.httpClient) + if c.httpClient != c2.httpClient { + t.Error("SetHTTPClient did not replace client") + } + + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 7}} + c.SetRetry(rc) + if c.retry.MaxRetries != 7 { + t.Errorf("expected MaxRetries=7, got %d", c.retry.MaxRetries) + } +} diff --git a/client/adapters/bedrock_test.go b/client/adapters/bedrock_test.go new file mode 100644 index 0000000..23c294a --- /dev/null +++ b/client/adapters/bedrock_test.go @@ -0,0 +1,567 @@ +package adapters + +import ( + "bytes" + "context" + "encoding/binary" + "encoding/json" + "hash/crc32" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewBedrockClient(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "session", "us-east-1") + if c == nil { + t.Fatal("NewBedrockClient returned nil") + } + if c.accessKeyID != "AKID" { + t.Errorf("accessKeyID = %q", c.accessKeyID) + } + if string(c.secretAccessKey) != "secret" { + t.Errorf("secretAccessKey = %q", string(c.secretAccessKey)) + } + if c.sessionToken != "session" { + t.Errorf("sessionToken = %q", c.sessionToken) + } + if c.region != "us-east-1" { + t.Errorf("region = %q", c.region) + } +} + +func TestBedrockClient_Name(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + if c.Name() != "anthropic-bedrock" { + t.Errorf("Name() = %q", c.Name()) + } +} + +func TestBedrockClient_String(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "my-secret-key-here", "", "us-east-1") + s := c.String() + if !strings.Contains(s, "AKID") { + t.Error("expected access key in string") + } + if !strings.Contains(s, "my-s") { + t.Error("expected partial secret in string") + } + if strings.Contains(s, "my-secret-key-here") { + t.Error("full secret should not appear") + } +} + +func TestBedrockClient_String_ShortSecret(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "abc", "", "us-east-1") + s := c.String() + if strings.Contains(s, "abc") { + t.Error("short secret should be fully masked") + } +} + +func TestBedrockClient_Chat_Success(t *testing.T) { + t.Parallel() + var capturedReq *http.Request + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + capturedReq = req + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_bedrock_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "Hello from Bedrock!"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewBedrockClient("AKID", "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "anthropic.claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from Bedrock!" { + t.Errorf("content = %q", resp.Content) + } + if capturedReq == nil { + t.Fatal("request not captured") + } + if capturedReq.Header.Get("Authorization") == "" { + t.Error("missing Authorization header (SigV4)") + } +} + +func TestBedrockClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestBedrockClient_Chat_EmptyRegion(t *testing.T) { + t.Parallel() + c := NewBedrockClient("", "", "", "") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error for empty credentials") + } +} + +func TestBedrockClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusForbidden, map[string]any{ + "error": map[string]string{"message": "AccessDenied"}, + }), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestBedrockClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestBedrockClient_Ping_IncompleteCredentials(t *testing.T) { + t.Parallel() + c := NewBedrockClient("", "", "", "") + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected error for incomplete credentials") + } +} + +func TestBedrockClient_Ping_Unauthorized(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestBedrockClient_modelURL(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + url := c.ModelURL("anthropic.claude-sonnet-4-20250514") + expected := "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-sonnet-4-20250514/invoke" + if url != expected { + t.Errorf("modelURL = %q, want %q", url, expected) + } +} + +func TestBedrockClient_BuildBody(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + body, err := c.BuildBody([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{Model: "claude", MaxTokens: 256}) + if err != nil { + t.Fatalf("BuildBody: %v", err) + } + var parsed map[string]interface{} + if err := json.Unmarshal(body, &parsed); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if parsed["anthropic_version"] != "bedrock-2023-05-31" { + t.Errorf("anthropic_version = %v", parsed["anthropic_version"]) + } +} + +func TestBedrockClient_HTTPClientAndRetry(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + if c.HTTPClient() == nil { + t.Error("HTTPClient is nil") + } + if c.Retry().MaxRetries != 3 { + t.Errorf("default MaxRetries = %d", c.Retry().MaxRetries) + } + if c.Region() != "us-east-1" { + t.Errorf("Region = %q", c.Region()) + } +} + +func TestCanonicalAWSHeaders(t *testing.T) { + t.Parallel() + h := http.Header{} + h.Set("Host", "bedrock-runtime.us-east-1.amazonaws.com") + h.Set("X-Amz-Date", "20250101T000000Z") + canonical, signed := CanonicalAWSHeaders(h) + if canonical == "" || signed == "" { + t.Error("expected non-empty canonical headers") + } + if !strings.Contains(signed, "host") { + t.Error("expected host in signed headers") + } + if !strings.Contains(signed, "x-amz-date") { + t.Error("expected x-amz-date in signed headers") + } +} + +func TestAWSCanonicalURI(t *testing.T) { + t.Parallel() + if AWSCanonicalURI("") != "/" { + t.Errorf(`AWSCanonicalURI("") = %q`, AWSCanonicalURI("")) + } + if AWSCanonicalURI("/path/to/model") != "/path/to/model" { + t.Errorf(`AWSCanonicalURI("/path") = %q`, AWSCanonicalURI("/path/to/model")) + } +} + +func TestSha256Hex(t *testing.T) { + t.Parallel() + hash := Sha256Hex([]byte("test")) + if len(hash) != 64 { + t.Errorf("hash length = %d", len(hash)) + } + if hash != "9f86d081884c7d659a2feaa0c55ad015a3bf4f1b2b0b822cd15d6c15b0f00a08" { + t.Errorf("hash = %q", hash) + } +} + +func TestAWSSigningKey(t *testing.T) { + t.Parallel() + key := AWSSigningKey("wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", "20250101", "us-east-1", "bedrock") + if len(key) != 32 { + t.Errorf("key length = %d", len(key)) + } +} + +func TestBedrockClient_Sign(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", "", "us-east-1") + req, _ := http.NewRequest("POST", "https://bedrock-runtime.us-east-1.amazonaws.com/model/claude/invoke", nil) + err := c.Sign(req, []byte(`{"test":true}`), time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC)) + if err != nil { + t.Fatalf("Sign: %v", err) + } + if req.Header.Get("Authorization") == "" { + t.Error("expected Authorization header") + } + if req.Header.Get("X-Amz-Date") == "" { + t.Error("expected X-Amz-Date header") + } + if req.Header.Get("X-Amz-Content-Sha256") == "" { + t.Error("expected X-Amz-Content-Sha256 header") + } +} + +func TestEventStreamReader(t *testing.T) { + t.Parallel() + frame := buildEventStreamFrame("test-payload") + reader := newEventStreamReader(readerFromBytes(frame)) + evt, err := reader.ReadEvent() + if err != nil { + t.Fatalf("ReadEvent: %v", err) + } + if string(evt.Payload) != "test-payload" { + t.Errorf("payload = %q", string(evt.Payload)) + } +} + +func TestEventStreamReader_InvalidCRC(t *testing.T) { + t.Parallel() + reader := newEventStreamReader(strings.NewReader("garbage")) + _, err := reader.ReadEvent() + if err == nil { + t.Fatal("expected error for garbage data") + } +} + +func TestEventStreamReader_WithHeaders(t *testing.T) { + t.Parallel() + payload := `{"type":"content_block_delta","delta":{"text":"hi"}}` + frame := buildEventStreamFrameWithHeaders(payload, map[string]string{":content-type": "application/json"}) + reader := newEventStreamReader(readerFromBytes(frame)) + evt, err := reader.ReadEvent() + if err != nil { + t.Fatalf("ReadEvent: %v", err) + } + if string(evt.Payload) != payload { + t.Errorf("payload = %q", string(evt.Payload)) + } + if evt.Headers[":content-type"] != "application/json" { + t.Errorf("content-type header = %q", evt.Headers[":content-type"]) + } +} + +func TestEventStreamReader_InvalidTotalLen(t *testing.T) { + t.Parallel() + frame := make([]byte, 12) // totalLen = 0 + binary.BigEndian.PutUint32(frame[8:12], crc32.ChecksumIEEE(frame[:8])) + reader := newEventStreamReader(readerFromBytes(frame)) + _, err := reader.ReadEvent() + if err == nil { + t.Fatal("expected error for invalid total length") + } +} + +func TestEventStreamReader_HeadersExceedTotal(t *testing.T) { + t.Parallel() + totalLen := uint32(20) + headersLen := uint32(100) // exceeds total + prelude := make([]byte, 12) + binary.BigEndian.PutUint32(prelude[0:4], totalLen) + binary.BigEndian.PutUint32(prelude[4:8], headersLen) + binary.BigEndian.PutUint32(prelude[8:12], crc32.ChecksumIEEE(prelude[:8])) + reader := newEventStreamReader(readerFromBytes(prelude)) + _, err := reader.ReadEvent() + if err == nil { + t.Fatal("expected error for headers exceeding total") + } +} + +// buildEventStreamFrameWithHeaders constructs a valid EventStream frame with string headers. +func buildEventStreamFrameWithHeaders(payload string, headers map[string]string) []byte { + payloadBytes := []byte(payload) + headerBytes := buildEventStreamHeaders(headers) + headersLen := len(headerBytes) + totalLen := 12 + headersLen + len(payloadBytes) + 4 + prelude := make([]byte, 12) + binary.BigEndian.PutUint32(prelude[0:4], uint32(totalLen)) + binary.BigEndian.PutUint32(prelude[4:8], uint32(headersLen)) + binary.BigEndian.PutUint32(prelude[8:12], crc32.ChecksumIEEE(prelude[:8])) + + crcInput := make([]byte, 0, totalLen-4) + crcInput = append(crcInput, prelude[:8]...) + crcInput = append(crcInput, headerBytes...) + crcInput = append(crcInput, payloadBytes...) + msgCRC := crc32.ChecksumIEEE(crcInput) + + result := make([]byte, 0, totalLen) + result = append(result, prelude...) + result = append(result, headerBytes...) + result = append(result, payloadBytes...) + result = append(result, byte(msgCRC>>24), byte(msgCRC>>16), byte(msgCRC>>8), byte(msgCRC)) + return result +} + +// buildEventStreamHeaders encodes string headers in EventStream binary format. +func buildEventStreamHeaders(headers map[string]string) []byte { + var buf []byte + for name, value := range headers { + buf = append(buf, byte(len(name))) + buf = append(buf, []byte(name)...) + buf = append(buf, 7) // string type + strLen := len(value) + buf = append(buf, byte(strLen>>8), byte(strLen)) + buf = append(buf, []byte(value)...) + } + return buf +} + +func TestBedrockClient_StreamChat_Success(t *testing.T) { + t.Parallel() + chunks := []string{ + `{"type":"message_start","message":{"usage":{"input_tokens":5}}}`, + `{"type":"content_block_start","index":0,"content_block":{"type":"text"}}`, + `{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}`, + `{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":" from Bedrock stream!"}}`, + `{"type":"content_block_stop","index":0}`, + `{"type":"message_delta","delta":{"stop_reason":"end_turn"}}`, + } + var eventStreamData []byte + for _, c := range chunks { + eventStreamData = append(eventStreamData, buildEventStreamFrame(c)...) + } + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/vnd.amazon.eventstream"}}, + Body: io.NopCloser(bytes.NewReader(eventStreamData)), + }, nil + }) + c := NewBedrockClient("AKID", "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "anthropic.claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "Hello from Bedrock stream!" { + t.Errorf("content = %q", content) + } +} + +func TestBedrockClient_StreamChat_Error(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusForbidden, map[string]any{ + "error": map[string]string{"message": "AccessDenied"}, + }), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestBedrockClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +// buildEventStreamFrame constructs a valid Amazon EventStream frame containing the given payload. +func buildEventStreamFrame(payload string) []byte { + payloadBytes := []byte(payload) + headersLen := 0 + totalLen := 12 + headersLen + len(payloadBytes) + 4 + prelude := make([]byte, 12) + binary.BigEndian.PutUint32(prelude[0:4], uint32(totalLen)) + binary.BigEndian.PutUint32(prelude[4:8], uint32(headersLen)) + + preludeCRC := crc32.ChecksumIEEE(prelude[:8]) + binary.BigEndian.PutUint32(prelude[8:12], preludeCRC) + + crcInput := make([]byte, 0, totalLen-4) + crcInput = append(crcInput, prelude[:8]...) + crcInput = append(crcInput, payloadBytes...) + msgCRC := crc32.ChecksumIEEE(crcInput) + + result := make([]byte, 0, totalLen) + result = append(result, prelude...) + result = append(result, payloadBytes...) + result = append(result, byte(msgCRC>>24), byte(msgCRC>>16), byte(msgCRC>>8), byte(msgCRC)) + return result +} + +// readerFromBytes returns an io.Reader from a byte slice. +func readerFromBytes(b []byte) *strings.Reader { + return strings.NewReader(string(b)) +} + +func TestBedrockClient_Chat_DefaultMaxTokens(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["max_tokens"] != float64(4096) { + t.Errorf("max_tokens = %v, expected 4096", body["max_tokens"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_bedrock_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "anthropic.claude-sonnet-4-20250514"}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestBedrockClient_Chat_SystemPrompt(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + sys, ok := body["system"].(string) + if !ok { + t.Fatalf("system is not a string: %T", body["system"]) + } + if !strings.Contains(sys, "You are a helpful assistant") { + t.Errorf("system = %q, expected to contain custom system", sys) + } + if !strings.Contains(sys, "system from message") { + t.Errorf("system = %q, expected to contain message system", sys) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_bedrock_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{ + {Role: "system", Content: "system from message"}, + {Role: "user", Content: "Hi"}, + }, core.ChatOptions{Model: "anthropic.claude-sonnet-4-20250514", MaxTokens: 256, System: "You are a helpful assistant"}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestBedrockClient_Chat_SystemOnlyOpts(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + sys, ok := body["system"].(string) + if !ok { + t.Fatalf("system is not a string: %T", body["system"]) + } + if sys != "only from opts" { + t.Errorf("system = %q, expected 'only from opts'", sys) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_bedrock_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "ok"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + c := NewBedrockClient("AKID", "secret", "", "us-east-1") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "anthropic.claude-sonnet-4-20250514", MaxTokens: 256, System: "only from opts"}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} diff --git a/client/adapters/deepseek_test.go b/client/adapters/deepseek_test.go new file mode 100644 index 0000000..b9f0c17 --- /dev/null +++ b/client/adapters/deepseek_test.go @@ -0,0 +1,191 @@ +package adapters + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewDeepSeekClient_WithAnthropicFallback(t *testing.T) { + t.Parallel() + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "https://api.deepseek.com/anthropic", nil) + if client == nil { + t.Fatal("NewDeepSeekClient returned nil") + } + if client.Name() != "deepseek" { + t.Errorf("expected name 'deepseek', got %q", client.Name()) + } + if client.router.OpenAI == nil { + t.Fatal("expected OpenAI client") + } + if client.router.Anthropic == nil { + t.Fatal("expected Anthropic client for fallback") + } +} + +func TestNewDeepSeekClient_WithoutAnthropicFallback(t *testing.T) { + t.Parallel() + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "", nil) + if client.router.Anthropic != nil { + t.Fatal("expected no Anthropic client when anthropicBase is empty") + } +} + +func TestDeepSeekClient_Name(t *testing.T) { + t.Parallel() + client := NewDeepSeekClient("key", "https://api.deepseek.com/v1", "", nil) + if client.Name() != "deepseek" { + t.Errorf("expected 'deepseek', got %q", client.Name()) + } +} + +func TestDeepSeekClient_ChatOpenAISuccess(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-1", + "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}}, + "usage": map[string]int{"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "", nil) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + resp, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "deepseek-chat", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello!" { + t.Errorf("content = %q, want Hello!", resp.Content) + } +} + +func TestDeepSeekClient_ChatFallbackToAnthropic(t *testing.T) { + t.Parallel() + var paths []string + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + paths = append(paths, req.URL.Path) + if strings.HasSuffix(req.URL.Path, "/chat/completions") { + return jsonResponse(http.StatusServiceUnavailable, map[string]any{ + "error": map[string]string{"message": "service unavailable"}, + }), nil + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "Hello from Anthropic!"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 1, "output_tokens": 2}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "https://api.deepseek.com/anthropic", nil) + // Disable retries to speed up fallback test + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + resp, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "deepseek-chat", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from Anthropic!" { + t.Errorf("content = %q, want Hello from Anthropic!", resp.Content) + } + if len(paths) < 2 { + t.Fatalf("expected fallback to anthropic, got %v paths", len(paths)) + } +} + +func TestDeepSeekClient_StreamChatFallbackToAnthropic(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + if strings.HasSuffix(req.URL.Path, "/chat/completions") { + return jsonResponse(http.StatusServiceUnavailable, map[string]any{ + "error": map[string]string{"message": "service unavailable"}, + }), nil + } + if strings.HasSuffix(req.URL.Path, "/messages") { + body := "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":10}}}\n\nevent: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello stream!\"}}\n\nevent: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-1", "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "OK"}, "finish_reason": "stop"}}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "https://api.deepseek.com/anthropic", nil) + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + result, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "deepseek-chat", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "error" { + t.Fatalf("unexpected stream error: %s", event.Error) + } + if event.Type == "content" { + content += event.Content + } + } + if content != "Hello stream!" { + t.Errorf("content = %q, want Hello stream!", content) + } +} + +func TestDeepSeekClient_PingOpenAISuccess(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]string{"status": "ok"}), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "", nil) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestDeepSeekClient_PingFallbackToAnthropic(t *testing.T) { + t.Parallel() + var paths []string + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + paths = append(paths, req.URL.Path) + // Only fail the OpenAI Ping, not the Anthropic Ping + if strings.HasSuffix(req.URL.Path, "/models") && strings.Contains(req.URL.String(), "api.deepseek.com/v1") { + return nil, &transportError{msg: "connection refused"} + } + return jsonResponse(http.StatusOK, map[string]string{"status": "ok"}), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "https://api.deepseek.com/anthropic", nil) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping fallback: %v", err) + } + if len(paths) < 2 { + t.Fatalf("expected anthropic fallback, got %v paths", len(paths)) + } +} diff --git a/client/adapters/dynamic_test.go b/client/adapters/dynamic_test.go new file mode 100644 index 0000000..34be594 --- /dev/null +++ b/client/adapters/dynamic_test.go @@ -0,0 +1,158 @@ +package adapters + +import ( + "os" + "testing" +) + +func TestDynamicProviderEnabled_EnvNotSet(t *testing.T) { + os.Unsetenv(DynamicProviderEnvVar) + if DynamicProviderEnabled() { + t.Error("expected false when env var is not set") + } +} + +func TestDynamicProviderEnabled_EnvSetTo1(t *testing.T) { + os.Setenv(DynamicProviderEnvVar, "1") + defer os.Unsetenv(DynamicProviderEnvVar) + if !DynamicProviderEnabled() { + t.Error("expected true when env var is '1'") + } +} + +func TestDynamicProviderEnabled_EnvSetToTrue(t *testing.T) { + os.Setenv(DynamicProviderEnvVar, "true") + defer os.Unsetenv(DynamicProviderEnvVar) + if !DynamicProviderEnabled() { + t.Error("expected true when env var is 'true'") + } +} + +func TestDynamicProviderEnabled_EnvSetToYes(t *testing.T) { + os.Setenv(DynamicProviderEnvVar, "yes") + defer os.Unsetenv(DynamicProviderEnvVar) + if !DynamicProviderEnabled() { + t.Error("expected true when env var is 'yes'") + } +} + +func TestDynamicProviderEnabled_EnvSetToNo(t *testing.T) { + os.Setenv(DynamicProviderEnvVar, "no") + defer os.Unsetenv(DynamicProviderEnvVar) + if DynamicProviderEnabled() { + t.Error("expected false when env var is 'no'") + } +} + +func TestFreezeRegistry(t *testing.T) { + FreezeRegistry() + if !registryFrozen.Load() { + t.Error("expected registry to be frozen after FreezeRegistry") + } + // Reset for other tests + registryFrozen.Store(false) +} + +func TestRegisterDynamicProvider_Success(t *testing.T) { + registryFrozen.Store(false) + // Save and restore the map + saved := OpenAICompatibleProviders + OpenAICompatibleProviders = make(map[string]ProviderRegistryConfig) + defer func() { OpenAICompatibleProviders = saved }() + + err := RegisterDynamicProvider("my-provider", "https://my-api.example.com", "MY_API_KEY") + if err != nil { + t.Fatalf("RegisterDynamicProvider failed: %v", err) + } + p, ok := OpenAICompatibleProviders["my-provider"] + if !ok { + t.Fatal("expected my-provider to be registered") + } + if p.Type != ProviderTypeOpenAICompatible { + t.Errorf("expected type openai-compatible, got %s", p.Type) + } + if p.BaseURL != "https://my-api.example.com" { + t.Errorf("expected base URL https://my-api.example.com, got %s", p.BaseURL) + } + if p.EnvKey != "MY_API_KEY" { + t.Errorf("expected env key MY_API_KEY, got %s", p.EnvKey) + } +} + +func TestRegisterDynamicProvider_Frozen(t *testing.T) { + registryFrozen.Store(true) + defer registryFrozen.Store(false) + + err := RegisterDynamicProvider("test", "https://example.com", "KEY") + if err == nil { + t.Fatal("expected error when registry is frozen") + } +} + +func TestRegisterDynamicProvider_EmptyBaseURL(t *testing.T) { + registryFrozen.Store(false) + err := RegisterDynamicProvider("test", "", "KEY") + if err == nil { + t.Fatal("expected error for empty baseURL") + } +} + +func TestRegisterDynamicProvider_InvalidURL(t *testing.T) { + registryFrozen.Store(false) + err := RegisterDynamicProvider("test", "not-a-url", "KEY") + if err == nil { + t.Fatal("expected error for invalid URL") + } +} + +func TestRegisterDynamicProvider_NoScheme(t *testing.T) { + registryFrozen.Store(false) + err := RegisterDynamicProvider("test", "example.com/api", "KEY") + if err == nil { + t.Fatal("expected error for URL without scheme") + } +} + +func TestOpenAIBaseFallbackURL_APIBASE(t *testing.T) { + os.Setenv("OPENAI_API_BASE", "https://api.example.com/v1") + defer os.Unsetenv("OPENAI_API_BASE") + os.Unsetenv("OPENAI_BASE_URL") + + u := OpenAIBaseFallbackURL() + if u != "https://api.example.com/v1" { + t.Errorf("expected OPENAI_API_BASE value, got %q", u) + } +} + +func TestOpenAIBaseFallbackURL_BASEURL(t *testing.T) { + os.Unsetenv("OPENAI_API_BASE") + os.Setenv("OPENAI_BASE_URL", "https://alt.example.com") + defer os.Unsetenv("OPENAI_BASE_URL") + + u := OpenAIBaseFallbackURL() + if u != "https://alt.example.com" { + t.Errorf("expected OPENAI_BASE_URL value, got %q", u) + } +} + +func TestOpenAIBaseFallbackURL_PrefersAPIBASE(t *testing.T) { + os.Setenv("OPENAI_API_BASE", "https://api.example.com") + defer os.Unsetenv("OPENAI_API_BASE") + os.Setenv("OPENAI_BASE_URL", "https://alt.example.com") + defer os.Unsetenv("OPENAI_BASE_URL") + + u := OpenAIBaseFallbackURL() + if u != "https://api.example.com" { + t.Errorf("expected OPENAI_API_BASE to take priority, got %q", u) + } +} + +func TestOpenAIBaseFallbackURL_NotSet(t *testing.T) { + os.Unsetenv("OPENAI_API_BASE") + os.Unsetenv("OPENAI_BASE_URL") + + u := OpenAIBaseFallbackURL() + if u != "" { + t.Errorf("expected empty string, got %q", u) + } +} diff --git a/client/adapters/gemini_test.go b/client/adapters/gemini_test.go new file mode 100644 index 0000000..064d6ca --- /dev/null +++ b/client/adapters/gemini_test.go @@ -0,0 +1,935 @@ +package adapters + +import ( + "context" + "encoding/json" + "io" + "net/http" + "os" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewGeminiClient(t *testing.T) { + t.Parallel() + c := NewGeminiClient("AIza-test", "https://custom.example/v1beta") + if c == nil { + t.Fatal("NewGeminiClient returned nil") + } + if c.apiKey != "AIza-test" { + t.Errorf("apiKey = %q", c.apiKey) + } + if c.baseURL != "https://custom.example/v1beta" { + t.Errorf("baseURL = %q", c.baseURL) + } +} + +func TestNewGeminiClient_EmptyBaseURL(t *testing.T) { + t.Parallel() + c := NewGeminiClient("AIza-test", "") + if c.baseURL != "https://generativelanguage.googleapis.com/v1beta" { + t.Errorf("baseURL = %q", c.baseURL) + } +} + +func TestGeminiClient_Name(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + if c.Name() != "gemini" { + t.Errorf("Name() = %q", c.Name()) + } +} + +func TestGeminiClient_VertexDetection(t *testing.T) { + t.Parallel() + vertex := NewGeminiClient("token", "https://us-central1-aiplatform.googleapis.com/v1beta") + nonVertex := NewGeminiClient("key", "https://generativelanguage.googleapis.com/v1beta") + if !vertex.isVertex() { + t.Error("expected Vertex detection for aiplatform URL") + } + if nonVertex.isVertex() { + t.Error("expected non-Vertex for generativelanguage URL") + } +} + +func TestGeminiClient_Chat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := `{"candidates":[{"content":{"role":"model","parts":[{"text":"Hello from Gemini!"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":10,"totalTokenCount":15}}` + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("AIza-key", "https://generativelanguage.googleapis.com/v1beta") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from Gemini!" { + t.Errorf("content = %q", resp.Content) + } + if resp.FinishReason != "end_turn" { + t.Errorf("finishReason = %q", resp.FinishReason) + } +} + +func TestGeminiClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestGeminiClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusForbidden, map[string]any{ + "error": map[string]string{"message": "API key not valid"}, + }), nil + }) + c := NewGeminiClient("bad-key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestGeminiClient_StreamChat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"Hello \"}]}}]}\n\ndata: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"stream!\"}]},\"finishReason\":\"STOP\"}],\"usageMetadata\":{\"promptTokenCount\":2,\"candidatesTokenCount\":5,\"totalTokenCount\":7}}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "Hello stream!" { + t.Errorf("content = %q", content) + } +} + +func TestGeminiClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestGeminiClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{"models": []map[string]any{}}), nil + }) + c := NewGeminiClient("AIza-key", "") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestGeminiClient_Ping_AuthError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + }) + c := NewGeminiClient("bad-key", "") + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestGeminiClient_HTTPClientAndRetry(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + if c.HTTPClient() == nil { + t.Error("HTTPClient is nil") + } + rc := c.Retry() + if rc.MaxRetries != 3 { + t.Errorf("default MaxRetries = %d", rc.MaxRetries) + } +} + +func TestGeminiClient_ParseResponse(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + data, _ := json.Marshal(map[string]any{ + "candidates": []map[string]any{ + { + "content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "Response"}}}, + "finishReason": "STOP", + }, + }, + "usageMetadata": map[string]int{ + "promptTokenCount": 10, "candidatesTokenCount": 20, "totalTokenCount": 30, + }, + }) + resp, err := c.parseResponse(data, "req_id") + if err != nil { + t.Fatalf("parseResponse: %v", err) + } + if resp.Content != "Response" { + t.Errorf("Content = %q", resp.Content) + } + if resp.FinishReason != "end_turn" { + t.Errorf("FinishReason = %q", resp.FinishReason) + } +} + +func TestGeminiClient_ParseResponse_NoCandidates(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + data, _ := json.Marshal(map[string]any{}) + _, err := c.parseResponse(data, "req_id") + if err == nil { + t.Fatal("expected error for no candidates") + } +} + +func TestGeminiClient_ParseResponse_Blocked(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "") + data, _ := json.Marshal(map[string]any{ + "promptFeedback": map[string]string{"blockReason": "SAFETY", "blockReasonMessage": "Content blocked"}, + }) + _, err := c.parseResponse(data, "req_id") + if err == nil { + t.Fatal("expected error for blocked prompt") + } +} + +func TestMapGeminiFinishReason(t *testing.T) { + t.Parallel() + tests := []struct{ in, want string }{ + {"STOP", "end_turn"}, + {"MAX_TOKENS", "max_tokens"}, + {"SAFETY", "content_filter"}, + {"OTHER", "OTHER"}, + {"", ""}, + } + for _, tt := range tests { + got := mapGeminiFinishReason(tt.in) + if got != tt.want { + t.Errorf("mapGeminiFinishReason(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestGeminiSharedParserEnabled(t *testing.T) { + t.Parallel() + os.Unsetenv(geminiSharedParserEnvVar) + if !geminiSharedParserEnabled() { + t.Error("expected enabled by default") + } + os.Setenv(geminiSharedParserEnvVar, "0") + if geminiSharedParserEnabled() { + t.Error("expected disabled for '0'") + } + os.Setenv(geminiSharedParserEnvVar, "false") + if geminiSharedParserEnabled() { + t.Error("expected disabled for 'false'") + } + os.Setenv(geminiSharedParserEnvVar, "no") + if geminiSharedParserEnabled() { + t.Error("expected disabled for 'no'") + } + os.Unsetenv(geminiSharedParserEnvVar) +} + +func TestProcessGeminiStream(t *testing.T) { + t.Parallel() + sseEvents := make(chan core.SSEEvent, 3) + ctx := context.Background() + events := ProcessGeminiStream(ctx, sseEvents, testLogger(t)) + sseEvents <- core.SSEEvent{Event: "data", Data: `{"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}`} + sseEvents <- core.SSEEvent{Event: "data", Data: `{"candidates":[{"content":{"parts":[{"text":" world"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":5,"totalTokenCount":7}}`} + close(sseEvents) + + var content string + for evt := range events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "Hello world" { + t.Errorf("content = %q", content) + } +} + +func TestProcessGeminiStream_ToolCall(t *testing.T) { + t.Parallel() + sseEvents := make(chan core.SSEEvent, 2) + ctx := context.Background() + events := ProcessGeminiStream(ctx, sseEvents, testLogger(t)) + sseEvents <- core.SSEEvent{Event: "data", Data: `{"candidates":[{"content":{"parts":[{"functionCall":{"name":"get_weather","args":{"city":"NYC"}}}]}}]}`} + close(sseEvents) + + var toolCalls int + for evt := range events { + if evt.Type == "tool_call" && evt.ToolCall != nil { + toolCalls++ + if evt.ToolCall.Name != "get_weather" { + t.Errorf("tool name = %q", evt.ToolCall.Name) + } + } + } + if toolCalls != 1 { + t.Errorf("tool calls = %d", toolCalls) + } +} + +func TestProcessGeminiStream_ToolCallWithUsage(t *testing.T) { + t.Parallel() + sseEvents := make(chan core.SSEEvent, 2) + ctx := context.Background() + events := ProcessGeminiStream(ctx, sseEvents, testLogger(t)) + sseEvents <- core.SSEEvent{Event: "data", Data: `{"candidates":[{"content":{"parts":[{"functionCall":{"name":"get_weather","args":{"city":"NYC"},"id":"call_1"}}]}}],"usageMetadata":{"promptTokenCount":1,"candidatesTokenCount":2,"totalTokenCount":3}}`} + close(sseEvents) + + var toolCalls int + var usage *core.EyrieUsage + for evt := range events { + if evt.Type == "tool_call" && evt.ToolCall != nil { + toolCalls++ + if evt.ToolCall.ID != "call_1" { + t.Errorf("tool call id = %q", evt.ToolCall.ID) + } + } + if evt.Type == "done" && evt.Usage != nil { + usage = evt.Usage + } + } + if toolCalls != 1 { + t.Errorf("tool calls = %d", toolCalls) + } + if usage == nil || usage.TotalTokens != 3 { + t.Errorf("usage = %+v", usage) + } +} + +func TestProcessGeminiStream_SSEError(t *testing.T) { + t.Parallel() + sseEvents := make(chan core.SSEEvent, 1) + ctx := context.Background() + events := ProcessGeminiStream(ctx, sseEvents, testLogger(t)) + sseEvents <- core.SSEEvent{Event: "error", Data: "parse error"} + close(sseEvents) + + gotError := false + for evt := range events { + if evt.Type == "error" { + gotError = true + } + } + if !gotError { + t.Error("expected error event") + } +} + +func TestGeminiClient_Chat_WithSystem(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["systemInstruction"] == nil { + t.Error("expected systemInstruction in request body") + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{ + {Role: "system", Content: "Be helpful"}, + {Role: "user", Content: "Hi"}, + }, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "ok" { + t.Errorf("content = %q", resp.Content) + } +} + +func TestGeminiClient_Chat_WithTools(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if _, ok := body["tools"]; !ok { + t.Error("expected tools in request body") + } + if _, ok := body["toolConfig"]; !ok { + t.Error("expected toolConfig in request body") + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"functionCall": map[string]any{"name": "get_weather", "args": map[string]string{"city": "NYC"}}}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Weather?"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", MaxTokens: 256, + Tools: []core.EyrieTool{{Name: "get_weather", Description: "Get weather", Parameters: map[string]interface{}{"type": "object"}}}, + ToolChoice: &core.ToolChoiceOption{Type: "any"}, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if len(resp.ToolCalls) != 1 || resp.ToolCalls[0].Name != "get_weather" { + t.Errorf("tool calls = %+v", resp.ToolCalls) + } +} + +func TestGeminiClient_Chat_WithResponseFormat(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + gc, ok := body["generationConfig"].(map[string]interface{}) + if !ok { + t.Fatal("expected generationConfig") + } + if gc["responseMimeType"] != "application/json" { + t.Errorf("responseMimeType = %v", gc["responseMimeType"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "json response"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "JSON please"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_schema", Schema: `{"type":"object"}`}, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "json response" { + t.Errorf("content = %q", resp.Content) + } +} + +func TestGeminiClient_Chat_WithTopLogProbs(t *testing.T) { + t.Parallel() + logprobs := 5 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + gc, ok := body["generationConfig"].(map[string]interface{}) + if !ok { + t.Fatal("expected generationConfig") + } + if gc["logprobs"] != float64(5) { + t.Errorf("logprobs = %v", gc["logprobs"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "logprobs done"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", + MaxTokens: 256, + TopLogProbs: &logprobs, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_WithPenalties(t *testing.T) { + t.Parallel() + penalty := 0.5 + seed := 42 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + gc, ok := body["generationConfig"].(map[string]interface{}) + if !ok { + t.Fatal("expected generationConfig") + } + if gc["presencePenalty"] != 0.5 { + t.Errorf("presencePenalty = %v", gc["presencePenalty"]) + } + if gc["seed"] != float64(42) { + t.Errorf("seed = %v", gc["seed"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", MaxTokens: 256, + PresencePenalty: &penalty, + Seed: &seed, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_VertexAuth(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.Header.Get("Authorization") != "Bearer token" { + t.Errorf("Authorization = %q", req.Header.Get("Authorization")) + } + if req.Header.Get("x-goog-api-key") != "" { + t.Error("unexpected x-goog-api-key for Vertex") + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "vertex"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("token", "https://us-central1-aiplatform.googleapis.com/v1beta") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "vertex" { + t.Errorf("content = %q", resp.Content) + } +} + +func TestGeminiClient_Chat_ToolResults(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + contents := body["contents"].([]interface{}) + parts := contents[0].(map[string]interface{})["parts"].([]interface{}) + fr := parts[0].(map[string]interface{})["functionResponse"].(map[string]interface{}) + if fr["name"] != "get_weather" { + t.Errorf("functionResponse name = %v", fr["name"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "done"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{ + {Role: "user", ToolResults: []core.ToolResult{{ToolUseID: "get_weather", Content: "72°F"}}}, + }, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_StreamChat_LegacyParser(t *testing.T) { + os.Setenv(geminiSharedParserEnvVar, "0") + defer os.Unsetenv(geminiSharedParserEnvVar) + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"legacy \"}]}}]}\ndata: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"stream\"}]},\"finishReason\":\"STOP\"}],\"usageMetadata\":{\"promptTokenCount\":1,\"candidatesTokenCount\":1,\"totalTokenCount\":2}}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "legacy stream" { + t.Errorf("content = %q", content) + } +} + +func TestGeminiClient_StreamChat_Legacy_InvalidJSON(t *testing.T) { + os.Setenv(geminiSharedParserEnvVar, "0") + defer os.Unsetenv(geminiSharedParserEnvVar) + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: not valid json\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var gotDone bool + for evt := range result.Events { + if evt.Type == "done" { + gotDone = true + } + } + if !gotDone { + t.Error("expected done event") + } +} + +func TestGeminiClient_StreamChat_Legacy_NoCandidates(t *testing.T) { + os.Setenv(geminiSharedParserEnvVar, "0") + defer os.Unsetenv(geminiSharedParserEnvVar) + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"candidates\":[]}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var gotDone bool + for evt := range result.Events { + if evt.Type == "done" { + gotDone = true + } + } + if !gotDone { + t.Error("expected done event") + } +} + +func TestGeminiClient_StreamChat_Legacy_FunctionCall(t *testing.T) { + os.Setenv(geminiSharedParserEnvVar, "0") + defer os.Unsetenv(geminiSharedParserEnvVar) + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"candidates\":[{\"content\":{\"parts\":[{\"functionCall\":{\"name\":\"get_weather\",\"args\":{\"city\":\"NYC\"},\"id\":\"call_1\"}}]}}]}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var toolCalls int + for evt := range result.Events { + if evt.Type == "tool_call" && evt.ToolCall != nil { + toolCalls++ + if evt.ToolCall.ID != "call_1" { + t.Errorf("tool call id = %q", evt.ToolCall.ID) + } + } + } + if toolCalls != 1 { + t.Errorf("tool calls = %d", toolCalls) + } +} + +func TestGeminiClient_StreamChat_Legacy_NoUsage(t *testing.T) { + os.Setenv(geminiSharedParserEnvVar, "0") + defer os.Unsetenv(geminiSharedParserEnvVar) + + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"hello\"}]}}]}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var gotContent, gotDone bool + for evt := range result.Events { + if evt.Type == "content" { + gotContent = true + } + if evt.Type == "done" { + gotDone = true + } + } + if !gotContent { + t.Error("expected content event") + } + if !gotDone { + t.Error("expected done event (from streamLoop fallback)") + } +} + +func TestGeminiClient_Chat_AssistantRole(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + contents := body["contents"].([]interface{}) + if len(contents) < 1 { + t.Fatal("no contents") + } + msg := contents[0].(map[string]interface{}) + if msg["role"] != "model" { + t.Errorf("role = %v", msg["role"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "assistant", Content: "Hello"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_UnknownRole(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + contents := body["contents"].([]interface{}) + if len(contents) < 1 { + t.Fatal("no contents") + } + msg := contents[0].(map[string]interface{}) + if msg["role"] != "user" { + t.Errorf("role = %v", msg["role"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "function", Content: "result"}}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_ContentParts(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + contents := body["contents"].([]interface{}) + msg := contents[0].(map[string]interface{}) + parts := msg["parts"].([]interface{}) + if len(parts) < 3 { + t.Fatalf("expected 3 parts, got %d", len(parts)) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{ + Role: "user", + ContentParts: []core.ContentPart{ + {Type: "text", Text: "hello"}, + {Type: "image_url", ImageURL: &core.ImageURLPart{URL: "https://example.com/img.png"}}, + {Type: "input_audio", InputAudio: &core.InputAudioPart{Data: "base64data", Format: "wav"}}, + }, + }}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_LegacyImages(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + contents := body["contents"].([]interface{}) + msg := contents[0].(map[string]interface{}) + parts := msg["parts"].([]interface{}) + if len(parts) < 1 { + t.Fatalf("expected parts") + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{ + Role: "user", + Images: []string{"base64imagedata"}, + }}, core.ChatOptions{Model: "gemini-2.0-flash", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_ToolChoiceNone(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["toolConfig"] == nil { + t.Fatal("expected toolConfig") + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", + MaxTokens: 256, + ToolChoice: &core.ToolChoiceOption{Type: "none"}, + Tools: []core.EyrieTool{{Name: "test", Description: "a test tool"}}, + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func TestGeminiClient_Chat_PenaltiesOnly(t *testing.T) { + t.Parallel() + freq := 0.5 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var body map[string]interface{} + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode: %v", err) + } + gc, ok := body["generationConfig"].(map[string]interface{}) + if !ok { + t.Fatal("expected generationConfig") + } + if gc["frequencyPenalty"] != 0.5 { + t.Errorf("frequencyPenalty = %v", gc["frequencyPenalty"]) + } + return jsonResponse(http.StatusOK, map[string]any{ + "candidates": []map[string]any{{"content": map[string]any{"role": "model", "parts": []map[string]any{{"text": "ok"}}}, "finishReason": "STOP"}}, + "usageMetadata": map[string]int{"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }), nil + }) + c := NewGeminiClient("key", "") + c.retry = core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}} + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "gemini-2.0-flash", + MaxTokens: 256, + FrequencyPenalty: &freq, + N: ptrInt(2), + LogProbs: ptrBool(true), + Seed: ptrInt(42), + }) + if err != nil { + t.Fatalf("Chat: %v", err) + } +} + +func ptrInt(v int) *int { return &v } +func ptrBool(v bool) *bool { return &v } diff --git a/client/adapters/mimo_test.go b/client/adapters/mimo_test.go index d2d9346..8f131c2 100644 --- a/client/adapters/mimo_test.go +++ b/client/adapters/mimo_test.go @@ -2,10 +2,14 @@ package adapters import ( "context" + "fmt" + "io" "net/http" + "strings" "testing" "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" ) func TestMiMoClientChatFallsBackToAnthropicOnParamIncorrect(t *testing.T) { @@ -72,3 +76,186 @@ func TestMiMoClientPreservesProtocolBaseURLs(t *testing.T) { t.Fatalf("Anthropic base URL = %q", got) } } + +func TestMiMoClient_Name(t *testing.T) { + t.Parallel() + client := NewMiMoClient("key", "https://oai.example/v1", "https://ant.example", &XiaomiCompat, "providerA") + if client.Name() != "providerA" { + t.Errorf("Name() = %q, want providerA", client.Name()) + } +} + +func TestMiMoClient_Name_NoOpenAI(t *testing.T) { + t.Parallel() + c := &MiMoClient{providerID: "bare-id"} + if c.Name() != "bare-id" { + t.Errorf("Name() = %q, want bare-id", c.Name()) + } +} + +func TestMiMoClient_ProviderID(t *testing.T) { + t.Parallel() + client := NewMiMoClient("key", "https://oai.example/v1", "", &XiaomiCompat, "my_provider") + if client.ProviderID() != "my_provider" { + t.Errorf("ProviderID = %q", client.ProviderID()) + } +} + +func TestMiMoClient_Ping_SuccessViaOpenAI(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{"data": []map[string]any{}}), nil + }) + client := NewMiMoClient("key", "https://oai.example/v1", "https://ant.example", &XiaomiCompat, "p") + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestMiMoClient_Ping_FallbackToAnthropic(t *testing.T) { + t.Parallel() + anthropicPinged := false + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return nil, fmt.Errorf("connection refused") + }) + anthropicTransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + anthropicPinged = true + return jsonResponse(http.StatusOK, nil), nil + }) + client := NewMiMoClient("key", "https://oai.example/v1", "https://ant.example", &XiaomiCompat, "p") + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + client.router.Anthropic.httpClient = &http.Client{Transport: anthropicTransport} + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } + if !anthropicPinged { + t.Fatal("expected Anthropic ping fallback") + } +} + +func TestMiMoClient_Ping_NoFallbackOnNonRetryable(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return nil, fmt.Errorf("some unknown error") + }) + client := NewMiMoClient("key", "https://oai.example/v1", "https://ant.example", &XiaomiCompat, "p") + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + err := client.Ping(context.Background()) + if err == nil { + t.Fatal("expected ping error when no fallback") + } +} + +func TestMiMoClient_Ping_NoAnthropicNoFallback(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return nil, fmt.Errorf("some unknown error") + }) + client := NewMiMoClient("key", "https://oai.example/v1", "", &XiaomiCompat, "p") + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + err := client.Ping(context.Background()) + if err == nil { + t.Fatal("expected error without anthropic fallback client") + } +} + +func TestMiMoClient_StreamChat_FallbackToAnthropic(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"message": "param incorrect"}}), nil + }) + anthropicTransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1}}}\n\nevent: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"stream ok\"}}\n\nevent: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + + client := NewMiMoClient("key", "https://oai.example/v1", "https://ant.example", &XiaomiCompat, "p") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.Anthropic.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + client.router.Anthropic.httpClient = &http.Client{Transport: anthropicTransport} + + result, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{Model: "mimo-pro", MaxTokens: 64}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "stream ok" { + t.Errorf("content = %q, want stream ok", content) + } +} + +func TestMimoFallbackChatError(t *testing.T) { + tests := []struct { + err error + want bool + }{ + {nil, false}, + {fmt.Errorf("param incorrect"), true}, + {fmt.Errorf("invalid format"), true}, + {fmt.Errorf("reasoning_content"), true}, + {fmt.Errorf("HTTP 400 xiaomi"), true}, + {fmt.Errorf("HTTP 400"), false}, + {fmt.Errorf("something else"), false}, + {fmt.Errorf(""), false}, + } + for _, tt := range tests { + got := MimoFallbackChatError(tt.err) + if got != tt.want { + t.Errorf("MimoFallbackChatError(%v) = %v, want %v", tt.err, got, tt.want) + } + } +} + +func TestMimoRetryableChatError_HTTPStatus(t *testing.T) { + err := fmt.Errorf("HTTP 401") + if !MimoRetryableChatError(err) { + t.Error("expected 401 to be retryable") + } +} + +func TestMimoRetryableChatError_TransientMessage(t *testing.T) { + err := fmt.Errorf("connection refused") + if !MimoRetryableChatError(err) { + t.Error("expected connection refused to be retryable") + } +} + +func TestMimoRetryableChatError_NonTransient(t *testing.T) { + err := fmt.Errorf("some unknown error") + if MimoRetryableChatError(err) { + t.Error("expected unknown error to NOT be retryable") + } +} + +func TestParseHTTPStatusFromError(t *testing.T) { + tests := []struct { + msg string + want int + }{ + {"HTTP 404 Not Found", 404}, + {"status 500 internal", 500}, + {"error (403) forbidden", 403}, + {"no status here", 0}, + {"", 0}, + } + for _, tt := range tests { + got := parseHTTPStatusFromError(tt.msg) + if got != tt.want { + t.Errorf("parseHTTPStatusFromError(%q) = %d, want %d", tt.msg, got, tt.want) + } + } +} diff --git a/client/adapters/openai_embedding_test.go b/client/adapters/openai_embedding_test.go new file mode 100644 index 0000000..8008116 --- /dev/null +++ b/client/adapters/openai_embedding_test.go @@ -0,0 +1,165 @@ +package adapters + +import ( + "context" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" +) + +func TestCreateEmbedding_Validation(t *testing.T) { + t.Parallel() + client := &OpenAIClient{providerName: "test"} + _, err := client.CreateEmbedding(context.Background(), core.EmbeddingRequest{Model: "", Input: []string{"hello"}}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestCreateEmbedding_Success(t *testing.T) { + t.Parallel() + var gotPath, gotMethod string + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + gotPath = req.URL.Path + gotMethod = req.Method + return jsonResponse(http.StatusOK, map[string]any{ + "object": "list", + "data": []map[string]any{ + {"object": "embedding", "index": 0, "embedding": []float64{0.1, 0.2, 0.3}}, + }, + "model": "text-embedding-3-small", + "usage": map[string]int{"prompt_tokens": 1, "total_tokens": 1}, + }), nil + }) + + client := &OpenAIClient{ + providerName: "openai", + baseURL: "https://api.openai.com/v1", + httpClient: &http.Client{Transport: transport}, + logger: testLogger(t), + } + + resp, err := client.CreateEmbedding(context.Background(), core.EmbeddingRequest{ + Model: "text-embedding-3-small", + Input: []string{"hello world"}, + }) + if err != nil { + t.Fatalf("CreateEmbedding: %v", err) + } + if !strings.HasSuffix(gotPath, "/embeddings") { + t.Errorf("path = %q, want suffix /embeddings", gotPath) + } + if gotMethod != "POST" { + t.Errorf("method = %q, want POST", gotMethod) + } + if resp.Model != "text-embedding-3-small" { + t.Errorf("model = %q, want text-embedding-3-small", resp.Model) + } + if len(resp.Embeddings) != 1 || len(resp.Embeddings[0]) != 3 { + t.Fatalf("expected 1 embedding of dim 3, got %d embeddings of dim %d", len(resp.Embeddings), len(resp.Embeddings[0])) + } + if resp.Usage == nil || resp.Usage.PromptTokens != 1 { + t.Errorf("expected prompt_tokens=1, got %+v", resp.Usage) + } +} + +func TestCreateEmbedding_MultipleInputs(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "object": "list", + "data": []map[string]any{ + {"object": "embedding", "index": 0, "embedding": []float64{1.0, 0.0}}, + {"object": "embedding", "index": 1, "embedding": []float64{0.0, 1.0}}, + }, + "model": "text-embedding-3-small", + "usage": map[string]int{"prompt_tokens": 2, "total_tokens": 2}, + }), nil + }) + + client := &OpenAIClient{ + providerName: "openai", + baseURL: "https://api.openai.com/v1", + httpClient: &http.Client{Transport: transport}, + logger: testLogger(t), + } + + resp, err := client.CreateEmbedding(context.Background(), core.EmbeddingRequest{ + Model: "text-embedding-3-small", + Input: []string{"a", "b"}, + }) + if err != nil { + t.Fatalf("CreateEmbedding: %v", err) + } + if len(resp.Embeddings) != 2 { + t.Fatalf("expected 2 embeddings, got %d", len(resp.Embeddings)) + } + if resp.Embeddings[0][0] != 1.0 || resp.Embeddings[1][1] != 1.0 { + t.Errorf("unexpected embeddings: %v", resp.Embeddings) + } +} + +func TestCreateEmbedding_WithExtraParams(t *testing.T) { + t.Parallel() + var sentBody struct { + Model string `json:"model"` + Input []string `json:"input"` + Dimensions string `json:"dimensions"` + } + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + jsonDecodeRequest(req, &sentBody) + return jsonResponse(http.StatusOK, map[string]any{ + "object": "list", + "data": []map[string]any{{"object": "embedding", "index": 0, "embedding": []float64{0.5}}}, + "model": "text-embedding-3-small", + "usage": map[string]int{"prompt_tokens": 1, "total_tokens": 1}, + }), nil + }) + + client := &OpenAIClient{ + providerName: "openai", + baseURL: "https://api.openai.com/v1", + httpClient: &http.Client{Transport: transport}, + logger: testLogger(t), + } + + _, err := client.CreateEmbedding(context.Background(), core.EmbeddingRequest{ + Model: "text-embedding-3-small", + Input: []string{"hello"}, + Params: map[string]string{ + "dimensions": "256", + }, + }) + if err != nil { + t.Fatalf("CreateEmbedding: %v", err) + } + if sentBody.Dimensions != "256" { + t.Errorf("expected dimensions=256, got %q", sentBody.Dimensions) + } +} + +func TestCreateEmbedding_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"message": "Invalid API key"}, + }), nil + }) + + client := &OpenAIClient{ + providerName: "openai", + baseURL: "https://api.openai.com/v1", + httpClient: &http.Client{Transport: transport}, + logger: testLogger(t), + } + + _, err := client.CreateEmbedding(context.Background(), core.EmbeddingRequest{ + Model: "text-embedding-3-small", + Input: []string{"hello"}, + }) + if err == nil { + t.Fatal("expected error for unauthorized request") + } +} diff --git a/client/adapters/openai_test.go b/client/adapters/openai_test.go new file mode 100644 index 0000000..20ac817 --- /dev/null +++ b/client/adapters/openai_test.go @@ -0,0 +1,476 @@ +package adapters + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewOpenAIClient(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("sk-test", "https://custom.proxy.com/v1", &OpenAICompat) + if c == nil { + t.Fatal("NewOpenAIClient returned nil") + } + if c.apiKey != "sk-test" { + t.Errorf("apiKey = %q", c.apiKey) + } + if c.baseURL != "https://custom.proxy.com/v1" { + t.Errorf("baseURL = %q", c.baseURL) + } + if c.providerName != "openai" { + t.Errorf("providerName = %q", c.providerName) + } +} + +func TestNewOpenAIClient_EmptyBaseURL(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("sk-test", "", nil) + if c.baseURL != "https://api.openai.com/v1" { + t.Errorf("baseURL = %q", c.baseURL) + } + if c.compat == nil { + t.Error("expected default compat config") + } +} + +func TestOpenAIClient_Name(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("key", "", nil) + if c.Name() != "openai" { + t.Errorf("Name() = %q", c.Name()) + } +} + +func TestOpenAIClient_Chat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-123", + "choices": []map[string]any{ + { + "message": map[string]any{ + "content": "Hello from OpenAI!", + }, + "finish_reason": "stop", + }, + }, + "usage": map[string]int{"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}, + }), nil + }) + c := NewOpenAIClient("sk-test", "https://api.openai.com/v1", nil) + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from OpenAI!" { + t.Errorf("content = %q", resp.Content) + } + if resp.FinishReason != "stop" { + t.Errorf("finishReason = %q", resp.FinishReason) + } +} + +func TestOpenAIClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("key", "", nil) + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestOpenAIClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusTooManyRequests, map[string]any{ + "error": map[string]string{"message": "rate limited"}, + }), nil + }) + c := NewOpenAIClient("sk-test", "https://api.openai.com/v1", nil) + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestOpenAIClient_StreamChat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"},\"index\":0}]}\n\ndata: {\"choices\":[{\"delta\":{\"content\":\" stream\"},\"index\":0}]}\n\ndata: [DONE]\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + c := NewOpenAIClient("sk-test", "https://api.openai.com/v1", nil) + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + var content string + for evt := range result.Events { + if evt.Type == "content" { + content += evt.Content + } + } + if content != "Hello stream" { + t.Errorf("content = %q", content) + } +} + +func TestOpenAIClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("key", "", nil) + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestOpenAIClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{"data": []map[string]any{}}), nil + }) + c := NewOpenAIClient("sk-test", "https://api.openai.com/v1", nil) + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestOpenAIClient_Ping_AuthError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + }) + c := NewOpenAIClient("bad-key", "https://api.openai.com/v1", nil) + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestOpenAIClient_SetHTTPClientAndRetry(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("key", "", nil) + c2 := NewOpenAIClient("key2", "", nil) + c.SetHTTPClient(c2.httpClient) + if c.httpClient != c2.httpClient { + t.Error("SetHTTPClient did not replace client") + } + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 3}}) + if c.retry.MaxRetries != 3 { + t.Errorf("MaxRetries=%d", c.retry.MaxRetries) + } +} + +func TestBuildRequestBase(t *testing.T) { + t.Parallel() + temp := 0.7 + topP := 0.9 + req := BuildRequestBase([]core.EyrieMessage{ + {Role: "user", Content: "hello"}, + }, core.ChatOptions{ + Model: "gpt-4o", MaxTokens: 256, Temperature: &temp, + TopP: &topP, StopSequences: []string{"\n"}, + }, false, &OpenAICompat) + if req.Model != "gpt-4o" { + t.Errorf("Model = %q", req.Model) + } + if req.Temperature == nil || *req.Temperature != 0.7 { + t.Errorf("Temperature = %v", req.Temperature) + } + if req.MaxCompletionTokens == nil || *req.MaxCompletionTokens != 256 { + t.Errorf("MaxCompletionTokens = %v", req.MaxCompletionTokens) + } +} + +func TestBuildRequestBase_MaxCompletionTokens(t *testing.T) { + t.Parallel() + compat := &OpenAICompatConfig{MaxTokensField: "max_completion_tokens"} + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{Model: "o1", MaxTokens: 512}, false, compat) + if req.MaxCompletionTokens == nil || *req.MaxCompletionTokens != 512 { + t.Errorf("MaxCompletionTokens = %v", req.MaxCompletionTokens) + } + if req.MaxTokens != nil { + t.Errorf("MaxTokens should be nil") + } +} + +func TestBuildRequestBase_StreamOptions(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}, true, nil) + if req.StreamOptions == nil || !req.StreamOptions.IncludeUsage { + t.Error("expected StreamOptions with IncludeUsage for streaming") + } +} + +func TestBuildRequestBase_NoStreamOptionsForIncompatible(t *testing.T) { + t.Parallel() + compat := &OpenAICompatConfig{SupportsUsageInStreaming: false} + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{Model: "deepseek", MaxTokens: 256}, true, compat) + if req.StreamOptions != nil { + t.Error("expected no StreamOptions for incompatible compat") + } +} + +func TestBuildRequestBase_ToolChoice(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{ + Model: "gpt-4o", MaxTokens: 256, + ToolChoice: &core.ToolChoiceOption{Type: "tool", Name: "get_weather"}, + Tools: []core.EyrieTool{{Name: "get_weather", Description: "Get weather", Parameters: map[string]interface{}{"type": "object"}}}, + }, false, &OpenAICompat) + if req.ToolChoice == nil { + t.Fatal("expected ToolChoice") + } + tc, ok := req.ToolChoice.(map[string]interface{}) + if !ok { + t.Fatalf("ToolChoice type = %T", req.ToolChoice) + } + if tc["type"] != "function" { + t.Errorf("ToolChoice type = %v", tc["type"]) + } +} + +func TestBuildRequestBase_ResponseFormat(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{ + Model: "gpt-4o", MaxTokens: 256, + ResponseFormat: &core.ResponseFormat{Type: "json_object"}, + }, false, &OpenAICompat) + if req.ResponseFormat == nil { + t.Fatal("expected ResponseFormat") + } +} + +func TestBuildRequestBase_ReasoningEffort(t *testing.T) { + t.Parallel() + compat := &OpenAICompatConfig{SupportsReasoningEffort: true} + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{ + Model: "o3", MaxTokens: 1024, ReasoningEffort: "high", + }, false, compat) + if req.ReasoningEffort != "high" { + t.Errorf("ReasoningEffort = %q", req.ReasoningEffort) + } +} + +func TestOpenAIToolChoice(t *testing.T) { + t.Parallel() + tc := OpenAIToolChoice(nil) + if tc != nil { + t.Errorf("expected nil, got %v", tc) + } + if got := OpenAIToolChoice(&core.ToolChoiceOption{Type: "auto"}); got != "auto" { + t.Errorf("auto = %v", got) + } + if got := OpenAIToolChoice(&core.ToolChoiceOption{Type: "none"}); got != "none" { + t.Errorf("none = %v", got) + } + if got := OpenAIToolChoice(&core.ToolChoiceOption{Type: "any"}); got != "required" { + t.Errorf("any = %v", got) + } + if got := OpenAIToolChoice(&core.ToolChoiceOption{Type: "tool"}); got != "required" { + t.Errorf("tool no name = %v", got) + } + if got := OpenAIToolChoice(&core.ToolChoiceOption{Type: "custom"}); got != "custom" { + t.Errorf("custom = %v", got) + } + withName := OpenAIToolChoice(&core.ToolChoiceOption{Type: "tool", Name: "get_weather"}) + m, ok := withName.(map[string]interface{}) + if !ok { + t.Fatalf("tool with name type = %T", withName) + } + if m["type"] != "function" { + t.Errorf("type = %v", m["type"]) + } +} + +func TestOpenAIClient_MimoAuthRetry(t *testing.T) { + t.Parallel() + calls := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + calls++ + if calls == 1 { + if req.Header.Get("api-key") == "" { + t.Error("expected api-key header on first call") + } + return jsonResponse(http.StatusUnauthorized, map[string]any{"error": "invalid"}), nil + } + if req.Header.Get("Authorization") != "Bearer tp-key" { + t.Errorf("expected Bearer on retry, got %q", req.Header.Get("Authorization")) + } + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-456", + "choices": []map[string]any{ + {"message": map[string]any{"content": "retried"}, "finish_reason": "stop"}, + }, + }), nil + }) + c := NewOpenAIClient("tp-key", "", nil) + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.SetMimoAuth() + c.httpClient = &http.Client{Transport: transport} + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "retried" { + t.Errorf("content = %q", resp.Content) + } + if calls != 2 { + t.Errorf("calls = %d, want 2", calls) + } +} + +func TestBuildRequestBase_ToolResults(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{ + {Role: "user", ToolResults: []core.ToolResult{{ToolUseID: "call_1", Content: "42"}}}, + }, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}, false, nil) + if len(req.Messages) != 1 { + t.Fatalf("messages = %d", len(req.Messages)) + } + if req.Messages[0]["role"] != "tool" { + t.Errorf("role = %v", req.Messages[0]["role"]) + } + if req.Messages[0]["tool_call_id"] != "call_1" { + t.Errorf("tool_call_id = %v", req.Messages[0]["tool_call_id"]) + } +} + +func TestBuildRequestBase_ToolUse(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{ + {Role: "assistant", Content: "Let me check", ToolUse: []core.ToolCall{{ID: "call_1", Name: "get_weather", Arguments: map[string]interface{}{"city": "NYC"}}}}, + }, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}, false, nil) + if len(req.Messages) != 1 { + t.Fatalf("messages = %d", len(req.Messages)) + } + toolCalls := req.Messages[0]["tool_calls"].([]map[string]interface{}) + if len(toolCalls) != 1 { + t.Fatalf("tool_calls = %d", len(toolCalls)) + } + if toolCalls[0]["id"] != "call_1" { + t.Errorf("tool_call id = %v", toolCalls[0]["id"]) + } +} + +func TestBuildRequestBase_ContentParts(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{ + {Role: "user", ContentParts: []core.ContentPart{ + {Type: "text", Text: "desc"}, + {Type: "image_url", ImageURL: &core.ImageURLPart{URL: "https://example.com/img.png", Detail: "high"}}, + }}, + }, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}, false, nil) + content := req.Messages[0]["content"].([]map[string]interface{}) + if len(content) != 2 { + t.Fatalf("content blocks = %d", len(content)) + } + if content[1]["type"] != "image_url" { + t.Errorf("block type = %v", content[1]["type"]) + } +} + +func TestBuildRequestBase_LegacyImages(t *testing.T) { + t.Parallel() + req := BuildRequestBase([]core.EyrieMessage{ + {Role: "user", Content: "Check", Images: []string{"https://example.com/img.png"}}, + }, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256}, false, nil) + content := req.Messages[0]["content"].([]map[string]interface{}) + if len(content) != 2 { + t.Fatalf("content blocks = %d", len(content)) + } +} + +func TestBuildRequestBase_CacheRole(t *testing.T) { + t.Parallel() + compat := &OpenAICompatConfig{SupportsCacheRole: true} + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{ + Model: "moonshot-v1", MaxTokens: 256, + KimiContextCacheID: "cache_abc", KimiCacheResetTTL: true, + }, false, compat) + if len(req.Messages) != 2 { + t.Fatalf("messages = %d", len(req.Messages)) + } + if req.Messages[0]["role"] != "cache" { + t.Errorf("first msg role = %v", req.Messages[0]["role"]) + } + if req.Messages[0]["reset_ttl"] != true { + t.Errorf("expected reset_ttl") + } +} + +func TestBuildRequestBase_ZAIThinking(t *testing.T) { + t.Parallel() + enabled := true + compat := &OpenAICompatConfig{ThinkingFormat: "zai"} + req := BuildRequestBase([]core.EyrieMessage{{Role: "user", Content: "hi"}}, core.ChatOptions{ + Model: "glm-4", MaxTokens: 256, + GLMThinkingEnabled: &enabled, + }, false, compat) + if req.Thinking == nil || req.Thinking["type"] != "enabled" { + t.Errorf("Thinking = %v", req.Thinking) + } +} + +func TestOpenAIClient_Ping_MimoAuthRetry(t *testing.T) { + t.Parallel() + calls := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + calls++ + if calls == 1 { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + } + return jsonResponse(http.StatusOK, map[string]any{"data": []map[string]any{}}), nil + }) + c := NewOpenAIClient("tp-key", "", nil) + c.SetMimoAuth() + c.httpClient = &http.Client{Transport: transport} + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } + if calls != 2 { + t.Errorf("calls = %d, want 2", calls) + } +} + +func TestOpenAIClient_BuildOpenAIRequest(t *testing.T) { + t.Parallel() + c := NewOpenAIClient("sk-test", "https://api.openai.com/v1", nil) + req, body, err := c.BuildOpenAIRequest(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "gpt-4o", MaxTokens: 256, Temperature: float64Ptr(0.5)}, false) + if err != nil { + t.Fatalf("BuildOpenAIRequest: %v", err) + } + if req == nil { + t.Fatal("expected non-nil request") + } + if req.Method != "POST" { + t.Errorf("method = %q", req.Method) + } + if req.Header.Get("Authorization") != "Bearer sk-test" { + t.Errorf("Authorization = %q", req.Header.Get("Authorization")) + } + if body == nil { + t.Error("expected non-nil body") + } +} diff --git a/client/adapters/opencodego_test.go b/client/adapters/opencodego_test.go index 3f25d65..7acd06c 100644 --- a/client/adapters/opencodego_test.go +++ b/client/adapters/opencodego_test.go @@ -2,12 +2,14 @@ package adapters import ( "context" + "fmt" "io" "net/http" "strings" "testing" "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" ) func TestOpenCodeGoClientRoutesMiniMaxToAnthropic(t *testing.T) { @@ -232,3 +234,68 @@ func TestOpenCodeGoClientStreamMiniMaxReasoningOnlyFallsBackToChat(t *testing.T) t.Fatalf("expected /messages then /chat/completions, got %v", paths) } } + +func TestOpenCodeGoClient_Name(t *testing.T) { + t.Parallel() + client := NewOpenCodeGoClient("key", "https://opencode.example/zen/go/v1") + if client.Name() != "opencodego" { + t.Errorf("Name() = %q, want opencodego", client.Name()) + } +} + +func TestOpenCodeGoClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{}), nil + }) + client := NewOpenCodeGoClient("key", "https://opencode.example/zen/go/v1") + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + if err := client.Ping(context.Background()); err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestOpenCodeGoClient_Ping_FallbackToAnthropic(t *testing.T) { + t.Parallel() + callCount := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + callCount++ + if callCount == 1 { + return jsonResponse(http.StatusUnauthorized, map[string]any{}), nil + } + return jsonResponse(http.StatusOK, map[string]any{}), nil + }) + client := NewOpenCodeGoClient("key", "https://openai.example/v1") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.Anthropic.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + if err := client.Ping(context.Background()); err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestOACompatUnsupportedError(t *testing.T) { + t.Parallel() + tests := []struct { + name string + err error + want bool + }{ + {"nil error", nil, false}, + {"401 status", fmt.Errorf("status=401"), true}, + {"http 401", fmt.Errorf("http 401"), true}, + {"oa-compat", fmt.Errorf("oa-compat unsupported"), true}, + {"not supported", fmt.Errorf("not supported"), true}, + {"other error", fmt.Errorf("something else"), false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := OACompatUnsupportedError(tt.err) + if got != tt.want { + t.Errorf("OACompatUnsupportedError(%v) = %v, want %v", tt.err, got, tt.want) + } + }) + } +} diff --git a/client/adapters/poolside_ext_test.go b/client/adapters/poolside_ext_test.go new file mode 100644 index 0000000..d354308 --- /dev/null +++ b/client/adapters/poolside_ext_test.go @@ -0,0 +1,82 @@ +package adapters + +import ( + "context" + "net/http" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" +) + +func TestPoolsideClient_Name(t *testing.T) { + t.Parallel() + c := NewPoolsideClient("psk", "https://poolside.example") + if c.Name() != "poolside" { + t.Errorf("Name() = %q, want poolside", c.Name()) + } +} + +func TestPoolsideClient_Chat(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-1", "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "Poolside!"}, "finish_reason": "stop"}}, + }), nil + }) + + c := NewPoolsideClient("psk", "https://poolside.example") + c.openAI.httpClient = &http.Client{Transport: transport} + + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "poolside/laguna-m.1", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Poolside!" { + t.Errorf("content = %q, want Poolside!", resp.Content) + } +} + +func TestPoolsideClient_Ping(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]string{"status": "ok"}), nil + }) + + c := NewPoolsideClient("psk", "https://poolside.example") + c.openAI.httpClient = &http.Client{Transport: transport} + + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestPoolsideClient_StreamChatContentful(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-1", "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}}, + }), nil + }) + + c := NewPoolsideClient("psk", "https://poolside.example") + c.openAI.httpClient = &http.Client{Transport: transport} + + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "poolside/laguna-m.1", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "content" { + content += event.Content + } + } + if content != "Answer" { + t.Errorf("content = %q, want Answer", content) + } +} diff --git a/client/adapters/protocol_router_test.go b/client/adapters/protocol_router_test.go index bad6055..007c867 100644 --- a/client/adapters/protocol_router_test.go +++ b/client/adapters/protocol_router_test.go @@ -151,3 +151,68 @@ func TestProtocolRouterNoFallbackWhenNil(t *testing.T) { t.Fatal("expected error without fallback") } } + +func TestStreamResultFromChat_FullResponse(t *testing.T) { + t.Parallel() + resp := &core.EyrieResponse{ + Thinking: "Let me think...", + Content: "Hello there!", + ToolCalls: []core.ToolCall{{Name: "get_weather", Arguments: map[string]interface{}{"city": "NYC"}}}, + Usage: &core.EyrieUsage{TotalTokens: 42}, + FinishReason: "", + } + result := streamResultFromChat(resp) + defer result.Close() + + var thinking, content string + var toolCalls int + var gotUsage, gotDone bool + var stopReason string + for evt := range result.Events { + switch evt.Type { + case "thinking": + thinking = evt.Thinking + case "content": + content = evt.Content + case "tool_call": + toolCalls++ + case "usage": + gotUsage = evt.Usage != nil + case "done": + gotDone = true + stopReason = evt.StopReason + } + } + if thinking != "Let me think..." { + t.Errorf("thinking = %q", thinking) + } + if content != "Hello there!" { + t.Errorf("content = %q", content) + } + if toolCalls != 1 { + t.Errorf("tool calls = %d", toolCalls) + } + if !gotUsage { + t.Error("expected usage event") + } + if !gotDone { + t.Error("expected done event") + } + if stopReason != "stop" { + t.Errorf("stop_reason = %q, expected 'stop'", stopReason) + } +} + +func TestStreamResultFromChat_NilResponse(t *testing.T) { + t.Parallel() + result := streamResultFromChat(nil) + defer result.Close() + + var eventCount int + for range result.Events { + eventCount++ + } + if eventCount != 0 { + t.Errorf("expected 0 events for nil response, got %d", eventCount) + } +} diff --git a/client/adapters/test_helpers_test.go b/client/adapters/test_helpers_test.go index dcef8ba..4f7a2cb 100644 --- a/client/adapters/test_helpers_test.go +++ b/client/adapters/test_helpers_test.go @@ -4,9 +4,21 @@ import ( "bytes" "encoding/json" "io" + "log/slog" "net/http" + "testing" ) +type transportError struct{ msg string } + +func (e *transportError) Error() string { return e.msg } +func (e *transportError) Timeout() bool { return false } +func (e *transportError) Temporary() bool { return true } + +func testLogger(t *testing.T) *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + type roundTripFunc func(*http.Request) (*http.Response, error) func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { @@ -35,3 +47,5 @@ func jsonDecodeRequest(req *http.Request, value any) error { } return json.Unmarshal(body, value) } + +func float64Ptr(v float64) *float64 { return &v } diff --git a/client/adapters/vertex_test.go b/client/adapters/vertex_test.go new file mode 100644 index 0000000..2395222 --- /dev/null +++ b/client/adapters/vertex_test.go @@ -0,0 +1,237 @@ +package adapters + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewVertexClient(t *testing.T) { + t.Parallel() + c := NewVertexClient("my-project", "us-central1", "ya29.token") + if c == nil { + t.Fatal("NewVertexClient returned nil") + } + if c.projectID != "my-project" { + t.Errorf("projectID = %q", c.projectID) + } + if c.region != "us-central1" { + t.Errorf("region = %q", c.region) + } + if c.token != "ya29.token" { + t.Errorf("token = %q", c.token) + } +} + +func TestVertexClient_Name(t *testing.T) { + t.Parallel() + c := NewVertexClient("p", "us-central1", "tok") + if c.Name() != "anthropic-vertex" { + t.Errorf("Name() = %q, want anthropic-vertex", c.Name()) + } +} + +func TestVertexClient_BaseURL(t *testing.T) { + t.Parallel() + c := NewVertexClient("my-project", "us-central1", "tok") + expected := "https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/anthropic/models" + if c.BaseURL() != expected { + t.Errorf("BaseURL = %q, want %q", c.BaseURL(), expected) + } +} + +func TestVertexClient_RegionAndProject(t *testing.T) { + t.Parallel() + c := NewVertexClient("proj", "europe-west4", "tok") + if c.Region() != "europe-west4" { + t.Errorf("Region = %q", c.Region()) + } + if c.ProjectID() != "proj" { + t.Errorf("ProjectID = %q", c.ProjectID()) + } +} + +func TestVertexClient_Chat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_vertex_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "Hello Vertex!"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 5, "output_tokens": 10}, + }), nil + }) + + c := NewVertexClient("proj", "us-central1", "tok") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + + resp, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello Vertex!" { + t.Errorf("content = %q, want Hello Vertex!", resp.Content) + } +} + +func TestVertexClient_Chat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewVertexClient("proj", "us-central1", "tok") + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestVertexClient_Chat_APIError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusForbidden, map[string]any{ + "error": map[string]string{"message": "permission denied"}, + }), nil + }) + + c := NewVertexClient("proj", "us-central1", "bad-tok") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + + _, err := c.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error for forbidden") + } +} + +func TestVertexClient_StreamChat_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + body := "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":10}}}\n\nevent: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello Vertex stream!\"}}\n\nevent: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }) + + c := NewVertexClient("proj", "us-central1", "tok") + c.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + c.httpClient = &http.Client{Transport: transport} + + result, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "claude-sonnet-4-20250514", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "error" { + t.Fatalf("unexpected error: %s", event.Error) + } + if event.Type == "content" { + content += event.Content + } + } + if content != "Hello Vertex stream!" { + t.Errorf("content = %q, want Hello Vertex stream!", content) + } +} + +func TestVertexClient_StreamChat_EmptyModel(t *testing.T) { + t.Parallel() + c := NewVertexClient("proj", "us-central1", "tok") + _, err := c.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: ""}) + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestVertexClient_Ping_Success(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "models": []map[string]any{{"name": "claude-sonnet-4-20250514"}}, + }), nil + }) + + c := NewVertexClient("proj", "us-central1", "tok") + c.httpClient = &http.Client{Transport: transport} + + err := c.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestVertexClient_Ping_AuthError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"message": "invalid credentials"}, + }), nil + }) + + c := NewVertexClient("proj", "us-central1", "bad-tok") + c.httpClient = &http.Client{Transport: transport} + + err := c.Ping(context.Background()) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestVertexClient_SetHTTPClientAndRetry(t *testing.T) { + t.Parallel() + c := NewVertexClient("proj", "us-central1", "tok") + c2 := NewVertexClient("proj", "us-central1", "tok") + + c.SetHTTPClient(c2.httpClient) + if c.httpClient != c2.httpClient { + t.Error("SetHTTPClient did not replace client") + } + if c.HTTPClient() != c2.httpClient { + t.Error("HTTPClient getter mismatch") + } + + rc := core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 7}} + c.SetRetry(rc) + if c.retry.MaxRetries != 7 { + t.Errorf("expected MaxRetries=7, got %d", c.retry.MaxRetries) + } +} + +func TestVertexClient_buildBody(t *testing.T) { + t.Parallel() + c := NewVertexClient("proj", "us-central1", "tok") + + body, err := c.BuildBody([]core.EyrieMessage{ + {Role: "user", Content: "hello"}, + }, core.ChatOptions{ + Model: "claude-sonnet-4-20250514", + MaxTokens: 256, + System: "Be helpful", + }, false) + if err != nil { + t.Fatalf("BuildBody: %v", err) + } + + var parsed map[string]interface{} + if err := jsonDecodeRequest(&http.Request{Body: io.NopCloser(strings.NewReader(string(body)))}, &parsed); err != nil { + t.Fatalf("parse body: %v", err) + } + + if parsed["anthropic_version"] != "vertex-2023-10-16" { + t.Errorf("expected anthropic_version, got %v", parsed["anthropic_version"]) + } + if parsed["model"] != "claude-sonnet-4-20250514" { + t.Errorf("model = %v", parsed["model"]) + } + if parsed["max_tokens"] != float64(256) { + t.Errorf("max_tokens = %v", parsed["max_tokens"]) + } +} diff --git a/client/adapters/zai_test.go b/client/adapters/zai_test.go new file mode 100644 index 0000000..49c51c9 --- /dev/null +++ b/client/adapters/zai_test.go @@ -0,0 +1,248 @@ +package adapters + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "testing" + + "github.com/GrayCodeAI/eyrie/client/core" + "github.com/GrayCodeAI/eyrie/types" +) + +func TestNewZAIClient_WithAnthropicFallback(t *testing.T) { + t.Parallel() + client := NewZAIClient("zai-key", "https://zai.example/paas/v4", "https://zai.example/api/anthropic", nil, "zai_payg") + if client == nil { + t.Fatal("NewZAIClient returned nil") + } + if client.router.OpenAI == nil { + t.Fatal("expected OpenAI client") + } + if client.router.Anthropic == nil { + t.Fatal("expected Anthropic client for fallback") + } +} + +func TestNewZAIClient_WithoutAnthropicFallback(t *testing.T) { + t.Parallel() + client := NewZAIClient("zai-key", "https://zai.example/paas/v4", "", nil, "zai_coding") + if client.router.Anthropic != nil { + t.Fatal("expected no Anthropic client when anthropicBase is empty") + } +} + +func TestZAIClient_Name(t *testing.T) { + t.Parallel() + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + if client.Name() == "" { + t.Fatal("expected non-empty Name") + } +} + +func TestZAIClient_ChatOpenAISuccess(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "chatcmpl-1", "object": "chat.completion", + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "Z.AI Hello!"}, "finish_reason": "stop"}}, + "usage": map[string]int{"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3}, + }), nil + }) + + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + resp, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "glm-4", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Z.AI Hello!" { + t.Errorf("content = %q, want Z.AI Hello!", resp.Content) + } +} + +func TestZAIClient_ChatFallbackToAnthropic(t *testing.T) { + t.Parallel() + var paths []string + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + paths = append(paths, req.URL.Path) + if strings.HasSuffix(req.URL.Path, "/chat/completions") { + return jsonResponse(http.StatusServiceUnavailable, map[string]any{ + "error": map[string]string{"message": "overloaded"}, + }), nil + } + if strings.HasSuffix(req.URL.Path, "/messages") { + return jsonResponse(http.StatusOK, map[string]any{ + "id": "msg_1", "type": "message", "role": "assistant", + "content": []map[string]string{{"type": "text", "text": "Hello from Z.AI Anthropic!"}}, + "stop_reason": "end_turn", + "usage": map[string]int{"input_tokens": 1, "output_tokens": 2}, + }), nil + } + return jsonResponse(http.StatusOK, map[string]any{ + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "OK"}, "finish_reason": "stop"}}, + }), nil + }) + + client := NewZAIClient("key", "https://zai.example/paas/v4", "https://zai.example/api/anthropic", nil, "zai_payg") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + resp, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "glm-4", MaxTokens: 256}) + if err != nil { + t.Fatalf("Chat: %v", err) + } + if resp.Content != "Hello from Z.AI Anthropic!" { + t.Errorf("content = %q, want Hello from Z.AI Anthropic!", resp.Content) + } + if len(paths) < 2 { + t.Fatalf("expected anthropic fallback, got %v paths", len(paths)) + } +} + +func TestZAIClient_StreamChatFallbackToAnthropic(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + if strings.HasSuffix(req.URL.Path, "/chat/completions") { + return jsonResponse(http.StatusServiceUnavailable, map[string]any{ + "error": map[string]string{"message": "overloaded"}, + }), nil + } + if strings.HasSuffix(req.URL.Path, "/messages") { + body := "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":10}}}\n\nevent: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello Z.AI stream!\"}}\n\nevent: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil + } + return jsonResponse(http.StatusOK, map[string]any{ + "choices": []map[string]any{{"message": map[string]string{"role": "assistant", "content": "OK"}, "finish_reason": "stop"}}, + }), nil + }) + + client := NewZAIClient("key", "https://zai.example/paas/v4", "https://zai.example/api/anthropic", nil, "zai_payg") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + result, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "glm-4", MaxTokens: 256}) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "error" { + t.Fatalf("unexpected stream error: %s", event.Error) + } + if event.Type == "content" { + content += event.Content + } + } + if content != "Hello Z.AI stream!" { + t.Errorf("content = %q, want Hello Z.AI stream!", content) + } +} + +func TestZAIClient_PingOpenAISuccess(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusOK, map[string]string{"status": "ok"}), nil + }) + + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping: %v", err) + } +} + +func TestZAIClient_PingFallbackToAnthropic(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + if strings.HasSuffix(req.URL.Path, "/models") && strings.Contains(req.URL.String(), "paas") { + return nil, &transportError{msg: "connection refused"} + } + return jsonResponse(http.StatusOK, map[string]string{"status": "ok"}), nil + }) + + client := NewZAIClient("key", "https://zai.example/paas/v4", "https://zai.example/api/anthropic", nil, "zai_payg") + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err != nil { + t.Fatalf("Ping fallback: %v", err) + } +} + +func TestZaiFallbackChatError(t *testing.T) { + t.Parallel() + tests := []struct { + name string + err error + want bool + }{ + {"nil error", nil, false}, + {"retriable eyrie error", &core.EyrieError{StatusCode: 503}, true}, + {"non-retriable eyrie error", &core.EyrieError{StatusCode: 400}, false}, + {"param incorrect", errors.New("param incorrect"), true}, + {"invalid format", errors.New("invalid format"), true}, + {"reasoning_content", errors.New("reasoning_content"), true}, + {"http 400 with zai", errors.New("http 400 zai error"), true}, + {"http 400 without zai", errors.New("http 400 generic"), false}, + {"generic error", errors.New("something else"), false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := zaiFallbackChatError(tt.err) + if got != tt.want { + t.Errorf("zaiFallbackChatError(%v) = %v, want %v", tt.err, got, tt.want) + } + }) + } +} + +func TestZaiRetryableChatError(t *testing.T) { + t.Parallel() + tests := []struct { + name string + err error + want bool + }{ + {"nil error", nil, false}, + {"http 500 in message", errors.New("HTTP 500"), true}, + {"http 401 in message", errors.New("HTTP 401"), true}, + {"http 403 in message", errors.New("HTTP 403"), true}, + {"http 400 in message", errors.New("HTTP 400"), false}, + {"http 200 in message", errors.New("HTTP 200"), false}, + {"retriable eyrie error", &core.EyrieError{StatusCode: 503}, true}, + {"non-retriable eyrie error", &core.EyrieError{StatusCode: 400}, false}, + {"transient error", &transportError{msg: "timeout"}, true}, + {"generic error", errors.New("something"), false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := zaiRetryableChatError(tt.err) + if got != tt.want { + t.Errorf("zaiRetryableChatError(%v) = %v, want %v", tt.err, got, tt.want) + } + }) + } +} + +func TestZAIClient_Name_NilOpenAI(t *testing.T) { + t.Parallel() + client := &ZAIClient{providerID: "zai_custom"} + if client.Name() != "zai_custom" { + t.Errorf("Name() = %q, want zai_custom", client.Name()) + } +} diff --git a/conversation/engine_test.go b/conversation/engine_test.go index adc7277..dcb34c1 100644 --- a/conversation/engine_test.go +++ b/conversation/engine_test.go @@ -168,7 +168,7 @@ func (m *maxTokensMockProvider) StreamChat(_ context.Context, msgs []client.Eyri sr := &client.StreamResult{Events: ch} // Wrap Close so we can count invocations. return &client.StreamResult{ - Events: sr.Events, + Events: sr.Events, RequestID: sr.RequestID, }, nil } @@ -258,10 +258,8 @@ func TestConversationEngine_ContextCancelClosesStream(t *testing.T) { } // blockingMockProvider returns a StreamResult whose Events channel blocks -// until the context is cancelled. It signals via a channel when Close is called. -type blockingMockProvider struct { - closed chan struct{} -} +// until the context is cancelled. +type blockingMockProvider struct{} func (b *blockingMockProvider) Name() string { return "blocking-mock" diff --git a/engine/convert_test.go b/engine/convert_test.go index 68b649d..e3463ab 100644 --- a/engine/convert_test.go +++ b/engine/convert_test.go @@ -222,9 +222,9 @@ func TestToClientOptions_NoOutputSchemaLeavesResponseFormatNil(t *testing.T) { func TestToClientOptions_ClonesSlicesAndMaps(t *testing.T) { // Verify that mutating the request after conversion does not affect the options. req := llm.GenerateRequest{ - Tools: []llm.EyrieTool{{Name: "a"}, {Name: "b"}}, - Options: llm.GenerationOptions{StopSequences: []string{"x", "y"}}, - OutputSchema: "orig", + Tools: []llm.EyrieTool{{Name: "a"}, {Name: "b"}}, + Options: llm.GenerationOptions{StopSequences: []string{"x", "y"}}, + OutputSchema: "orig", } route := Route{Provider: "test", Model: "test/model"} opts := toClientOptions(req, route, false) diff --git a/go.mod b/go.mod index aa9abf2..a6f0aea 100644 --- a/go.mod +++ b/go.mod @@ -3,7 +3,7 @@ module github.com/GrayCodeAI/eyrie go 1.26.5 require ( - github.com/GrayCodeAI/hawk-core-contracts v0.1.8 + github.com/GrayCodeAI/hawk-core-contracts v0.1.9 github.com/google/uuid v1.6.0 github.com/tiktoken-go/tokenizer v0.8.0 github.com/zalando/go-keyring v0.2.8 diff --git a/go.sum b/go.sum index 78e324e..1e8d6e2 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,5 @@ -github.com/GrayCodeAI/hawk-core-contracts v0.1.8 h1:SkDsGZJXL+3DYG0Fi3NXvNe/NlhP/KZn+Feofnx35Zc= -github.com/GrayCodeAI/hawk-core-contracts v0.1.8/go.mod h1:BXbh68YrCf+s9HVqND5F8DAvl2MnE5NcOwZZZB56HGA= +github.com/GrayCodeAI/hawk-core-contracts v0.1.9 h1:uXX/gtNM+3kxSEzu+rZkHykzcEaAbASn1lmPyOGMXvc= +github.com/GrayCodeAI/hawk-core-contracts v0.1.9/go.mod h1:BXbh68YrCf+s9HVqND5F8DAvl2MnE5NcOwZZZB56HGA= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/danieljoos/wincred v1.2.3 h1:v7dZC2x32Ut3nEfRH+vhoZGvN72+dQ/snVXo/vMFLdQ= diff --git a/operationsgraph/operations_graph.go b/operationsgraph/operations_graph.go index 3da70c1..1baf537 100644 --- a/operationsgraph/operations_graph.go +++ b/operationsgraph/operations_graph.go @@ -33,11 +33,11 @@ type OperationEdge struct { // OperationsGraph represents a graph of operations for eyrie. type OperationsGraph struct { mu sync.RWMutex - ID string `json:"id"` - Name string `json:"name"` + ID string `json:"id"` + Name string `json:"name"` Nodes map[string]*OperationNode `json:"nodes"` - Edges []OperationEdge `json:"edges"` - Attrs map[string]interface{} `json:"attrs,omitempty"` + Edges []OperationEdge `json:"edges"` + Attrs map[string]interface{} `json:"attrs,omitempty"` } // NewOperationsGraph creates a new operations graph. @@ -126,7 +126,7 @@ func (g *OperationsGraph) ToGraphSpec() *graphcontracts.GraphSpec { nodes = append(nodes, graphcontracts.NodeSpec{ ID: id, - Type: graphcontracts.NodeTypeOperations, + Type: graphcontracts.NodeTypeFunction, Name: node.Name, Config: config, }) @@ -142,9 +142,9 @@ func (g *OperationsGraph) ToGraphSpec() *graphcontracts.GraphSpec { } return &graphcontracts.GraphSpec{ - ID: g.ID, - Name: g.Name, - Nodes: nodes, - Edges: edges, + ID: g.ID, + Name: g.Name, + Nodes: nodes, + Edges: edges, } } From bdd12c221f6be84ff307c60280afc0358d67c3e9 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Mon, 27 Jul 2026 06:36:40 +0530 Subject: [PATCH 2/4] test: add coverage for ResolveEnvSecret, ResolveProviderModelEnvOverride, isReasoningOnlyStreamDiagnostic MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add tests for provider_registry.go and protocol_router.go functions: - ResolveEnvSecret: 0% → 100% (found/not found) - ResolveProviderModelEnvOverride: 0% → 83.3% (with/without model env) - isReasoningOnlyStreamDiagnostic: 0% → 100% (6 subtests) Overall coverage: 84.6% → 85.1% --- client/adapters/protocol_router_test.go | 24 +++++++++ client/adapters/provider_registry_test.go | 59 +++++++++++++++++++++++ 2 files changed, 83 insertions(+) create mode 100644 client/adapters/provider_registry_test.go diff --git a/client/adapters/protocol_router_test.go b/client/adapters/protocol_router_test.go index 007c867..9ef64d8 100644 --- a/client/adapters/protocol_router_test.go +++ b/client/adapters/protocol_router_test.go @@ -216,3 +216,27 @@ func TestStreamResultFromChat_NilResponse(t *testing.T) { t.Errorf("expected 0 events for nil response, got %d", eventCount) } } + +func TestIsReasoningOnlyStreamDiagnostic(t *testing.T) { + t.Parallel() + tests := []struct { + name string + message string + want bool + }{ + {"error_only_reasoning", "error_only_reasoning", true}, + {"reasoning tokens but no answer", "reasoning tokens but no answer", true}, + {"case insensitive", "ERROR_ONLY_REASONING", true}, + {"with whitespace", " error_only_reasoning ", true}, + {"normal message", "Hello world", false}, + {"empty", "", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := isReasoningOnlyStreamDiagnostic(tt.message) + if got != tt.want { + t.Errorf("isReasoningOnlyStreamDiagnostic(%q) = %v, want %v", tt.message, got, tt.want) + } + }) + } +} diff --git a/client/adapters/provider_registry_test.go b/client/adapters/provider_registry_test.go new file mode 100644 index 0000000..835eb0e --- /dev/null +++ b/client/adapters/provider_registry_test.go @@ -0,0 +1,59 @@ +package adapters + +import ( + "testing" + + "github.com/GrayCodeAI/eyrie/credentials" +) + +func TestResolveEnvSecret(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "openai_api_key": "sk-test-key", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := ResolveEnvSecret("OPENAI_API_KEY") + if got != "sk-test-key" { + t.Errorf("ResolveEnvSecret = %q, want sk-test-key", got) + } +} + +func TestResolveEnvSecret_NotFound(t *testing.T) { + store := &credentials.MapStore{Data: map[string]string{}} + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := ResolveEnvSecret("NONEXISTENT_KEY") + if got != "" { + t.Errorf("ResolveEnvSecret = %q, want empty", got) + } +} + +func TestResolveProviderModelEnvOverride(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "openai_model": "gpt-4o", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := ResolveProviderModelEnvOverride("openai") + if got != "gpt-4o" { + t.Errorf("ResolveProviderModelEnvOverride = %q, want gpt-4o", got) + } +} + +func TestResolveProviderModelEnvOverride_Empty(t *testing.T) { + store := &credentials.MapStore{Data: map[string]string{}} + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := ResolveProviderModelEnvOverride("openai") + if got != "" { + t.Errorf("ResolveProviderModelEnvOverride = %q, want empty", got) + } +} From 288f23afc6fe6bf57a6d9ff61943ddfa44b8ae1a Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Mon, 27 Jul 2026 06:57:06 +0530 Subject: [PATCH 3/4] test: add coverage for DetectProvider, newStreamWithReasoningFallback, fallback paths Add tests for: - DetectProvider: 0% -> 100% (9 subtests covering all providers, priority order, defaults) - newStreamWithReasoningFallback: 53.8% -> 92.3% (error fallback, fallback fails, non-reasoning error, normal content, tool calls) - reasoningOnlyFallbackChat: 80% -> 100% (MaxTokens bump from <512 to 512) - NewOpenCodeGoClient: 80% -> 100% (default base URL) - DeepSeekClient Chat/StreamChat/Ping: 80% -> 100% (no fallback when Anthropic nil, non-retriable errors) - MiMoClient Chat/StreamChat: 80% -> 100% (no fallback on non-retriable) - ZAIClient Chat/StreamChat/Ping: 80% -> 100% (no fallback on non-retriable) Overall coverage: 85.1% -> 88.0% --- client/adapters/deepseek_test.go | 55 +++++++ client/adapters/mimo_test.go | 30 ++++ client/adapters/opencodego_test.go | 11 ++ client/adapters/poolside_test.go | 55 +++++++ client/adapters/protocol_router_test.go | 183 ++++++++++++++++++++++ client/adapters/provider_registry_test.go | 135 ++++++++++++++++ client/adapters/zai_test.go | 45 ++++++ 7 files changed, 514 insertions(+) diff --git a/client/adapters/deepseek_test.go b/client/adapters/deepseek_test.go index b9f0c17..5e90a53 100644 --- a/client/adapters/deepseek_test.go +++ b/client/adapters/deepseek_test.go @@ -189,3 +189,58 @@ func TestDeepSeekClient_PingFallbackToAnthropic(t *testing.T) { t.Fatalf("expected anthropic fallback, got %v paths", len(paths)) } } + +func TestDeepSeekClient_StreamChatNoFallbackWhenAnthropicNil(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusServiceUnavailable, map[string]any{ + "error": map[string]string{"message": "service unavailable"}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "", nil) + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + _, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "deepseek-chat", MaxTokens: 256}) + if err == nil { + t.Fatal("expected error when no fallback available") + } +} + +func TestDeepSeekClient_PingNoFallbackWhenAnthropicNil(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"message": "unauthorized"}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "", nil) + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err == nil { + t.Fatal("expected error when no fallback available") + } +} + +func TestDeepSeekClient_PingNonRetriableError(t *testing.T) { + t.Parallel() + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{ + "error": map[string]string{"message": "unauthorized"}, + }), nil + }) + + client := NewDeepSeekClient("ds-key", "https://api.deepseek.com/v1", "https://api.deepseek.com/anthropic", nil) + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: transport} + client.router.Anthropic.httpClient = &http.Client{Transport: transport} + + err := client.Ping(context.Background()) + if err == nil { + t.Fatal("expected error for non-retriable error") + } +} diff --git a/client/adapters/mimo_test.go b/client/adapters/mimo_test.go index 8f131c2..62d9830 100644 --- a/client/adapters/mimo_test.go +++ b/client/adapters/mimo_test.go @@ -162,6 +162,36 @@ func TestMiMoClient_Ping_NoAnthropicNoFallback(t *testing.T) { } } +func TestMiMoClient_Chat_NoFallbackOnNonRetriable(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"message": "some other error"}}), nil + }) + client := NewMiMoClient("key", "https://oai.example/v1", "", &XiaomiCompat, "p") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + + _, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "mimo"}) + if err == nil { + t.Fatal("expected error without fallback") + } +} + +func TestMiMoClient_StreamChat_NoFallbackOnNonRetriable(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"message": "some other error"}}), nil + }) + client := NewMiMoClient("key", "https://oai.example/v1", "", &XiaomiCompat, "p") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + + _, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "mimo"}) + if err == nil { + t.Fatal("expected error without fallback") + } +} + func TestMiMoClient_StreamChat_FallbackToAnthropic(t *testing.T) { t.Parallel() openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { diff --git a/client/adapters/opencodego_test.go b/client/adapters/opencodego_test.go index 7acd06c..11cbe8a 100644 --- a/client/adapters/opencodego_test.go +++ b/client/adapters/opencodego_test.go @@ -299,3 +299,14 @@ func TestOACompatUnsupportedError(t *testing.T) { }) } } + +func TestNewOpenCodeGoClient_DefaultBaseURL(t *testing.T) { + t.Parallel() + client := NewOpenCodeGoClient("key", "") + if client == nil { + t.Fatal("expected non-nil client") + } + if client.Name() != "opencodego" { + t.Errorf("Name = %q, want opencodego", client.Name()) + } +} diff --git a/client/adapters/poolside_test.go b/client/adapters/poolside_test.go index cee2c10..d6d6aa4 100644 --- a/client/adapters/poolside_test.go +++ b/client/adapters/poolside_test.go @@ -70,3 +70,58 @@ func TestPoolsideClientReasoningOnlyStreamFallsBackToChat(t *testing.T) { t.Fatalf("requests = %d, want stream plus chat fallback", requests) } } + +func TestPoolsideClientReasoningOnlyStreamFallbackBumpsMaxTokens(t *testing.T) { + t.Parallel() + var requests int + var fallbackMaxTokens int + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + requests++ + var body struct { + MaxTokens int `json:"max_tokens"` + } + if err := jsonDecodeRequest(req, &body); err != nil { + t.Fatalf("decode request: %v", err) + } + if requests == 2 { + fallbackMaxTokens = body.MaxTokens + } + if requests == 1 { + return jsonResponse(http.StatusOK, map[string]any{ + "choices": []map[string]any{{ + "message": map[string]string{"role": "assistant", "reasoning_content": "thinking"}, + "finish_reason": "stop", + }}, + }), nil + } + return jsonResponse(http.StatusOK, map[string]any{ + "choices": []map[string]any{{ + "message": map[string]string{"role": "assistant", "content": "Hi"}, + "finish_reason": "stop", + }}, + }), nil + }) + + client := NewPoolsideClient("poolside-test-key", "https://poolside.example/v1") + client.openAI.httpClient = &http.Client{Transport: transport} + result, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{ + Model: "poolside/laguna-m.1", MaxTokens: 128, + }) + if err != nil { + t.Fatalf("StreamChat: %v", err) + } + defer result.Close() + + var content string + for event := range result.Events { + if event.Type == "content" { + content += event.Content + } + } + if content != "Hi" { + t.Fatalf("content = %q, want Hi", content) + } + if fallbackMaxTokens != 512 { + t.Fatalf("fallback max_tokens = %d, want 512 (bumped from 128)", fallbackMaxTokens) + } +} diff --git a/client/adapters/protocol_router_test.go b/client/adapters/protocol_router_test.go index 9ef64d8..6897d2d 100644 --- a/client/adapters/protocol_router_test.go +++ b/client/adapters/protocol_router_test.go @@ -217,6 +217,189 @@ func TestStreamResultFromChat_NilResponse(t *testing.T) { } } +func TestNewStreamWithReasoningFallbackErrorFallback(t *testing.T) { + t.Parallel() + primaryEvents := make(chan core.EyrieStreamEvent, 4) + primaryEvents <- core.EyrieStreamEvent{Type: "thinking", Thinking: "internal reasoning"} + primaryEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "end_turn"} + close(primaryEvents) + primary := llm.NewStreamResult(primaryEvents, "", func() {}) + + fallbackEvents := make(chan core.EyrieStreamEvent, 4) + fallbackEvents <- core.EyrieStreamEvent{Type: "content", Content: "fallback content"} + fallbackEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "stop"} + close(fallbackEvents) + + fallback := protocolStreamFallback{ + chat: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.EyrieResponse, error) { + return nil, fmt.Errorf("chat failed") + }, + stream: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.StreamResult, error) { + return llm.NewStreamResult(fallbackEvents, "", func() {}), nil + }, + } + + result := newStreamWithReasoningFallback(context.Background(), nil, core.ChatOptions{}, primary, fallback) + var content string + var gotError bool + for event := range result.Events { + if event.Type == "content" { + content += event.Content + } + if event.Type == "error" { + gotError = true + } + } + if content != "fallback content" { + t.Fatalf("content = %q, want fallback content", content) + } + if gotError { + t.Fatal("did not expect error event") + } +} + +func TestNewStreamWithReasoningFallbackErrorFallbackFails(t *testing.T) { + t.Parallel() + primaryEvents := make(chan core.EyrieStreamEvent, 4) + primaryEvents <- core.EyrieStreamEvent{Type: "thinking", Thinking: "internal reasoning"} + primaryEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "end_turn"} + close(primaryEvents) + primary := llm.NewStreamResult(primaryEvents, "", func() {}) + + fallback := protocolStreamFallback{ + chat: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.EyrieResponse, error) { + return nil, fmt.Errorf("chat failed") + }, + stream: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.StreamResult, error) { + return nil, fmt.Errorf("stream also failed") + }, + } + + result := newStreamWithReasoningFallback(context.Background(), nil, core.ChatOptions{}, primary, fallback) + var gotError bool + var errorMsg string + for event := range result.Events { + if event.Type == "error" { + gotError = true + errorMsg = event.Error + } + } + if !gotError { + t.Fatal("expected error event when both fallbacks fail") + } + if errorMsg != "stream also failed" { + t.Errorf("error = %q, want stream also failed", errorMsg) + } +} + +func TestNewStreamWithReasoningFallbackNonReasoningError(t *testing.T) { + t.Parallel() + primaryEvents := make(chan core.EyrieStreamEvent, 4) + primaryEvents <- core.EyrieStreamEvent{Type: "thinking", Thinking: "internal reasoning"} + primaryEvents <- core.EyrieStreamEvent{Type: "error", Error: "connection refused"} + primaryEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "end_turn"} + close(primaryEvents) + primary := llm.NewStreamResult(primaryEvents, "", func() {}) + + var fallbackCalled bool + fallback := protocolStreamFallback{ + chat: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.EyrieResponse, error) { + fallbackCalled = true + return &core.EyrieResponse{Content: "fallback"}, nil + }, + stream: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.StreamResult, error) { + fallbackCalled = true + return nil, fmt.Errorf("should not reach") + }, + } + + result := newStreamWithReasoningFallback(context.Background(), nil, core.ChatOptions{}, primary, fallback) + var gotError bool + for event := range result.Events { + if event.Type == "error" { + gotError = true + } + } + if !gotError { + t.Fatal("expected error event for non-reasoning error") + } + if fallbackCalled { + t.Fatal("fallback should not be called for non-reasoning error") + } +} + +func TestNewStreamWithReasoningFallbackNormalContent(t *testing.T) { + t.Parallel() + primaryEvents := make(chan core.EyrieStreamEvent, 4) + primaryEvents <- core.EyrieStreamEvent{Type: "thinking", Thinking: "internal reasoning"} + primaryEvents <- core.EyrieStreamEvent{Type: "content", Content: "actual answer"} + primaryEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "stop"} + close(primaryEvents) + primary := llm.NewStreamResult(primaryEvents, "", func() {}) + + var fallbackCalled bool + fallback := protocolStreamFallback{ + chat: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.EyrieResponse, error) { + fallbackCalled = true + return &core.EyrieResponse{Content: "fallback"}, nil + }, + stream: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.StreamResult, error) { + fallbackCalled = true + return nil, fmt.Errorf("should not reach") + }, + } + + result := newStreamWithReasoningFallback(context.Background(), nil, core.ChatOptions{}, primary, fallback) + var content string + for event := range result.Events { + if event.Type == "content" { + content += event.Content + } + } + if content != "actual answer" { + t.Fatalf("content = %q, want actual answer", content) + } + if fallbackCalled { + t.Fatal("fallback should not be called when content is present") + } +} + +func TestNewStreamWithReasoningFallbackToolCalls(t *testing.T) { + t.Parallel() + primaryEvents := make(chan core.EyrieStreamEvent, 4) + primaryEvents <- core.EyrieStreamEvent{Type: "thinking", Thinking: "internal reasoning"} + primaryEvents <- core.EyrieStreamEvent{Type: "tool_call", ToolCall: &core.ToolCall{Name: "test_tool"}} + primaryEvents <- core.EyrieStreamEvent{Type: "done", StopReason: "tool_calls"} + close(primaryEvents) + primary := llm.NewStreamResult(primaryEvents, "", func() {}) + + var fallbackCalled bool + fallback := protocolStreamFallback{ + chat: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.EyrieResponse, error) { + fallbackCalled = true + return &core.EyrieResponse{Content: "fallback"}, nil + }, + stream: func(context.Context, []core.EyrieMessage, core.ChatOptions) (*core.StreamResult, error) { + fallbackCalled = true + return nil, fmt.Errorf("should not reach") + }, + } + + result := newStreamWithReasoningFallback(context.Background(), nil, core.ChatOptions{}, primary, fallback) + var gotToolCall bool + for event := range result.Events { + if event.Type == "tool_call" { + gotToolCall = true + } + } + if !gotToolCall { + t.Fatal("expected tool_call event to be forwarded") + } + if fallbackCalled { + t.Fatal("fallback should not be called when tool calls are present") + } +} + func TestIsReasoningOnlyStreamDiagnostic(t *testing.T) { t.Parallel() tests := []struct { diff --git a/client/adapters/provider_registry_test.go b/client/adapters/provider_registry_test.go index 835eb0e..7e1fcab 100644 --- a/client/adapters/provider_registry_test.go +++ b/client/adapters/provider_registry_test.go @@ -57,3 +57,138 @@ func TestResolveProviderModelEnvOverride_Empty(t *testing.T) { t.Errorf("ResolveProviderModelEnvOverride = %q, want empty", got) } } + +func TestDetectProvider(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "openai_api_key": "sk-test-key", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "openai" { + t.Errorf("DetectProvider = %q, want openai", got) + } +} + +func TestDetectProvider_NoProvider(t *testing.T) { + store := &credentials.MapStore{Data: map[string]string{}} + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "anthropic" { + t.Errorf("DetectProvider = %q, want anthropic (default)", got) + } +} + +func TestDetectProvider_PriorityOrder(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "anthropic_api_key": "sk-ant-test", + "openai_api_key": "sk-test-key", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "anthropic" { + t.Errorf("DetectProvider = %q, want anthropic (priority order)", got) + } +} + +func TestDetectProvider_Ollama(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "ollama_base_url": "http://localhost:11434", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "ollama" { + t.Errorf("DetectProvider = %q, want ollama", got) + } +} + +func TestDetectProvider_Azure(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "azure_openai_api_key": "azure-key", + "azure_openai_endpoint": "https://test.openai.azure.com", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "azure" { + t.Errorf("DetectProvider = %q, want azure", got) + } +} + +func TestDetectProvider_Bedrock(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "aws_access_key_id": "AKIA-TEST", + "aws_secret_access_key": "secret", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "bedrock" { + t.Errorf("DetectProvider = %q, want bedrock", got) + } +} + +func TestDetectProvider_Vertex(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "vertex_project_id": "test-project", + "vertex_access_token": "token", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "vertex" { + t.Errorf("DetectProvider = %q, want vertex", got) + } +} + +func TestDetectProvider_Grok(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "xai_api_key": "xai-test", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "grok" { + t.Errorf("DetectProvider = %q, want grok", got) + } +} + +func TestDetectProvider_Gemini(t *testing.T) { + store := &credentials.MapStore{ + Data: map[string]string{ + "gemini_api_key": "gemini-key", + }, + } + credentials.SetDefaultStore(store) + t.Cleanup(func() { credentials.SetDefaultStore(nil) }) + + got := DetectProvider() + if got != "gemini" { + t.Errorf("DetectProvider = %q, want gemini", got) + } +} diff --git a/client/adapters/zai_test.go b/client/adapters/zai_test.go index 49c51c9..a35a107 100644 --- a/client/adapters/zai_test.go +++ b/client/adapters/zai_test.go @@ -246,3 +246,48 @@ func TestZAIClient_Name_NilOpenAI(t *testing.T) { t.Errorf("Name() = %q, want zai_custom", client.Name()) } } + +func TestZAIClient_Chat_NoFallbackOnNonRetriable(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"message": "some other error"}}), nil + }) + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + + _, err := client.Chat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "glm-4"}) + if err == nil { + t.Fatal("expected error without fallback") + } +} + +func TestZAIClient_StreamChat_NoFallbackOnNonRetriable(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"message": "some other error"}}), nil + }) + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + + _, err := client.StreamChat(context.Background(), []core.EyrieMessage{{Role: "user", Content: "Hi"}}, core.ChatOptions{Model: "glm-4"}) + if err == nil { + t.Fatal("expected error without fallback") + } +} + +func TestZAIClient_Ping_NoAnthropicNoFallback(t *testing.T) { + t.Parallel() + openAITransport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return jsonResponse(http.StatusUnauthorized, map[string]any{"error": map[string]string{"message": "unauthorized"}}), nil + }) + client := NewZAIClient("key", "https://zai.example/paas/v4", "", nil, "zai_payg") + client.router.OpenAI.SetRetry(core.RetryConfig{RetryConfig: types.RetryConfig{MaxRetries: 0}}) + client.router.OpenAI.httpClient = &http.Client{Transport: openAITransport} + + err := client.Ping(context.Background()) + if err == nil { + t.Fatal("expected error without anthropic fallback client") + } +} From afc123d759822147f9c8f2322912b92cc25cefeb Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Mon, 27 Jul 2026 07:14:47 +0530 Subject: [PATCH 4/4] test: add coverage for processStreamChunk (old bespoke parser path) Add tests for the old bespoke Gemini stream parser path (streamLoop/processStreamChunk): - processStreamChunk: 90.0% -> 95.0% (content, tool_call, usage/done, context cancelled, empty candidates, invalid JSON) Overall coverage: 88.0% -> 88.1% --- client/adapters/gemini_test.go | 102 +++++++++++++++++++++++++++++++++ 1 file changed, 102 insertions(+) diff --git a/client/adapters/gemini_test.go b/client/adapters/gemini_test.go index 064d6ca..f10c5ba 100644 --- a/client/adapters/gemini_test.go +++ b/client/adapters/gemini_test.go @@ -933,3 +933,105 @@ func TestGeminiClient_Chat_PenaltiesOnly(t *testing.T) { func ptrInt(v int) *int { return &v } func ptrBool(v bool) *bool { return &v } + +func TestProcessStreamChunk_Content(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent, 1) + data := `{"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}` + cont := c.processStreamChunk(context.Background(), data, events) + if cont { + t.Fatal("processStreamChunk returned true (done), expected false") + } + select { + case evt := <-events: + if evt.Type != "content" || evt.Content != "Hello" { + t.Errorf("event = %+v", evt) + } + default: + t.Fatal("expected content event") + } +} + +func TestProcessStreamChunk_ToolCall(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent, 1) + data := `{"candidates":[{"content":{"parts":[{"functionCall":{"name":"get_weather","args":{"city":"NYC"},"id":"call_1"}}]}}]}` + cont := c.processStreamChunk(context.Background(), data, events) + if cont { + t.Fatal("processStreamChunk returned true (done), expected false") + } + select { + case evt := <-events: + if evt.Type != "tool_call" || evt.ToolCall == nil || evt.ToolCall.Name != "get_weather" { + t.Errorf("event = %+v", evt) + } + default: + t.Fatal("expected tool_call event") + } +} + +func TestProcessStreamChunk_UsageDone(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent, 2) + data := `{"candidates":[{"content":{"parts":[{"text":"Hi"}],"finishReason":"STOP"}}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":5,"totalTokenCount":7}}` + cont := c.processStreamChunk(context.Background(), data, events) + if !cont { + t.Fatal("processStreamChunk returned false, expected true (done)") + } + var gotDone bool + for i := 0; i < 2; i++ { + select { + case evt := <-events: + if evt.Type == "done" { + gotDone = true + } + default: + return + } + } + if !gotDone { + t.Fatal("expected done event") + } +} + +func TestProcessStreamChunk_ContextCancelled(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + data := `{"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}` + cont := c.processStreamChunk(ctx, data, events) + if cont { + t.Fatal("processStreamChunk returned true, expected false") + } +} + +func TestProcessStreamChunk_EmptyCandidates(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent, 1) + data := `{"candidates":[]}` + cont := c.processStreamChunk(context.Background(), data, events) + if cont { + t.Fatal("processStreamChunk returned true, expected false") + } + select { + case <-events: + t.Fatal("expected no events") + default: + } +} + +func TestProcessStreamChunk_InvalidJSON(t *testing.T) { + t.Parallel() + c := NewGeminiClient("key", "https://gemini.example") + events := make(chan core.EyrieStreamEvent, 1) + cont := c.processStreamChunk(context.Background(), "not json", events) + if cont { + t.Fatal("processStreamChunk returned true, expected false") + } +}