diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index c05fb0ad4..9eadb8e79 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -10,6 +10,7 @@ import ( "strings" "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/providers/openai_compat" ) // createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store. @@ -84,12 +85,16 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if apiBase == "" { apiBase = getDefaultAPIBase(protocol) } + // The factory strips the outer protocol prefix before calling the HTTP + // provider, so pass an explicit hint to preserve the requested + // OpenAI-specific /responses-first behavior. return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout( cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField, cfg.RequestTimeout, + openai_compat.WithResponsesPreferred(), ), modelID, nil case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia", diff --git a/pkg/providers/factory_provider_test.go b/pkg/providers/factory_provider_test.go index 78389f331..928c2b3f2 100644 --- a/pkg/providers/factory_provider_test.go +++ b/pkg/providers/factory_provider_test.go @@ -8,6 +8,7 @@ package providers import ( "net/http" "net/http/httptest" + "reflect" "strings" "testing" "time" @@ -99,6 +100,56 @@ func TestCreateProviderFromConfig_OpenAI(t *testing.T) { } } +func TestCreateProviderFromConfig_OpenAIUsesResponsesFirst(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"status":"completed","output":[{"type":"message","content":[{"type":"output_text","text":"from responses"}]}]}`)) + case "/chat/completions": + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"from chat completions"},"finish_reason":"stop"}]}`)) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + cfg := &config.ModelConfig{ + ModelName: "test-openai", + Model: "openai/gpt-4o", + APIKey: "test-key", + APIBase: server.URL, + } + + provider, modelID, err := CreateProviderFromConfig(cfg) + if err != nil { + t.Fatalf("CreateProviderFromConfig() error = %v", err) + } + + out, err := provider.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + modelID, + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "from responses" { + t.Fatalf("Content = %q, want %q", out.Content, "from responses") + } + if !reflect.DeepEqual(paths, []string{"/responses"}) { + t.Fatalf("paths = %v, want [/responses]", paths) + } +} + func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) { tests := []struct { name string diff --git a/pkg/providers/http_provider.go b/pkg/providers/http_provider.go index 5c328f418..a46c868a0 100644 --- a/pkg/providers/http_provider.go +++ b/pkg/providers/http_provider.go @@ -17,9 +17,11 @@ type HTTPProvider struct { delegate *openai_compat.Provider } -func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider { +// NewHTTPProvider forwards optional provider-specific compatibility flags +// without changing the shared HTTP provider interface. +func NewHTTPProvider(apiKey, apiBase, proxy string, opts ...openai_compat.Option) *HTTPProvider { return &HTTPProvider{ - delegate: openai_compat.NewProvider(apiKey, apiBase, proxy), + delegate: openai_compat.NewProvider(apiKey, apiBase, proxy, opts...), } } @@ -30,15 +32,18 @@ func NewHTTPProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField st func NewHTTPProviderWithMaxTokensFieldAndRequestTimeout( apiKey, apiBase, proxy, maxTokensField string, requestTimeoutSeconds int, + opts ...openai_compat.Option, ) *HTTPProvider { + // Apply the legacy defaults first, then append any protocol-specific + // behavior switches such as OpenAI's /responses preference. + providerOpts := []openai_compat.Option{ + openai_compat.WithMaxTokensField(maxTokensField), + openai_compat.WithRequestTimeout(time.Duration(requestTimeoutSeconds) * time.Second), + } + providerOpts = append(providerOpts, opts...) + return &HTTPProvider{ - delegate: openai_compat.NewProvider( - apiKey, - apiBase, - proxy, - openai_compat.WithMaxTokensField(maxTokensField), - openai_compat.WithRequestTimeout(time.Duration(requestTimeoutSeconds)*time.Second), - ), + delegate: openai_compat.NewProvider(apiKey, apiBase, proxy, providerOpts...), } } diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index caee6f9ce..3e45fbc9a 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -5,6 +5,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "log" @@ -30,10 +31,11 @@ type ( ) type Provider struct { - apiKey string - apiBase string - maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models) - httpClient *http.Client + apiKey string + apiBase string + maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models) + httpClient *http.Client + preferResponses bool // Prefer /responses for OpenAI-native models selected via the factory. } type Option func(*Provider) @@ -54,6 +56,14 @@ func WithRequestTimeout(timeout time.Duration) Option { } } +// WithResponsesPreferred marks this provider instance as OpenAI-native so it +// prefers /responses even after the factory strips the outer "openai/" prefix. +func WithResponsesPreferred() Option { + return func(p *Provider) { + p.preferResponses = true + } +} + func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider { client := &http.Client{ Timeout: defaultRequestTimeout, @@ -113,8 +123,62 @@ func (p *Provider) Chat( return nil, fmt.Errorf("API base not configured") } - model = normalizeModel(model, p.apiBase) + normalizedModel := normalizeModel(model, p.apiBase) + // Keep the legacy chat/completions path for histories that already depend on + // reasoning_content, because Responses represents reasoning state differently. + if shouldPreferResponses(model, normalizedModel, p.preferResponses) && !hasReasoningContentHistory(messages) { + out, err := p.chatResponses(ctx, messages, tools, normalizedModel, options) + if err == nil { + return out, nil + } + if ctx.Err() != nil { + return nil, err + } + log.Printf("openai_compat: /responses failed for %q, falling back to /chat/completions: %v", normalizedModel, err) + fallbackOut, fallbackErr := p.chatCompletions(ctx, messages, tools, normalizedModel, options) + if fallbackErr != nil { + return nil, fmt.Errorf("responses request failed: %w; fallback chat/completions failed: %v", err, fallbackErr) + } + return fallbackOut, nil + } + + return p.chatCompletions(ctx, messages, tools, normalizedModel, options) +} + +func (p *Provider) chatCompletions( + ctx context.Context, + messages []Message, + tools []ToolDefinition, + model string, + options map[string]any, +) (*LLMResponse, error) { + requestBody := buildChatCompletionsRequestBody(messages, tools, model, options, p.maxTokensField, p.apiBase) + return p.doRequest(ctx, "/chat/completions", requestBody, parseResponse) +} + +func (p *Provider) chatResponses( + ctx context.Context, + messages []Message, + tools []ToolDefinition, + model string, + options map[string]any, +) (*LLMResponse, error) { + requestBody, err := buildResponsesRequestBody(messages, tools, model, options, p.apiBase) + if err != nil { + return nil, err + } + return p.doRequest(ctx, "/responses", requestBody, parseResponsesResponse) +} + +func buildChatCompletionsRequestBody( + messages []Message, + tools []ToolDefinition, + model string, + options map[string]any, + maxTokensField string, + apiBase string, +) map[string]any { requestBody := map[string]any{ "model": model, "messages": serializeMessages(messages), @@ -126,10 +190,10 @@ func (p *Provider) Chat( } if maxTokens, ok := asInt(options["max_tokens"]); ok { - // Use configured maxTokensField if specified, otherwise fallback to model-based detection - fieldName := p.maxTokensField + // Use configured maxTokensField if specified, otherwise fallback to model-based detection. + fieldName := maxTokensField if fieldName == "" { - // Fallback: detect from model name for backward compatibility + // Fallback: detect from model name for backward compatibility. lowerModel := strings.ToLower(model) if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") || strings.Contains(lowerModel, "gpt-5") { @@ -141,34 +205,257 @@ func (p *Provider) Chat( requestBody[fieldName] = maxTokens } - if temperature, ok := asFloat(options["temperature"]); ok { - lowerModel := strings.ToLower(model) - // Kimi k2 models only support temperature=1. - if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") { - requestBody["temperature"] = 1.0 - } else { - requestBody["temperature"] = temperature - } + if temperature, ok := requestTemperature(model, options); ok { + requestBody["temperature"] = temperature } // Prompt caching: pass a stable cache key so OpenAI can bucket requests // with the same key and reuse prefix KV cache across calls. - // The key is typically the agent ID — stable per agent, shared across requests. + // The key is typically the agent ID - stable per agent, shared across requests. // See: https://platform.openai.com/docs/guides/prompt-caching // Prompt caching is only supported by OpenAI-native endpoints. // Gemini and other providers reject unknown fields, so skip for non-OpenAI APIs. if cacheKey, ok := options["prompt_cache_key"].(string); ok && cacheKey != "" { - if !strings.Contains(p.apiBase, "generativelanguage.googleapis.com") { + if !strings.Contains(apiBase, "generativelanguage.googleapis.com") { requestBody["prompt_cache_key"] = cacheKey } } + return requestBody +} + +// buildResponsesRequestBody keeps the option handling close to the legacy +// chat/completions path so the new route can reuse the existing compatibility +// knobs with minimal behavioral drift. +func buildResponsesRequestBody( + messages []Message, + tools []ToolDefinition, + model string, + options map[string]any, + apiBase string, +) (map[string]any, error) { + input, err := buildResponsesInput(messages) + if err != nil { + return nil, err + } + + requestBody := map[string]any{ + "model": model, + "input": input, + } + + if len(tools) > 0 { + requestBody["tools"] = serializeResponseTools(tools) + requestBody["tool_choice"] = "auto" + } + + if maxTokens, ok := asInt(options["max_tokens"]); ok { + requestBody["max_output_tokens"] = maxTokens + } + + if temperature, ok := requestTemperature(model, options); ok { + requestBody["temperature"] = temperature + } + + // Prompt caching follows the same compatibility rule as chat/completions: + // send the key only to endpoints that are expected to understand it. + if cacheKey, ok := options["prompt_cache_key"].(string); ok && cacheKey != "" { + if !strings.Contains(apiBase, "generativelanguage.googleapis.com") { + requestBody["prompt_cache_key"] = cacheKey + } + } + + return requestBody, nil +} + +// buildResponsesInput translates the existing conversation format into the +// item-based Responses input shape while preserving tool call history. +func buildResponsesInput(messages []Message) ([]any, error) { + input := make([]any, 0, len(messages)) + + for _, m := range messages { + switch m.Role { + case "system", "user": + input = append(input, map[string]any{ + "type": "message", + "role": m.Role, + "content": serializeResponsesMessageContent(m), + }) + case "assistant": + if strings.TrimSpace(m.Content) != "" || strings.TrimSpace(m.ReasoningContent) != "" || len(m.Media) > 0 || len(m.ToolCalls) == 0 { + input = append(input, map[string]any{ + "type": "message", + "role": m.Role, + "content": serializeResponsesMessageContent(m), + }) + } + + for _, tc := range m.ToolCalls { + name, args, ok := resolveResponseToolCall(tc) + if !ok { + log.Printf("openai_compat: skipping invalid assistant tool call in responses history: id=%q", tc.ID) + continue + } + input = append(input, map[string]any{ + "type": "function_call", + "call_id": tc.ID, + "name": name, + "arguments": args, + }) + } + case "tool": + if strings.TrimSpace(m.ToolCallID) == "" { + return nil, fmt.Errorf("tool message missing tool_call_id") + } + input = append(input, map[string]any{ + "type": "function_call_output", + "call_id": m.ToolCallID, + "output": m.Content, + }) + default: + return nil, fmt.Errorf("unsupported message role: %s", m.Role) + } + } + + return input, nil +} + +// serializeResponsesMessageContent converts plain text and inline image data +// into the content format expected by the Responses API. +func serializeResponsesMessageContent(m Message) any { + effectiveText := m.Content + if effectiveText == "" { + effectiveText = m.ReasoningContent + } + + if len(m.Media) == 0 { + return effectiveText + } + + parts := make([]map[string]any, 0, 1+len(m.Media)) + if effectiveText != "" { + parts = append(parts, map[string]any{ + "type": "input_text", + "text": effectiveText, + }) + } + + for _, mediaURL := range m.Media { + if strings.HasPrefix(mediaURL, "data:image/") { + parts = append(parts, map[string]any{ + "type": "input_image", + "image_url": mediaURL, + }) + } + } + + if len(parts) == 0 { + return effectiveText + } + + return parts +} + +// serializeResponseTools maps the existing OpenAI-compatible tool schema to the +// smaller function-tool shape accepted by the Responses API. +func serializeResponseTools(tools []ToolDefinition) []map[string]any { + result := make([]map[string]any, 0, len(tools)) + for _, tool := range tools { + if tool.Type != "" && tool.Type != "function" { + continue + } + + entry := map[string]any{ + "type": "function", + "name": tool.Function.Name, + "parameters": tool.Function.Parameters, + } + if entry["parameters"] == nil { + entry["parameters"] = map[string]any{"type": "object", "properties": map[string]any{}} + } + if tool.Function.Description != "" { + entry["description"] = tool.Function.Description + } + + result = append(result, entry) + } + return result +} + +// resolveResponseToolCall rebuilds the assistant-side tool call record into the +// stringified argument form required by Responses conversation history. +func resolveResponseToolCall(tc ToolCall) (name string, arguments string, ok bool) { + name = tc.Name + if name == "" && tc.Function != nil { + name = tc.Function.Name + } + if name == "" { + return "", "", false + } + + if len(tc.Arguments) > 0 { + argsJSON, err := json.Marshal(tc.Arguments) + if err != nil { + return "", "", false + } + return name, string(argsJSON), true + } + + if tc.Function != nil && tc.Function.Arguments != "" { + return name, tc.Function.Arguments, true + } + + return name, "{}", true +} + +func requestTemperature(model string, options map[string]any) (float64, bool) { + temperature, ok := asFloat(options["temperature"]) + if !ok { + return 0, false + } + + lowerModel := strings.ToLower(model) + if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") { + return 1.0, true + } + return temperature, true +} + +// shouldPreferResponses centralizes the opt-in rule so OpenAI-native configs +// and gpt-5 models can try /responses first while other compat backends keep +// their existing chat/completions behavior. +func shouldPreferResponses(rawModel, normalizedModel string, preferOpenAIModels bool) bool { + rawModel = strings.ToLower(strings.TrimSpace(rawModel)) + normalizedModel = strings.ToLower(strings.TrimSpace(normalizedModel)) + + return preferOpenAIModels || strings.HasPrefix(rawModel, "openai/") || + strings.HasPrefix(rawModel, "gpt-5") || + strings.HasPrefix(normalizedModel, "gpt-5") +} + +// hasReasoningContentHistory detects histories that already rely on the legacy +// reasoning_content field so they can stay on the older wire format. +func hasReasoningContentHistory(messages []Message) bool { + for _, message := range messages { + if strings.TrimSpace(message.ReasoningContent) != "" { + return true + } + } + return false +} + +func (p *Provider) doRequest( + ctx context.Context, + path string, + requestBody map[string]any, + parse func(io.Reader) (*LLMResponse, error), +) (*LLMResponse, error) { jsonData, err := json.Marshal(requestBody) if err != nil { return nil, fmt.Errorf("failed to marshal request: %w", err) } - req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData)) + req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+path, bytes.NewReader(jsonData)) if err != nil { return nil, fmt.Errorf("failed to create request: %w", err) } @@ -185,7 +472,6 @@ func (p *Provider) Chat( defer resp.Body.Close() contentType := resp.Header.Get("Content-Type") - // Non-200: read a prefix to tell HTML error page apart from JSON error body. if resp.StatusCode != http.StatusOK { body, readErr := io.ReadAll(io.LimitReader(resp.Body, 256)) @@ -212,7 +498,7 @@ func (p *Provider) Chat( return nil, wrapHTMLResponseError(resp.StatusCode, prefix, contentType, p.apiBase) } - out, err := parseResponse(reader) + out, err := parse(reader) if err != nil { return nil, fmt.Errorf("failed to parse JSON response: %w", err) } @@ -361,6 +647,162 @@ func parseResponse(body io.Reader) (*LLMResponse, error) { }, nil } +// parseResponsesResponse maps the Responses API envelope back to the legacy +// provider response shape used by the rest of the codebase. +func parseResponsesResponse(body io.Reader) (*LLMResponse, error) { + var apiResponse struct { + Status string `json:"status"` + Error *struct { + Message string `json:"message"` + } `json:"error"` + Output []struct { + ID string `json:"id"` + Type string `json:"type"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` + Summary []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"summary"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + Refusal string `json:"refusal"` + } `json:"content"` + } `json:"output"` + IncompleteDetails *struct { + Reason string `json:"reason"` + } `json:"incomplete_details"` + Usage *struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` + } + + if err := json.NewDecoder(body).Decode(&apiResponse); err != nil { + return nil, fmt.Errorf("failed to decode response: %w", err) + } + if strings.TrimSpace(apiResponse.Status) == "" && len(apiResponse.Output) == 0 { + return nil, errors.New("openai responses returned unexpected response shape") + } + + var content strings.Builder + var reasoning strings.Builder + var reasoningContent strings.Builder + reasoningDetails := make([]ReasoningDetail, 0) + toolCalls := make([]ToolCall, 0) + for _, item := range apiResponse.Output { + switch item.Type { + case "message": + for _, part := range item.Content { + if part.Text != "" { + content.WriteString(part.Text) + continue + } + if part.Refusal != "" { + content.WriteString(part.Refusal) + } + } + case "reasoning": + for _, part := range item.Summary { + if part.Text == "" { + continue + } + if reasoning.Len() > 0 { + reasoning.WriteString("\n") + } + reasoning.WriteString(part.Text) + reasoningDetails = append(reasoningDetails, ReasoningDetail{ + Format: "text", + Index: len(reasoningDetails), + Type: part.Type, + Text: part.Text, + }) + } + for _, part := range item.Content { + if part.Text == "" { + continue + } + if reasoningContent.Len() > 0 { + reasoningContent.WriteString("\n") + } + reasoningContent.WriteString(part.Text) + reasoningDetails = append(reasoningDetails, ReasoningDetail{ + Format: "text", + Index: len(reasoningDetails), + Type: part.Type, + Text: part.Text, + }) + } + case "function_call": + arguments := make(map[string]any) + if item.Arguments != "" { + if err := json.Unmarshal([]byte(item.Arguments), &arguments); err != nil { + log.Printf("openai_compat: failed to decode responses tool call arguments for %q: %v", item.Name, err) + arguments["raw"] = item.Arguments + } + } + + toolCalls = append(toolCalls, ToolCall{ + ID: firstNonEmpty(item.CallID, item.ID), + Name: item.Name, + Arguments: arguments, + }) + } + } + + if apiResponse.Status == "failed" { + if apiResponse.Error != nil && apiResponse.Error.Message != "" { + return nil, errors.New(apiResponse.Error.Message) + } + return nil, errors.New("openai responses request failed") + } + + finishReason := "stop" + if len(toolCalls) > 0 { + finishReason = "tool_calls" + } else if apiResponse.Status == "incomplete" { + finishReason = "length" + if apiResponse.IncompleteDetails != nil && apiResponse.IncompleteDetails.Reason != "" && apiResponse.IncompleteDetails.Reason != "max_output_tokens" { + finishReason = apiResponse.IncompleteDetails.Reason + } + } else if apiResponse.Status == "failed" { + finishReason = "error" + } + + var usage *UsageInfo + if apiResponse.Usage != nil { + usage = &UsageInfo{ + PromptTokens: apiResponse.Usage.InputTokens, + CompletionTokens: apiResponse.Usage.OutputTokens, + TotalTokens: apiResponse.Usage.TotalTokens, + } + } + + return &LLMResponse{ + Content: content.String(), + ReasoningContent: reasoningContent.String(), + Reasoning: reasoning.String(), + ReasoningDetails: reasoningDetails, + ToolCalls: toolCalls, + FinishReason: finishReason, + Usage: usage, + }, nil +} + +// firstNonEmpty prefers call_id but falls back to the raw item id when the +// response item omits it. +func firstNonEmpty(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" { + return value + } + } + return "" +} + // openaiMessage is the wire-format message for OpenAI-compatible APIs. // It mirrors protocoltypes.Message but omits SystemParts, which is an // internal field that would be unknown to third-party endpoints. diff --git a/pkg/providers/openai_compat/provider_test.go b/pkg/providers/openai_compat/provider_test.go index 5c4dcd1b0..ecf9922d8 100644 --- a/pkg/providers/openai_compat/provider_test.go +++ b/pkg/providers/openai_compat/provider_test.go @@ -8,6 +8,7 @@ import ( "net/http" "net/http/httptest" "net/url" + "reflect" "strings" "testing" "time" @@ -15,6 +16,546 @@ import ( "github.com/sipeed/picoclaw/pkg/providers/protocoltypes" ) +func TestProviderChat_PrefersResponsesForOpenAIPrefixedModel(t *testing.T) { + var paths []string + var responsesBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + if err := json.NewDecoder(r.Body).Decode(&responsesBody); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{ + { + "type": "message", + "content": []map[string]any{ + {"type": "output_text", "text": "from responses"}, + }, + }, + }, + "usage": map[string]any{ + "input_tokens": 12, + "output_tokens": 3, + "total_tokens": 15, + }, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "from chat completions"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + "openai/gpt-4o", + map[string]any{"max_tokens": 256}, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "from responses" { + t.Fatalf("Content = %q, want %q", out.Content, "from responses") + } + if !reflect.DeepEqual(paths, []string{"/responses"}) { + t.Fatalf("paths = %v, want [/responses]", paths) + } + if responsesBody["model"] != "openai/gpt-4o" { + t.Fatalf("model = %v, want openai/gpt-4o", responsesBody["model"]) + } + if _, ok := responsesBody["input"]; !ok { + t.Fatalf("expected responses request body to contain input") + } + if _, ok := responsesBody["messages"]; ok { + t.Fatalf("did not expect messages in responses request body") + } + if responsesBody["max_output_tokens"] != float64(256) { + t.Fatalf("max_output_tokens = %v, want 256", responsesBody["max_output_tokens"]) + } +} + +func TestProviderChat_FallsBackToChatCompletionsWhenResponsesFails(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":"responses not supported"}`)) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "fallback chat completion"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "fallback chat completion" { + t.Fatalf("Content = %q, want %q", out.Content, "fallback chat completion") + } + if !reflect.DeepEqual(paths, []string{"/responses", "/chat/completions"}) { + t.Fatalf("paths = %v, want [/responses /chat/completions]", paths) + } +} + +func TestProviderChat_ParsesToolCallsFromResponses(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{ + { + "type": "function_call", + "id": "fc_1", + "call_id": "call_1", + "name": "get_weather", + "arguments": "{\"city\":\"SF\"}", + }, + }, + "usage": map[string]any{ + "input_tokens": 9, + "output_tokens": 4, + "total_tokens": 13, + }, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "from chat completions"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "weather?"}}, + []ToolDefinition{{ + Type: "function", + Function: ToolFunctionDefinition{ + Name: "get_weather", + Description: "Get weather", + Parameters: map[string]any{ + "type": "object", + "properties": map[string]any{ + "city": map[string]any{"type": "string"}, + }, + }, + }, + }}, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.FinishReason != "tool_calls" { + t.Fatalf("FinishReason = %q, want tool_calls", out.FinishReason) + } + if len(out.ToolCalls) != 1 { + t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls)) + } + if out.ToolCalls[0].ID != "call_1" { + t.Fatalf("ToolCalls[0].ID = %q, want call_1", out.ToolCalls[0].ID) + } + if out.ToolCalls[0].Name != "get_weather" { + t.Fatalf("ToolCalls[0].Name = %q, want get_weather", out.ToolCalls[0].Name) + } + if out.ToolCalls[0].Arguments["city"] != "SF" { + t.Fatalf("ToolCalls[0].Arguments[city] = %v, want SF", out.ToolCalls[0].Arguments["city"]) + } +} + +func TestProviderChat_FallsBackWhenResponsesStatusFailed(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "failed", + "error": map[string]any{ + "message": "responses failed", + }, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "fallback after failed status"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "fallback after failed status" { + t.Fatalf("Content = %q, want %q", out.Content, "fallback after failed status") + } + if !reflect.DeepEqual(paths, []string{"/responses", "/chat/completions"}) { + t.Fatalf("paths = %v, want [/responses /chat/completions]", paths) + } +} + +func TestProviderChat_ParsesReasoningContentFromResponses(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{ + { + "type": "reasoning", + "summary": []map[string]any{ + {"type": "summary_text", "text": "brief reasoning"}, + }, + "content": []map[string]any{ + {"type": "reasoning_text", "text": "step by step"}, + }, + }, + { + "type": "message", + "content": []map[string]any{ + {"type": "output_text", "text": "final answer"}, + }, + }, + }, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "chat fallback"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "why?"}}, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "final answer" { + t.Fatalf("Content = %q, want %q", out.Content, "final answer") + } + if out.Reasoning != "brief reasoning" { + t.Fatalf("Reasoning = %q, want %q", out.Reasoning, "brief reasoning") + } + if out.ReasoningContent != "step by step" { + t.Fatalf("ReasoningContent = %q, want %q", out.ReasoningContent, "step by step") + } +} + +func TestProviderChat_FallsBackWhenResponsesReturnsUnexpected200Body(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "wrong envelope"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "fallback after invalid responses body"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "fallback after invalid responses body" { + t.Fatalf("Content = %q, want %q", out.Content, "fallback after invalid responses body") + } + if !reflect.DeepEqual(paths, []string{"/responses", "/chat/completions"}) { + t.Fatalf("paths = %v, want [/responses /chat/completions]", paths) + } +} + +func TestProviderChat_SkipsResponsesWhenHistoryHasReasoningContent(t *testing.T) { + var paths []string + var requestBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{{ + "type": "message", + "content": []map[string]any{{"type": "output_text", "text": "responses path"}}, + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "chat path"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{ + {Role: "user", Content: "1+1?"}, + {Role: "assistant", Content: "2", ReasoningContent: "internal reasoning"}, + {Role: "user", Content: "2+2?"}, + }, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "chat path" { + t.Fatalf("Content = %q, want %q", out.Content, "chat path") + } + if !reflect.DeepEqual(paths, []string{"/chat/completions"}) { + t.Fatalf("paths = %v, want [/chat/completions]", paths) + } + + reqMessages, ok := requestBody["messages"].([]any) + if !ok { + t.Fatalf("messages is not []any: %T", requestBody["messages"]) + } + assistantMsg, ok := reqMessages[1].(map[string]any) + if !ok { + t.Fatalf("assistant message is not map[string]any: %T", reqMessages[1]) + } + if assistantMsg["reasoning_content"] != "internal reasoning" { + t.Fatalf("reasoning_content = %v, want internal reasoning", assistantMsg["reasoning_content"]) + } +} + +func TestProviderChat_DoesNotPreferResponsesForNestedOpenAINamespace(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{{ + "type": "message", + "content": []map[string]any{{"type": "output_text", "text": "responses path"}}, + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "chat path"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "hi"}}, + nil, + "groq/openai/gpt-oss-120b", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "chat path" { + t.Fatalf("Content = %q, want %q", out.Content, "chat path") + } + if !reflect.DeepEqual(paths, []string{"/chat/completions"}) { + t.Fatalf("paths = %v, want [/chat/completions]", paths) + } +} + +func TestProviderChat_ParsesRefusalFromResponses(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/responses": + resp := map[string]any{ + "status": "completed", + "output": []map[string]any{{ + "type": "message", + "content": []map[string]any{{"type": "refusal", "refusal": "I can't help with that."}}, + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + case "/chat/completions": + resp := map[string]any{ + "choices": []map[string]any{{ + "message": map[string]any{"content": "chat fallback"}, + "finish_reason": "stop", + }}, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + default: + http.Error(w, "not found", http.StatusNotFound) + } + })) + defer server.Close() + + p := NewProvider("key", server.URL, "") + out, err := p.Chat( + t.Context(), + []Message{{Role: "user", Content: "unsafe request"}}, + nil, + "gpt-5.2", + nil, + ) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + if out.Content != "I can't help with that." { + t.Fatalf("Content = %q, want %q", out.Content, "I can't help with that.") + } +} + func TestProviderChat_UsesMaxCompletionTokensForGLM(t *testing.T) { var requestBody map[string]any