From ba37217bfc89fa6ad52d8462ef70c21dad903f6c Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Fri, 26 Dec 2025 20:04:34 +0100 Subject: [PATCH 1/4] feat: added Groq provider --- .env.template | 4 + cmd/gomodel/main.go | 1 + config/config.go | 6 + internal/providers/groq/groq.go | 420 ++++++++++++++++++++++++++++++++ 4 files changed, 431 insertions(+) create mode 100644 internal/providers/groq/groq.go diff --git a/.env.template b/.env.template index af69acd1c..4b2f0fa98 100644 --- a/.env.template +++ b/.env.template @@ -37,3 +37,7 @@ CACHE_TYPE=local # xAI (Grok) # XAI_API_KEY=... + +# Groq +# GROQ_API_KEY=gsk_... + diff --git a/cmd/gomodel/main.go b/cmd/gomodel/main.go index 390f07ce5..93c50547e 100644 --- a/cmd/gomodel/main.go +++ b/cmd/gomodel/main.go @@ -17,6 +17,7 @@ import ( // Import provider packages to trigger their init() registration _ "gomodel/internal/providers/anthropic" _ "gomodel/internal/providers/gemini" + _ "gomodel/internal/providers/groq" _ "gomodel/internal/providers/openai" _ "gomodel/internal/providers/xai" "gomodel/internal/server" diff --git a/config/config.go b/config/config.go index 516f04d7a..4519b883e 100644 --- a/config/config.go +++ b/config/config.go @@ -144,6 +144,12 @@ func Load() (*Config, error) { APIKey: apiKey, } } + if apiKey := viper.GetString("GROQ_API_KEY"); apiKey != "" { + cfg.Providers["groq-primary"] = ProviderConfig{ + Type: "groq", + APIKey: apiKey, + } + } } return &cfg, nil diff --git a/internal/providers/groq/groq.go b/internal/providers/groq/groq.go new file mode 100644 index 000000000..493979581 --- /dev/null +++ b/internal/providers/groq/groq.go @@ -0,0 +1,420 @@ +// Package groq provides Groq API integration for the LLM gateway. +package groq + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "gomodel/internal/core" + "gomodel/internal/pkg/llmclient" + "gomodel/internal/providers" +) + +const ( + defaultBaseURL = "https://api.groq.com/openai/v1" +) + +func init() { + // Self-register with the factory + providers.RegisterProvider("groq", New) +} + +// Provider implements the core.Provider interface for Groq +type Provider struct { + client *llmclient.Client + apiKey string +} + +// New creates a new Groq provider +func New(apiKey string) *Provider { + p := &Provider{apiKey: apiKey} + cfg := llmclient.DefaultConfig("groq", defaultBaseURL) + // Apply global hooks if available + cfg.Hooks = providers.GetGlobalHooks() + p.client = llmclient.New(cfg, p.setHeaders) + return p +} + +// NewWithHTTPClient creates a new Groq provider with a custom HTTP client +func NewWithHTTPClient(apiKey string, httpClient *http.Client) *Provider { + p := &Provider{apiKey: apiKey} + cfg := llmclient.DefaultConfig("groq", defaultBaseURL) + // Apply global hooks if available + cfg.Hooks = providers.GetGlobalHooks() + p.client = llmclient.NewWithHTTPClient(httpClient, cfg, p.setHeaders) + return p +} + +// SetBaseURL allows configuring a custom base URL for the provider +func (p *Provider) SetBaseURL(url string) { + p.client.SetBaseURL(url) +} + +// setHeaders sets the required headers for Groq API requests +func (p *Provider) setHeaders(req *http.Request) { + req.Header.Set("Authorization", "Bearer "+p.apiKey) +} + +// ChatCompletion sends a chat completion request to Groq +func (p *Provider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { + var resp core.ChatResponse + err := p.client.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/chat/completions", + Body: req, + }, &resp) + if err != nil { + return nil, err + } + return &resp, nil +} + +// StreamChatCompletion returns a raw response body for streaming (caller must close) +func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { + req.Stream = true + return p.client.DoStream(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/chat/completions", + Body: req, + }) +} + +// ListModels retrieves the list of available models from Groq +func (p *Provider) ListModels(ctx context.Context) (*core.ModelsResponse, error) { + var resp core.ModelsResponse + err := p.client.Do(ctx, llmclient.Request{ + Method: http.MethodGet, + Endpoint: "/models", + }, &resp) + if err != nil { + return nil, err + } + return &resp, nil +} + +// convertResponsesRequestToChat converts a ResponsesRequest to ChatRequest for Groq +func convertResponsesRequestToChat(req *core.ResponsesRequest) *core.ChatRequest { + chatReq := &core.ChatRequest{ + Model: req.Model, + Messages: make([]core.Message, 0), + Temperature: req.Temperature, + Stream: req.Stream, + } + + if req.MaxOutputTokens != nil { + chatReq.MaxTokens = req.MaxOutputTokens + } + + // Add system instruction if provided + if req.Instructions != "" { + chatReq.Messages = append(chatReq.Messages, core.Message{ + Role: "system", + Content: req.Instructions, + }) + } + + // Convert input to messages + switch input := req.Input.(type) { + case string: + chatReq.Messages = append(chatReq.Messages, core.Message{ + Role: "user", + Content: input, + }) + case []interface{}: + for _, item := range input { + if msgMap, ok := item.(map[string]interface{}); ok { + role, _ := msgMap["role"].(string) + content := extractContentFromInput(msgMap["content"]) + if role != "" && content != "" { + chatReq.Messages = append(chatReq.Messages, core.Message{ + Role: role, + Content: content, + }) + } + } + } + } + + return chatReq +} + +// extractContentFromInput extracts text content from responses input +func extractContentFromInput(content interface{}) string { + switch c := content.(type) { + case string: + return c + case []interface{}: + // Array of content parts - extract text + var texts []string + for _, part := range c { + if partMap, ok := part.(map[string]interface{}); ok { + if text, ok := partMap["text"].(string); ok { + texts = append(texts, text) + } + } + } + return strings.Join(texts, " ") + } + return "" +} + +// convertChatResponseToResponses converts a ChatResponse to ResponsesResponse +func convertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResponse { + content := "" + if len(resp.Choices) > 0 { + content = resp.Choices[0].Message.Content + } + + return &core.ResponsesResponse{ + ID: resp.ID, + Object: "response", + CreatedAt: resp.Created, + Model: resp.Model, + Status: "completed", + Output: []core.ResponsesOutputItem{ + { + ID: fmt.Sprintf("msg_%d", time.Now().UnixNano()), + Type: "message", + Role: "assistant", + Status: "completed", + Content: []core.ResponsesContentItem{ + { + Type: "output_text", + Text: content, + Annotations: []string{}, + }, + }, + }, + }, + Usage: &core.ResponsesUsage{ + InputTokens: resp.Usage.PromptTokens, + OutputTokens: resp.Usage.CompletionTokens, + TotalTokens: resp.Usage.TotalTokens, + }, + } +} + +// Responses sends a Responses API request to Groq (converted to chat format) +func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { + // Convert ResponsesRequest to ChatRequest + chatReq := convertResponsesRequestToChat(req) + + // Use the existing ChatCompletion method + chatResp, err := p.ChatCompletion(ctx, chatReq) + if err != nil { + return nil, err + } + + return convertChatResponseToResponses(chatResp), nil +} + +// StreamResponses returns a raw response body for streaming Responses API (caller must close) +func (p *Provider) StreamResponses(ctx context.Context, req *core.ResponsesRequest) (io.ReadCloser, error) { + // Convert ResponsesRequest to ChatRequest + chatReq := convertResponsesRequestToChat(req) + chatReq.Stream = true + + // Get the streaming response from chat completions + stream, err := p.StreamChatCompletion(ctx, chatReq) + if err != nil { + return nil, err + } + + // Wrap the stream to convert chat completion format to Responses API format + return newGroqResponsesStreamConverter(stream, req.Model), nil +} + +// groqResponsesStreamConverter wraps a chat completion stream and converts it to Responses API format +type groqResponsesStreamConverter struct { + reader io.ReadCloser + model string + responseID string + buffer []byte + lineBuffer []byte + closed bool + sentCreate bool + sentDone bool +} + +func newGroqResponsesStreamConverter(reader io.ReadCloser, model string) *groqResponsesStreamConverter { + return &groqResponsesStreamConverter{ + reader: reader, + model: model, + responseID: "resp_" + time.Now().Format("20060102150405"), + buffer: make([]byte, 0, 4096), + lineBuffer: make([]byte, 0, 1024), + } +} + +func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { + if sc.closed { + return 0, io.EOF + } + + // If we have buffered data, return it first + if len(sc.buffer) > 0 { + n = copy(p, sc.buffer) + sc.buffer = sc.buffer[n:] + return n, nil + } + + // Send response.created event first + if !sc.sentCreate { + sc.sentCreate = true + createdEvent := map[string]interface{}{ + "type": "response.created", + "response": map[string]interface{}{ + "id": sc.responseID, + "object": "response", + "status": "in_progress", + "model": sc.model, + "created_at": time.Now().Unix(), + }, + } + jsonData, err := json.Marshal(createdEvent) + if err != nil { + slog.Error("failed to marshal response.created event", "error", err, "response_id", sc.responseID) + return 0, nil + } + created := fmt.Sprintf("event: response.created\ndata: %s\n\n", jsonData) + sc.buffer = append(sc.buffer, []byte(created)...) + n = copy(p, sc.buffer) + sc.buffer = sc.buffer[n:] + return n, nil + } + + // Read from the underlying stream + tempBuf := make([]byte, 1024) + nr, readErr := sc.reader.Read(tempBuf) + if nr > 0 { + sc.lineBuffer = append(sc.lineBuffer, tempBuf[:nr]...) + + // Process complete lines + for { + idx := bytes.Index(sc.lineBuffer, []byte("\n")) + if idx == -1 { + break + } + + line := sc.lineBuffer[:idx] + sc.lineBuffer = sc.lineBuffer[idx+1:] + + line = bytes.TrimSpace(line) + if len(line) == 0 { + continue + } + + if bytes.HasPrefix(line, []byte("data: ")) { + data := bytes.TrimPrefix(line, []byte("data: ")) + if bytes.Equal(data, []byte("[DONE]")) { + // Send done event + if !sc.sentDone { + sc.sentDone = true + doneEvent := map[string]interface{}{ + "type": "response.done", + "response": map[string]interface{}{ + "id": sc.responseID, + "object": "response", + "status": "completed", + "model": sc.model, + "created_at": time.Now().Unix(), + }, + } + jsonData, err := json.Marshal(doneEvent) + if err != nil { + slog.Error("failed to marshal response.done event", "error", err, "response_id", sc.responseID) + continue + } + doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) + sc.buffer = append(sc.buffer, []byte(doneMsg)...) + } + continue + } + + // Parse the chat completion chunk + var chunk map[string]interface{} + if err := json.Unmarshal(data, &chunk); err != nil { + continue + } + + // Extract content delta + if choices, ok := chunk["choices"].([]interface{}); ok && len(choices) > 0 { + if choice, ok := choices[0].(map[string]interface{}); ok { + if delta, ok := choice["delta"].(map[string]interface{}); ok { + if content, ok := delta["content"].(string); ok && content != "" { + deltaEvent := map[string]interface{}{ + "type": "response.output_text.delta", + "delta": content, + } + jsonData, err := json.Marshal(deltaEvent) + if err != nil { + slog.Error("failed to marshal content delta event", "error", err, "response_id", sc.responseID) + continue + } + sc.buffer = append(sc.buffer, []byte(fmt.Sprintf("event: response.output_text.delta\ndata: %s\n\n", jsonData))...) + } + } + } + } + } + } + } + + if readErr != nil { + if readErr == io.EOF { + // Send final done event if we haven't already + if !sc.sentDone { + sc.sentDone = true + doneEvent := map[string]interface{}{ + "type": "response.done", + "response": map[string]interface{}{ + "id": sc.responseID, + "object": "response", + "status": "completed", + "model": sc.model, + "created_at": time.Now().Unix(), + }, + } + jsonData, err := json.Marshal(doneEvent) + if err != nil { + slog.Error("failed to marshal final response.done event", "error", err, "response_id", sc.responseID) + } else { + doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) + sc.buffer = append(sc.buffer, []byte(doneMsg)...) + } + } + + if len(sc.buffer) > 0 { + n = copy(p, sc.buffer) + sc.buffer = sc.buffer[n:] + return n, nil + } + + sc.closed = true + _ = sc.reader.Close() + return 0, io.EOF + } + return 0, readErr + } + + if len(sc.buffer) > 0 { + n = copy(p, sc.buffer) + sc.buffer = sc.buffer[n:] + return n, nil + } + + // No data yet, try again + return 0, nil +} + +func (sc *groqResponsesStreamConverter) Close() error { + sc.closed = true + return sc.reader.Close() +} From ecba6a68c21c7322de08ece756a241823d3c014e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Fri, 26 Dec 2025 22:33:52 +0100 Subject: [PATCH 2/4] fix: groq adjusted to the standard syntax used for streaming --- go.mod | 1 + go.sum | 2 ++ internal/providers/anthropic/anthropic.go | 6 ++++-- internal/providers/gemini/gemini.go | 6 ++++-- internal/providers/groq/groq.go | 26 ++++++++++------------- 5 files changed, 22 insertions(+), 19 deletions(-) diff --git a/go.mod b/go.mod index ef774b4e5..294b4ad3c 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module gomodel go 1.24.0 require ( + github.com/google/uuid v1.6.0 github.com/joho/godotenv v1.5.1 github.com/labstack/echo/v4 v4.14.0 github.com/prometheus/client_golang v1.23.2 diff --git a/go.sum b/go.sum index 952ccca7e..a13f4e5b0 100644 --- a/go.sum +++ b/go.sum @@ -18,6 +18,8 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0= github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= diff --git a/internal/providers/anthropic/anthropic.go b/internal/providers/anthropic/anthropic.go index 6cbf93c97..5eeb95f9a 100644 --- a/internal/providers/anthropic/anthropic.go +++ b/internal/providers/anthropic/anthropic.go @@ -13,6 +13,8 @@ import ( "strings" "time" + "github.com/google/uuid" + "gomodel/internal/core" "gomodel/internal/pkg/llmclient" "gomodel/internal/providers" @@ -508,7 +510,7 @@ func convertAnthropicResponseToResponses(resp *anthropicResponse, model string) Status: "completed", Output: []core.ResponsesOutputItem{ { - ID: fmt.Sprintf("msg_%d", time.Now().UnixNano()), + ID: "msg_" + uuid.New().String(), Type: "message", Role: "assistant", Status: "completed", @@ -580,7 +582,7 @@ func newResponsesStreamConverter(body io.ReadCloser, model string) *responsesStr reader: bufio.NewReader(body), body: body, model: model, - responseID: fmt.Sprintf("resp_%d", time.Now().UnixNano()), + responseID: "resp_" + uuid.New().String(), buffer: make([]byte, 0, 1024), } } diff --git a/internal/providers/gemini/gemini.go b/internal/providers/gemini/gemini.go index f863dff96..d937f2c33 100644 --- a/internal/providers/gemini/gemini.go +++ b/internal/providers/gemini/gemini.go @@ -12,6 +12,8 @@ import ( "strings" "time" + "github.com/google/uuid" + "gomodel/internal/core" "gomodel/internal/pkg/llmclient" "gomodel/internal/providers" @@ -262,7 +264,7 @@ func convertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResp Status: "completed", Output: []core.ResponsesOutputItem{ { - ID: fmt.Sprintf("msg_%d", time.Now().UnixNano()), + ID: "msg_" + uuid.New().String(), Type: "message", Role: "assistant", Status: "completed", @@ -328,7 +330,7 @@ func newGeminiResponsesStreamConverter(reader io.ReadCloser, model string) *gemi return &geminiResponsesStreamConverter{ reader: reader, model: model, - responseID: "resp_" + time.Now().Format("20060102150405"), + responseID: "resp_" + uuid.New().String(), buffer: make([]byte, 0, 4096), lineBuffer: make([]byte, 0, 1024), } diff --git a/internal/providers/groq/groq.go b/internal/providers/groq/groq.go index 493979581..753246039 100644 --- a/internal/providers/groq/groq.go +++ b/internal/providers/groq/groq.go @@ -7,11 +7,12 @@ import ( "encoding/json" "fmt" "io" - "log/slog" "net/http" "strings" "time" + "github.com/google/uuid" + "gomodel/internal/core" "gomodel/internal/pkg/llmclient" "gomodel/internal/providers" @@ -78,11 +79,10 @@ func (p *Provider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (* // StreamChatCompletion returns a raw response body for streaming (caller must close) func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { - req.Stream = true return p.client.DoStream(ctx, llmclient.Request{ Method: http.MethodPost, Endpoint: "/chat/completions", - Body: req, + Body: req.WithStreaming(), }) } @@ -180,7 +180,7 @@ func convertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResp Status: "completed", Output: []core.ResponsesOutputItem{ { - ID: fmt.Sprintf("msg_%d", time.Now().UnixNano()), + ID: "msg_" + uuid.New().String(), Type: "message", Role: "assistant", Status: "completed", @@ -247,7 +247,7 @@ func newGroqResponsesStreamConverter(reader io.ReadCloser, model string) *groqRe return &groqResponsesStreamConverter{ reader: reader, model: model, - responseID: "resp_" + time.Now().Format("20060102150405"), + responseID: "resp_" + uuid.New().String(), buffer: make([]byte, 0, 4096), lineBuffer: make([]byte, 0, 1024), } @@ -280,8 +280,7 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(createdEvent) if err != nil { - slog.Error("failed to marshal response.created event", "error", err, "response_id", sc.responseID) - return 0, nil + return 0, fmt.Errorf("failed to marshal response.created event: %w", err) } created := fmt.Sprintf("event: response.created\ndata: %s\n\n", jsonData) sc.buffer = append(sc.buffer, []byte(created)...) @@ -329,8 +328,7 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(doneEvent) if err != nil { - slog.Error("failed to marshal response.done event", "error", err, "response_id", sc.responseID) - continue + return 0, fmt.Errorf("failed to marshal response.done event: %w", err) } doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) sc.buffer = append(sc.buffer, []byte(doneMsg)...) @@ -355,8 +353,7 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(deltaEvent) if err != nil { - slog.Error("failed to marshal content delta event", "error", err, "response_id", sc.responseID) - continue + return 0, fmt.Errorf("failed to marshal content delta event: %w", err) } sc.buffer = append(sc.buffer, []byte(fmt.Sprintf("event: response.output_text.delta\ndata: %s\n\n", jsonData))...) } @@ -384,11 +381,10 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(doneEvent) if err != nil { - slog.Error("failed to marshal final response.done event", "error", err, "response_id", sc.responseID) - } else { - doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) - sc.buffer = append(sc.buffer, []byte(doneMsg)...) + return 0, fmt.Errorf("failed to marshal final response.done event: %w", err) } + doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) + sc.buffer = append(sc.buffer, []byte(doneMsg)...) } if len(sc.buffer) > 0 { From 203a198e5ea14bf0d094a362f851d966da2187e1 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Fri, 26 Dec 2025 22:46:15 +0100 Subject: [PATCH 3/4] chore: added default config template for Groq --- config/config.yaml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/config/config.yaml b/config/config.yaml index 529631f22..28ac57623 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -48,6 +48,10 @@ providers: type: "xai" api_key: "${XAI_API_KEY}" + groq-primary: + type: "groq" + api_key: "${GROQ_API_KEY}" + # Example: Groq (OpenAI-compatible) # groq: # type: "openai" From 0d288bca91a66a38a5c12df50f32a503852778ff Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Fri, 26 Dec 2025 23:01:25 +0100 Subject: [PATCH 4/4] fix: fixed error raporting + tests --- internal/providers/groq/groq.go | 17 +- internal/providers/groq/groq_test.go | 1004 ++++++++++++++++++++++++++ 2 files changed, 1015 insertions(+), 6 deletions(-) create mode 100644 internal/providers/groq/groq_test.go diff --git a/internal/providers/groq/groq.go b/internal/providers/groq/groq.go index 753246039..85b9b02f1 100644 --- a/internal/providers/groq/groq.go +++ b/internal/providers/groq/groq.go @@ -7,6 +7,7 @@ import ( "encoding/json" "fmt" "io" + "log/slog" "net/http" "strings" "time" @@ -280,7 +281,8 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(createdEvent) if err != nil { - return 0, fmt.Errorf("failed to marshal response.created event: %w", err) + slog.Error("failed to marshal response.created event", "error", err, "response_id", sc.responseID) + return 0, nil } created := fmt.Sprintf("event: response.created\ndata: %s\n\n", jsonData) sc.buffer = append(sc.buffer, []byte(created)...) @@ -328,7 +330,8 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(doneEvent) if err != nil { - return 0, fmt.Errorf("failed to marshal response.done event: %w", err) + slog.Error("failed to marshal response.done event", "error", err, "response_id", sc.responseID) + continue } doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) sc.buffer = append(sc.buffer, []byte(doneMsg)...) @@ -353,7 +356,8 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(deltaEvent) if err != nil { - return 0, fmt.Errorf("failed to marshal content delta event: %w", err) + slog.Error("failed to marshal content delta event", "error", err, "response_id", sc.responseID) + continue } sc.buffer = append(sc.buffer, []byte(fmt.Sprintf("event: response.output_text.delta\ndata: %s\n\n", jsonData))...) } @@ -381,10 +385,11 @@ func (sc *groqResponsesStreamConverter) Read(p []byte) (n int, err error) { } jsonData, err := json.Marshal(doneEvent) if err != nil { - return 0, fmt.Errorf("failed to marshal final response.done event: %w", err) + slog.Error("failed to marshal final response.done event", "error", err, "response_id", sc.responseID) + } else { + doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) + sc.buffer = append(sc.buffer, []byte(doneMsg)...) } - doneMsg := fmt.Sprintf("event: response.done\ndata: %s\n\ndata: [DONE]\n\n", jsonData) - sc.buffer = append(sc.buffer, []byte(doneMsg)...) } if len(sc.buffer) > 0 { diff --git a/internal/providers/groq/groq_test.go b/internal/providers/groq/groq_test.go new file mode 100644 index 000000000..dd35fd719 --- /dev/null +++ b/internal/providers/groq/groq_test.go @@ -0,0 +1,1004 @@ +package groq + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "gomodel/internal/core" +) + +func TestNew(t *testing.T) { + apiKey := "test-api-key" + provider := New(apiKey) + + if provider.apiKey != apiKey { + t.Errorf("apiKey = %q, want %q", provider.apiKey, apiKey) + } + if provider.client == nil { + t.Error("client should not be nil") + } +} + +func TestChatCompletion(t *testing.T) { + tests := []struct { + name string + statusCode int + responseBody string + expectedError bool + checkResponse func(*testing.T, *core.ChatResponse) + }{ + { + name: "successful request", + statusCode: http.StatusOK, + responseBody: `{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama-3.3-70b-versatile", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30 + } + }`, + expectedError: false, + checkResponse: func(t *testing.T, resp *core.ChatResponse) { + if resp.ID != "chatcmpl-123" { + t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") + } + if resp.Model != "llama-3.3-70b-versatile" { + t.Errorf("Model = %q, want %q", resp.Model, "llama-3.3-70b-versatile") + } + if len(resp.Choices) != 1 { + t.Fatalf("len(Choices) = %d, want 1", len(resp.Choices)) + } + if resp.Choices[0].Message.Content != "Hello! How can I help you today?" { + t.Errorf("Message content = %q, want %q", resp.Choices[0].Message.Content, "Hello! How can I help you today?") + } + if resp.Usage.PromptTokens != 10 { + t.Errorf("PromptTokens = %d, want 10", resp.Usage.PromptTokens) + } + if resp.Usage.CompletionTokens != 20 { + t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens) + } + if resp.Usage.TotalTokens != 30 { + t.Errorf("TotalTokens = %d, want 30", resp.Usage.TotalTokens) + } + }, + }, + { + name: "API error", + statusCode: http.StatusUnauthorized, + responseBody: `{"error": {"message": "Invalid API key"}}`, + expectedError: true, + }, + { + name: "rate limit error", + statusCode: http.StatusTooManyRequests, + responseBody: `{"error": {"message": "Rate limit exceeded"}}`, + expectedError: true, + }, + { + name: "server error", + statusCode: http.StatusInternalServerError, + responseBody: `{"error": {"message": "Internal server error"}}`, + expectedError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify request headers + if r.Header.Get("Content-Type") != "application/json" { + t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") + } + authHeader := r.Header.Get("Authorization") + if !strings.HasPrefix(authHeader, "Bearer ") { + t.Errorf("Authorization header should start with 'Bearer '") + } + + // Verify request body + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var req core.ChatRequest + if err := json.Unmarshal(body, &req); err != nil { + t.Fatalf("failed to unmarshal request: %v", err) + } + + w.WriteHeader(tt.statusCode) + _, _ = w.Write([]byte(tt.responseBody)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + req := &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{ + {Role: "user", Content: "Hello"}, + }, + } + + resp, err := provider.ChatCompletion(context.Background(), req) + + if tt.expectedError { + if err == nil { + t.Error("expected error, got nil") + } + } else { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if tt.checkResponse != nil { + tt.checkResponse(t, resp) + } + } + }) + } +} + +func TestStreamChatCompletion(t *testing.T) { + tests := []struct { + name string + statusCode int + responseBody string + expectedError bool + }{ + { + name: "successful streaming request", + statusCode: http.StatusOK, + responseBody: `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} + +data: [DONE] +`, + expectedError: false, + }, + { + name: "API error", + statusCode: http.StatusUnauthorized, + responseBody: `{"error": {"message": "Invalid API key"}}`, + expectedError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify request headers + if r.Header.Get("Content-Type") != "application/json" { + t.Errorf("Content-Type = %q, want %q", r.Header.Get("Content-Type"), "application/json") + } + authHeader := r.Header.Get("Authorization") + if !strings.HasPrefix(authHeader, "Bearer ") { + t.Errorf("Authorization header should start with 'Bearer '") + } + + // Verify stream is set in request body + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var req core.ChatRequest + if err := json.Unmarshal(body, &req); err != nil { + t.Fatalf("failed to unmarshal request: %v", err) + } + if !req.Stream { + t.Error("Stream should be true in request") + } + + w.WriteHeader(tt.statusCode) + _, _ = w.Write([]byte(tt.responseBody)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + req := &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{ + {Role: "user", Content: "Hello"}, + }, + } + + body, err := provider.StreamChatCompletion(context.Background(), req) + + if tt.expectedError { + if err == nil { + t.Error("expected error, got nil") + } + } else { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if body == nil { + t.Fatal("body should not be nil") + } + defer func() { _ = body.Close() }() + + // Read and verify the streaming response + respBody, err := io.ReadAll(body) + if err != nil { + t.Fatalf("failed to read response body: %v", err) + } + if string(respBody) != tt.responseBody { + t.Errorf("response body = %q, want %q", string(respBody), tt.responseBody) + } + } + }) + } +} + +func TestListModels(t *testing.T) { + tests := []struct { + name string + statusCode int + responseBody string + expectedError bool + checkResponse func(*testing.T, *core.ModelsResponse) + }{ + { + name: "successful request", + statusCode: http.StatusOK, + responseBody: `{ + "object": "list", + "data": [ + { + "id": "llama-3.3-70b-versatile", + "object": "model", + "created": 1687882411, + "owned_by": "groq" + }, + { + "id": "mixtral-8x7b-32768", + "object": "model", + "created": 1687882410, + "owned_by": "groq" + } + ] + }`, + expectedError: false, + checkResponse: func(t *testing.T, resp *core.ModelsResponse) { + if resp.Object != "list" { + t.Errorf("Object = %q, want %q", resp.Object, "list") + } + if len(resp.Data) != 2 { + t.Fatalf("len(Data) = %d, want 2", len(resp.Data)) + } + if resp.Data[0].ID != "llama-3.3-70b-versatile" { + t.Errorf("Data[0].ID = %q, want %q", resp.Data[0].ID, "llama-3.3-70b-versatile") + } + if resp.Data[0].OwnedBy != "groq" { + t.Errorf("Data[0].OwnedBy = %q, want %q", resp.Data[0].OwnedBy, "groq") + } + }, + }, + { + name: "API error", + statusCode: http.StatusUnauthorized, + responseBody: `{"error": {"message": "Invalid API key"}}`, + expectedError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify request method and path + if r.Method != http.MethodGet { + t.Errorf("Method = %q, want %q", r.Method, http.MethodGet) + } + if r.URL.Path != "/models" { + t.Errorf("Path = %q, want %q", r.URL.Path, "/models") + } + + // Verify authorization header + authHeader := r.Header.Get("Authorization") + if !strings.HasPrefix(authHeader, "Bearer ") { + t.Errorf("Authorization header should start with 'Bearer '") + } + + w.WriteHeader(tt.statusCode) + _, _ = w.Write([]byte(tt.responseBody)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + resp, err := provider.ListModels(context.Background()) + + if tt.expectedError { + if err == nil { + t.Error("expected error, got nil") + } + } else { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if tt.checkResponse != nil { + tt.checkResponse(t, resp) + } + } + }) + } +} + +func TestChatCompletionWithContext(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Simulate a slow response + <-r.Context().Done() + w.WriteHeader(http.StatusRequestTimeout) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // Cancel immediately + + req := &core.ChatRequest{ + Model: "llama-3.3-70b-versatile", + Messages: []core.Message{ + {Role: "user", Content: "Hello"}, + }, + } + + _, err := provider.ChatCompletion(ctx, req) + if err == nil { + t.Error("expected error when context is cancelled, got nil") + } +} + +func TestResponses(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify request path for chat completions (Groq converts Responses to chat) + if r.URL.Path != "/chat/completions" { + t.Errorf("Path = %q, want %q", r.URL.Path, "/chat/completions") + } + + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama-3.3-70b-versatile", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30 + } + }`)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + req := &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + } + + resp, err := provider.Responses(context.Background(), req) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if resp.ID != "chatcmpl-123" { + t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") + } + if resp.Object != "response" { + t.Errorf("Object = %q, want %q", resp.Object, "response") + } + if resp.Model != "llama-3.3-70b-versatile" { + t.Errorf("Model = %q, want %q", resp.Model, "llama-3.3-70b-versatile") + } + if resp.Status != "completed" { + t.Errorf("Status = %q, want %q", resp.Status, "completed") + } + if len(resp.Output) != 1 { + t.Fatalf("len(Output) = %d, want 1", len(resp.Output)) + } + if len(resp.Output[0].Content) != 1 { + t.Fatalf("len(Output[0].Content) = %d, want 1", len(resp.Output[0].Content)) + } + if resp.Output[0].Content[0].Text != "Hello! How can I help you today?" { + t.Errorf("Output text = %q, want %q", resp.Output[0].Content[0].Text, "Hello! How can I help you today?") + } + if resp.Usage == nil { + t.Fatal("Usage should not be nil") + } + if resp.Usage.InputTokens != 10 { + t.Errorf("InputTokens = %d, want 10", resp.Usage.InputTokens) + } + if resp.Usage.OutputTokens != 20 { + t.Errorf("OutputTokens = %d, want 20", resp.Usage.OutputTokens) + } +} + +func TestResponsesWithArrayInput(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify request body is converted to chat format + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + + var req map[string]interface{} + if err := json.Unmarshal(body, &req); err != nil { + t.Fatalf("failed to unmarshal request: %v", err) + } + + // Verify messages array exists (converted from input) + messages, ok := req["messages"].([]interface{}) + if !ok { + t.Fatal("messages should be an array") + } + // Should have system message + 2 input messages + if len(messages) != 3 { + t.Errorf("len(messages) = %d, want 3", len(messages)) + } + + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama-3.3-70b-versatile", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello!" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15 + } + }`)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + req := &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello", + }, + map[string]interface{}{ + "role": "assistant", + "content": "Hi there!", + }, + }, + Instructions: "Be helpful", + } + + resp, err := provider.Responses(context.Background(), req) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if resp.ID != "chatcmpl-123" { + t.Errorf("ID = %q, want %q", resp.ID, "chatcmpl-123") + } +} + +func TestStreamResponses(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Verify stream is set in request body + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var req core.ChatRequest + if err := json.Unmarshal(body, &req); err != nil { + t.Fatalf("failed to unmarshal request: %v", err) + } + if !req.Stream { + t.Error("Stream should be true in request") + } + + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":null}]} + +data: [DONE] +`)) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + req := &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + } + + body, err := provider.StreamResponses(context.Background(), req) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if body == nil { + t.Fatal("body should not be nil") + } + defer func() { _ = body.Close() }() + + respBody, err := io.ReadAll(body) + if err != nil { + t.Fatalf("failed to read response body: %v", err) + } + + responseStr := string(respBody) + if !strings.Contains(responseStr, "response.created") { + t.Error("response should contain response.created event") + } + if !strings.Contains(responseStr, "response.output_text.delta") { + t.Error("response should contain response.output_text.delta event") + } + if !strings.Contains(responseStr, "[DONE]") { + t.Error("response should end with [DONE]") + } +} + +func TestResponsesWithContext(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Simulate a slow response + <-r.Context().Done() + w.WriteHeader(http.StatusRequestTimeout) + })) + defer server.Close() + + provider := New("test-api-key") + provider.SetBaseURL(server.URL) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // Cancel immediately + + req := &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + } + + _, err := provider.Responses(ctx, req) + if err == nil { + t.Error("expected error when context is cancelled, got nil") + } +} + +func TestConvertResponsesRequestToChat(t *testing.T) { + temp := 0.7 + maxTokens := 1024 + + tests := []struct { + name string + input *core.ResponsesRequest + checkFn func(*testing.T, *core.ChatRequest) + }{ + { + name: "string input", + input: &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + }, + checkFn: func(t *testing.T, req *core.ChatRequest) { + if req.Model != "llama-3.3-70b-versatile" { + t.Errorf("Model = %q, want %q", req.Model, "llama-3.3-70b-versatile") + } + if len(req.Messages) != 1 { + t.Errorf("len(Messages) = %d, want 1", len(req.Messages)) + } + if req.Messages[0].Role != "user" { + t.Errorf("Messages[0].Role = %q, want %q", req.Messages[0].Role, "user") + } + if req.Messages[0].Content != "Hello" { + t.Errorf("Messages[0].Content = %q, want %q", req.Messages[0].Content, "Hello") + } + }, + }, + { + name: "with instructions", + input: &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + Instructions: "Be helpful", + }, + checkFn: func(t *testing.T, req *core.ChatRequest) { + if len(req.Messages) < 2 { + t.Fatalf("len(Messages) = %d, want at least 2", len(req.Messages)) + } + if req.Messages[0].Role != "system" { + t.Errorf("Messages[0].Role = %q, want %q", req.Messages[0].Role, "system") + } + if req.Messages[0].Content != "Be helpful" { + t.Errorf("Messages[0].Content = %q, want %q", req.Messages[0].Content, "Be helpful") + } + }, + }, + { + name: "with parameters", + input: &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + Temperature: &temp, + MaxOutputTokens: &maxTokens, + }, + checkFn: func(t *testing.T, req *core.ChatRequest) { + if req.Temperature == nil || *req.Temperature != 0.7 { + t.Errorf("Temperature = %v, want 0.7", req.Temperature) + } + if req.MaxTokens == nil || *req.MaxTokens != 1024 { + t.Errorf("MaxTokens = %v, want 1024", req.MaxTokens) + } + }, + }, + { + name: "with streaming enabled", + input: &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: "Hello", + Stream: true, + }, + checkFn: func(t *testing.T, req *core.ChatRequest) { + if !req.Stream { + t.Error("Stream should be true") + } + }, + }, + { + name: "array input with messages", + input: &core.ResponsesRequest{ + Model: "llama-3.3-70b-versatile", + Input: []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello", + }, + map[string]interface{}{ + "role": "assistant", + "content": "Hi there!", + }, + }, + }, + checkFn: func(t *testing.T, req *core.ChatRequest) { + if len(req.Messages) != 2 { + t.Fatalf("len(Messages) = %d, want 2", len(req.Messages)) + } + if req.Messages[0].Role != "user" { + t.Errorf("Messages[0].Role = %q, want %q", req.Messages[0].Role, "user") + } + if req.Messages[0].Content != "Hello" { + t.Errorf("Messages[0].Content = %q, want %q", req.Messages[0].Content, "Hello") + } + if req.Messages[1].Role != "assistant" { + t.Errorf("Messages[1].Role = %q, want %q", req.Messages[1].Role, "assistant") + } + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := convertResponsesRequestToChat(tt.input) + tt.checkFn(t, result) + }) + } +} + +func TestConvertChatResponseToResponses(t *testing.T) { + resp := &core.ChatResponse{ + ID: "chatcmpl-123", + Object: "chat.completion", + Model: "llama-3.3-70b-versatile", + Created: 1677652288, + Choices: []core.Choice{ + { + Index: 0, + Message: core.Message{ + Role: "assistant", + Content: "Hello! How can I help you today?", + }, + FinishReason: "stop", + }, + }, + Usage: core.Usage{ + PromptTokens: 10, + CompletionTokens: 20, + TotalTokens: 30, + }, + } + + result := convertChatResponseToResponses(resp) + + if result.ID != "chatcmpl-123" { + t.Errorf("ID = %q, want %q", result.ID, "chatcmpl-123") + } + if result.Object != "response" { + t.Errorf("Object = %q, want %q", result.Object, "response") + } + if result.Model != "llama-3.3-70b-versatile" { + t.Errorf("Model = %q, want %q", result.Model, "llama-3.3-70b-versatile") + } + if result.Status != "completed" { + t.Errorf("Status = %q, want %q", result.Status, "completed") + } + if len(result.Output) != 1 { + t.Fatalf("len(Output) = %d, want 1", len(result.Output)) + } + if result.Output[0].Type != "message" { + t.Errorf("Output[0].Type = %q, want %q", result.Output[0].Type, "message") + } + if result.Output[0].Role != "assistant" { + t.Errorf("Output[0].Role = %q, want %q", result.Output[0].Role, "assistant") + } + if result.Output[0].Status != "completed" { + t.Errorf("Output[0].Status = %q, want %q", result.Output[0].Status, "completed") + } + if len(result.Output[0].Content) != 1 { + t.Fatalf("len(Output[0].Content) = %d, want 1", len(result.Output[0].Content)) + } + if result.Output[0].Content[0].Type != "output_text" { + t.Errorf("Content[0].Type = %q, want %q", result.Output[0].Content[0].Type, "output_text") + } + if result.Output[0].Content[0].Text != "Hello! How can I help you today?" { + t.Errorf("Content[0].Text = %q, want %q", result.Output[0].Content[0].Text, "Hello! How can I help you today?") + } + if result.Usage == nil { + t.Fatal("Usage should not be nil") + } + if result.Usage.InputTokens != 10 { + t.Errorf("InputTokens = %d, want 10", result.Usage.InputTokens) + } + if result.Usage.OutputTokens != 20 { + t.Errorf("OutputTokens = %d, want 20", result.Usage.OutputTokens) + } + if result.Usage.TotalTokens != 30 { + t.Errorf("TotalTokens = %d, want 30", result.Usage.TotalTokens) + } +} + +func TestConvertChatResponseToResponses_EmptyChoices(t *testing.T) { + resp := &core.ChatResponse{ + ID: "chatcmpl-123", + Object: "chat.completion", + Model: "llama-3.3-70b-versatile", + Created: 1677652288, + Choices: []core.Choice{}, + Usage: core.Usage{ + PromptTokens: 10, + CompletionTokens: 0, + TotalTokens: 10, + }, + } + + result := convertChatResponseToResponses(resp) + + if len(result.Output) != 1 { + t.Fatalf("len(Output) = %d, want 1", len(result.Output)) + } + // Content should be empty string when no choices + if result.Output[0].Content[0].Text != "" { + t.Errorf("Content[0].Text = %q, want empty string", result.Output[0].Content[0].Text) + } +} + +func TestExtractContentFromInput(t *testing.T) { + tests := []struct { + name string + input interface{} + expected string + }{ + { + name: "string input", + input: "Hello world", + expected: "Hello world", + }, + { + name: "array with text parts", + input: []interface{}{ + map[string]interface{}{ + "type": "text", + "text": "Hello", + }, + map[string]interface{}{ + "type": "text", + "text": "world", + }, + }, + expected: "Hello world", + }, + { + name: "nil input", + input: nil, + expected: "", + }, + { + name: "unsupported type", + input: 12345, + expected: "", + }, + { + name: "array with non-text parts", + input: []interface{}{ + map[string]interface{}{ + "type": "image", + "url": "http://example.com/image.png", + }, + }, + expected: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := extractContentFromInput(tt.input) + if result != tt.expected { + t.Errorf("extractContentFromInput(%v) = %q, want %q", tt.input, result, tt.expected) + } + }) + } +} + +func TestGroqResponsesStreamConverter(t *testing.T) { + // Test the stream converter with mock chat completion stream + mockStream := `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":null}]} + +data: [DONE] +` + + reader := io.NopCloser(strings.NewReader(mockStream)) + converter := newGroqResponsesStreamConverter(reader, "llama-3.3-70b-versatile") + + // Read all data from converter + data, err := io.ReadAll(converter) + if err != nil { + t.Fatalf("failed to read from converter: %v", err) + } + + result := string(data) + + // Check that the stream contains expected events + if !strings.Contains(result, "response.created") { + t.Error("stream should contain response.created event") + } + if !strings.Contains(result, "response.output_text.delta") { + t.Error("stream should contain response.output_text.delta event") + } + if !strings.Contains(result, "Hello") { + t.Error("stream should contain 'Hello' content") + } + if !strings.Contains(result, " world") { + t.Error("stream should contain ' world' content") + } + if !strings.Contains(result, "response.done") { + t.Error("stream should contain response.done event") + } + if !strings.Contains(result, "[DONE]") { + t.Error("stream should contain [DONE] marker") + } +} + +func TestGroqResponsesStreamConverter_Close(t *testing.T) { + reader := io.NopCloser(strings.NewReader("data: [DONE]\n")) + converter := newGroqResponsesStreamConverter(reader, "test-model") + + err := converter.Close() + if err != nil { + t.Errorf("Close() returned error: %v", err) + } + + // Subsequent reads should return EOF + buf := make([]byte, 100) + n, err := converter.Read(buf) + if n != 0 || err != io.EOF { + t.Errorf("Read after Close: n=%d, err=%v, want n=0, err=EOF", n, err) + } +} + +func TestGroqResponsesStreamConverter_EmptyDelta(t *testing.T) { + // Test that empty deltas are not emitted + mockStream := `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":""},"finish_reason":null}]} + +data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"llama-3.3-70b-versatile","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} + +data: [DONE] +` + + reader := io.NopCloser(strings.NewReader(mockStream)) + converter := newGroqResponsesStreamConverter(reader, "llama-3.3-70b-versatile") + + data, err := io.ReadAll(converter) + if err != nil { + t.Fatalf("failed to read from converter: %v", err) + } + + result := string(data) + + // Count delta event lines - should only have one with "Hello" + // Each event has "event: response.output_text.delta\n" line + deltaCount := strings.Count(result, "event: response.output_text.delta") + if deltaCount != 1 { + t.Errorf("expected 1 delta event line, got %d", deltaCount) + } + + // Verify the Hello content is present + if !strings.Contains(result, `"delta":"Hello"`) { + t.Error("expected delta with Hello content") + } +} + +func TestNewWithHTTPClient(t *testing.T) { + customClient := &http.Client{} + apiKey := "test-api-key" + + provider := NewWithHTTPClient(apiKey, customClient) + + if provider.apiKey != apiKey { + t.Errorf("apiKey = %q, want %q", provider.apiKey, apiKey) + } + if provider.client == nil { + t.Error("client should not be nil") + } +} + +func TestSetBaseURL(t *testing.T) { + provider := New("test-api-key") + customURL := "https://custom.groq.api.com/v1" + + provider.SetBaseURL(customURL) + + // We can't directly check the baseURL as it's encapsulated in llmclient + // but we can verify the provider still works by making a test request + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) + })) + defer server.Close() + + provider.SetBaseURL(server.URL) + _, err := provider.ListModels(context.Background()) + if err != nil { + t.Errorf("SetBaseURL should allow using custom URL: %v", err) + } +}