feat(llm): support OpenAI Responses API via per-provider api_type
This commit is contained in:
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user