feat(llm): support OpenAI Responses API via per-provider api_type
This commit is contained in:
@@ -8,6 +8,7 @@ model/provider switching, and local tools.
|
|||||||
|
|
||||||
- TUI-first chat experience with a light theme by default.
|
- TUI-first chat experience with a light theme by default.
|
||||||
- OpenAI-compatible `/v1/chat/completions` provider support.
|
- OpenAI-compatible `/v1/chat/completions` provider support.
|
||||||
|
- OpenAI Responses API (`/v1/responses`) provider support.
|
||||||
- Multiple providers from `~/.agentu/config.yaml`.
|
- Multiple providers from `~/.agentu/config.yaml`.
|
||||||
- Session-only provider, model, and thinking-level switching.
|
- Session-only provider, model, and thinking-level switching.
|
||||||
- Read-only project tools enabled by default.
|
- Read-only project tools enabled by default.
|
||||||
@@ -48,6 +49,7 @@ Use `--config <path>` to load another file.
|
|||||||
```yaml
|
```yaml
|
||||||
providers:
|
providers:
|
||||||
loveuer:
|
loveuer:
|
||||||
|
api_type: chat
|
||||||
base_url: https://ai.loveuer.com
|
base_url: https://ai.loveuer.com
|
||||||
api_key: ${AGENTU_API_KEY}
|
api_key: ${AGENTU_API_KEY}
|
||||||
models:
|
models:
|
||||||
@@ -64,6 +66,19 @@ Provider names are the keys under `providers`. At startup, agentu selects the
|
|||||||
first provider by name. Use `/model provider <name>` to switch providers for the
|
first provider by name. Use `/model provider <name>` to switch providers for the
|
||||||
current session.
|
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 <name>` is restricted to that list. Each
|
`models` is required. `/model model <name>` is restricted to that list. Each
|
||||||
entry is an object with:
|
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
|
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
|
## 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
|
- Provider/model/thinking changes are session-only; agentu does not write them
|
||||||
back to `~/.agentu/config.yaml`.
|
back to `~/.agentu/config.yaml`.
|
||||||
|
- `api_type` is per provider and cannot be changed at runtime; use
|
||||||
|
`/model provider <name>` to switch to a provider configured with a different
|
||||||
|
API type.
|
||||||
- The config schema is intentionally strict during MVP development. Deprecated
|
- The config schema is intentionally strict during MVP development. Deprecated
|
||||||
fields such as `provider` or `active_provider` are rejected.
|
fields such as `provider` or `active_provider` are rejected.
|
||||||
|
|||||||
@@ -1,5 +1,8 @@
|
|||||||
providers:
|
providers:
|
||||||
loveuer:
|
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
|
base_url: https://ai.loveuer.com
|
||||||
api_key: ${AGENTU_API_KEY}
|
api_key: ${AGENTU_API_KEY}
|
||||||
models:
|
models:
|
||||||
|
|||||||
+8
-1
@@ -91,10 +91,17 @@ func loadConfig(configPath string) (*config.Config, error) {
|
|||||||
func defaultProvider(cfg *config.Config, httpClient *http.Client) (string, config.ProviderConfig, llm.Provider) {
|
func defaultProvider(cfg *config.Config, httpClient *http.Client) (string, config.ProviderConfig, llm.Provider) {
|
||||||
providerName := cfg.DefaultProviderName()
|
providerName := cfg.DefaultProviderName()
|
||||||
providerConfig := cfg.Providers[providerName]
|
providerConfig := cfg.Providers[providerName]
|
||||||
provider := llm.NewOpenAICompatibleClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient)
|
provider := newProviderClient(providerConfig, httpClient)
|
||||||
return providerName, providerConfig, provider
|
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 {
|
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{
|
return agent.New(agent.Options{
|
||||||
Provider: provider,
|
Provider: provider,
|
||||||
|
|||||||
@@ -204,10 +204,18 @@ func (m *Manager) CurrentThinking() string {
|
|||||||
return m.provider.Thinking
|
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 {
|
func (m *Manager) Info() string {
|
||||||
var b strings.Builder
|
var b strings.Builder
|
||||||
fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider())
|
fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider())
|
||||||
fmt.Fprintf(&b, "model: %s\n", m.CurrentModel())
|
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, "thinking: %s\n", m.CurrentThinking())
|
||||||
fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", "))
|
fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", "))
|
||||||
fmt.Fprintf(&b, "models: %s\n", strings.Join(m.provider.Models, ", "))
|
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.providerName = name
|
||||||
m.provider = provider
|
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
|
return "Switched provider.\n" + m.summary(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -257,7 +265,14 @@ func (m *Manager) SetThinking(level string) (string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) summary() string {
|
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 {
|
func (m *Manager) maxContextTokens() int {
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ package session
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"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) {
|
func TestManagerExposesModelCommandCandidates(t *testing.T) {
|
||||||
manager, _, _ := testManager(t)
|
manager, _, _ := testManager(t)
|
||||||
models := manager.AvailableModels()
|
models := manager.AvailableModels()
|
||||||
|
|||||||
@@ -14,6 +14,11 @@ import (
|
|||||||
|
|
||||||
const DefaultPath = "~/.agentu/config.yaml"
|
const DefaultPath = "~/.agentu/config.yaml"
|
||||||
|
|
||||||
|
const (
|
||||||
|
APITypeChat = "chat"
|
||||||
|
APITypeResponses = "responses"
|
||||||
|
)
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Providers map[string]ProviderConfig `yaml:"providers"`
|
Providers map[string]ProviderConfig `yaml:"providers"`
|
||||||
Agent AgentConfig `yaml:"agent"`
|
Agent AgentConfig `yaml:"agent"`
|
||||||
@@ -21,6 +26,7 @@ type Config struct {
|
|||||||
|
|
||||||
type ProviderConfig struct {
|
type ProviderConfig struct {
|
||||||
Name string
|
Name string
|
||||||
|
APIType string
|
||||||
BaseURL string
|
BaseURL string
|
||||||
Model string
|
Model string
|
||||||
Models []string
|
Models []string
|
||||||
@@ -56,6 +62,7 @@ type rawConfig struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type rawProviderConfig struct {
|
type rawProviderConfig struct {
|
||||||
|
APIType string `yaml:"api_type"`
|
||||||
BaseURL string `yaml:"base_url"`
|
BaseURL string `yaml:"base_url"`
|
||||||
Models []rawModelConfig `yaml:"models"`
|
Models []rawModelConfig `yaml:"models"`
|
||||||
APIKey string `yaml:"api_key"`
|
APIKey string `yaml:"api_key"`
|
||||||
@@ -153,6 +160,9 @@ func (c *Config) Validate() error {
|
|||||||
for name, provider := range c.Providers {
|
for name, provider := range c.Providers {
|
||||||
var missing []string
|
var missing []string
|
||||||
prefix := "providers." + name
|
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) == "" {
|
if strings.TrimSpace(provider.BaseURL) == "" {
|
||||||
missing = append(missing, prefix+".base_url")
|
missing = append(missing, prefix+".base_url")
|
||||||
}
|
}
|
||||||
@@ -220,6 +230,7 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
|
|||||||
models, modelConfigs := normalizeModelConfigs(raw.Models)
|
models, modelConfigs := normalizeModelConfigs(raw.Models)
|
||||||
provider := ProviderConfig{
|
provider := ProviderConfig{
|
||||||
Name: name,
|
Name: name,
|
||||||
|
APIType: normalizeAPIType(raw.APIType),
|
||||||
BaseURL: raw.BaseURL,
|
BaseURL: raw.BaseURL,
|
||||||
Models: models,
|
Models: models,
|
||||||
ModelConfigs: modelConfigs,
|
ModelConfigs: modelConfigs,
|
||||||
@@ -235,6 +246,28 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
|
|||||||
return provider
|
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) {
|
func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) {
|
||||||
models := make([]string, 0, len(rawModels))
|
models := make([]string, 0, len(rawModels))
|
||||||
configs := make(map[string]ModelConfig, len(rawModels))
|
configs := make(map[string]ModelConfig, len(rawModels))
|
||||||
|
|||||||
@@ -370,3 +370,74 @@ providers:
|
|||||||
t.Fatalf("err = %v", err)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user