feat(providers): prefer OpenAI responses API with fallback
This commit is contained in:
parent
4768edc67b
commit
1655561ac8
5 changed files with 1074 additions and 30 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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...),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue