diff --git a/README.md b/README.md index 0b2e862..beaaf17 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,7 @@ model/provider switching, and local tools. - TUI-first chat experience with a light theme by default. - OpenAI-compatible `/v1/chat/completions` provider support. +- OpenAI Responses API (`/v1/responses`) provider support. - Multiple providers from `~/.agentu/config.yaml`. - Session-only provider, model, and thinking-level switching. - Read-only project tools enabled by default. @@ -48,6 +49,7 @@ Use `--config ` to load another file. ```yaml providers: loveuer: + api_type: chat base_url: https://ai.loveuer.com api_key: ${AGENTU_API_KEY} models: @@ -64,6 +66,19 @@ Provider names are the keys under `providers`. At startup, agentu selects the first provider by name. Use `/model provider ` to switch providers for the current session. +`api_type` selects which API protocol the provider speaks: + +```text +chat OpenAI-compatible /v1/chat/completions (default) +responses OpenAI Responses API /v1/responses +``` + +The Responses API client converts the conversation history to Responses input +items (`user`/`assistant`/`system` messages, `function_call`, and +`function_call_output`), streams `response.output_text.delta` events, and +supports function calling through `response.output_item.added`, +`response.function_call_arguments.delta`, and `response.output_item.done`. + `models` is required. `/model model ` is restricted to that list. Each entry is an object with: @@ -83,7 +98,10 @@ none, middle, high, xhigh, max ``` When configured or changed with `/model thinking ...`, the value is sent as a -top-level OpenAI-compatible request field named `thinking`. +top-level request field named `thinking`. For OpenAI reasoning models on the +Responses API, set `thinking_param: reasoning` so the level is sent as the +official `reasoning: {"effort": ...}` field; agentu maps its `middle` level to +the official `medium` value and passes `none`, `high`, `xhigh`, `max` through. ## CLI Flags @@ -174,5 +192,8 @@ Local config and build outputs are intentionally ignored by git. - Provider/model/thinking changes are session-only; agentu does not write them back to `~/.agentu/config.yaml`. +- `api_type` is per provider and cannot be changed at runtime; use + `/model provider ` to switch to a provider configured with a different + API type. - The config schema is intentionally strict during MVP development. Deprecated fields such as `provider` or `active_provider` are rejected. diff --git a/agentu.example.yaml b/agentu.example.yaml index e6a5a52..e1655c2 100644 --- a/agentu.example.yaml +++ b/agentu.example.yaml @@ -1,5 +1,8 @@ providers: loveuer: + # API protocol: "chat" (OpenAI-compatible /v1/chat/completions, default) + # or "responses" (OpenAI Responses API /v1/responses). + api_type: chat base_url: https://ai.loveuer.com api_key: ${AGENTU_API_KEY} models: diff --git a/internal/app/app.go b/internal/app/app.go index 37e3e99..698c4af 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -91,10 +91,17 @@ func loadConfig(configPath string) (*config.Config, error) { func defaultProvider(cfg *config.Config, httpClient *http.Client) (string, config.ProviderConfig, llm.Provider) { providerName := cfg.DefaultProviderName() providerConfig := cfg.Providers[providerName] - provider := llm.NewOpenAICompatibleClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient) + provider := newProviderClient(providerConfig, httpClient) return providerName, providerConfig, provider } +func newProviderClient(providerConfig config.ProviderConfig, httpClient *http.Client) llm.Provider { + if providerConfig.APIType == config.APITypeResponses { + return llm.NewOpenAIResponsesClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient) + } + return llm.NewOpenAICompatibleClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient) +} + func newAssistant(cfg *config.Config, providerName string, providerConfig config.ProviderConfig, provider llm.Provider, yolo bool, httpClient *http.Client) *agent.Agent { return agent.New(agent.Options{ Provider: provider, diff --git a/internal/session/manager.go b/internal/session/manager.go index 1c7ce6f..ba581a6 100644 --- a/internal/session/manager.go +++ b/internal/session/manager.go @@ -204,10 +204,18 @@ func (m *Manager) CurrentThinking() string { return m.provider.Thinking } +func (m *Manager) CurrentAPIType() string { + if m.provider.APIType == "" { + return config.APITypeChat + } + return m.provider.APIType +} + func (m *Manager) Info() string { var b strings.Builder fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider()) fmt.Fprintf(&b, "model: %s\n", m.CurrentModel()) + fmt.Fprintf(&b, "api: %s\n", m.CurrentAPIType()) fmt.Fprintf(&b, "thinking: %s\n", m.CurrentThinking()) fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", ")) fmt.Fprintf(&b, "models: %s\n", strings.Join(m.provider.Models, ", ")) @@ -223,7 +231,7 @@ func (m *Manager) SetProvider(name string) (string, error) { } m.providerName = name m.provider = provider - m.syncAgent(llm.NewOpenAICompatibleClient(provider.BaseURL, provider.APIKey, m.httpClient)) + m.syncAgent(providerClient(provider, m.httpClient)) return "Switched provider.\n" + m.summary(), nil } @@ -257,7 +265,14 @@ func (m *Manager) SetThinking(level string) (string, error) { } func (m *Manager) summary() string { - return fmt.Sprintf("provider: %s\nmodel: %s\nthinking: %s", m.CurrentProvider(), m.CurrentModel(), m.CurrentThinking()) + return fmt.Sprintf("provider: %s\nmodel: %s\napi: %s\nthinking: %s", m.CurrentProvider(), m.CurrentModel(), m.CurrentAPIType(), m.CurrentThinking()) +} + +func providerClient(provider config.ProviderConfig, httpClient *http.Client) llm.Provider { + if provider.APIType == config.APITypeResponses { + return llm.NewOpenAIResponsesClient(provider.BaseURL, provider.APIKey, httpClient) + } + return llm.NewOpenAICompatibleClient(provider.BaseURL, provider.APIKey, httpClient) } func (m *Manager) maxContextTokens() int { diff --git a/internal/session/manager_test.go b/internal/session/manager_test.go index 71cdd7b..de5186a 100644 --- a/internal/session/manager_test.go +++ b/internal/session/manager_test.go @@ -2,6 +2,8 @@ package session import ( "context" + "net/http" + "net/http/httptest" "path/filepath" "strings" "testing" @@ -109,6 +111,64 @@ func TestManagerRejectsUnknownModelAndThinking(t *testing.T) { } } +func TestManagerSwitchesToResponsesProvider(t *testing.T) { + var hitPath string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hitPath = r.URL.Path + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n")) + _, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n")) + })) + defer server.Close() + + cfg := &config.Config{ + Providers: map[string]config.ProviderConfig{ + "chat": { + Name: "chat", + APIType: config.APITypeChat, + BaseURL: "https://chat.example.com", + APIKey: "sk-test", + Model: "model-a", + Models: []string{"model-a"}, + }, + "responses": { + Name: "responses", + APIType: config.APITypeResponses, + BaseURL: server.URL, + APIKey: "sk-test", + Model: "model-b", + Models: []string{"model-b"}, + }, + }, + } + assistant := agent.New(agent.Options{Provider: &recordingProvider{}, Model: "model-a"}) + store, err := NewStore(filepath.Join(t.TempDir(), "sessions")) + if err != nil { + t.Fatal(err) + } + manager, err := NewManager(cfg, assistant, server.Client(), store) + if err != nil { + t.Fatal(err) + } + + if _, err := manager.SetProvider("responses"); err != nil { + t.Fatal(err) + } + if manager.CurrentAPIType() != config.APITypeResponses { + t.Fatalf("api type = %q", manager.CurrentAPIType()) + } + var out strings.Builder + if err := assistant.RunTurn(context.Background(), "hi", &out, nil); err != nil { + t.Fatal(err) + } + if hitPath != "/v1/responses" { + t.Fatalf("path = %q, want /v1/responses", hitPath) + } + if out.String() != "ok" { + t.Fatalf("output = %q", out.String()) + } +} + func TestManagerExposesModelCommandCandidates(t *testing.T) { manager, _, _ := testManager(t) models := manager.AvailableModels() diff --git a/pkg/config/config.go b/pkg/config/config.go index 5dc0ea7..57257b9 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -14,6 +14,11 @@ import ( const DefaultPath = "~/.agentu/config.yaml" +const ( + APITypeChat = "chat" + APITypeResponses = "responses" +) + type Config struct { Providers map[string]ProviderConfig `yaml:"providers"` Agent AgentConfig `yaml:"agent"` @@ -21,6 +26,7 @@ type Config struct { type ProviderConfig struct { Name string + APIType string BaseURL string Model string Models []string @@ -56,6 +62,7 @@ type rawConfig struct { } type rawProviderConfig struct { + APIType string `yaml:"api_type"` BaseURL string `yaml:"base_url"` Models []rawModelConfig `yaml:"models"` APIKey string `yaml:"api_key"` @@ -153,6 +160,9 @@ func (c *Config) Validate() error { for name, provider := range c.Providers { var missing []string prefix := "providers." + name + if !IsAPIType(provider.APIType) { + return fmt.Errorf("%s.api_type must be one of: %s", prefix, strings.Join(APITypes(), ", ")) + } if strings.TrimSpace(provider.BaseURL) == "" { missing = append(missing, prefix+".base_url") } @@ -220,6 +230,7 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig { models, modelConfigs := normalizeModelConfigs(raw.Models) provider := ProviderConfig{ Name: name, + APIType: normalizeAPIType(raw.APIType), BaseURL: raw.BaseURL, Models: models, ModelConfigs: modelConfigs, @@ -235,6 +246,28 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig { return provider } +func normalizeAPIType(value string) string { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return APITypeChat + } + return value +} + +func APITypes() []string { + return []string{APITypeChat, APITypeResponses} +} + +func IsAPIType(value string) bool { + value = strings.ToLower(strings.TrimSpace(value)) + for _, apiType := range APITypes() { + if value == apiType { + return true + } + } + return false +} + func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) { models := make([]string, 0, len(rawModels)) configs := make(map[string]ModelConfig, len(rawModels)) diff --git a/pkg/config/config_test.go b/pkg/config/config_test.go index 5e6d7bd..2f54c7a 100644 --- a/pkg/config/config_test.go +++ b/pkg/config/config_test.go @@ -370,3 +370,74 @@ providers: t.Fatalf("err = %v", err) } } + +func TestLoadDefaultsAPITypeToChat(t *testing.T) { + dir := t.TempDir() + t.Setenv("AGENTU_TEST_KEY", "sk-test") + path := filepath.Join(dir, "agentu.yaml") + if err := os.WriteFile(path, []byte(` +providers: + test: + base_url: https://example.com + api_key: ${AGENTU_TEST_KEY} + models: + - id: test-model +`), 0o644); err != nil { + t.Fatal(err) + } + + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if got := cfg.Providers["test"].APIType; got != APITypeChat { + t.Fatalf("api type = %q, want %q", got, APITypeChat) + } +} + +func TestLoadParsesResponsesAPIType(t *testing.T) { + dir := t.TempDir() + t.Setenv("AGENTU_TEST_KEY", "sk-test") + path := filepath.Join(dir, "agentu.yaml") + if err := os.WriteFile(path, []byte(` +providers: + openai: + api_type: RESPONSES + base_url: https://api.openai.com + api_key: ${AGENTU_TEST_KEY} + models: + - id: gpt-test +`), 0o644); err != nil { + t.Fatal(err) + } + + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if got := cfg.Providers["openai"].APIType; got != APITypeResponses { + t.Fatalf("api type = %q, want %q", got, APITypeResponses) + } +} + +func TestLoadRejectsInvalidAPIType(t *testing.T) { + dir := t.TempDir() + t.Setenv("AGENTU_TEST_KEY", "sk-test") + path := filepath.Join(dir, "agentu.yaml") + if err := os.WriteFile(path, []byte(` +providers: + test: + api_type: completions + base_url: https://example.com + api_key: ${AGENTU_TEST_KEY} + models: + - id: test-model +`), 0o644); err != nil { + t.Fatal(err) + } + + _, err := Load(path) + if err == nil || !strings.Contains(err.Error(), "api_type") { + t.Fatalf("err = %v", err) + } +} diff --git a/pkg/llm/responses.go b/pkg/llm/responses.go new file mode 100644 index 0000000..1c56008 --- /dev/null +++ b/pkg/llm/responses.go @@ -0,0 +1,460 @@ +package llm + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "sort" + "strings" +) + +const responsesEndpoint = "/v1/responses" + +// OpenAIResponsesClient implements llm.Provider against the OpenAI Responses +// API (/v1/responses). It maps the agent's chat-style conversation history to +// Responses input items (messages, function_call, function_call_output) and +// converts Responses streaming events back into the shared StreamEvent shape. +type OpenAIResponsesClient struct { + baseURL string + apiKey string + httpClient *http.Client +} + +func NewOpenAIResponsesClient(baseURL, apiKey string, httpClient *http.Client) *OpenAIResponsesClient { + if httpClient == nil { + httpClient = http.DefaultClient + } + return &OpenAIResponsesClient{ + baseURL: strings.TrimRight(baseURL, "/"), + apiKey: apiKey, + httpClient: httpClient, + } +} + +func (c *OpenAIResponsesClient) ChatStream(ctx context.Context, req ChatRequest, emit func(StreamEvent) error) error { + body, err := json.Marshal(buildResponsesRequest(req)) + if err != nil { + return fmt.Errorf("marshal responses request: %w", err) + } + state := &responsesStreamState{} + return postResponsesSSE(ctx, c.baseURL, c.apiKey, c.httpClient, body, func(payload string) error { + return state.handle(payload, emit) + }) +} + +func postResponsesSSE(ctx context.Context, baseURL, apiKey string, httpClient *http.Client, body []byte, handle func(payload string) error) error { + var lastErr error + for attempt := 1; attempt <= maxChatAttempts; attempt++ { + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+responsesEndpoint, bytes.NewReader(body)) + if err != nil { + return fmt.Errorf("create responses request: %w", err) + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "text/event-stream") + if apiKey != "" { + httpReq.Header.Set("Authorization", "Bearer "+apiKey) + } + + resp, err := httpClient.Do(httpReq) + if err != nil { + lastErr = fmt.Errorf("send responses request: %w", err) + if !shouldRetryRequest(ctx, attempt, 0) { + return lastErr + } + if waitErr := waitBeforeRetry(ctx); waitErr != nil { + return waitErr + } + continue + } + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + limit := io.LimitReader(resp.Body, 4096) + data, _ := io.ReadAll(limit) + _ = resp.Body.Close() + lastErr = fmt.Errorf("responses request failed: status=%d body=%s", resp.StatusCode, strings.TrimSpace(string(data))) + if !shouldRetryRequest(ctx, attempt, resp.StatusCode) { + return lastErr + } + if waitErr := waitBeforeRetry(ctx); waitErr != nil { + return waitErr + } + continue + } + + defer resp.Body.Close() + return readResponsesSSE(resp.Body, handle) + } + return lastErr +} + +func readResponsesSSE(r io.Reader, handle func(payload string) error) error { + reader := bufio.NewReader(r) + for { + line, err := reader.ReadString('\n') + if err != nil && !errors.Is(err, io.EOF) { + return fmt.Errorf("read responses stream: %w", err) + } + + line = strings.TrimSpace(line) + if strings.HasPrefix(line, "data:") { + payload := strings.TrimSpace(strings.TrimPrefix(line, "data:")) + if payload == "[DONE]" { + return nil + } + if handleErr := handle(payload); handleErr != nil { + return handleErr + } + } + + if errors.Is(err, io.EOF) { + return nil + } + } +} + +// buildResponsesRequest converts a ChatRequest into the /v1/responses body. +func buildResponsesRequest(req ChatRequest) map[string]any { + payload := map[string]any{ + "model": req.Model, + "input": buildResponsesInput(req.Messages), + "stream": true, + } + if len(req.Tools) > 0 { + payload["tools"] = buildResponsesTools(req.Tools) + } + if req.ToolChoice != "" { + payload["tool_choice"] = req.ToolChoice + } + for key, value := range req.Extra { + if _, exists := payload[key]; exists { + continue + } + if key == "reasoning" { + if level, ok := value.(string); ok { + payload[key] = map[string]any{"effort": mapReasoningEffort(level)} + continue + } + } + payload[key] = value + } + return payload +} + +// buildResponsesInput maps chat-style messages to Responses input items. +// Assistant tool calls become function_call items and tool results become +// function_call_output items, preserving their position in the conversation. +func buildResponsesInput(messages []Message) []any { + items := make([]any, 0, len(messages)+4) + for _, msg := range messages { + switch msg.Role { + case RoleSystem: + items = append(items, responsesMessageInput{Role: "system", Content: msg.Content}) + case RoleUser: + items = append(items, responsesMessageInput{Role: "user", Content: msg.Content}) + case RoleAssistant: + if msg.Content != "" { + items = append(items, responsesMessageInput{Role: "assistant", Content: msg.Content}) + } + for _, call := range msg.ToolCalls { + items = append(items, responsesFunctionCallInput{ + Type: "function_call", + CallID: call.ID, + Name: call.Function.Name, + Arguments: call.Function.Arguments, + }) + } + case RoleTool: + items = append(items, responsesFunctionCallOutput{ + Type: "function_call_output", + CallID: msg.ToolCallID, + Output: msg.Content, + }) + } + } + return items +} + +func buildResponsesTools(tools []Tool) []responsesTool { + out := make([]responsesTool, 0, len(tools)) + for _, tool := range tools { + out = append(out, responsesTool{ + Type: "function", + Name: tool.Function.Name, + Description: tool.Function.Description, + Parameters: tool.Function.Parameters, + }) + } + return out +} + +func mapReasoningEffort(level string) string { + level = strings.ToLower(strings.TrimSpace(level)) + if level == "middle" { + return "medium" + } + return level +} + +type responsesMessageInput struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type responsesFunctionCallInput struct { + Type string `json:"type"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type responsesFunctionCallOutput struct { + Type string `json:"type"` + CallID string `json:"call_id"` + Output string `json:"output"` +} + +type responsesTool struct { + Type string `json:"type"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + Parameters json.RawMessage `json:"parameters,omitempty"` +} + +// responsesStreamState tracks in-flight output items while parsing SSE events. +type responsesStreamState struct { + calls map[int]*responsesCallBuilder + done bool +} + +type responsesCallBuilder struct { + index int + id string + callType string + name string + arguments strings.Builder + finalized bool +} + +func (s *responsesStreamState) handle(payload string, emit func(StreamEvent) error) error { + if s.done { + return nil + } + var event responsesStreamEvent + if err := json.Unmarshal([]byte(payload), &event); err != nil { + return fmt.Errorf("parse responses stream payload: %w", err) + } + + switch event.Type { + case "error": + s.done = true + return event.err() + case "response.failed": + s.done = true + if event.Response != nil && event.Response.Error != nil { + return fmt.Errorf("provider error: %s: %s", event.Response.Error.Code, event.Response.Error.Message) + } + return fmt.Errorf("provider error: response failed") + case "response.incomplete": + s.done = true + return s.finish(emit, "length", nil) + case "response.completed": + s.done = true + var usage *Usage + if event.Response != nil && event.Response.Usage != nil { + usage = &Usage{ + PromptTokens: event.Response.Usage.InputTokens, + CompletionTokens: event.Response.Usage.OutputTokens, + TotalTokens: event.Response.Usage.TotalTokens, + } + } + return s.finish(emit, "stop", usage) + case "response.output_text.delta": + if event.Delta != "" { + return emit(StreamEvent{Content: event.Delta}) + } + case "response.output_item.added": + if event.Item != nil && event.Item.Type == "function_call" { + s.builder(event.OutputIndex).applyAdded(*event.Item) + } + case "response.function_call_arguments.delta": + if event.Delta != "" { + s.builder(event.OutputIndex).arguments.WriteString(event.Delta) + } + case "response.function_call_arguments.done": + builder := s.builder(event.OutputIndex) + if event.Name != "" { + builder.name = event.Name + } + if event.Arguments != "" { + builder.arguments.Reset() + builder.arguments.WriteString(event.Arguments) + } + case "response.output_item.done": + if event.Item != nil && event.Item.Type == "function_call" { + builder := s.builder(event.OutputIndex) + builder.applyDone(*event.Item) + return s.emitCall(builder, emit) + } + } + return nil +} + +// finish flushes any not-yet-finalized function calls and reports the +// terminal finish reason. +func (s *responsesStreamState) finish(emit func(StreamEvent) error, finishReason string, usage *Usage) error { + if len(s.calls) == 0 { + return emit(StreamEvent{FinishReason: finishReason, Usage: usage}) + } + indexes := make([]int, 0, len(s.calls)) + for index := range s.calls { + indexes = append(indexes, index) + } + sort.Ints(indexes) + var deltas []ToolCallDelta + for _, index := range indexes { + builder := s.calls[index] + if builder.finalized { + continue + } + builder.finalized = true + deltas = append(deltas, builder.delta()) + } + if len(deltas) > 0 { + if err := emit(StreamEvent{ToolCalls: deltas}); err != nil { + return err + } + } + return emit(StreamEvent{FinishReason: finishReason, Usage: usage}) +} + +func (s *responsesStreamState) emitCall(builder *responsesCallBuilder, emit func(StreamEvent) error) error { + if builder.finalized { + return nil + } + builder.finalized = true + return emit(StreamEvent{ToolCalls: []ToolCallDelta{builder.delta()}}) +} + +func (s *responsesStreamState) builder(index int) *responsesCallBuilder { + if s.calls == nil { + s.calls = make(map[int]*responsesCallBuilder) + } + builder := s.calls[index] + if builder == nil { + builder = &responsesCallBuilder{index: index} + s.calls[index] = builder + } + return builder +} + +func (b *responsesCallBuilder) applyAdded(item responsesStreamItem) { + if b.id == "" { + b.id = item.CallID + if b.id == "" { + b.id = item.ID + } + } + if b.name == "" { + b.name = item.Name + } + if item.Arguments != "" { + b.arguments.Reset() + b.arguments.WriteString(item.Arguments) + } +} + +func (b *responsesCallBuilder) applyDone(item responsesStreamItem) { + if item.CallID != "" { + b.id = item.CallID + } else if item.ID != "" { + b.id = item.ID + } + if item.Name != "" { + b.name = item.Name + } + if item.Arguments != "" { + b.arguments.Reset() + b.arguments.WriteString(item.Arguments) + } +} + +func (b *responsesCallBuilder) delta() ToolCallDelta { + id := b.id + if id == "" { + id = fmt.Sprintf("call_%d", b.index) + } + callType := b.callType + if callType == "" { + callType = "function" + } + return ToolCallDelta{ + Index: b.index, + ID: id, + Type: callType, + Name: b.name, + Arguments: b.arguments.String(), + } +} + +type responsesStreamEvent struct { + Type string `json:"type"` + OutputIndex int `json:"output_index"` + Delta string `json:"delta"` + Arguments string `json:"arguments"` + Name string `json:"name"` + Code string `json:"code"` + Message string `json:"message"` + Param string `json:"param"` + Item *responsesStreamItem `json:"item"` + Response *responsesStreamResponse `json:"response"` + Error *responsesStreamError `json:"error"` +} + +type responsesStreamItem struct { + Type string `json:"type"` + ID string `json:"id"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type responsesStreamResponse struct { + Status string `json:"status"` + Error *responsesStreamError `json:"error"` + Usage *responsesStreamUsage `json:"usage"` +} + +type responsesStreamUsage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` +} + +type responsesStreamError struct { + Type string `json:"type"` + Code string `json:"code"` + Message string `json:"message"` + Param string `json:"param"` +} + +func (e responsesStreamEvent) err() error { + code, message, param := e.Code, e.Message, e.Param + if code == "" && message == "" && e.Error != nil { + code, message, param = e.Error.Code, e.Error.Message, e.Error.Param + } + if message == "" { + if param != "" { + return fmt.Errorf("provider error: %s", param) + } + return fmt.Errorf("provider error") + } + if code != "" { + return fmt.Errorf("provider error: %s: %s", code, message) + } + return fmt.Errorf("provider error: %s", message) +} diff --git a/pkg/llm/responses_test.go b/pkg/llm/responses_test.go new file mode 100644 index 0000000..31e9335 --- /dev/null +++ b/pkg/llm/responses_test.go @@ -0,0 +1,305 @@ +package llm + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestBuildResponsesRequest(t *testing.T) { + req := ChatRequest{ + Model: "gpt-test", + Messages: []Message{ + {Role: RoleSystem, Content: "be concise"}, + {Role: RoleUser, Content: "hi"}, + { + Role: RoleAssistant, + Content: "let me check", + ToolCalls: []ToolCall{{ + ID: "call_1", + Type: "function", + Function: FunctionCall{ + Name: "file_read", + Arguments: `{"path":"README.md"}`, + }, + }}, + }, + {Role: RoleTool, ToolCallID: "call_1", Content: "contents"}, + {Role: RoleAssistant, Content: "done"}, + }, + Tools: []Tool{{ + Type: "function", + Function: ToolFunction{ + Name: "file_read", + Description: "Read a local file", + Parameters: json.RawMessage(`{"type":"object","properties":{}}`), + }, + }}, + ToolChoice: "auto", + Extra: map[string]any{ + "reasoning": "middle", + "thinking": "high", + "temperature": 0.2, + }, + } + + payload := buildResponsesRequest(req) + data, err := json.Marshal(payload) + if err != nil { + t.Fatal(err) + } + var raw map[string]any + if err := json.Unmarshal(data, &raw); err != nil { + t.Fatal(err) + } + + if raw["model"] != "gpt-test" || raw["stream"] != true { + t.Fatalf("model/stream = %#v / %#v", raw["model"], raw["stream"]) + } + if raw["tool_choice"] != "auto" { + t.Fatalf("tool_choice = %#v", raw["tool_choice"]) + } + + input, ok := raw["input"].([]any) + if !ok || len(input) != 6 { + t.Fatalf("input = %#v", raw["input"]) + } + checkInputItem := func(index int, wantType, wantKey, wantValue string) { + t.Helper() + item, ok := input[index].(map[string]any) + if !ok { + t.Fatalf("input[%d] = %#v", index, input[index]) + } + if item[wantKey] != wantValue { + t.Fatalf("input[%d].%s = %#v, want %q", index, wantKey, item[wantKey], wantValue) + } + if wantType != "" && item["type"] != wantType { + t.Fatalf("input[%d].type = %#v, want %q", index, item["type"], wantType) + } + } + checkInputItem(0, "", "role", "system") + checkInputItem(1, "", "role", "user") + checkInputItem(2, "", "role", "assistant") + checkInputItem(3, "function_call", "call_id", "call_1") + checkInputItem(4, "function_call_output", "call_id", "call_1") + checkInputItem(5, "", "role", "assistant") + if item := input[3].(map[string]any); item["name"] != "file_read" || item["arguments"] != `{"path":"README.md"}` { + t.Fatalf("function_call item = %#v", item) + } + if item := input[4].(map[string]any); item["output"] != "contents" { + t.Fatalf("function_call_output item = %#v", item) + } + + tools, ok := raw["tools"].([]any) + if !ok || len(tools) != 1 { + t.Fatalf("tools = %#v", raw["tools"]) + } + tool := tools[0].(map[string]any) + if tool["type"] != "function" || tool["name"] != "file_read" || tool["description"] != "Read a local file" { + t.Fatalf("tool = %#v", tool) + } + if _, ok := tool["parameters"].(map[string]any); !ok { + t.Fatalf("tool parameters = %#v", tool["parameters"]) + } + + reasoning, ok := raw["reasoning"].(map[string]any) + if !ok || reasoning["effort"] != "medium" { + t.Fatalf("reasoning = %#v", raw["reasoning"]) + } + if raw["thinking"] != "high" { + t.Fatalf("thinking = %#v", raw["thinking"]) + } + if raw["temperature"] != 0.2 { + t.Fatalf("temperature = %#v", raw["temperature"]) + } +} + +func TestOpenAIResponsesClientStreamsContentAndToolCalls(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/responses" { + t.Fatalf("path = %s", r.URL.Path) + } + if got := r.Header.Get("Authorization"); got != "Bearer sk-test" { + t.Fatalf("authorization = %q", got) + } + body := mustReadBody(t, r) + var raw map[string]any + if err := json.NewDecoder(strings.NewReader(body)).Decode(&raw); err != nil { + t.Fatal(err) + } + if raw["stream"] != true { + t.Fatal("request did not enable stream") + } + if raw["tool_choice"] != "auto" { + t.Fatalf("tool_choice = %#v", raw["tool_choice"]) + } + input, ok := raw["input"].([]any) + if !ok || len(input) == 0 { + t.Fatalf("input = %#v", raw["input"]) + } + + w.Header().Set("Content-Type", "text/event-stream") + events := []string{ + `{"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg_1","role":"assistant","content":[]}}`, + `{"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"Hel"}`, + `{"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"lo"}`, + `{"type":"response.output_item.added","output_index":1,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"file_read","arguments":""}}`, + `{"type":"response.function_call_arguments.delta","item_id":"fc_1","output_index":1,"delta":"{\"path\""}`, + `{"type":"response.function_call_arguments.delta","item_id":"fc_1","output_index":1,"delta":":\"README.md\"}"}`, + `{"type":"response.function_call_arguments.done","item_id":"fc_1","output_index":1,"name":"file_read","arguments":"{\"path\":\"README.md\"}"}`, + `{"type":"response.output_item.done","output_index":1,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"file_read","arguments":"{\"path\":\"README.md\"}"}}`, + `{"type":"response.output_item.done","output_index":0,"item":{"type":"message","id":"msg_1","role":"assistant","content":[{"type":"output_text","text":"Hello"}]}}`, + `{"type":"response.completed","response":{"id":"resp_1","status":"completed","usage":{"input_tokens":10,"output_tokens":20,"total_tokens":30}}}`, + } + for _, event := range events { + _, _ = w.Write([]byte("data: " + event + "\n\n")) + } + })) + defer server.Close() + + client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client()) + var content strings.Builder + var calls []ToolCallDelta + var finishes []string + var usage *Usage + err := client.ChatStream(context.Background(), ChatRequest{ + Model: "gpt-test", + Messages: []Message{{Role: RoleUser, Content: "hi"}}, + ToolChoice: "auto", + Tools: []Tool{{ + Type: "function", + Function: ToolFunction{ + Name: "file_read", + Description: "Read a local file", + Parameters: json.RawMessage(`{"type":"object"}`), + }, + }}, + }, func(event StreamEvent) error { + content.WriteString(event.Content) + calls = append(calls, event.ToolCalls...) + if event.Usage != nil { + usage = event.Usage + } + if event.FinishReason != "" { + finishes = append(finishes, event.FinishReason) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + if content.String() != "Hello" { + t.Fatalf("content = %q", content.String()) + } + if len(calls) != 1 { + t.Fatalf("calls = %#v", calls) + } + call := calls[0] + if call.Index != 1 || call.ID != "call_1" || call.Name != "file_read" { + t.Fatalf("call = %#v", call) + } + if call.Arguments != `{"path":"README.md"}` { + t.Fatalf("arguments = %q", call.Arguments) + } + if len(finishes) != 1 || finishes[0] != "stop" { + t.Fatalf("finishes = %#v", finishes) + } + if usage == nil || usage.PromptTokens != 10 || usage.CompletionTokens != 20 || usage.TotalTokens != 30 { + t.Fatalf("usage = %#v", usage) + } +} + +func TestOpenAIResponsesClientErrorEvent(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"type":"error","code":"invalid_request_error","message":"bad request"}` + "\n\n")) + })) + defer server.Close() + + client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client()) + err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(StreamEvent) error { + return nil + }) + if err == nil || !strings.Contains(err.Error(), "bad request") { + t.Fatalf("err = %v", err) + } +} + +func TestOpenAIResponsesClientFailedEvent(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_error","message":"boom"}}}` + "\n\n")) + })) + defer server.Close() + + client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client()) + err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(StreamEvent) error { + return nil + }) + if err == nil || !strings.Contains(err.Error(), "boom") { + t.Fatalf("err = %v", err) + } +} + +func TestOpenAIResponsesClientIncomplete(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"partial"}` + "\n\n")) + _, _ = w.Write([]byte(`data: {"type":"response.incomplete","response":{"id":"resp_1","status":"incomplete","incomplete_details":{"reason":"max_tokens"}}}` + "\n\n")) + })) + defer server.Close() + + client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client()) + var content strings.Builder + var finishes []string + err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(event StreamEvent) error { + content.WriteString(event.Content) + if event.FinishReason != "" { + finishes = append(finishes, event.FinishReason) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + if content.String() != "partial" { + t.Fatalf("content = %q", content.String()) + } + if len(finishes) != 1 || finishes[0] != "length" { + t.Fatalf("finishes = %#v", finishes) + } +} + +func TestOpenAIResponsesClientRetriesTransientHTTPError(t *testing.T) { + attempts := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts++ + if attempts == 1 { + http.Error(w, "temporary", http.StatusBadGateway) + return + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n")) + _, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n")) + })) + defer server.Close() + + client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client()) + var content strings.Builder + err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(event StreamEvent) error { + content.WriteString(event.Content) + return nil + }) + if err != nil { + t.Fatal(err) + } + if attempts != 2 { + t.Fatalf("attempts = %d, want 2", attempts) + } + if content.String() != "ok" { + t.Fatalf("content = %q", content.String()) + } +}