diff --git a/agent-schema.json b/agent-schema.json index 076905f91..9561d9f54 100644 --- a/agent-schema.json +++ b/agent-schema.json @@ -224,7 +224,7 @@ "properties": { "provider": { "type": "string", - "description": "The underlying provider type. Defaults to \"openai\" when not set. Supported values: openai, anthropic, google, amazon-bedrock, dmr, and any built-in alias (requesty, openrouter, azure, xai, ollama, mistral, baseten, ovhcloud, groq, fireworks, deepseek, cerebras, together, huggingface, moonshot, vercel, cloudflare-workers-ai, cloudflare-ai-gateway, nvidia, github-copilot, chatgpt, etc.).", + "description": "The underlying provider type. Defaults to \"openai\" when not set. Supported values: openai, anthropic, google, amazon-bedrock, dmr, and any built-in alias (requesty, openrouter, azure, xai, ollama, mistral, baseten, ovhcloud, groq, fireworks-ai, deepseek, cerebras, togetherai, huggingface, moonshotai, vercel, cloudflare-workers-ai, cloudflare-ai-gateway, nvidia, github-copilot, chatgpt, opencode, opencode-go, etc.).", "examples": [ "openai", "anthropic", diff --git a/cmd/root/models.go b/cmd/root/models.go index 8ce52382d..03ac56b37 100644 --- a/cmd/root/models.go +++ b/cmd/root/models.go @@ -46,8 +46,8 @@ const listTimeout = 5 * time.Second // from the snapshot) prevents surprising side effects like `docker agent models // --provider ollama` issuing a real GET against localhost. var liveFetchProviders = map[string]bool{ - "opencode-zen": true, - "opencode-go": true, + "opencode": true, + "opencode-go": true, } // modelRow represents a single model entry for display or serialization. @@ -129,11 +129,22 @@ func (f *modelsListFlags) runModelsListCommand(cmd *cobra.Command, args []string out := cli.NewPrinter(cmd.OutOrStdout()) env := f.runConfig.EnvProvider() - // Normalize the provider filter to lowercase so case-sensitive map lookups - // in db.Providers and IsCatalogProvider all match the same way - // strings.EqualFold does in the outer row filter below. + isCustomProvider := func(name string) bool { + for customName := range f.runConfig.Providers { + if strings.EqualFold(customName, name) { + return true + } + } + return false + } + + // Custom names take precedence over built-in aliases, regardless of case. + customFilter := isCustomProvider(f.providerFilter) if f.providerFilter != "" { f.providerFilter = strings.ToLower(f.providerFilter) + if !customFilter { + f.providerFilter = modelsdev.CanonicalProviderID(f.providerFilter) + } } // Determine which model auto-selection would pick. DMR discovery is left @@ -158,10 +169,15 @@ func (f *modelsListFlags) runModelsListCommand(cmd *cobra.Command, args []string rows = f.collectModels(ctx, env, availableProviders, autoModel) } - // Apply provider filter if f.providerFilter != "" { rows = slices.DeleteFunc(rows, func(r modelRow) bool { - return !strings.EqualFold(r.Provider, f.providerFilter) + if strings.EqualFold(r.Provider, f.providerFilter) { + return false + } + if customFilter || isCustomProvider(r.Provider) { + return true + } + return modelsdev.CanonicalProviderID(strings.ToLower(r.Provider)) != f.providerFilter }) } diff --git a/cmd/root/models_test.go b/cmd/root/models_test.go index 91456adc2..61b879f57 100644 --- a/cmd/root/models_test.go +++ b/cmd/root/models_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "fmt" "maps" "net/http" "net/http/httptest" @@ -18,6 +19,7 @@ import ( "github.com/docker/docker-agent/pkg/config" "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider" "github.com/docker/docker-agent/pkg/modelsdev" ) @@ -805,3 +807,176 @@ func TestModelsListCommand_AliasCredentialsListAliasModels(t *testing.T) { assert.Contains(t, buf.String(), "grok-4", "an alias credential must surface the alias's catalog models without --all") } + +func TestModelsListCommand_LegacyProviderFilter(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + canonical := modelsdev.CanonicalProviderID(legacy) + alias, ok := provider.LookupAlias(canonical) + require.True(t, ok) + var buf bytes.Buffer + cmd := newModelsCmd(func(rc *config.RuntimeConfig) { + rc.EnvProviderForTests = environment.NewMapEnvProvider(map[string]string{alias.TokenEnvVar: "test-key"}) + rc.Providers = map[string]latest.ProviderConfig{} + rc.ModelsDevStoreOverride = modelsdev.NewDatabaseStore(&modelsdev.Database{Providers: map[string]modelsdev.Provider{ + canonical: {Models: map[string]modelsdev.Model{"catalog-only": {Modalities: modelsdev.Modalities{Output: []string{"text"}}}}}, + }}) + }) + cmd.SetOut(&buf) + cmd.SetErr(&buf) + cmd.SetArgs([]string{"--provider", legacy, "--format", "json"}) + require.NoError(t, cmd.Execute()) + var rows []modelRow + require.NoError(t, json.Unmarshal(buf.Bytes(), &rows)) + require.NotEmpty(t, rows) + var catalogFound bool + for _, row := range rows { + assert.Equal(t, canonical, row.Provider) + catalogFound = catalogFound || row.Model == "catalog-only" + } + assert.True(t, catalogFound) + }) + } +} + +func TestModelsListCommand_LegacyNamedCustomProviderFilter(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + for _, name := range []string{legacy, strings.ToUpper(legacy[:1]) + legacy[1:]} { + t.Run(name, func(t *testing.T) { + t.Parallel() + server, _ := newCustomProviderServer(t, []string{"custom-model"}) + for _, filter := range []string{legacy, strings.ToUpper(legacy[:1]) + legacy[1:], strings.ToUpper(legacy)} { + t.Run(filter, func(t *testing.T) { + t.Parallel() + var buf bytes.Buffer + cmd := newModelsCmd( + withTestConfig(map[string]string{"MYPROVIDER_API_KEY": "custom-key"}), + withProviders(map[string]latest.ProviderConfig{ + name: {BaseURL: server.URL, TokenKey: "MYPROVIDER_API_KEY"}, + }), + ) + cmd.SetOut(&buf) + cmd.SetErr(&buf) + cmd.SetArgs([]string{"--provider", filter, "--format", "json"}) + require.NoError(t, cmd.Execute()) + var rows []modelRow + require.NoError(t, json.Unmarshal(buf.Bytes(), &rows)) + require.Len(t, rows, 1) + assert.Equal(t, name, rows[0].Provider) + assert.Equal(t, "custom-model", rows[0].Model) + }) + } + }) + } + }) + } +} + +func TestModelsListCommand_GatewayLegacyProviderFilter(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + canonical := modelsdev.CanonicalProviderID(legacy) + for _, servedProvider := range []string{legacy, canonical} { + t.Run(servedProvider, func(t *testing.T) { + t.Parallel() + servedID := servedProvider + "/served-model" + gw, _ := newGatewayServer(t, fmt.Sprintf(`{"object":"list","data":[{"id":%q},{"id":"openai/other-model"}]}`, servedID)) + for _, filter := range []string{legacy, canonical} { + t.Run(filter, func(t *testing.T) { + t.Parallel() + var buf bytes.Buffer + cmd := newModelsCmd( + withTestConfig(gatewayTestEnv(nil)), + withProviders(map[string]latest.ProviderConfig{}), + withCatalog(&modelsdev.Database{}), + ) + cmd.SetOut(&buf) + cmd.SetErr(&buf) + cmd.SetArgs([]string{"--models-gateway", gw.URL, "--provider", filter, "--format", "json"}) + require.NoError(t, cmd.Execute()) + var rows []modelRow + require.NoError(t, json.Unmarshal(buf.Bytes(), &rows)) + require.Len(t, rows, 1) + assert.Equal(t, servedProvider, rows[0].Provider) + assert.Equal(t, "served-model", rows[0].Model) + assert.Equal(t, servedID, rows[0].Provider+"/"+rows[0].Model) + }) + } + }) + } + }) + } +} + +func TestModelsListCommand_LegacyCustomProviderPrecedence(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + canonical := modelsdev.CanonicalProviderID(legacy) + name := strings.ToUpper(legacy[:1]) + legacy[1:] + custom, _ := newCustomProviderServer(t, []string{"custom-model"}) + gw, _ := newGatewayServer(t, fmt.Sprintf(`{"object":"list","data":[{"id":%q},{"id":%q}]}`, canonical+"/gateway-model", legacy+"/legacy-gateway-model")) + for _, tt := range []struct { + filter string + want []modelRow + }{ + {legacy, []modelRow{{Provider: name, Model: "custom-model"}, {Provider: legacy, Model: "legacy-gateway-model"}}}, + {canonical, []modelRow{{Provider: canonical, Model: "gateway-model"}}}, + } { + t.Run(tt.filter, func(t *testing.T) { + t.Parallel() + env := gatewayTestEnv(nil) + env["MYPROVIDER_API_KEY"] = "custom-key" + var buf bytes.Buffer + cmd := newModelsCmd( + withTestConfig(env), + withCatalog(&modelsdev.Database{}), + withProviders(map[string]latest.ProviderConfig{ + name: {BaseURL: custom.URL, TokenKey: "MYPROVIDER_API_KEY"}, + }), + ) + cmd.SetOut(&buf) + cmd.SetErr(&buf) + cmd.SetArgs([]string{"--models-gateway", gw.URL, "--provider", tt.filter, "--format", "json"}) + require.NoError(t, cmd.Execute()) + var rows []modelRow + require.NoError(t, json.Unmarshal(buf.Bytes(), &rows)) + assert.Equal(t, tt.want, rows) + }) + } + }) + } +} + +func TestModelsListCommand_CaseDuplicateCustomProviderFilter(t *testing.T) { + t.Parallel() + upper, _ := newCustomProviderServer(t, []string{"upper-model"}) + lower, _ := newCustomProviderServer(t, []string{"lower-model"}) + var buf bytes.Buffer + cmd := newModelsCmd( + withTestConfig(map[string]string{"MYPROVIDER_API_KEY": "custom-key"}), + withProviders(map[string]latest.ProviderConfig{ + "Fireworks": {BaseURL: upper.URL, TokenKey: "MYPROVIDER_API_KEY"}, + "fireworks": {BaseURL: lower.URL, TokenKey: "MYPROVIDER_API_KEY"}, + }), + ) + cmd.SetOut(&buf) + cmd.SetErr(&buf) + cmd.SetArgs([]string{"--provider", "FiReWoRkS", "--format", "json"}) + require.NoError(t, cmd.Execute()) + var rows []modelRow + require.NoError(t, json.Unmarshal(buf.Bytes(), &rows)) + require.Len(t, rows, 2) + assert.Equal(t, "Fireworks", rows[0].Provider) + assert.Equal(t, "upper-model", rows[0].Model) + assert.Equal(t, "fireworks", rows[1].Provider) + assert.Equal(t, "lower-model", rows[1].Model) +} diff --git a/docs/configuration/models/index.md b/docs/configuration/models/index.md index 2a82cdc0e..9950243c5 100644 --- a/docs/configuration/models/index.md +++ b/docs/configuration/models/index.md @@ -18,7 +18,7 @@ models: first_available: [list] # Optional: candidate model refs, tried in order by available credentials. # Mutually exclusive with other model settings. provider: string # Required unless using first_available. One of: openai, anthropic, google, amazon-bedrock, - # dmr, mistral, xai, nebius, nvidia, minimax, baseten, ovhcloud, groq, fireworks, deepseek, cerebras, together, huggingface, moonshot, vercel, cloudflare-workers-ai, cloudflare-ai-gateway, requesty, openrouter, + # dmr, mistral, xai, nebius, nvidia, minimax, baseten, ovhcloud, groq, fireworks-ai, deepseek, cerebras, togetherai, huggingface, moonshotai, vercel, cloudflare-workers-ai, cloudflare-ai-gateway, requesty, openrouter, # azure, ollama, github-copilot, or a named provider defined # under the top-level `providers:` section. model: string # Required: model identifier @@ -60,7 +60,7 @@ models: | Property | Type | Required | Description | | --------------------- | ---------- | -------- | ------------------------------------------------------------------------------------- | | `first_available` | array | ✗ | Candidate model references tried in order; selects the first whose credentials are configured. Mutually exclusive with other model settings. | -| `provider` | string | ✓/✗ | Required for regular model definitions; omitted for `first_available` selectors. Provider: `openai`, `anthropic`, `google`, `amazon-bedrock`, `dmr`, `mistral`, `xai`, `nebius`, `nvidia`, `minimax`, `baseten`, `ovhcloud`, `groq`, `fireworks`, `deepseek`, `cerebras`, `together`, `huggingface`, `moonshot`, `vercel`, `cloudflare-workers-ai`, `cloudflare-ai-gateway`, `requesty`, `openrouter`, `azure`, `ollama`, `github-copilot`, `chatgpt`, or any [named provider](../../providers/custom/index.md). | +| `provider` | string | ✓/✗ | Required for regular model definitions; omitted for `first_available` selectors. Provider: `openai`, `anthropic`, `google`, `amazon-bedrock`, `dmr`, `mistral`, `xai`, `nebius`, `nvidia`, `minimax`, `baseten`, `ovhcloud`, `groq`, `fireworks-ai`, `deepseek`, `cerebras`, `togetherai`, `huggingface`, `moonshotai`, `vercel`, `cloudflare-workers-ai`, `cloudflare-ai-gateway`, `requesty`, `openrouter`, `azure`, `ollama`, `github-copilot`, `chatgpt`, `opencode`, `opencode-go`, or any [named provider](../../providers/custom/index.md). | | `model` | string | ✓/✗ | Required for regular model definitions; omitted for `first_available` selectors. Model name (e.g., `gpt-4o`, `claude-sonnet-4-5`, `gemini-3.5-flash`) | | `description` | string | ✗ | Informational, human-readable summary of the model's purpose or strengths (e.g., "fast and cheap, good for summaries"). Not sent to the model. Can be combined with `first_available` (a selector's description is kept when it resolves). | | `temperature` | float | ✗ | Sampling randomness. Range is provider-dependent — typically `0.0–2.0` (Anthropic caps at `1.0`). `0.0` is deterministic. | @@ -84,6 +84,12 @@ models: | `compaction_threshold` | float | ✗ | Fraction of the context window at which proactive auto-compaction triggers for agents running this model. Must be greater than `0` and at most `1`. Takes precedence over the agent-level `compaction_threshold`. Cannot be combined with `first_available`. Default: `0.9`. See the [Context & Compaction guide](../../guides/compaction/index.md). | | `bypass_models_gateway` | boolean | ✗ | When `true`, this model connects directly to its provider even when a models gateway (`--models-gateway` / `DOCKER_AGENT_MODELS_GATEWAY`) is configured. Implied by a custom `base_url`. See [Gateway Bypass](#gateway-bypass). | +Built-in provider IDs use models.dev names: `fireworks-ai`, `togetherai`, +`moonshotai`, and `opencode`. Existing configurations using `fireworks`, +`together`, `moonshot`, or `opencode-zen` continue to work without warnings. +New configurations using canonical IDs require a version that supports them. +Explicitly named custom providers still take precedence. + ## Attachment Capability Overrides For custom OpenAI-compatible providers, local models (Ollama, DMR), and any @@ -550,7 +556,7 @@ for a complete local-server configuration. ## Custom HTTP Headers For OpenAI-compatible providers (`openai`, `github-copilot`, `mistral`, `xai`, -`nebius`, `nvidia`, `minimax`, `baseten`, `ovhcloud`, `groq`, `fireworks`, `deepseek`, `cerebras`, `together`, `huggingface`, `moonshot`, `vercel`, `cloudflare-workers-ai`, `cloudflare-ai-gateway`, `requesty`, `openrouter`, `ollama`, and any custom provider using the OpenAI API), +`nebius`, `nvidia`, `minimax`, `baseten`, `ovhcloud`, `groq`, `fireworks-ai`, `deepseek`, `cerebras`, `togetherai`, `huggingface`, `moonshotai`, `vercel`, `cloudflare-workers-ai`, `cloudflare-ai-gateway`, `requesty`, `openrouter`, `ollama`, and any custom provider using the OpenAI API), `provider_opts.http_headers` adds arbitrary HTTP headers to every outgoing request: diff --git a/docs/features/cli/index.md b/docs/features/cli/index.md index 89e59f13f..5753c63e9 100644 --- a/docs/features/cli/index.md +++ b/docs/features/cli/index.md @@ -230,6 +230,8 @@ $ docker agent models --provider openai $ docker agent models --format json | jq ``` +Provider filters are case-insensitive. For built-in providers, legacy and canonical names match the same models (`fireworks` / `fireworks-ai`, `together` / `togetherai`, `moonshot` / `moonshotai`, and `opencode-zen` / `opencode`), including gateway listings. Gateway model references retain the prefix returned by the gateway. A configured custom provider name takes precedence over a built-in alias and is matched by its own name without alias expansion. + When a models gateway is configured (`--models-gateway`, `DOCKER_AGENT_MODELS_GATEWAY`, or the user config), the command first queries the gateway's `/v1/models` endpoint. A non-empty response is authoritative for the models routed through the gateway: the listing shows the models the gateway serves (`--provider` filters within it), alongside any custom providers you have configured, which serve their models from their own endpoints rather than through the gateway. If the gateway cannot be queried or serves no usable model (endpoint not implemented, empty list, invalid response, timeout, missing authentication), the command falls back to the providers you have configured directly — provider API keys, provider aliases, and custom providers — plus the model catalog; a failure of one source never prevents the others from being listed. Docker Desktop authentication is required only for HTTPS `docker.com` gateways. An available Docker Desktop token may also be sent to trusted loopback gateways, but is never sent to third-party gateways. ### `docker agent toolsets` diff --git a/docs/providers/fireworks/index.md b/docs/providers/fireworks/index.md index 65cd1fe15..264d70f10 100644 --- a/docs/providers/fireworks/index.md +++ b/docs/providers/fireworks/index.md @@ -15,6 +15,8 @@ models, serving Kimi, Qwen, DeepSeek, GLM and others through an OpenAI-compatible API. Docker Agent includes built-in support for Fireworks AI as an alias provider. +Use the provider ID `fireworks-ai`. The legacy ID `fireworks` is also accepted. + ## Setup 1. Create an API key from the [Fireworks dashboard](https://fireworks.ai/account/api-keys). @@ -33,7 +35,7 @@ The simplest way to use Fireworks AI: ```yaml agents: root: - model: fireworks/accounts/fireworks/models/kimi-k3 + model: fireworks-ai/accounts/fireworks/models/kimi-k3 description: Assistant using Fireworks AI instruction: You are a helpful assistant. ``` @@ -45,7 +47,7 @@ For more control over parameters: ```yaml models: fireworks_model: - provider: fireworks + provider: fireworks-ai model: accounts/fireworks/models/kimi-k3 temperature: 0.7 max_tokens: 8192 @@ -93,7 +95,7 @@ messages into a single one for this provider. ```yaml agents: coder: - model: fireworks/accounts/fireworks/models/kimi-k2p7-code + model: fireworks-ai/accounts/fireworks/models/kimi-k2p7-code description: Code assistant using Kimi K2.7 Code on Fireworks AI instruction: | You are an expert programmer. diff --git a/docs/providers/moonshot/index.md b/docs/providers/moonshot/index.md index 6c1cf162d..c193d1439 100644 --- a/docs/providers/moonshot/index.md +++ b/docs/providers/moonshot/index.md @@ -15,6 +15,8 @@ OpenAI-compatible API. The Kimi K2 models have strong momentum for coding and agentic tasks. Docker Agent includes built-in support for Moonshot AI as an alias provider. +Use the provider ID `moonshotai`. The legacy ID `moonshot` is also accepted. + ## Setup 1. Create an API key from the [Moonshot AI console](https://platform.moonshot.ai/console/api-keys). @@ -33,7 +35,7 @@ The simplest way to use Moonshot AI: ```yaml agents: root: - model: moonshot/kimi-k3 + model: moonshotai/kimi-k3 description: Assistant using Moonshot AI instruction: You are a helpful assistant. ``` @@ -45,7 +47,7 @@ For more control over parameters: ```yaml models: moonshot_model: - provider: moonshot + provider: moonshotai model: kimi-k3 temperature: 0.7 max_tokens: 8192 @@ -85,7 +87,7 @@ Moonshot AI is implemented as a built-in alias in Docker Agent: ```yaml agents: coder: - model: moonshot/kimi-k3 + model: moonshotai/kimi-k3 description: Code assistant using Kimi K2 instruction: | You are an expert programmer. diff --git a/docs/providers/opencode-zen/index.md b/docs/providers/opencode-zen/index.md index 033887df7..48b216372 100644 --- a/docs/providers/opencode-zen/index.md +++ b/docs/providers/opencode-zen/index.md @@ -14,6 +14,8 @@ _Use OpenCode Zen models with Docker Agent._ Docker Agent includes built-in support for OpenCode Zen as an alias provider for OpenAI-compatible models. Anthropic and Google models are supported via custom provider definitions. +Use the provider ID `opencode`. The legacy ID `opencode-zen` is also accepted. + ## Setup 1. Sign in to [OpenCode Zen](https://opencode.ai/auth), add billing information, and copy your API key @@ -38,7 +40,7 @@ The simplest way to use OpenCode Zen with a free model: ```yaml agents: root: - model: opencode-zen/deepseek-v4-flash-free + model: opencode/deepseek-v4-flash-free description: Assistant using OpenCode Zen (free) instruction: You are a helpful assistant. ``` @@ -50,7 +52,7 @@ For more control over parameters: ```yaml models: zen_model: - provider: opencode-zen + provider: opencode model: gpt-5.5 temperature: 0.7 max_tokens: 16384 @@ -80,7 +82,7 @@ These models are available at no cost: ### OpenAI-Compatible (Chat Completions) -These models use the `/v1/chat/completions` endpoint and work directly with the `opencode-zen` alias: +These models use the `/v1/chat/completions` endpoint and work directly with the `opencode` provider: | Model | Description | | --------------------- | ---------------------------------- | diff --git a/docs/providers/overview/index.md b/docs/providers/overview/index.md index 7771c6017..2ad4c8ff5 100644 --- a/docs/providers/overview/index.md +++ b/docs/providers/overview/index.md @@ -44,7 +44,7 @@ Use this table to find a built-in provider's config key and authentication metho | [AWS Bedrock](../bedrock/index.md) | `amazon-bedrock` | `AWS_BEARER_TOKEN_BEDROCK` or the standard AWS credentials chain | | [Docker Model Runner](../dmr/index.md) | `dmr` | None (local) | | ChatGPT (OpenAI account) | [`chatgpt`](../chatgpt/index.md) | None (sign in via `docker agent setup`) | -| OpenCode Zen | `opencode-zen` | `OPENCODE_API_KEY` | +| OpenCode Zen | `opencode` | `OPENCODE_API_KEY` | | OpenCode Go | `opencode-go` | `OPENCODE_API_KEY` | | Mistral | `mistral` | `MISTRAL_API_KEY` | | xAI (Grok) | `xai` | `XAI_API_KEY` | @@ -54,13 +54,13 @@ Use this table to find a built-in provider's config key and authentication metho | Baseten | `baseten` | `BASETEN_API_KEY` | | OVHcloud | `ovhcloud` | `OVH_AI_ENDPOINTS_ACCESS_TOKEN` | | Groq | `groq` | `GROQ_API_KEY` | -| Fireworks AI | `fireworks` | `FIREWORKS_API_KEY` | +| Fireworks AI | `fireworks-ai` | `FIREWORKS_API_KEY` | | DeepSeek | `deepseek` | `DEEPSEEK_API_KEY` | | Cerebras | `cerebras` | `CEREBRAS_API_KEY` | -| Together AI | `together` | `TOGETHER_API_KEY` | +| Together AI | `togetherai` | `TOGETHER_API_KEY` | | Hugging Face | `huggingface` | `HF_TOKEN` | | Cloudflare Workers AI | `cloudflare-workers-ai` | `CLOUDFLARE_API_TOKEN` + `CLOUDFLARE_ACCOUNT_ID` | -| Moonshot AI | `moonshot` | `MOONSHOT_API_KEY` | +| Moonshot AI | `moonshotai` | `MOONSHOT_API_KEY` | | Vercel AI Gateway | `vercel` | `AI_GATEWAY_API_KEY` | | Cloudflare AI Gateway | `cloudflare-ai-gateway` | `CLOUDFLARE_API_TOKEN` + `CLOUDFLARE_ACCOUNT_ID` + `CLOUDFLARE_GATEWAY_ID` | | Requesty | `requesty` | `REQUESTY_API_KEY` | @@ -111,3 +111,10 @@ agents: helper: model: local # helper runs locally for free ``` + +Built-in provider IDs follow the models.dev catalogue. The legacy IDs +`fireworks`, `together`, `moonshot`, and `opencode-zen` remain accepted as +`fireworks-ai`, `togetherai`, `moonshotai`, and `opencode`, respectively. New +configurations use the canonical IDs; older Docker Agent binaries may not +recognize them. API model names, URLs, and environment variables are unchanged. +Custom provider definitions take precedence over built-in IDs. diff --git a/docs/providers/together/index.md b/docs/providers/together/index.md index 8dda71006..4cd5acd3d 100644 --- a/docs/providers/together/index.md +++ b/docs/providers/together/index.md @@ -15,6 +15,8 @@ models, serving Llama, Qwen, DeepSeek, Kimi, GLM and others through an OpenAI-compatible API. Docker Agent includes built-in support for Together AI as an alias provider. +Use the provider ID `togetherai`. The legacy ID `together` is also accepted. + ## Setup 1. Create an API key from the [Together AI settings](https://api.together.ai/settings/api-keys). @@ -33,7 +35,7 @@ The simplest way to use Together AI: ```yaml agents: root: - model: together/meta-llama/Llama-3.3-70B-Instruct-Turbo + model: togetherai/meta-llama/Llama-3.3-70B-Instruct-Turbo description: Assistant using Together AI instruction: You are a helpful assistant. ``` @@ -45,7 +47,7 @@ For more control over parameters: ```yaml models: together_model: - provider: together + provider: togetherai model: meta-llama/Llama-3.3-70B-Instruct-Turbo temperature: 0.7 max_tokens: 8192 @@ -89,7 +91,7 @@ system messages into a single one for this provider. ```yaml agents: coder: - model: together/Qwen/Qwen3-235B-A22B-Instruct-2507-tput + model: togetherai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput description: Code assistant using Qwen3 on Together AI instruction: | You are an expert programmer. diff --git a/examples/README.md b/examples/README.md index 1f6196985..afea7f346 100644 --- a/examples/README.md +++ b/examples/README.md @@ -193,6 +193,7 @@ remote MCP endpoints. | File | What it shows | |------|---------------| | [`custom_provider.yaml`](custom_provider.yaml) | Talking to any OpenAI-compatible endpoint via a custom provider. | +| [`opencode-zen.yaml`](opencode-zen.yaml) | OpenCode Zen gateway (`opencode`). | | [`compose-secrets.yaml`](compose-secrets.yaml) | Reading API keys from Docker Compose / Swarm secrets. | | [`env_placeholders.yaml`](env_placeholders.yaml) | `${env.VAR}` substitution inside the YAML. | | [`model_env_substitution.yaml`](model_env_substitution.yaml) | `${env.VAR}` substitution in a model's `model` / `base_url`. | @@ -200,12 +201,12 @@ remote MCP endpoints. | [`baseten.yaml`](baseten.yaml) | Baseten cloud provider. | | [`ovhcloud.yaml`](ovhcloud.yaml) | OVHcloud AI Endpoints provider. | | [`groq.yaml`](groq.yaml) | Groq fast-inference provider. | -| [`fireworks.yaml`](fireworks.yaml) | Fireworks AI open-model inference provider. | +| [`fireworks.yaml`](fireworks.yaml) | Fireworks AI open-model inference (`fireworks-ai`). | | [`deepseek.yaml`](deepseek.yaml) | DeepSeek chat and reasoning provider. | | [`cerebras.yaml`](cerebras.yaml) | Cerebras fast-inference provider. | -| [`together.yaml`](together.yaml) | Together AI open-model inference provider. | +| [`together.yaml`](together.yaml) | Together AI open-model inference (`togetherai`). | | [`huggingface.yaml`](huggingface.yaml) | Hugging Face Inference Providers open-model router. | -| [`moonshot.yaml`](moonshot.yaml) | Moonshot AI (Kimi K2) provider. | +| [`moonshot.yaml`](moonshot.yaml) | Moonshot AI Kimi provider (`moonshotai`). | | [`vercel.yaml`](vercel.yaml) | Vercel AI Gateway multi-provider router. | | [`cloudflare-workers-ai.yaml`](cloudflare-workers-ai.yaml) | Cloudflare Workers AI edge-hosted open models. | | [`cloudflare-ai-gateway.yaml`](cloudflare-ai-gateway.yaml) | Cloudflare AI Gateway multi-provider router. | @@ -225,6 +226,10 @@ remote MCP endpoints. | [`thinking_budget.yaml`](thinking_budget.yaml) | Reasoning/thinking budgets across OpenAI, Anthropic and Google. | | [`task_budget.yaml`](task_budget.yaml) | Anthropic `task_budget`: cap total tokens spent across a multi-step agentic task. | +Provider examples use the canonical IDs `fireworks-ai`, `togetherai`, +`moonshotai`, and `opencode`. The old IDs `fireworks`, `together`, `moonshot`, +and `opencode-zen` remain accepted for existing configurations. + --- ## Permissions, redaction & sandboxing diff --git a/examples/fireworks.yaml b/examples/fireworks.yaml index 77d4f7206..83cee749c 100644 --- a/examples/fireworks.yaml +++ b/examples/fireworks.yaml @@ -2,7 +2,7 @@ models: fireworks_model: - provider: fireworks + provider: fireworks-ai model: accounts/fireworks/models/kimi-k3 agents: diff --git a/examples/moonshot.yaml b/examples/moonshot.yaml index 20f447a1b..27af5a555 100644 --- a/examples/moonshot.yaml +++ b/examples/moonshot.yaml @@ -2,7 +2,7 @@ models: moonshot_model: - provider: moonshot + provider: moonshotai model: kimi-k3 agents: diff --git a/examples/opencode-zen.yaml b/examples/opencode-zen.yaml index e8249c1ea..5bf3cbd3a 100644 --- a/examples/opencode-zen.yaml +++ b/examples/opencode-zen.yaml @@ -10,7 +10,7 @@ agents: root: - model: opencode-zen/deepseek-v4-flash-free + model: opencode/deepseek-v4-flash-free description: A helpful AI assistant powered by OpenCode Zen (free) instruction: | You are a helpful AI assistant powered by OpenCode Zen. diff --git a/examples/together.yaml b/examples/together.yaml index 6654f3d2b..f4ec76d6e 100644 --- a/examples/together.yaml +++ b/examples/together.yaml @@ -2,7 +2,7 @@ models: together_model: - provider: together + provider: togetherai model: meta-llama/Llama-3.3-70B-Instruct-Turbo agents: diff --git a/pkg/config/auto.go b/pkg/config/auto.go index 79db37706..84b537706 100644 --- a/pkg/config/auto.go +++ b/pkg/config/auto.go @@ -42,7 +42,7 @@ type providerConfig struct { // The first provider with a configured API key will be selected by AutoModelConfig. // DMR is always appended as the final fallback (not listed here). // -// opencode-zen is ordered before opencode-go because both share OPENCODE_API_KEY: +// opencode is ordered before opencode-go because both share OPENCODE_API_KEY: // when the key is set, Zen wins auto-selection. A subscriber who only uses Go // should set the provider explicitly (e.g. `--model opencode-go/...`) rather than // relying on auto; see docs/providers/opencode-go for details. @@ -65,12 +65,12 @@ var cloudProviders = []providerConfig{ {"baseten", []string{"BASETEN_API_KEY"}, "BASETEN_API_KEY", "BASETEN_API_KEY"}, {"ovhcloud", []string{"OVH_AI_ENDPOINTS_ACCESS_TOKEN"}, "OVH_AI_ENDPOINTS_ACCESS_TOKEN", "OVH_AI_ENDPOINTS_ACCESS_TOKEN"}, {"groq", []string{"GROQ_API_KEY"}, "GROQ_API_KEY", "GROQ_API_KEY"}, - {"fireworks", []string{"FIREWORKS_API_KEY"}, "FIREWORKS_API_KEY", "FIREWORKS_API_KEY"}, + {"fireworks-ai", []string{"FIREWORKS_API_KEY"}, "FIREWORKS_API_KEY", "FIREWORKS_API_KEY"}, {"deepseek", []string{"DEEPSEEK_API_KEY"}, "DEEPSEEK_API_KEY", "DEEPSEEK_API_KEY"}, {"cerebras", []string{"CEREBRAS_API_KEY"}, "CEREBRAS_API_KEY", "CEREBRAS_API_KEY"}, - {"together", []string{"TOGETHER_API_KEY"}, "TOGETHER_API_KEY", "TOGETHER_API_KEY"}, + {"togetherai", []string{"TOGETHER_API_KEY"}, "TOGETHER_API_KEY", "TOGETHER_API_KEY"}, {"huggingface", []string{"HF_TOKEN"}, "HF_TOKEN", "HF_TOKEN"}, - {"moonshot", []string{"MOONSHOT_API_KEY"}, "MOONSHOT_API_KEY", "MOONSHOT_API_KEY"}, + {"moonshotai", []string{"MOONSHOT_API_KEY"}, "MOONSHOT_API_KEY", "MOONSHOT_API_KEY"}, {"vercel", []string{"AI_GATEWAY_API_KEY"}, "AI_GATEWAY_API_KEY", "AI_GATEWAY_API_KEY"}, {"amazon-bedrock", []string{ "AWS_BEARER_TOKEN_BEDROCK", @@ -78,7 +78,7 @@ var cloudProviders = []providerConfig{ "AWS_PROFILE", "AWS_ROLE_ARN", }, "AWS_ACCESS_KEY_ID (or AWS_PROFILE, AWS_ROLE_ARN, AWS_BEARER_TOKEN_BEDROCK)", ""}, - {"opencode-zen", []string{"OPENCODE_API_KEY"}, "OPENCODE_API_KEY", "OPENCODE_API_KEY"}, + {"opencode", []string{"OPENCODE_API_KEY"}, "OPENCODE_API_KEY", "OPENCODE_API_KEY"}, {"opencode-go", []string{"OPENCODE_API_KEY"}, "OPENCODE_API_KEY", "OPENCODE_API_KEY"}, } @@ -161,16 +161,16 @@ var DefaultModels = map[string]string{ "baseten": "deepseek-ai/DeepSeek-V4-Pro", "ovhcloud": "Qwen3.5-397B-A17B", "groq": "llama-3.3-70b-versatile", - "fireworks": "accounts/fireworks/models/kimi-k3", + "fireworks-ai": "accounts/fireworks/models/kimi-k3", "deepseek": "deepseek-v4-pro", "cerebras": "gpt-oss-120b", - "together": "meta-llama/Llama-3.3-70B-Instruct-Turbo", + "togetherai": "meta-llama/Llama-3.3-70B-Instruct-Turbo", "huggingface": "meta-llama/Llama-3.3-70B-Instruct", - "moonshot": "kimi-k3", + "moonshotai": "kimi-k3", "vercel": "openai/gpt-5.6-sol", "amazon-bedrock": "global.anthropic.claude-sonnet-5", "opencode-go": "deepseek-v4-flash", - "opencode-zen": "deepseek-v4-flash-free", + "opencode": "deepseek-v4-flash-free", } // nonForwardableTokenEnvVars lists provider token env vars that are NOT safe to diff --git a/pkg/config/auto_test.go b/pkg/config/auto_test.go index f2db1d298..1b05e58a9 100644 --- a/pkg/config/auto_test.go +++ b/pkg/config/auto_test.go @@ -13,6 +13,7 @@ import ( "github.com/docker/docker-agent/pkg/chatgpt" "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider" "github.com/docker/docker-agent/pkg/modelsdev" ) @@ -99,7 +100,7 @@ func TestAvailableProviders_NoGateway(t *testing.T) { envVars: map[string]string{ "FIREWORKS_API_KEY": "test-key", }, - expectedProvider: "fireworks", + expectedProvider: "fireworks-ai", }, { name: "deepseek api key present", @@ -120,7 +121,7 @@ func TestAvailableProviders_NoGateway(t *testing.T) { envVars: map[string]string{ "TOGETHER_API_KEY": "test-key", }, - expectedProvider: "together", + expectedProvider: "togetherai", }, { name: "huggingface token present", @@ -134,7 +135,7 @@ func TestAvailableProviders_NoGateway(t *testing.T) { envVars: map[string]string{ "MOONSHOT_API_KEY": "test-key", }, - expectedProvider: "moonshot", + expectedProvider: "moonshotai", }, { name: "vercel ai gateway key present", @@ -337,7 +338,7 @@ func TestAutoModelConfig(t *testing.T) { envVars: map[string]string{ "FIREWORKS_API_KEY": "test-key", }, - expectedProvider: "fireworks", + expectedProvider: "fireworks-ai", expectedModel: "accounts/fireworks/models/kimi-k3", expectedMaxTokens: 32000, }, @@ -364,7 +365,7 @@ func TestAutoModelConfig(t *testing.T) { envVars: map[string]string{ "TOGETHER_API_KEY": "test-key", }, - expectedProvider: "together", + expectedProvider: "togetherai", expectedModel: "meta-llama/Llama-3.3-70B-Instruct-Turbo", expectedMaxTokens: 32000, }, @@ -382,7 +383,7 @@ func TestAutoModelConfig(t *testing.T) { envVars: map[string]string{ "MOONSHOT_API_KEY": "test-key", }, - expectedProvider: "moonshot", + expectedProvider: "moonshotai", expectedModel: "kimi-k3", expectedMaxTokens: 32000, }, @@ -477,7 +478,7 @@ func TestDefaultModels(t *testing.T) { t.Parallel() // Test that DefaultModels map has all expected providers - expectedProviders := []string{"openai", "anthropic", "google", "dmr", "mistral", "openrouter", "baseten", "ovhcloud", "groq", "fireworks", "deepseek", "cerebras", "together", "huggingface", "moonshot", "vercel", "amazon-bedrock", "opencode-zen", "opencode-go", "github-copilot"} + expectedProviders := []string{"openai", "anthropic", "google", "dmr", "mistral", "openrouter", "baseten", "ovhcloud", "groq", "fireworks-ai", "deepseek", "cerebras", "togetherai", "huggingface", "moonshotai", "vercel", "amazon-bedrock", "opencode", "opencode-go", "github-copilot"} for _, provider := range expectedProviders { t.Run(provider, func(t *testing.T) { @@ -498,23 +499,23 @@ func TestDefaultModels(t *testing.T) { assert.Equal(t, "deepseek-ai/DeepSeek-V4-Pro", DefaultModels["baseten"]) assert.Equal(t, "Qwen3.5-397B-A17B", DefaultModels["ovhcloud"]) assert.Equal(t, "llama-3.3-70b-versatile", DefaultModels["groq"]) - assert.Equal(t, "accounts/fireworks/models/kimi-k3", DefaultModels["fireworks"]) + assert.Equal(t, "accounts/fireworks/models/kimi-k3", DefaultModels["fireworks-ai"]) assert.Equal(t, "deepseek-v4-pro", DefaultModels["deepseek"]) assert.Equal(t, "gpt-oss-120b", DefaultModels["cerebras"]) - assert.Equal(t, "meta-llama/Llama-3.3-70B-Instruct-Turbo", DefaultModels["together"]) + assert.Equal(t, "meta-llama/Llama-3.3-70B-Instruct-Turbo", DefaultModels["togetherai"]) assert.Equal(t, "meta-llama/Llama-3.3-70B-Instruct", DefaultModels["huggingface"]) - assert.Equal(t, "kimi-k3", DefaultModels["moonshot"]) + assert.Equal(t, "kimi-k3", DefaultModels["moonshotai"]) assert.Equal(t, "openai/gpt-5.6-sol", DefaultModels["vercel"]) assert.Equal(t, "global.anthropic.claude-sonnet-5", DefaultModels["amazon-bedrock"]) assert.Equal(t, "deepseek-v4-flash", DefaultModels["opencode-go"]) - assert.Equal(t, "deepseek-v4-flash-free", DefaultModels["opencode-zen"]) + assert.Equal(t, "deepseek-v4-flash-free", DefaultModels["opencode"]) } func TestAutoModelConfig_IntegrationWithDefaultModels(t *testing.T) { t.Parallel() // Verify that AutoModelConfig always returns a model from DefaultModels - providers := []string{"openai", "anthropic", "google", "mistral", "openrouter", "baseten", "ovhcloud", "groq", "fireworks", "deepseek", "cerebras", "together", "huggingface", "moonshot", "vercel", "opencode-zen", "github-copilot"} + providers := []string{"openai", "anthropic", "google", "mistral", "openrouter", "baseten", "ovhcloud", "groq", "fireworks-ai", "deepseek", "cerebras", "togetherai", "huggingface", "moonshotai", "vercel", "opencode", "github-copilot"} for _, provider := range providers { t.Run(provider, func(t *testing.T) { @@ -542,21 +543,21 @@ func TestAutoModelConfig_IntegrationWithDefaultModels(t *testing.T) { envVars["OVH_AI_ENDPOINTS_ACCESS_TOKEN"] = "test-token" case "groq": envVars["GROQ_API_KEY"] = "test-key" - case "fireworks": + case "fireworks-ai": envVars["FIREWORKS_API_KEY"] = "test-key" case "deepseek": envVars["DEEPSEEK_API_KEY"] = "test-key" case "cerebras": envVars["CEREBRAS_API_KEY"] = "test-key" - case "together": + case "togetherai": envVars["TOGETHER_API_KEY"] = "test-key" case "huggingface": envVars["HF_TOKEN"] = "test-token" - case "moonshot": + case "moonshotai": envVars["MOONSHOT_API_KEY"] = "test-key" case "vercel": envVars["AI_GATEWAY_API_KEY"] = "test-key" - case "opencode-zen": + case "opencode": envVars["OPENCODE_API_KEY"] = "test-key" } @@ -691,7 +692,7 @@ func TestAvailableProviders_PrecedenceOrder(t *testing.T) { "DEEPSEEK_API_KEY": "test-key", }) providers = AvailableProviders(t.Context(), "", env) - assert.Equal(t, "fireworks", providers[0]) + assert.Equal(t, "fireworks-ai", providers[0]) // deepseek wins over cerebras env = environment.NewMapEnvProvider(map[string]string{ @@ -715,7 +716,7 @@ func TestAvailableProviders_PrecedenceOrder(t *testing.T) { "HF_TOKEN": "test-token", }) providers = AvailableProviders(t.Context(), "", env) - assert.Equal(t, "together", providers[0]) + assert.Equal(t, "togetherai", providers[0]) // huggingface wins over moonshot env = environment.NewMapEnvProvider(map[string]string{ @@ -731,7 +732,7 @@ func TestAvailableProviders_PrecedenceOrder(t *testing.T) { "AI_GATEWAY_API_KEY": "test-key", }) providers = AvailableProviders(t.Context(), "", env) - assert.Equal(t, "moonshot", providers[0]) + assert.Equal(t, "moonshotai", providers[0]) // vercel wins over amazon-bedrock env = environment.NewMapEnvProvider(map[string]string{ @@ -746,7 +747,7 @@ func TestAvailableProviders_PrecedenceOrder(t *testing.T) { "OPENCODE_API_KEY": "test-key", }) providers = AvailableProviders(t.Context(), "", env) - assert.Equal(t, "opencode-zen", providers[0]) + assert.Equal(t, "opencode", providers[0]) // No keys at all - dmr should be selected env = environment.NewNoEnvProvider() @@ -1182,34 +1183,33 @@ func TestCloudProviderEnvVars(t *testing.T) { assert.Equal(t, []string{"GITHUB_TOKEN", "GH_TOKEN"}, providers[copilotIdx].EnvVars) } -// TestDefaultModelsExistInModelsDev is the regression test for issue #4133: -// DefaultModels must reference models that actually exist in the models.dev -// catalog, since AutoModelConfig hands them straight to real users with no -// other validation. modelsDevAbsentProviders and modelsDevCatalogProviders -// (both defined in examples_test.go) are reused so providers legitimately -// absent, or aliased under a different id, in the catalog don't produce -// false failures. +// Defaults must resolve through the production lookup against the committed snapshot. func TestDefaultModelsExistInModelsDev(t *testing.T) { t.Parallel() - - modelsStore, err := modelsdev.NewStore() - require.NoError(t, err) - - for provider, model := range DefaultModels { - t.Run(provider, func(t *testing.T) { + store := modelsdev.NewDatabaseStore(modelsdev.EmbeddedSnapshot()) + for providerID, model := range DefaultModels { + t.Run(providerID, func(t *testing.T) { t.Parallel() - - if modelsDevAbsentProviders[provider] { - t.Skipf("provider %q is not expected to exist in the models.dev catalog", provider) + require.True(t, provider.IsKnownProvider(providerID)) + require.Equal(t, providerID, modelsdev.CanonicalProviderID(providerID)) + if providerID == "dmr" { + return // Local models are not catalogued. } - - catalogProvider := provider - if id, ok := modelsDevCatalogProviders[provider]; ok { - catalogProvider = id + if providerID == "chatgpt" { + // ChatGPT is distinct; only its model names are validated against OpenAI. + providerID = "openai" } - - _, err := modelsStore.GetModel(t.Context(), modelsdev.NewID(catalogProvider, model)) - require.NoError(t, err, "DefaultModels[%q] = %q must exist in the models.dev catalog", provider, model) + _, err := store.GetModel(t.Context(), modelsdev.NewID(providerID, model)) + require.NoError(t, err) }) } } + +func TestCloudProvidersCanonical(t *testing.T) { + t.Parallel() + for _, cfg := range cloudProviders { + assert.Equal(t, cfg.name, modelsdev.CanonicalProviderID(cfg.name)) + assert.True(t, provider.IsKnownProvider(cfg.name), cfg.name) + assert.Contains(t, DefaultModels, cfg.name) + } +} diff --git a/pkg/config/examples_test.go b/pkg/config/examples_test.go index 800eff15b..f6f74d685 100644 --- a/pkg/config/examples_test.go +++ b/pkg/config/examples_test.go @@ -18,16 +18,10 @@ import ( "github.com/docker/docker-agent/pkg/modelsdev" ) -// modelsDevCatalogProviders maps a docker-agent provider name to the id -// models.dev actually catalogs it under, for providers where the two -// diverge. Resolving through this map (instead of skipping validation -// outright) is what caught the stale Fireworks model reference in #4132. +// modelsDevCatalogProviders validates ChatGPT example names against OpenAI's +// catalog without treating the subscription backend as a full metadata alias. var modelsDevCatalogProviders = map[string]string{ - "fireworks": "fireworks-ai", // models.dev catalogs Fireworks under the "fireworks-ai" id, not "fireworks" - "together": "togetherai", // models.dev catalogs Together AI under the "togetherai" id, not "together" - "moonshot": "moonshotai", // models.dev catalogs Moonshot AI under the "moonshotai" id, not "moonshot" - "chatgpt": "openai", // ChatGPT subscription backend; models.dev catalogs its models under the "openai" id - "opencode-zen": "opencode", // models.dev catalogs the OpenCode Zen router under the "opencode" id + "chatgpt": "openai", } // modelsDevAbsentProviders lists providers that are valid at runtime but @@ -36,7 +30,6 @@ var modelsDevCatalogProviders = map[string]string{ // lookups for these to avoid false failures. var modelsDevAbsentProviders = map[string]bool{ "dmr": true, // Docker Model Runner (local, not in catalog) - "ovhcloud": true, // models.dev lower-cases OVHcloud model ids (e.g. "qwen3.5-397b-a17b"); the provider API is case-sensitive and takes "Qwen3.5-397B-A17B" "cloudflare-workers-ai": true, // example uses an @cf/... model id not present in the models.dev snapshot (only variant ids like -fp8 are listed) "cloudflare-ai-gateway": true, // multi-provider router; example model ids use the gateway's provider/model form, not guaranteed to match a models.dev id } @@ -69,9 +62,8 @@ func collectExamples(t *testing.T) []string { // from the environment's credentials, routed models span multiple // providers, custom providers are self-contained (already validated via // cfg.Providers), modelsDevAbsentProviders lists providers models.dev -// deliberately does not catalog, and modelsDevCatalogProviders resolves -// the remaining providers to the id models.dev actually catalogs them -// under, when it diverges from the docker-agent provider name. +// deliberately does not catalog, and modelsDevCatalogProviders validates +// ChatGPT model names against OpenAI's catalog. func catalogModelRefs(cfg *latest.Config) []modelsdev.ID { var ids []modelsdev.ID for _, model := range cfg.Models { @@ -240,3 +232,42 @@ func TestHCLExamplesMatchYAML(t *testing.T) { }) } } + +func TestExampleProvidersCanonical(t *testing.T) { + t.Parallel() + for _, example := range collectExamples(t) { + t.Run(example, func(t *testing.T) { + t.Parallel() + cfg, err := Load(t.Context(), hcl.NewSource(NewFileSource(example))) + require.NoError(t, err) + data, err := yaml.Marshal(cfg) + require.NoError(t, err) + var document any + require.NoError(t, yaml.Unmarshal(data, &document)) + var check func(any) + check = func(value any) { + switch node := value.(type) { + case map[string]any: + for key, child := range node { + if key == "provider" { + if id, ok := child.(string); ok { + assert.Equal(t, id, modelsdev.CanonicalProviderID(id), "provider in %s", example) + } + } + check(child) + } + case []any: + for _, child := range node { + check(child) + } + case string: + // Inline refs also appear in fallback, routing and first_available lists. + if id, _, ok := strings.Cut(node, "/"); ok { + assert.Equal(t, id, modelsdev.CanonicalProviderID(id), "inline reference %s", node) + } + } + } + check(document) + }) + } +} diff --git a/pkg/model/provider/aliases.go b/pkg/model/provider/aliases.go index 49e612364..7327cc8cb 100644 --- a/pkg/model/provider/aliases.go +++ b/pkg/model/provider/aliases.go @@ -7,6 +7,7 @@ import ( "strings" "github.com/docker/docker-agent/pkg/chatgpt" + "github.com/docker/docker-agent/pkg/modelsdev" ) // Alias defines the configuration for a provider alias. @@ -93,7 +94,7 @@ var Aliases = map[string]Alias{ BaseURL: "https://api.groq.com/openai/v1", TokenEnvVar: "GROQ_API_KEY", }, - "fireworks": { + "fireworks-ai": { APIType: "openai", BaseURL: "https://api.fireworks.ai/inference/v1", TokenEnvVar: "FIREWORKS_API_KEY", @@ -108,7 +109,7 @@ var Aliases = map[string]Alias{ BaseURL: "https://api.cerebras.ai/v1", TokenEnvVar: "CEREBRAS_API_KEY", }, - "together": { + "togetherai": { APIType: "openai", BaseURL: "https://api.together.xyz/v1", TokenEnvVar: "TOGETHER_API_KEY", @@ -118,7 +119,7 @@ var Aliases = map[string]Alias{ BaseURL: "https://router.huggingface.co/v1", TokenEnvVar: "HF_TOKEN", }, - "moonshot": { + "moonshotai": { APIType: "openai", BaseURL: "https://api.moonshot.ai/v1", TokenEnvVar: "MOONSHOT_API_KEY", @@ -161,7 +162,7 @@ var Aliases = map[string]Alias{ BaseURL: "https://opencode.ai/zen/go/v1", TokenEnvVar: "OPENCODE_API_KEY", }, - "opencode-zen": { + "opencode": { APIType: "openai", BaseURL: "https://opencode.ai/zen/v1", TokenEnvVar: "OPENCODE_API_KEY", @@ -172,7 +173,7 @@ var Aliases = map[string]Alias{ // Lookup is case-sensitive; callers that need case-insensitive matching // should normalise the name first (e.g. [strings.ToLower]). func LookupAlias(name string) (Alias, bool) { - alias, ok := Aliases[name] + alias, ok := Aliases[modelsdev.CanonicalProviderID(name)] return alias, ok } diff --git a/pkg/model/provider/aliases_test.go b/pkg/model/provider/aliases_test.go index e0c48a1b5..2b4b4115c 100644 --- a/pkg/model/provider/aliases_test.go +++ b/pkg/model/provider/aliases_test.go @@ -6,6 +6,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/modelsdev" ) func TestLookupAlias(t *testing.T) { @@ -36,18 +38,18 @@ func TestCatalogAliases(t *testing.T) { t.Parallel() expected := map[string]Alias{ - "openrouter": {APIType: "openai", BaseURL: "https://openrouter.ai/api/v1", TokenEnvVar: "OPENROUTER_API_KEY"}, - "baseten": {APIType: "openai", BaseURL: "https://inference.baseten.co/v1", TokenEnvVar: "BASETEN_API_KEY"}, - "ovhcloud": {APIType: "openai", BaseURL: "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1", TokenEnvVar: "OVH_AI_ENDPOINTS_ACCESS_TOKEN"}, - "groq": {APIType: "openai", BaseURL: "https://api.groq.com/openai/v1", TokenEnvVar: "GROQ_API_KEY"}, - "deepseek": {APIType: "openai", BaseURL: "https://api.deepseek.com/v1", TokenEnvVar: "DEEPSEEK_API_KEY"}, - "cerebras": {APIType: "openai", BaseURL: "https://api.cerebras.ai/v1", TokenEnvVar: "CEREBRAS_API_KEY"}, - "fireworks": {APIType: "openai", BaseURL: "https://api.fireworks.ai/inference/v1", TokenEnvVar: "FIREWORKS_API_KEY"}, - "together": {APIType: "openai", BaseURL: "https://api.together.xyz/v1", TokenEnvVar: "TOGETHER_API_KEY"}, - "huggingface": {APIType: "openai", BaseURL: "https://router.huggingface.co/v1", TokenEnvVar: "HF_TOKEN"}, - "moonshot": {APIType: "openai", BaseURL: "https://api.moonshot.ai/v1", TokenEnvVar: "MOONSHOT_API_KEY"}, - "nvidia": {APIType: "openai", BaseURL: "https://integrate.api.nvidia.com/v1", TokenEnvVar: "NVIDIA_API_KEY"}, - "vercel": {APIType: "openai", BaseURL: "https://ai-gateway.vercel.sh/v1", TokenEnvVar: "AI_GATEWAY_API_KEY"}, + "openrouter": {APIType: "openai", BaseURL: "https://openrouter.ai/api/v1", TokenEnvVar: "OPENROUTER_API_KEY"}, + "baseten": {APIType: "openai", BaseURL: "https://inference.baseten.co/v1", TokenEnvVar: "BASETEN_API_KEY"}, + "ovhcloud": {APIType: "openai", BaseURL: "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1", TokenEnvVar: "OVH_AI_ENDPOINTS_ACCESS_TOKEN"}, + "groq": {APIType: "openai", BaseURL: "https://api.groq.com/openai/v1", TokenEnvVar: "GROQ_API_KEY"}, + "deepseek": {APIType: "openai", BaseURL: "https://api.deepseek.com/v1", TokenEnvVar: "DEEPSEEK_API_KEY"}, + "cerebras": {APIType: "openai", BaseURL: "https://api.cerebras.ai/v1", TokenEnvVar: "CEREBRAS_API_KEY"}, + "fireworks-ai": {APIType: "openai", BaseURL: "https://api.fireworks.ai/inference/v1", TokenEnvVar: "FIREWORKS_API_KEY"}, + "togetherai": {APIType: "openai", BaseURL: "https://api.together.xyz/v1", TokenEnvVar: "TOGETHER_API_KEY"}, + "huggingface": {APIType: "openai", BaseURL: "https://router.huggingface.co/v1", TokenEnvVar: "HF_TOKEN"}, + "moonshotai": {APIType: "openai", BaseURL: "https://api.moonshot.ai/v1", TokenEnvVar: "MOONSHOT_API_KEY"}, + "nvidia": {APIType: "openai", BaseURL: "https://integrate.api.nvidia.com/v1", TokenEnvVar: "NVIDIA_API_KEY"}, + "vercel": {APIType: "openai", BaseURL: "https://ai-gateway.vercel.sh/v1", TokenEnvVar: "AI_GATEWAY_API_KEY"}, } for name, want := range expected { @@ -118,3 +120,37 @@ func TestEachAlias_EarlyTermination(t *testing.T) { } assert.Equal(t, 1, count, "iteration should stop when consumer breaks out") } + +func TestProviderIDsMatchEmbeddedCatalog(t *testing.T) { + t.Parallel() + db := modelsdev.EmbeddedSnapshot() + absent := map[string]string{ + "chatgpt": "subscription backend, distinct from OpenAI", + "dmr": "local Docker Model Runner", + "ollama": "local Ollama server", + } + for _, name := range AllProviders() { + t.Run(name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, name, modelsdev.CanonicalProviderID(name)) + _, exists := db.Providers[name] + if reason, exempt := absent[name]; exempt { + assert.False(t, exists, "remove stale exemption: %s", reason) + } else { + assert.True(t, exists, "built-in provider must use its models.dev ID") + } + }) + } + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + canonical := modelsdev.CanonicalProviderID(legacy) + assert.NotEqual(t, legacy, canonical) + assert.NotContains(t, Aliases, legacy) + assert.NotContains(t, AllProviders(), legacy) + assert.NotContains(t, db.Providers, legacy, "legacy ID now collides with an upstream service") + require.Contains(t, Aliases, canonical) + got, ok := LookupAlias(legacy) + require.True(t, ok) + assert.Equal(t, Aliases[canonical], got) + assert.True(t, IsKnownProvider(legacy)) + } +} diff --git a/pkg/model/provider/canonical_provider_test.go b/pkg/model/provider/canonical_provider_test.go new file mode 100644 index 000000000..c64f254d8 --- /dev/null +++ b/pkg/model/provider/canonical_provider_test.go @@ -0,0 +1,49 @@ +package provider + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/modelsdev" +) + +func TestCanonicalProviderDefaults(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + canonical := modelsdev.CanonicalProviderID(legacy) + cfg := &latest.ModelConfig{Provider: legacy, Model: "Mixed/Model"} + got := applyProviderDefaults(cfg, nil) + want := applyProviderDefaults(&latest.ModelConfig{Provider: canonical, Model: cfg.Model}, nil) + assert.Equal(t, want, got) + assert.Equal(t, legacy, cfg.Provider) + assert.Equal(t, canonical, got.Provider) + assert.Equal(t, cfg.Model, got.Model) + assert.Equal(t, "openai", resolveProviderType(got), "transport is not a naming alias") + + custom := latest.ProviderConfig{BaseURL: "https://custom.invalid/v1", TokenKey: "CUSTOM_TOKEN", APIType: "openai_chatcompletions"} + providers := map[string]latest.ProviderConfig{legacy: custom, canonical: {BaseURL: "https://canonical-custom.invalid/v1", APIType: "openai_chatcompletions"}} + got = applyProviderDefaults(cfg, providers) + assert.Equal(t, legacy, got.Provider, "custom key wins before normalization") + assert.Equal(t, custom.BaseURL, got.BaseURL) + assert.Equal(t, custom.TokenKey, got.TokenKey) + assert.Equal(t, "openai_chatcompletions", resolveProviderType(got)) + assert.Equal(t, providers[canonical].BaseURL, applyProviderDefaults(&latest.ModelConfig{Provider: canonical, Model: cfg.Model}, providers).BaseURL) + + providers = map[string]latest.ProviderConfig{"mine": {Provider: legacy, BaseURL: custom.BaseURL, TokenKey: custom.TokenKey}} + got = applyProviderDefaults(&latest.ModelConfig{Provider: "mine", Model: cfg.Model}, providers) + assert.Equal(t, canonical, got.Provider) + assert.Equal(t, custom.BaseURL, got.BaseURL) + assert.Equal(t, custom.TokenKey, got.TokenKey) + + cfg.BaseURL, cfg.TokenKey = "https://override.invalid/v1", "OVERRIDE_TOKEN" + got = applyProviderDefaults(cfg, nil) + require.Equal(t, cfg.BaseURL, got.BaseURL) + assert.Equal(t, cfg.TokenKey, got.TokenKey) + }) + } +} diff --git a/pkg/model/provider/defaults.go b/pkg/model/provider/defaults.go index 0707d8bcc..e0666796e 100644 --- a/pkg/model/provider/defaults.go +++ b/pkg/model/provider/defaults.go @@ -12,6 +12,7 @@ import ( "github.com/docker/docker-agent/pkg/environment" "github.com/docker/docker-agent/pkg/model/provider/options" "github.com/docker/docker-agent/pkg/modelinfo" + "github.com/docker/docker-agent/pkg/modelsdev" ) // expandModelConfigEnv substitutes ${env.X} / ${X} references in the model @@ -157,11 +158,15 @@ func applyProviderDefaults(cfg *latest.ModelConfig, customProviders map[string]l "base_url", providerCfg.BaseURL, ) mergeFromProviderConfig(enhancedCfg, providerCfg) + if providerCfg.Provider != "" { + enhancedCfg.Provider = modelsdev.CanonicalProviderID(enhancedCfg.Provider) + } applyModelDefaults(enhancedCfg) return enhancedCfg } - if alias, exists := LookupAlias(cfg.Provider); exists { + enhancedCfg.Provider = modelsdev.CanonicalProviderID(enhancedCfg.Provider) + if alias, exists := LookupAlias(enhancedCfg.Provider); exists { applyAliasFallbacks(enhancedCfg, alias) } diff --git a/pkg/model/provider/openai/canonical_provider_test.go b/pkg/model/provider/openai/canonical_provider_test.go new file mode 100644 index 000000000..aff6c095c --- /dev/null +++ b/pkg/model/provider/openai/canonical_provider_test.go @@ -0,0 +1,28 @@ +package openai + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/modelsdev" +) + +func TestCanonicalProviderBehavior(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + canonical := modelsdev.CanonicalProviderID(legacy) + assert.Equal(t, autoSelectsResponsesAPI(legacy), autoSelectsResponsesAPI(canonical)) + assert.Equal(t, + shouldMergeConsecutiveMessages(&latest.ModelConfig{Provider: legacy}), + shouldMergeConsecutiveMessages(&latest.ModelConfig{Provider: canonical}), + ) + assert.NotContains(t, openModelHostProviders, legacy) + } + assert.True(t, autoSelectsResponsesAPI("opencode")) + assert.True(t, autoSelectsResponsesAPI("opencode-zen")) + assert.False(t, autoSelectsResponsesAPI("opencode-go")) + assert.True(t, shouldMergeConsecutiveMessages(&latest.ModelConfig{Provider: "fireworks-ai"})) + assert.True(t, shouldMergeConsecutiveMessages(&latest.ModelConfig{Provider: "togetherai"})) +} diff --git a/pkg/model/provider/openai/catalog_alias_test.go b/pkg/model/provider/openai/catalog_alias_test.go new file mode 100644 index 000000000..52a84e3e4 --- /dev/null +++ b/pkg/model/provider/openai/catalog_alias_test.go @@ -0,0 +1,176 @@ +package openai + +import ( + "encoding/json" + "io" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider/options" + "github.com/docker/docker-agent/pkg/modelsdev" + "github.com/docker/docker-agent/pkg/tools" +) + +type catalogRequestTransport struct { + http.RoundTripper + + requests chan *http.Request +} + +func (c catalogRequestTransport) RoundTrip(req *http.Request) (*http.Response, error) { + c.requests <- req.Clone(req.Context()) + return c.RoundTripper.RoundTrip(req) +} + +func TestChatCompletions_CatalogAliasesPreserveImages(t *testing.T) { + t.Parallel() + + for _, provider := range []struct{ configured, catalog, model string }{ + {"fireworks", "fireworks-ai", "accounts/fireworks/models/kimi-k3"}, + {"together", "togetherai", "Qwen/Qwen3.5-397B-A17B"}, + {"moonshot", "moonshotai", "kimi-k3"}, + {"opencode-zen", "opencode", "kimi-k3"}, + {"ovhcloud", "ovhcloud", "Qwen3.5-397B-A17B"}, + {"fireworks-ai", "fireworks-ai", "accounts/fireworks/models/kimi-k3"}, + {"togetherai", "togetherai", "Qwen/Qwen3.5-397B-A17B"}, + {"moonshotai", "moonshotai", "kimi-k3"}, + {"opencode", "opencode", "kimi-k3"}, + } { + for _, mode := range []string{"snapshot vision", "text only", "unknown", "override disables", "override enables", "direct text wins"} { + t.Run(provider.configured+"/"+mode, func(t *testing.T) { + t.Parallel() + model := provider.model + store := modelsdev.NewDatabaseStore(modelsdev.EmbeddedSnapshot()) + var override *latest.CapabilitiesConfig + wantImages := false + switch mode { + case "snapshot vision": + wantImages = true + case "text only", "override enables": + model = "text-only" + store = modelsdev.NewDatabaseStore(&modelsdev.Database{Providers: map[string]modelsdev.Provider{ + provider.catalog: {Models: map[string]modelsdev.Model{model: {Modalities: modelsdev.Modalities{Input: []string{"text"}}}}}, + }}) + if mode == "override enables" { + override = &latest.CapabilitiesConfig{Image: true} + wantImages = true + } + case "unknown": + model = "unknown-catalog-model" + case "override disables": + override = &latest.CapabilitiesConfig{Image: false} + case "direct text wins": + catalogModel := model + if provider.configured == "ovhcloud" { + catalogModel = "qwen3.5-397b-a17b" + } + providers := map[string]modelsdev.Provider{ + provider.catalog: {Models: map[string]modelsdev.Model{catalogModel: {Modalities: modelsdev.Modalities{Input: []string{"text", "image"}}}}}, + } + direct := providers[provider.configured] + if direct.Models == nil { + direct.Models = map[string]modelsdev.Model{} + } + direct.Models[model] = modelsdev.Model{Modalities: modelsdev.Modalities{Input: []string{"text"}}} + providers[provider.configured] = direct + store = modelsdev.NewDatabaseStore(&modelsdev.Database{Providers: providers}) + } + + server, capturedBody := captureRequestBody(t) + requests := make(chan *http.Request, 1) + maxTokens, wantMaxTokens := int64(32000), int64(32000) + if metadata, err := store.GetModel(t.Context(), modelsdev.NewID(provider.configured, model)); err == nil && metadata.Limit.Output > 0 { + maxTokens = metadata.Limit.Output + wantMaxTokens = maxTokens + if metadata.Limit.Context > 0 { + wantMaxTokens = min(maxTokens, max(int64(metadata.Limit.Context)-1024, 1)) + } + } + cfg := &latest.ModelConfig{ + Provider: provider.configured, Model: model, + BaseURL: server.URL + "/configured/v1", TokenKey: "CATALOG_TEST_TOKEN", + MaxTokens: &maxTokens, + ProviderOpts: map[string]any{"api_type": "openai_chatcompletions"}, Capabilities: override, + } + client, err := NewClient(t.Context(), cfg, environment.NewMapEnvProvider(map[string]string{ + "CATALOG_TEST_TOKEN": "fake-configured-token", "OPENAI_API_KEY": "fake-wrong-token", + }), options.WithModelsDevStore(store), options.WithHTTPTransportWrapper(func(base http.RoundTripper) http.RoundTripper { + return catalogRequestTransport{RoundTripper: base, requests: requests} + })) + require.NoError(t, err) + imagePart := func(name string, data byte) chat.MessagePart { + return chat.MessagePart{Type: chat.MessagePartTypeDocument, Document: &chat.Document{ + Name: name, MimeType: "image/png", Source: chat.DocumentSource{InlineData: []byte{data}}, + }} + } + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, MultiContent: []chat.MessagePart{ + {Type: chat.MessagePartTypeText, Text: "describe the attachment"}, imagePart("attachment.png", 1), + }}, + {Role: chat.MessageRoleAssistant, ToolCalls: []tools.ToolCall{{ID: "call-image", Type: "function", Function: tools.FunctionCall{Name: "read_file", Arguments: `{}`}}}}, + {Role: chat.MessageRoleTool, ToolCallID: "call-image", MultiContent: []chat.MessagePart{ + {Type: chat.MessagePartTypeText, Text: "image loaded"}, imagePart("tool.png", 2), + }}, + }, nil) + require.NoError(t, err) + defer stream.Close() + for { + _, err := stream.Recv() + if err != nil { + require.ErrorIs(t, err, io.EOF) + break + } + } + var req *http.Request + select { + case req = <-requests: + default: + t.Fatal("no request captured") + } + assert.Equal(t, server.URL+"/configured/v1/chat/completions", req.URL.String()) + assert.Equal(t, "Bearer fake-configured-token", req.Header.Get("Authorization")) + assert.Equal(t, *cfg, client.ModelConfig) + var payload struct { + Model string `json:"model"` + MaxTokens int64 `json:"max_tokens"` + Messages []struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` + } `json:"messages"` + } + require.NoError(t, json.Unmarshal(capturedBody(), &payload)) + assert.Equal(t, model, payload.Model) + assert.Equal(t, wantMaxTokens, payload.MaxTokens) + var imageURLs []string + for _, msg := range payload.Messages { + if msg.Role != "user" { + continue + } + var parts []struct { + Type string `json:"type"` + ImageURL struct { + URL string `json:"url"` + } `json:"image_url"` + } + require.NoError(t, json.Unmarshal(msg.Content, &parts)) + for _, part := range parts { + if part.Type == "image_url" { + imageURLs = append(imageURLs, part.ImageURL.URL) + } + } + } + if wantImages { + assert.Equal(t, []string{"data:image/png;base64,AQ==", "data:image/png;base64,Ag=="}, imageURLs) + } else { + assert.Empty(t, imageURLs) + } + }) + } + } +} diff --git a/pkg/model/provider/openai/client.go b/pkg/model/provider/openai/client.go index ad04de0bd..2d5ba87cb 100644 --- a/pkg/model/provider/openai/client.go +++ b/pkg/model/provider/openai/client.go @@ -30,6 +30,7 @@ import ( "github.com/docker/docker-agent/pkg/model/provider/options" "github.com/docker/docker-agent/pkg/model/provider/providerutil" "github.com/docker/docker-agent/pkg/modelinfo" + "github.com/docker/docker-agent/pkg/modelsdev" "github.com/docker/docker-agent/pkg/rag/prompts" "github.com/docker/docker-agent/pkg/rag/types" "github.com/docker/docker-agent/pkg/tools" @@ -256,16 +257,16 @@ func (c *Client) convertMessages(ctx context.Context, messages []chat.Message) [ // (openai, mistral, xai, minimax, github-copilot, opencode) tolerate multiple // system messages and are deliberately absent so their behavior is unchanged. var openModelHostProviders = map[string]bool{ - "baseten": true, - "ovhcloud": true, - "openrouter": true, - "nebius": true, - "nvidia": true, - "cerebras": true, - "fireworks": true, - "together": true, - "huggingface": true, - "vercel": true, + "baseten": true, + "ovhcloud": true, + "openrouter": true, + "nebius": true, + "nvidia": true, + "cerebras": true, + "fireworks-ai": true, + "togetherai": true, + "huggingface": true, + "vercel": true, // Cloudflare Workers AI serves open-weight models directly; the AI Gateway // fronts them (and other providers) through one endpoint. Both plausibly // reach models with strict single-system-message chat templates. @@ -301,7 +302,7 @@ func shouldMergeConsecutiveMessages(cfg *latest.ModelConfig) bool { if cfg.Provider == "openai" && cfg.BaseURL != "" { return true } - return openModelHostProviders[cfg.Provider] + return openModelHostProviders[modelsdev.CanonicalProviderID(cfg.Provider)] } // contextLimit returns this model's context window in tokens, preferring an @@ -1449,8 +1450,8 @@ func isCustomProvider(cfg *latest.ModelConfig) bool { // driven by modelinfo.SupportsResponsesAPI so new models are picked up by // naming convention rather than a hardcoded allow-list. func autoSelectsResponsesAPI(provider string) bool { - switch provider { - case "openai", "github-copilot", "opencode-zen", chatgpt.ProviderName: + switch modelsdev.CanonicalProviderID(provider) { + case "openai", "github-copilot", "opencode", chatgpt.ProviderName: return true } return false diff --git a/pkg/model/provider/openai_alias_providers_test.go b/pkg/model/provider/openai_alias_providers_test.go index cc0f3c03d..291fd2250 100644 --- a/pkg/model/provider/openai_alias_providers_test.go +++ b/pkg/model/provider/openai_alias_providers_test.go @@ -10,6 +10,7 @@ import ( "net/http" "net/http/httptest" "os" + "slices" "strings" "sync" "testing" @@ -21,11 +22,12 @@ import ( "github.com/docker/docker-agent/pkg/chat" "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/modelsdev" "github.com/docker/docker-agent/pkg/tools" ) // openAIAliasProvider describes a built-in OpenAI-compatible alias provider -// (deepseek, cerebras, fireworks, ...) for the shared wiring tests below. New +// (deepseek, cerebras, fireworks-ai, ...) for the shared wiring tests below. New // aliases of the same shape only need a row here rather than a fresh copy of // the whole end-to-end/live test. type openAIAliasProvider struct { @@ -161,7 +163,14 @@ var openAIAliasProviders = []openAIAliasProvider{ func TestOpenAIAliasProvider_EndToEndRequest(t *testing.T) { t.Parallel() + providers := slices.Clone(openAIAliasProviders) for _, p := range openAIAliasProviders { + if canonical := modelsdev.CanonicalProviderID(p.provider); canonical != p.provider { + p.provider = canonical + providers = append(providers, p) + } + } + for _, p := range providers { t.Run(p.provider, func(t *testing.T) { t.Parallel() diff --git a/pkg/modelsdev/store.go b/pkg/modelsdev/store.go index 415585524..6d9489993 100644 --- a/pkg/modelsdev/store.go +++ b/pkg/modelsdev/store.go @@ -254,6 +254,9 @@ func (s *Store) getProvider(ctx context.Context, providerID string) (*Provider, } provider, exists := db.Providers[providerID] + if !exists { + provider, exists = db.Providers[CanonicalProviderID(providerID)] + } if !exists { return nil, fmt.Errorf("provider %q not found", providerID) } @@ -261,6 +264,23 @@ func (s *Store) getProvider(ctx context.Context, providerID string) (*Provider, return &provider, nil } +// legacyProviderIDs contains shipped names for the same service in models.dev. +var legacyProviderIDs = map[string]string{ + "fireworks": "fireworks-ai", + "together": "togetherai", + "moonshot": "moonshotai", + "opencode-zen": "opencode", +} + +// CanonicalProviderID returns the models.dev ID for a built-in provider name. +// Unknown names and distinct services, including ChatGPT, are unchanged. +func CanonicalProviderID(providerID string) string { + if canonical, ok := legacyProviderIDs[providerID]; ok { + return canonical + } + return providerID +} + // GetModel returns a specific model by ID. The ID must carry both a // provider and a model component; pass the result of [NewID], [ParseID], // or a provider's [ID] method. @@ -269,27 +289,44 @@ func (s *Store) GetModel(ctx context.Context, id ID) (*Model, error) { return nil, fmt.Errorf("invalid model ID: %q", id.String()) } - provider, err := s.getProvider(ctx, id.Provider) + allowFetch := s.knownProvider == nil || s.knownProvider(id.Provider) + db, err := s.getDatabase(ctx, allowFetch) if err != nil { return nil, err } - model, exists := provider.Models[id.Model] + provider, exists := db.Providers[id.Provider] + if model, ok := lookupModel(provider, id); ok { + return &model, nil + } - // For amazon-bedrock, try stripping region/inference profile prefixes. - // Bedrock uses prefixes for cross-region inference profiles, - // but models.dev stores models without these prefixes. - if !exists && id.Provider == "amazon-bedrock" { - if prefix, after, ok := strings.Cut(id.Model, "."); ok && bedrockRegionPrefixes[prefix] { - model, exists = provider.Models[after] + if canonical := CanonicalProviderID(id.Provider); canonical != id.Provider { + aliasedProvider, aliasExists := db.Providers[canonical] + exists = exists || aliasExists + if model, ok := lookupModel(aliasedProvider, NewID(canonical, id.Model)); ok { + return &model, nil } } if !exists { - return nil, fmt.Errorf("model %q not found in provider %q", id.Model, id.Provider) + return nil, fmt.Errorf("provider %q not found", id.Provider) } + return nil, fmt.Errorf("model %q not found in provider %q", id.Model, id.Provider) +} - return &model, nil +func lookupModel(provider Provider, id ID) (Model, bool) { + model, exists := provider.Models[id.Model] + if !exists && id.Provider == "amazon-bedrock" { + // Cross-region inference profile prefixes are absent from the catalog. + if prefix, after, ok := strings.Cut(id.Model, "."); ok && bedrockRegionPrefixes[prefix] { + model, exists = provider.Models[after] + } + } + if !exists && id.Provider == "ovhcloud" { + // OVHcloud's API uses mixed-case IDs; its catalog uses lowercase IDs. + model, exists = provider.Models[strings.ToLower(id.Model)] + } + return model, exists } // loadDatabase loads the database from the local cache file or diff --git a/pkg/modelsdev/store_test.go b/pkg/modelsdev/store_test.go index f42958095..c59b8d77c 100644 --- a/pkg/modelsdev/store_test.go +++ b/pkg/modelsdev/store_test.go @@ -396,3 +396,213 @@ func TestDatePattern(t *testing.T) { }) } } + +func TestStore_GetModel_CatalogProviderAliases(t *testing.T) { + t.Parallel() + + catalogModel := Model{ + Name: "Vision model", Family: "vision", Reasoning: true, ToolCall: true, + Temperature: true, Attachment: true, OpenWeights: true, ReleaseDate: "2026-09-01", + Cost: &Cost{ + Input: 1.2, Output: 3.4, CacheRead: 0.1, CacheWrite: 0.2, + Tiers: []CostTier{{Rates: Rates{Input: 2.4, Output: 6.8}, Tier: TierSpec{Type: "context", Size: 100000}}}, + }, + Limit: Limit{Context: 262144, Output: 65536}, + Modalities: Modalities{Input: []string{"text", "image"}, Output: []string{"text"}}, + } + directModel := Model{Name: "Direct entry", Limit: Limit{Context: 1000, Output: 500}} + for _, alias := range []struct{ configured, catalog string }{ + {"fireworks", "fireworks-ai"}, + {"together", "togetherai"}, + {"moonshot", "moonshotai"}, + {"opencode-zen", "opencode"}, + } { + for _, mode := range []string{"catalog only", "direct entry wins", "direct provider missing model"} { + t.Run(alias.configured+"/"+mode, func(t *testing.T) { + t.Parallel() + providers := map[string]Provider{alias.catalog: {Models: map[string]Model{"MixedModel": catalogModel}}} + want := catalogModel + switch mode { + case "direct entry wins": + providers[alias.configured] = Provider{Models: map[string]Model{"MixedModel": directModel}} + want = directModel + case "direct provider missing model": + providers[alias.configured] = Provider{Models: map[string]Model{"other": directModel}} + } + store := NewDatabaseStore(&Database{Providers: providers}) + got, err := store.GetModel(t.Context(), NewID(alias.configured, "MixedModel")) + require.NoError(t, err) + assert.Equal(t, &want, got) + canonicalModel, err := store.GetModel(t.Context(), NewID(alias.catalog, "MixedModel")) + require.NoError(t, err) + assert.Equal(t, &catalogModel, canonicalModel) + _, err = store.GetModel(t.Context(), NewID(alias.configured, "missing")) + require.EqualError(t, err, `model "missing" not found in provider "`+alias.configured+`"`) + _, err = store.GetModel(t.Context(), NewID(alias.configured, "mixedmodel")) + require.EqualError(t, err, `model "mixedmodel" not found in provider "`+alias.configured+`"`) + }) + } + } +} + +func TestStore_GetModel_CatalogLookupBoundaries(t *testing.T) { + t.Parallel() + + lower := Model{Name: "Lowercase catalog model", Modalities: Modalities{Input: []string{"text", "image"}}} + direct := Model{Name: "Exact catalog model"} + for _, tc := range []struct { + name string + id ID + providers map[string]Provider + want *Model + wantErr string + }{ + { + name: "OVH lowercase fallback", id: NewID("ovhcloud", "Qwen3.5-397B-A17B"), + providers: map[string]Provider{"ovhcloud": {Models: map[string]Model{"qwen3.5-397b-a17b": lower}}}, want: &lower, + }, + { + name: "OVH direct precedence", id: NewID("ovhcloud", "Qwen3.5-397B-A17B"), + providers: map[string]Provider{"ovhcloud": {Models: map[string]Model{"qwen3.5-397b-a17b": lower, "Qwen3.5-397B-A17B": direct}}}, want: &direct, + }, + { + name: "OVH unknown preserves spelling", id: NewID("ovhcloud", "MissingModel"), + providers: map[string]Provider{"ovhcloud": {}}, wantErr: `model "MissingModel" not found in provider "ovhcloud"`, + }, + { + name: "other providers case sensitive", id: NewID("openai", "MixedModel"), + providers: map[string]Provider{"openai": {Models: map[string]Model{"mixedmodel": lower}}}, wantErr: `model "MixedModel" not found in provider "openai"`, + }, + { + name: "ChatGPT is not a full OpenAI alias", id: NewID("chatgpt", "vision"), + providers: map[string]Provider{"openai": {Models: map[string]Model{"vision": lower}}}, wantErr: `provider "chatgpt" not found`, + }, + { + name: "alias missing provider", id: NewID("fireworks", "MissingModel"), + providers: map[string]Provider{}, wantErr: `provider "fireworks" not found`, + }, + { + name: "custom provider missing", id: NewID("custom", "MissingModel"), + providers: map[string]Provider{}, wantErr: `provider "custom" not found`, + }, + { + name: "Bedrock direct precedence", id: NewID("amazon-bedrock", "us.model"), + providers: map[string]Provider{"amazon-bedrock": {Models: map[string]Model{"model": lower, "us.model": direct}}}, want: &direct, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got, err := NewDatabaseStore(&Database{Providers: tc.providers}).GetModel(t.Context(), tc.id) + if tc.wantErr != "" { + require.EqualError(t, err, tc.wantErr) + assert.Nil(t, got) + } else { + require.NoError(t, err) + assert.Equal(t, tc.want, got) + } + }) + } +} + +func TestStore_GetModel_CatalogAliasFetchPolicy(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + known bool + fail bool + cancel bool + missing bool + wantFetch int + }{ + {name: "known success", known: true, wantFetch: 1}, + {name: "known missing model", known: true, missing: true, wantFetch: 1}, + {name: "known fetch failure", known: true, fail: true, wantFetch: 1}, + {name: "known canceled fetch", known: true, cancel: true, wantFetch: 1}, + {name: "custom never fetches"}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + var calls int + var checked []string + model := "catalog-fetch-test" + store, err := NewStore( + WithCache(filepath.Join(t.TempDir(), CacheFileName)), + WithKnownProvider(func(provider string) bool { + checked = append(checked, provider) + return tc.known && provider == "moonshot" + }), + WithFetcher(func(ctx context.Context, _ string) (*Database, string, error) { + calls++ + if tc.cancel { + require.ErrorIs(t, ctx.Err(), context.Canceled) + return nil, "", ctx.Err() + } + if tc.fail { + return nil, "", errors.New("offline") + } + models := map[string]Model{} + if !tc.missing { + models[model] = Model{Name: "Fetched model"} + } + return &Database{Providers: map[string]Provider{"moonshotai": {Models: models}}}, "", nil + }), + ) + require.NoError(t, err) + ctx := t.Context() + if tc.cancel { + var cancel context.CancelFunc + ctx, cancel = context.WithCancel(ctx) + cancel() + } + provider := "moonshot" + if !tc.known { + provider = "custom" + } + got, err := store.GetModel(ctx, NewID(provider, model)) + if tc.known && !tc.fail && !tc.cancel && !tc.missing { + require.NoError(t, err) + assert.Equal(t, &Model{Name: "Fetched model"}, got) + } else { + require.Error(t, err) + assert.Contains(t, err.Error(), provider) + assert.Nil(t, got) + } + assert.Equal(t, tc.wantFetch, calls) + assert.Equal(t, []string{provider}, checked, "fetch policy uses only the configured provider") + }) + } +} + +func TestCanonicalProviderID(t *testing.T) { + t.Parallel() + for legacy, canonical := range legacyProviderIDs { + assert.Equal(t, canonical, CanonicalProviderID(legacy)) + assert.Equal(t, canonical, CanonicalProviderID(canonical)) + } + for _, id := range []string{"", "chatgpt", "openai", "custom", "FIREWORKS"} { + assert.Equal(t, id, CanonicalProviderID(id)) + } +} + +func TestStore_ResolveModelAlias_LegacyProviders(t *testing.T) { + t.Parallel() + for legacy, canonical := range legacyProviderIDs { + t.Run(legacy, func(t *testing.T) { + t.Parallel() + catalog := Provider{Models: map[string]Model{ + "latest": {Name: "Model (latest)"}, + "model-20260101": {Name: "Model"}, + }} + db := &Database{Providers: map[string]Provider{canonical: catalog}} + store := NewDatabaseStore(db) + assert.Equal(t, "model-20260101", store.ResolveModelAlias(t.Context(), legacy, "latest")) + assert.Equal(t, "model-20260101", store.ResolveModelAlias(t.Context(), canonical, "latest")) + assert.Equal(t, "unknown", store.ResolveModelAlias(t.Context(), legacy, "unknown")) + db.Providers[legacy] = Provider{Models: map[string]Model{"latest": {Name: "Direct model"}}} + assert.Equal(t, "latest", store.ResolveModelAlias(t.Context(), legacy, "latest"), "direct provider wins") + _, err := store.GetModel(t.Context(), NewID(canonical, "latest")) + require.NoError(t, err) + }) + } +} diff --git a/pkg/runtime/model_switcher.go b/pkg/runtime/model_switcher.go index 15c4e47f9..a1e0f74c8 100644 --- a/pkg/runtime/model_switcher.go +++ b/pkg/runtime/model_switcher.go @@ -744,7 +744,11 @@ func (r *LocalRuntime) buildCatalogChoices(ctx context.Context) []ModelChoice { for name, cfg := range r.modelSwitcherCfg.Models { existingRefs[name] = true if cfg.Provider != "" && cfg.Model != "" { - existingRefs[cfg.Provider+"/"+cfg.Model] = true + providerID := cfg.Provider + if _, custom := r.modelSwitcherCfg.Providers[providerID]; !custom { + providerID = modelsdev.CanonicalProviderID(providerID) + } + existingRefs[providerID+"/"+cfg.Model] = true } } diff --git a/pkg/runtime/model_switcher_test.go b/pkg/runtime/model_switcher_test.go index b36d739e1..2f900234b 100644 --- a/pkg/runtime/model_switcher_test.go +++ b/pkg/runtime/model_switcher_test.go @@ -3,6 +3,7 @@ package runtime import ( "context" "errors" + "fmt" "slices" "testing" @@ -1373,3 +1374,45 @@ func TestAgentThinkingConfigurationIsLocalAndConservative(t *testing.T) { }) } } + +func TestBuildCatalogChoices_CanonicalProviders(t *testing.T) { + t.Parallel() + for _, legacy := range []string{"fireworks", "together", "moonshot", "opencode-zen"} { + for _, custom := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/custom=%t", legacy, custom), func(t *testing.T) { + t.Parallel() + canonical := modelsdev.CanonicalProviderID(legacy) + alias, ok := provider.LookupAlias(canonical) + require.True(t, ok) + r := &LocalRuntime{ + modelsStore: &mockCatalogStore{db: &modelsdev.Database{Providers: map[string]modelsdev.Provider{ + canonical: {Models: map[string]modelsdev.Model{ + "configured": {Modalities: modelsdev.Modalities{Output: []string{"text"}}}, + "available": {Modalities: modelsdev.Modalities{Output: []string{"text"}}}, + }}, + }}}, + modelSwitcherCfg: &ModelSwitcherConfig{ + ProviderRegistry: testProviderRegistry(), + EnvProvider: environment.NewMapEnvProvider(map[string]string{alias.TokenEnvVar: "test-key"}), + Models: map[string]latest.ModelConfig{"mine": {Provider: legacy, Model: "configured"}}, + }, + } + if custom { + r.modelSwitcherCfg.Providers = map[string]latest.ProviderConfig{legacy: {BaseURL: "https://custom.invalid/v1"}} + } + choices := r.buildCatalogChoices(t.Context()) + wantLen := 1 + if custom { + wantLen = 2 + } + require.Len(t, choices, wantLen) + for _, choice := range choices { + assert.Equal(t, canonical, choice.Provider) + if !custom { + assert.Equal(t, canonical+"/available", choice.Ref) + } + } + }) + } + } +} diff --git a/pkg/runtime/transforms_catalog_alias_test.go b/pkg/runtime/transforms_catalog_alias_test.go new file mode 100644 index 000000000..12ea17c1b --- /dev/null +++ b/pkg/runtime/transforms_catalog_alias_test.go @@ -0,0 +1,114 @@ +package runtime + +import ( + "slices" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/agent" + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/model/provider/base" + "github.com/docker/docker-agent/pkg/modelsdev" + "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/team" + "github.com/docker/docker-agent/pkg/tools" +) + +func TestRunStream_CatalogAliasesFilterImages(t *testing.T) { + t.Parallel() + + for _, provider := range []struct{ configured, catalog, model string }{ + {"fireworks", "fireworks-ai", "accounts/fireworks/models/kimi-k3"}, + {"together", "togetherai", "Qwen/Qwen3.5-397B-A17B"}, + {"moonshot", "moonshotai", "kimi-k3"}, + {"opencode-zen", "opencode", "kimi-k3"}, + {"ovhcloud", "ovhcloud", "Qwen3.5-397B-A17B"}, + {"fireworks-ai", "fireworks-ai", "accounts/fireworks/models/kimi-k3"}, + {"togetherai", "togetherai", "Qwen/Qwen3.5-397B-A17B"}, + {"moonshotai", "moonshotai", "kimi-k3"}, + {"opencode", "opencode", "kimi-k3"}, + } { + for _, mode := range []string{"snapshot vision", "text only", "unknown", "override disables", "override enables", "direct text wins"} { + t.Run(provider.configured+"/"+mode, func(t *testing.T) { + t.Parallel() + model := provider.model + store := modelsdev.NewDatabaseStore(modelsdev.EmbeddedSnapshot()) + var override *latest.CapabilitiesConfig + wantImages := false + switch mode { + case "snapshot vision": + wantImages = true + case "text only", "override enables": + model = "text-only" + store = modelsdev.NewDatabaseStore(&modelsdev.Database{Providers: map[string]modelsdev.Provider{ + provider.catalog: {Models: map[string]modelsdev.Model{model: {Modalities: modelsdev.Modalities{Input: []string{"text"}}}}}, + }}) + if mode == "override enables" { + override = &latest.CapabilitiesConfig{Image: true} + wantImages = true + } + case "unknown": + model = "unknown-catalog-model" + case "override disables": + override = &latest.CapabilitiesConfig{Image: false} + case "direct text wins": + catalogModel := model + if provider.configured == "ovhcloud" { + catalogModel = "qwen3.5-397b-a17b" + } + providers := map[string]modelsdev.Provider{ + provider.catalog: {Models: map[string]modelsdev.Model{catalogModel: {Modalities: modelsdev.Modalities{Input: []string{"text", "image"}}}}}, + } + direct := providers[provider.configured] + if direct.Models == nil { + direct.Models = map[string]modelsdev.Model{} + } + direct.Models[model] = modelsdev.Model{Modalities: modelsdev.Modalities{Input: []string{"text"}}} + providers[provider.configured] = direct + store = modelsdev.NewDatabaseStore(&modelsdev.Database{Providers: providers}) + } + prov := &recordingMsgProvider{ + mockProvider: mockProvider{id: provider.configured + "/" + model, stream: &mockStream{}}, + baseConfig: base.Config{ModelConfig: latest.ModelConfig{Capabilities: override}}, + } + a := agent.New("root", "instructions", agent.WithModel(prov)) + tm := team.New(team.WithAgents(a)) + lazy := &lazyModelStore{} + lazy.once.Do(func() { lazy.st = store }) + rt, err := NewLocalRuntime(t.Context(), tm, WithModelStore(lazy)) + require.NoError(t, err) + defer rt.Close() + + user := mixedMediaMsg() + user.MultiContent = append(user.MultiContent, chat.MessagePart{Type: chat.MessagePartTypeDocument, Document: &chat.Document{ + Name: "attachment.png", MimeType: "image/png", Source: chat.DocumentSource{InlineData: []byte{1}}, + }}) + tool := chat.Message{Role: chat.MessageRoleTool, ToolCallID: "call-image", MultiContent: slices.Clone(user.MultiContent)} + assistant := chat.Message{Role: chat.MessageRoleAssistant, ToolCalls: []tools.ToolCall{{ + ID: "call-image", Type: "function", Function: tools.FunctionCall{Name: "read_file", Arguments: `{}`}, + }}} + sess := session.New(session.WithMessages([]session.Item{ + session.NewMessageItem(&session.Message{Message: user}), + session.NewMessageItem(&session.Message{Message: assistant}), + session.NewMessageItem(&session.Message{Message: tool}), + })) + for range rt.RunStream(t.Context(), sess) { + } + require.NotEmpty(t, prov.got) + for _, role := range []chat.MessageRole{chat.MessageRoleUser, chat.MessageRoleTool} { + idx := slices.IndexFunc(prov.got[0], func(m chat.Message) bool { return m.Role == role }) + require.NotEqual(t, -1, idx) + parts := prov.got[0][idx].MultiContent + assert.Equal(t, wantImages, slices.ContainsFunc(parts, func(p chat.MessagePart) bool { return p.Type == chat.MessagePartTypeImageURL })) + assert.Equal(t, wantImages, slices.ContainsFunc(parts, func(p chat.MessagePart) bool { + return p.Document != nil && p.Document.MimeType == "image/png" + })) + assert.True(t, slices.ContainsFunc(parts, func(p chat.MessagePart) bool { return p.Type == chat.MessagePartTypeText })) + } + }) + } + } +}