diff --git a/.env.template b/.env.template index ec4d28a3b..8c085895a 100644 --- a/.env.template +++ b/.env.template @@ -106,6 +106,11 @@ # MODEL_OVERRIDES_ENABLED=true # Hide provider models from GET /v1/models and expose only enabled aliases (default: false). # KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT=false +# How providers..models and [_SUFFIX]_MODELS affect provider inventory. +# fallback (default): use configured models only when upstream /models fails, is nil, or is empty. +# allowlist: expose only the configured models for providers that define a list, and skip their upstream /models calls. +# CONFIGURED_PROVIDER_MODELS_MODE=fallback +# Examples: OPENROUTER_MODELS=..., OPENROUTER_EU_MODELS=..., AZURE_MODELS=..., VLLM_MODELS=... # Fallback & Workflow Configuration # Default translated-route fallback mode: auto, manual, or off (default: auto) @@ -261,6 +266,8 @@ # OpenRouter (default base URL: https://openrouter.ai/api/v1) # OPENROUTER_API_KEY=sk-or-... # OPENROUTER_BASE_URL=https://openrouter.ai/api/v1 +# Optional configured model list; see CONFIGURED_PROVIDER_MODELS_MODE below +# OPENROUTER_MODELS=openai/gpt-oss-120b,anthropic/claude-sonnet-4 # OPENROUTER_SITE_URL=https://gomodel.enterpilot.io # OPENROUTER_APP_NAME=GoModel @@ -281,8 +288,7 @@ # Oracle # ORACLE_API_KEY=... # ORACLE_BASE_URL=https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/v1 -# Optional fallback model inventory when Oracle's /models endpoint is unavailable -# Comma-separated; whitespace around entries is ignored +# Optional configured model list; comma-separated, whitespace around entries is ignored # ORACLE_MODELS=openai.gpt-oss-120b,xai.grok-3 # Ollama (local LLM server) diff --git a/CLAUDE.md b/CLAUDE.md index 94cddfb34..cf9bf41ad 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -107,7 +107,7 @@ Full reference: `.env.template` and `config/config.yaml` - `ALLOW_PASSTHROUGH_V1_ALIAS` (true: Allow /p/{provider}/v1/... aliases while keeping /p/{provider}/... canonical) - `ENABLED_PASSTHROUGH_PROVIDERS` (openai,anthropic,openrouter,zai,vllm: Comma-separated list of enabled passthrough providers) - **Storage:** `STORAGE_TYPE` (sqlite), `SQLITE_PATH` (data/gomodel.db), `POSTGRES_URL`, `MONGODB_URL` -- **Models:** `MODELS_ENABLED_BY_DEFAULT` (true), `MODEL_OVERRIDES_ENABLED` (false), `KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT` (false); persisted overrides restrict/allow selectors with `user_paths`. When alias-only models listing is enabled, `GET /v1/models` returns only model aliases, not full concrete model specs, to operators. +- **Models:** `MODELS_ENABLED_BY_DEFAULT` (true), `MODEL_OVERRIDES_ENABLED` (true), `KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT` (false), `CONFIGURED_PROVIDER_MODELS_MODE` (`fallback` or `allowlist`, default `fallback`; `allowlist` skips upstream `/models` for providers with configured lists); persisted overrides restrict/allow selectors with `user_paths`. When alias-only models listing is enabled, `GET /v1/models` returns only model aliases, not full concrete model specs, to operators. - **Audit logging:** `LOGGING_ENABLED` (false), `LOGGING_LOG_BODIES` (false), `LOGGING_LOG_HEADERS` (false), `LOGGING_RETENTION_DAYS` (30) - **Usage tracking:** `USAGE_ENABLED` (true), `ENFORCE_RETURNING_USAGE_DATA` (true), `USAGE_RETENTION_DAYS` (90) - **Cache:** `CACHE_REFRESH_INTERVAL` (3600s), `REDIS_URL`, `REDIS_KEY_MODELS`, `REDIS_TTL_MODELS`. Exact response cache uses `cache.response.simple` in `config.yaml` (optional `enabled`); `REDIS_KEY_RESPONSES`, `REDIS_TTL_RESPONSES`, and `REDIS_URL` apply only when that block exists or when `RESPONSE_CACHE_SIMPLE_ENABLED=true`. Semantic response cache uses `cache.response.semantic` (optional `enabled`); when enabled, `embedder.provider` must name a key in the top-level `providers` map (no default embedder). At runtime that key is resolved against the same env-merged, credential-filtered provider set as routing (not YAML-only), so env-only credentials apply. `vector_store.type` must be set explicitly to one of `qdrant`, `pgvector`, `pinecone`, `weaviate` (each has its own nested config and `SEMANTIC_CACHE_*` env vars). Tuning via `SEMANTIC_CACHE_*` applies when the semantic block exists or `SEMANTIC_CACHE_ENABLED=true`. @@ -115,5 +115,5 @@ Full reference: `.env.template` and `config/config.yaml` - **Resilience:** Configured via `config/config.yaml` - global `resilience.retry.*` and `resilience.circuit_breaker.*` defaults with optional per-provider overrides under `providers..resilience.retry.*` and `providers..resilience.circuit_breaker.*`. Retry defaults: `max_retries` (3), `initial_backoff` (1s), `max_backoff` (30s), `backoff_factor` (2.0), `jitter_factor` (0.1). Circuit breaker defaults: `failure_threshold` (5), `success_threshold` (2), `timeout` (30s) - **Metrics:** `METRICS_ENABLED` (false), `METRICS_ENDPOINT` (/metrics) - **Guardrails:** Configured via `config/config.yaml` only (except `GUARDRAILS_ENABLED` env var) -- **Providers:** `OPENAI_API_KEY`, `ANTHROPIC_API_KEY`, `GEMINI_API_KEY`, `XAI_API_KEY`, `GROQ_API_KEY`, `OPENROUTER_API_KEY`, `ZAI_API_KEY`, `ZAI_BASE_URL` (optional Z.ai endpoint override), `MINIMAX_API_KEY`, `MINIMAX_BASE_URL` (optional MiniMax endpoint override), `AZURE_API_KEY`, `AZURE_BASE_URL` (Azure OpenAI deployment base URL), `AZURE_API_VERSION` (optional Azure API version), `ORACLE_API_KEY` (Oracle API key), `ORACLE_BASE_URL` (Oracle OpenAI-compatible base URL), `ORACLE_MODELS` (comma-separated Oracle fallback model inventory), `OLLAMA_BASE_URL`, `VLLM_BASE_URL`, `VLLM_API_KEY` (optional upstream vLLM bearer token) +- **Providers:** `OPENAI_API_KEY`, `ANTHROPIC_API_KEY`, `GEMINI_API_KEY`, `XAI_API_KEY`, `GROQ_API_KEY`, `OPENROUTER_API_KEY`, `ZAI_API_KEY`, `ZAI_BASE_URL` (optional Z.ai endpoint override), `MINIMAX_API_KEY`, `MINIMAX_BASE_URL` (optional MiniMax endpoint override), `AZURE_API_KEY`, `AZURE_BASE_URL` (Azure OpenAI deployment base URL), `AZURE_API_VERSION` (optional Azure API version), `ORACLE_API_KEY` (Oracle API key), `ORACLE_BASE_URL` (Oracle OpenAI-compatible base URL), `[_SUFFIX]_MODELS` (comma-separated configured model list for any provider type), `OLLAMA_BASE_URL`, `VLLM_BASE_URL`, `VLLM_API_KEY` (optional upstream vLLM bearer token) - **Provider model metadata:** `providers..models` accepts either model IDs (strings) or `{id, metadata}` objects. When `metadata` is supplied (`display_name`, `context_window`, `max_output_tokens`, `modes`, `capabilities`, `pricing`, …) it is merged onto the remote ai-model-list entry during enrichment, with operator values winning per-field. Primary use case: advertising context windows, capabilities, and pricing for local models (Ollama) and other custom endpoints whose IDs are not in the upstream registry. diff --git a/GETTING_STARTED.md b/GETTING_STARTED.md index cfab1bce4..9dfef7822 100644 --- a/GETTING_STARTED.md +++ b/GETTING_STARTED.md @@ -199,11 +199,19 @@ Provider credentials: | `AZURE_API_VERSION` | Azure OpenAI API version override (default: `2024-10-21`) | | `ORACLE_API_KEY` | Oracle | | `ORACLE_BASE_URL` | Oracle OpenAI-compatible base URL | -| `ORACLE_MODELS` | Oracle fallback model inventory (comma-separated, used when `/models` is unavailable) | | `OLLAMA_BASE_URL` | Ollama (default: `http://localhost:11434/v1`) | | `VLLM_BASE_URL` | vLLM OpenAI-compatible server (default: `http://localhost:8000/v1`) | | `VLLM_API_KEY` | vLLM bearer token, only when upstream vLLM was started with `--api-key` | +Model configuration: + +| Variable | Description | +| --------------------------------- | -------------------------------------------------------------------------------------------- | +| `_MODELS` | Optional configured model list for any provider type, for example `OPENROUTER_MODELS` | +| `CONFIGURED_PROVIDER_MODELS_MODE` | `fallback` by default; set `allowlist` to expose only configured models and skip upstream `/models` for configured lists | + +Configured model lists work for every provider via YAML `providers..models` or env vars like `OPENROUTER_MODELS`, `ORACLE_MODELS`, `AZURE_MODELS`, or `VLLM_MODELS`. In the default `fallback` mode, the list is used only when upstream `/models` fails, returns nil, or returns an empty list. In `allowlist` mode, providers with configured lists expose only those models and skip their upstream `/models` calls. + See `.env.template` for the full list of all configurable environment variables. --- @@ -227,7 +235,6 @@ Ollama requires no API key. Even with no YAML and no `OLLAMA_BASE_URL` set, an O **Oracle requires both key and base URL.** `ORACLE_API_KEY` alone is not enough for auto-discovery. Set `ORACLE_BASE_URL` to the Oracle OpenAI-compatible endpoint, otherwise the provider is ignored. -If your Oracle endpoint does not return a usable model list, set `ORACLE_MODELS` or configure `providers..models` in YAML to seed the router with explicit model IDs. **Azure ships with a pinned API version by default.** If you do not set `AZURE_API_VERSION`, the gateway sends `api-version=2024-10-21`. Override it only when you need a different Azure API version. diff --git a/README.md b/README.md index 0ab5fc1a6..36d0c457e 100644 --- a/README.md +++ b/README.md @@ -87,14 +87,20 @@ Example model identifiers are illustrative and subject to change; consult provid ✅ Supported ❌ Unsupported For Z.ai's GLM Coding Plan, set `ZAI_BASE_URL=https://api.z.ai/api/coding/paas/v4`. -For Oracle, set `ORACLE_MODELS=openai.gpt-oss-120b,xai.grok-3` when the -upstream `/models` endpoint is unavailable. +Configured model lists are available for every provider with +`_MODELS`, for example +`OPENROUTER_MODELS=openai/gpt-oss-120b,anthropic/claude-sonnet-4` or +`ORACLE_MODELS=openai.gpt-oss-120b,xai.grok-3`. By default, +`CONFIGURED_PROVIDER_MODELS_MODE=fallback` uses those lists only when upstream +`/models` is unavailable or empty. Set `CONFIGURED_PROVIDER_MODELS_MODE=allowlist` +to expose only configured models for providers that define a list, skipping +their upstream `/models` calls. For vLLM, set `VLLM_API_KEY` only if the upstream server was started with `--api-key`. To register multiple instances of the same provider type without `config.yaml`, use suffixed env vars such as `OPENAI_EAST_API_KEY` and -`OPENAI_EAST_BASE_URL`; this registers provider `openai-east` with type -`openai`. +`OPENAI_EAST_BASE_URL`; add `OPENAI_EAST_MODELS` to configure that instance's +model list. This registers provider `openai-east` with type `openai`. --- diff --git a/config/config.example.yaml b/config/config.example.yaml index 5298f3d2a..e6e6a3121 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -15,6 +15,7 @@ server: models: enabled_by_default: true # env: MODELS_ENABLED_BY_DEFAULT; when false, models stay unavailable until an override allows one or more user paths overrides_enabled: true # env: MODEL_OVERRIDES_ENABLED; load/enforce persisted model overrides and enable dashboard editing + configured_provider_models_mode: "fallback" # env: CONFIGURED_PROVIDER_MODELS_MODE; "fallback" uses configured lists only when upstream /models is unavailable/empty, "allowlist" exposes only configured models and skips upstream /models for configured lists cache: model: @@ -219,11 +220,18 @@ providers: # base_url: "https://api.groq.com/openai/v1" # api_key: "${GROQ_API_KEY}" - # Example: OpenRouter + # Example: OpenRouter with an explicit configured model list. + # In fallback mode (default), this list is used only if upstream /models is + # unavailable or empty. In allowlist mode, only these models are exposed and + # upstream /models is skipped for this provider. + # You can also set OPENROUTER_MODELS="openai/gpt-oss-120b,anthropic/claude-sonnet-4". # openrouter: # type: "openrouter" # base_url: "https://openrouter.ai/api/v1" # api_key: "${OPENROUTER_API_KEY}" + # models: + # - openai/gpt-oss-120b + # - anthropic/claude-sonnet-4 # Example: Azure OpenAI # azure: diff --git a/config/config.go b/config/config.go index bcca8131c..e762bf237 100644 --- a/config/config.go +++ b/config/config.go @@ -145,6 +145,44 @@ type ModelsConfig struct { // provider models and returns only alias-projected model entries. // Default: false. KeepOnlyAliasesAtModelsEndpoint bool `yaml:"keep_only_aliases_at_models_endpoint" env:"KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT"` + + // ConfiguredProviderModelsMode controls how providers..models and + // provider *_MODELS env vars affect the provider model inventory. + // Supported values: "fallback", "allowlist". Default: "fallback". + ConfiguredProviderModelsMode ConfiguredProviderModelsMode `yaml:"configured_provider_models_mode" env:"CONFIGURED_PROVIDER_MODELS_MODE"` +} + +// ConfiguredProviderModelsMode controls how explicitly configured provider +// model lists are applied to the discovered model inventory. +type ConfiguredProviderModelsMode string + +const ( + ConfiguredProviderModelsModeFallback ConfiguredProviderModelsMode = "fallback" + ConfiguredProviderModelsModeAllowlist ConfiguredProviderModelsMode = "allowlist" +) + +// Valid reports whether mode is one of the supported configured-provider-models modes. +func (m ConfiguredProviderModelsMode) Valid() bool { + switch NormalizeConfiguredProviderModelsMode(m) { + case ConfiguredProviderModelsModeFallback, ConfiguredProviderModelsModeAllowlist: + return true + default: + return false + } +} + +// NormalizeConfiguredProviderModelsMode canonicalizes a configured provider models mode. +func NormalizeConfiguredProviderModelsMode(mode ConfiguredProviderModelsMode) ConfiguredProviderModelsMode { + return ConfiguredProviderModelsMode(strings.ToLower(strings.TrimSpace(string(mode)))) +} + +// ResolveConfiguredProviderModelsMode canonicalizes mode and applies the process default. +func ResolveConfiguredProviderModelsMode(mode ConfiguredProviderModelsMode) ConfiguredProviderModelsMode { + mode = NormalizeConfiguredProviderModelsMode(mode) + if mode == "" { + return ConfiguredProviderModelsModeFallback + } + return mode } // FallbackConfig holds translated-route model fallback policy. @@ -890,6 +928,7 @@ func buildDefaultConfig() *Config { EnabledByDefault: true, OverridesEnabled: true, KeepOnlyAliasesAtModelsEndpoint: false, + ConfiguredProviderModelsMode: ConfiguredProviderModelsModeFallback, }, Cache: CacheConfig{ Model: ModelCacheConfig{ @@ -977,6 +1016,10 @@ func Load() (*LoadResult, error) { if err := applyEnvOverrides(cfg); err != nil { return nil, err } + cfg.Models.ConfiguredProviderModelsMode = ResolveConfiguredProviderModelsMode(cfg.Models.ConfiguredProviderModelsMode) + if !cfg.Models.ConfiguredProviderModelsMode.Valid() { + return nil, fmt.Errorf("models.configured_provider_models_mode must be one of: fallback, allowlist") + } if err := loadFallbackConfig(&cfg.Fallback); err != nil { return nil, err diff --git a/config/config_test.go b/config/config_test.go index f955698f1..a96829246 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -15,17 +15,17 @@ import ( func clearProviderEnvVars(t *testing.T) { t.Helper() for _, key := range []string{ - "OPENAI_API_KEY", "OPENAI_BASE_URL", - "ANTHROPIC_API_KEY", "ANTHROPIC_BASE_URL", - "GEMINI_API_KEY", "GEMINI_BASE_URL", - "XAI_API_KEY", "XAI_BASE_URL", - "GROQ_API_KEY", "GROQ_BASE_URL", - "OPENROUTER_API_KEY", "OPENROUTER_BASE_URL", "OPENROUTER_SITE_URL", "OPENROUTER_APP_NAME", - "ZAI_API_KEY", "ZAI_BASE_URL", - "AZURE_API_KEY", "AZURE_BASE_URL", "AZURE_API_VERSION", - "ORACLE_API_KEY", "ORACLE_BASE_URL", + "OPENAI_API_KEY", "OPENAI_BASE_URL", "OPENAI_MODELS", + "ANTHROPIC_API_KEY", "ANTHROPIC_BASE_URL", "ANTHROPIC_MODELS", + "GEMINI_API_KEY", "GEMINI_BASE_URL", "GEMINI_MODELS", + "XAI_API_KEY", "XAI_BASE_URL", "XAI_MODELS", + "GROQ_API_KEY", "GROQ_BASE_URL", "GROQ_MODELS", + "OPENROUTER_API_KEY", "OPENROUTER_BASE_URL", "OPENROUTER_MODELS", "OPENROUTER_SITE_URL", "OPENROUTER_APP_NAME", + "ZAI_API_KEY", "ZAI_BASE_URL", "ZAI_MODELS", + "AZURE_API_KEY", "AZURE_BASE_URL", "AZURE_API_VERSION", "AZURE_MODELS", + "ORACLE_API_KEY", "ORACLE_BASE_URL", "ORACLE_MODELS", "VLLM_API_KEY", "VLLM_BASE_URL", "VLLM_MODELS", - "OLLAMA_API_KEY", "OLLAMA_BASE_URL", + "OLLAMA_API_KEY", "OLLAMA_BASE_URL", "OLLAMA_MODELS", } { t.Setenv(key, "") os.Unsetenv(key) @@ -57,7 +57,7 @@ func clearAllConfigEnvVars(t *testing.T) { "USAGE_BUFFER_SIZE", "USAGE_FLUSH_INTERVAL", "USAGE_RETENTION_DAYS", "GUARDRAILS_ENABLED", "ENABLE_GUARDRAILS_FOR_BATCH_PROCESSING", "FEATURE_FALLBACK_MODE", "FALLBACK_MANUAL_RULES_PATH", - "MODEL_OVERRIDES_ENABLED", "MODELS_ENABLED_BY_DEFAULT", "KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT", + "MODEL_OVERRIDES_ENABLED", "MODELS_ENABLED_BY_DEFAULT", "KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT", "CONFIGURED_PROVIDER_MODELS_MODE", "HTTP_TIMEOUT", "HTTP_RESPONSE_HEADER_TIMEOUT", "WORKFLOW_REFRESH_INTERVAL", } { @@ -100,6 +100,9 @@ func TestBuildDefaultConfig(t *testing.T) { if got, want := cfg.Server.EnabledPassthroughProviders, []string{"openai", "anthropic", "openrouter", "zai", "vllm"}; !reflect.DeepEqual(got, want) { t.Errorf("expected Server.EnabledPassthroughProviders=%v, got %v", want, got) } + if cfg.Models.ConfiguredProviderModelsMode != ConfiguredProviderModelsModeFallback { + t.Errorf("expected Models.ConfiguredProviderModelsMode=fallback, got %q", cfg.Models.ConfiguredProviderModelsMode) + } if cfg.Cache.Model.Local != nil { t.Error("expected Cache.Model.Local to be nil in raw defaults") } @@ -232,6 +235,7 @@ models: enabled_by_default: false overrides_enabled: false keep_only_aliases_at_models_endpoint: true + configured_provider_models_mode: allowlist cache: model: redis: @@ -268,6 +272,9 @@ logging: if !cfg.Models.KeepOnlyAliasesAtModelsEndpoint { t.Error("expected Models.KeepOnlyAliasesAtModelsEndpoint=true from YAML") } + if cfg.Models.ConfiguredProviderModelsMode != ConfiguredProviderModelsModeAllowlist { + t.Errorf("expected Models.ConfiguredProviderModelsMode=allowlist from YAML, got %q", cfg.Models.ConfiguredProviderModelsMode) + } if cfg.Cache.Model.Redis == nil { t.Fatal("expected Cache.Model.Redis to be set") } @@ -374,6 +381,28 @@ fallback: }) } +func TestLoad_InvalidConfiguredProviderModelsMode(t *testing.T) { + clearAllConfigEnvVars(t) + + withTempDir(t, func(dir string) { + yaml := ` +models: + configured_provider_models_mode: strict +` + if err := os.WriteFile(filepath.Join(dir, "config.yaml"), []byte(yaml), 0644); err != nil { + t.Fatalf("Failed to write config.yaml: %v", err) + } + + _, err := Load() + if err == nil { + t.Fatal("expected Load() to fail for invalid configured provider models mode") + } + if !strings.Contains(err.Error(), "models.configured_provider_models_mode must be one of") { + t.Fatalf("Load() error = %v, want configured provider models mode validation error", err) + } + }) +} + func TestLoad_EmptyFallbackOverrideMode(t *testing.T) { clearAllConfigEnvVars(t) @@ -848,6 +877,7 @@ func TestLoad_EnvOverridesDefaults(t *testing.T) { t.Setenv("MODEL_OVERRIDES_ENABLED", "false") t.Setenv("MODELS_ENABLED_BY_DEFAULT", "false") t.Setenv("KEEP_ONLY_ALIASES_AT_MODELS_ENDPOINT", "true") + t.Setenv("CONFIGURED_PROVIDER_MODELS_MODE", "allowlist") t.Setenv("STORAGE_TYPE", "postgresql") t.Setenv("POSTGRES_URL", "postgres://localhost/test") t.Setenv("POSTGRES_MAX_CONNS", "20") @@ -870,6 +900,9 @@ func TestLoad_EnvOverridesDefaults(t *testing.T) { if !cfg.Models.KeepOnlyAliasesAtModelsEndpoint { t.Error("expected aliases-only models endpoint from env") } + if cfg.Models.ConfiguredProviderModelsMode != ConfiguredProviderModelsModeAllowlist { + t.Errorf("expected configured provider models mode allowlist from env, got %q", cfg.Models.ConfiguredProviderModelsMode) + } if cfg.Storage.Type != "postgresql" { t.Errorf("expected storage type postgresql, got %s", cfg.Storage.Type) } diff --git a/docs/advanced/config-yaml.mdx b/docs/advanced/config-yaml.mdx index e028aca54..5a4b22a5e 100644 --- a/docs/advanced/config-yaml.mdx +++ b/docs/advanced/config-yaml.mdx @@ -14,6 +14,7 @@ especially: - per-provider resilience overrides - custom provider instance names that do not fit the generated `-` env naming +- richer reviewable provider model lists, especially when using allowlist mode - larger nested config that is easier to review in one file For multiple provider instances, env vars support @@ -23,9 +24,12 @@ For multiple provider instances, env vars support `OLLAMA_A_BASE_URL` registers `ollama-a`. Azure also supports `__API_VERSION`. -For Oracle specifically, a single fallback model list can now stay in env via -`ORACLE_MODELS`. Use suffixed env vars such as `ORACLE_US_MODELS` for multiple -Oracle instances without YAML. +Configured provider model lists can stay in env via `_MODELS`, for +example `OPENROUTER_MODELS`, `ORACLE_MODELS`, `AZURE_MODELS`, or `VLLM_MODELS`. +Set `CONFIGURED_PROVIDER_MODELS_MODE=fallback` (default) to use those lists only +when upstream `/models` fails or is empty, or `allowlist` to expose only the +configured models for providers that define a list and skip their upstream +`/models` calls. ## Priority Order diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index fba4ab5a2..be4083699 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -144,7 +144,16 @@ Set these to automatically register providers. No YAML configuration required. | `VLLM_BASE_URL` | vLLM (no API key needed unless upstream requires) | Most providers can use a custom base URL via `_BASE_URL` (for example `OPENAI_BASE_URL`). OpenRouter defaults to `https://openrouter.ai/api/v1` and can be overridden with `OPENROUTER_BASE_URL`. Z.ai defaults to `https://api.z.ai/api/paas/v4`; set `ZAI_BASE_URL=https://api.z.ai/api/coding/paas/v4` for the GLM Coding Plan endpoint. vLLM defaults to `http://localhost:8000/v1` when `VLLM_API_KEY` is set, but keyless deployments should set `VLLM_BASE_URL` explicitly to register the provider. Azure uses `AZURE_BASE_URL` for its deployment base URL and accepts an optional `AZURE_API_VERSION` override; otherwise it defaults to `2024-10-21`. Oracle requires `ORACLE_BASE_URL` because its OpenAI-compatible endpoint is region-specific. -When using Oracle, set `ORACLE_MODELS` to a comma-separated list such as `openai.gpt-oss-120b,xai.grok-3` if the upstream endpoint does not expose a usable `/models` response. YAML `models:` remains available for custom provider names and larger Oracle provider blocks. + +Every provider type also accepts a comma-separated configured model list via +`_MODELS`, for example `OPENROUTER_MODELS`, `ORACLE_MODELS`, +`AZURE_MODELS`, or `VLLM_MODELS`. By default, +`CONFIGURED_PROVIDER_MODELS_MODE=fallback` uses configured lists only when +upstream `/models` fails, returns nil, or returns an empty list. Set +`CONFIGURED_PROVIDER_MODELS_MODE=allowlist` to expose only configured models for +providers that define a list and skip their upstream `/models` calls. YAML +`providers..models` provides the same model-list input for named provider +blocks. For OpenRouter, GoModel also sends default attribution headers unless the request already sets them. Override those defaults with `OPENROUTER_SITE_URL` and `OPENROUTER_APP_NAME`. @@ -252,7 +261,9 @@ export AZURE_API_KEY="..." # Registers "azure" provider when paired wi export AZURE_BASE_URL="https://your-resource.openai.azure.com/openai/deployments/your-deployment" export ORACLE_API_KEY="..." # Registers "oracle" provider when paired with ORACLE_BASE_URL export ORACLE_BASE_URL="https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/v1" -export ORACLE_MODELS="openai.gpt-oss-120b,xai.grok-3" # Optional fallback inventory for Oracle +export ORACLE_MODELS="openai.gpt-oss-120b,xai.grok-3" # Optional configured model list +export OPENROUTER_MODELS="openai/gpt-oss-120b,anthropic/claude-sonnet-4" +export CONFIGURED_PROVIDER_MODELS_MODE="fallback" # fallback or allowlist export OLLAMA_BASE_URL="http://localhost:11434/v1" # Registers "ollama" provider export VLLM_BASE_URL="http://localhost:8000/v1" # Registers keyless "vllm" provider # Optional: export VLLM_API_KEY="token-abc123" @@ -283,6 +294,11 @@ For more control (custom names, per-provider resilience, or larger structured settings), use the YAML file: ```yaml +models: + # fallback is the default. Use allowlist when configured provider model lists + # should hide upstream models and skip upstream /models calls. + configured_provider_models_mode: fallback + providers: # Override OpenAI base URL openai: @@ -313,7 +329,7 @@ providers: # api_key is optional; set it only when vllm serve uses --api-key. # api_key: "token-abc123" - # Restrict to specific models + # Configure a model list for fallback or allowlist mode gemini: type: gemini api_key: "..." @@ -323,15 +339,10 @@ providers: ``` - For Oracle, give GoModel a fallback inventory with `ORACLE_MODELS` or - `models:`. `ORACLE_MODELS` is enough for the default single-provider setup; - use suffixed env vars such as `ORACLE_US_MODELS` for env-only multi-provider - setups. Use YAML when you need custom names or larger provider blocks. - See the [Oracle guide](/guides/oracle) for the required OCI policy and a - tested configuration. Automatic model discovery is not yet a reliable, - validated path for this provider: GoModel can try Oracle's OpenAI-compatible - `/models` endpoint, but Oracle may not return a usable inventory there. - OCI-native Oracle model discovery is not integrated yet. + `models:` works for every provider block. In fallback mode it is a safety net + when upstream `/models` is unavailable or empty. In allowlist mode it becomes + the exposed inventory for that provider and skips upstream `/models`. For Oracle, see the [Oracle + guide](/guides/oracle) for the required OCI policy and a tested configuration. ### Ollama (Local Models) diff --git a/docs/guides/oracle.mdx b/docs/guides/oracle.mdx index bfeb1d95c..078d228ce 100644 --- a/docs/guides/oracle.mdx +++ b/docs/guides/oracle.mdx @@ -1,6 +1,6 @@ --- title: "GoModel & Oracle" -description: "Configure Oracle's OpenAI-compatible Generative AI endpoint in GoModel, including the required OCI policy and model fallback." +description: "Configure Oracle's OpenAI-compatible Generative AI endpoint in GoModel, including the required OCI policy and configured model lists." icon: "cloud" --- @@ -16,7 +16,8 @@ Flow: - Create an Oracle Generative AI API key. - Add an OCI IAM policy for `generativeaiapikey`. - Choose a supported Oracle region and model. -- Decide whether you want env-only `ORACLE_MODELS` or YAML `models:`. +- Decide whether you want env-only `ORACLE_MODELS` or YAML `models:`, and + whether configured lists should stay in fallback mode or act as an allowlist. ## 1. Add the OCI policy @@ -49,8 +50,10 @@ export ORACLE_MODELS="openai.gpt-oss-120b,xai.grok-3" ``` `ORACLE_MODELS` is a comma-separated list. GoModel trims whitespace around each -entry and uses the list as the fallback inventory when Oracle's `/models` -endpoint is unavailable. +entry. With the default `CONFIGURED_PROVIDER_MODELS_MODE=fallback`, GoModel uses +the list when Oracle's `/models` endpoint is unavailable, returns nil, or +returns an empty list. Set `CONFIGURED_PROVIDER_MODELS_MODE=allowlist` to expose +only the configured Oracle models and skip Oracle's upstream `/models` call. For multiple Oracle providers without YAML, use suffixed env vars such as `ORACLE_US_BASE_URL`, `ORACLE_US_API_KEY`, and `ORACLE_US_MODELS`. For @@ -82,11 +85,12 @@ Why `models:` matters: - Oracle inference works through `chat/completions` and `responses` - Oracle's `/models` endpoint may not be available for this API-key flow -- GoModel can fall back to the configured model list when `/models` is - unavailable +- GoModel can use the configured model list consistently with the global + `CONFIGURED_PROVIDER_MODELS_MODE` -If both are set, `ORACLE_MODELS` overrides YAML `models:` for the default -`oracle` provider. +If both are set, `ORACLE_MODELS` overrides YAML `models:` for the matching +Oracle provider. For multiple env-only Oracle instances, use suffixed variables +such as `ORACLE_US_MODELS`. ## Current status @@ -94,7 +98,8 @@ What is integrated today: - Oracle's OpenAI-compatible inference endpoints - manual model configuration through `ORACLE_MODELS` or `models:` -- GoModel `/v1/models` from the configured-model fallback +- GoModel `/v1/models` from the configured model list in fallback or allowlist + mode What is not yet validated as reliable: @@ -166,7 +171,7 @@ text. wrong, or the model is not available to the account. - `model registry has no models` Set `ORACLE_MODELS` or add `models:` to the Oracle provider config so GoModel - can use the fallback. + can use the configured model list. - OCI CLI works but Oracle bearer requests fail These are different auth flows. OCI CLI uses API signing keys; Oracle Generative AI inference uses the Generative AI bearer API key. diff --git a/internal/providers/azure/azure.go b/internal/providers/azure/azure.go index 95c3f97ac..7b3eac282 100644 --- a/internal/providers/azure/azure.go +++ b/internal/providers/azure/azure.go @@ -2,7 +2,6 @@ package azure import ( "context" - "log/slog" "net/http" "net/url" "strconv" @@ -31,7 +30,6 @@ type Provider struct { resourceProvider *openai.CompatibleProvider openAIResourceProvider *openai.CompatibleProvider apiVersion string - configuredModels []string } func New(providerCfg providers.ProviderConfig, opts providers.ProviderOptions) core.Provider { @@ -39,12 +37,10 @@ func New(providerCfg providers.ProviderConfig, opts providers.ProviderOptions) c apiVersion := providers.ResolveAPIVersion(providerCfg.APIVersion, defaultAPIVersion) p := &Provider{apiVersion: apiVersion} clientCfg := openai.CompatibleProviderConfig{ - ProviderName: "azure", - BaseURL: baseURL, - SetHeaders: setHeaders, - ConfiguredModels: opts.Models, + ProviderName: "azure", + BaseURL: baseURL, + SetHeaders: setHeaders, } - p.configuredModels = opts.Models p.CompatibleProvider = openai.NewCompatibleProvider(providerCfg.APIKey, opts, clientCfg) p.resourceProvider = openai.NewCompatibleProvider(providerCfg.APIKey, opts, clientCfg) p.openAIResourceProvider = openai.NewCompatibleProvider(providerCfg.APIKey, opts, clientCfg) @@ -84,31 +80,7 @@ func (p *Provider) ListModels(ctx context.Context) (*core.ModelsResponse, error) Method: http.MethodGet, Endpoint: "/openai/models", }, &resp); err != nil { - if len(p.configuredModels) == 0 { - return nil, err - } - - slog.Warn("azure upstream ListModels failed, using configured models fallback", - "error", err, - "configured_models", len(p.configuredModels), - ) - - data := make([]core.Model, 0, len(p.configuredModels)) - for _, modelID := range p.configuredModels { - modelID = strings.TrimSpace(modelID) - if modelID == "" { - continue - } - data = append(data, core.Model{ - ID: modelID, - Object: "model", - OwnedBy: "azure", - }) - } - return &core.ModelsResponse{ - Object: "list", - Data: data, - }, nil + return nil, err } return &resp, nil } diff --git a/internal/providers/config.go b/internal/providers/config.go index 749f482b1..28cbcc15d 100644 --- a/internal/providers/config.go +++ b/internal/providers/config.go @@ -149,9 +149,6 @@ func parseProviderEnvKey(prefix, key string, spec DiscoveryConfig) (string, prov if candidate.field == providerEnvFieldAPIVersion && !spec.SupportsAPIVersion { continue } - if candidate.field == providerEnvFieldModels && !spec.SupportsModelsEnv { - continue - } if rest == candidate.name { return "", candidate.field, true } diff --git a/internal/providers/config_test.go b/internal/providers/config_test.go index d397a5c38..c4980f8c2 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -48,8 +48,7 @@ var testDiscoveryConfigs = map[string]DiscoveryConfig{ SupportsAPIVersion: true, }, "oracle": { - RequireBaseURL: true, - SupportsModelsEnv: true, + RequireBaseURL: true, }, "ollama": { DefaultBaseURL: "http://localhost:11434/v1", @@ -511,13 +510,20 @@ func TestApplyProviderEnvVars_DiscoversVLLMFromAPIKeyWithDefaultBaseURL(t *testi } } -func TestApplyProviderEnvVars_IgnoresVLLMModelsEnv(t *testing.T) { +func TestApplyProviderEnvVars_DiscoversVLLMFromModelsEnv(t *testing.T) { t.Setenv("VLLM_MODELS", "meta-llama/Llama-3.1-8B-Instruct") got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, testDiscoveryConfigs) - if _, exists := got["vllm"]; exists { - t.Fatal("expected VLLM_MODELS not to discover vllm provider") + p, exists := got["vllm"] + if !exists { + t.Fatal("expected VLLM_MODELS to discover keyless vllm provider") + } + if p.Type != "vllm" { + t.Fatalf("Type = %q, want vllm", p.Type) + } + if len(p.Models) != 1 || p.Models[0].ID != "meta-llama/Llama-3.1-8B-Instruct" { + t.Fatalf("Models = %v, want [meta-llama/Llama-3.1-8B-Instruct]", p.Models) } } @@ -565,6 +571,7 @@ func TestApplyProviderEnvVars_DiscoversSuffixedProvidersForEveryRegisteredType(t for providerType, spec := range testDiscoveryConfigs { prefix := envPrefix(providerType) t.Setenv(prefix+"_EAST_API_KEY", "key-"+providerType) + t.Setenv(prefix+"_EAST_MODELS", "model-a-"+providerType+", model-b-"+providerType) if spec.RequireBaseURL { t.Setenv(prefix+"_EAST_BASE_URL", "https://"+providerType+".example.com/v1") } else { @@ -594,6 +601,9 @@ func TestApplyProviderEnvVars_DiscoversSuffixedProvidersForEveryRegisteredType(t } else if spec.DefaultBaseURL != "" && p.BaseURL != spec.DefaultBaseURL { t.Errorf("%s BaseURL = %q, want %q", name, p.BaseURL, spec.DefaultBaseURL) } + if len(p.Models) != 2 || p.Models[0].ID != "model-a-"+providerType || p.Models[1].ID != "model-b-"+providerType { + t.Errorf("%s Models = %v, want [model-a-%s model-b-%s]", name, p.Models, providerType, providerType) + } } } diff --git a/internal/providers/configured_models.go b/internal/providers/configured_models.go new file mode 100644 index 000000000..96bb4d7ef --- /dev/null +++ b/internal/providers/configured_models.go @@ -0,0 +1,212 @@ +package providers + +import ( + "sort" + "strings" + "time" + + "gomodel/config" + "gomodel/internal/core" +) + +type configuredProviderModelsApplyReason string + +const ( + configuredProviderModelsNotApplied configuredProviderModelsApplyReason = "" + configuredProviderModelsAllowlist configuredProviderModelsApplyReason = "allowlist" + configuredProviderModelsUpstreamError configuredProviderModelsApplyReason = "upstream_error" + configuredProviderModelsUpstreamNil configuredProviderModelsApplyReason = "upstream_nil" + configuredProviderModelsUpstreamEmpty configuredProviderModelsApplyReason = "upstream_empty" +) + +func normalizeConfiguredProviderModels(models []string) []string { + if len(models) == 0 { + return nil + } + + seen := make(map[string]struct{}, len(models)) + normalized := make([]string, 0, len(models)) + for _, model := range models { + model = strings.TrimSpace(model) + if model == "" { + continue + } + if _, exists := seen[model]; exists { + continue + } + seen[model] = struct{}{} + normalized = append(normalized, model) + } + if len(normalized) == 0 { + return nil + } + return normalized +} + +func applyConfiguredProviderModels( + providerName string, + providerType string, + mode config.ConfiguredProviderModelsMode, + configuredModels []string, + upstream *core.ModelsResponse, + upstreamErr error, + fallbackCreated int64, +) (*core.ModelsResponse, configuredProviderModelsApplyReason) { + if len(configuredModels) == 0 { + return upstream, configuredProviderModelsNotApplied + } + + mode = config.ResolveConfiguredProviderModelsMode(mode) + if mode == config.ConfiguredProviderModelsModeAllowlist { + return configuredProviderModelsResponse(providerName, providerType, configuredModels, upstream, fallbackCreated), configuredProviderModelsAllowlist + } + + if upstreamErr != nil { + return configuredProviderModelsResponse(providerName, providerType, configuredModels, upstream, fallbackCreated), configuredProviderModelsUpstreamError + } + if upstream == nil { + return configuredProviderModelsResponse(providerName, providerType, configuredModels, upstream, fallbackCreated), configuredProviderModelsUpstreamNil + } + if len(upstream.Data) == 0 { + return configuredProviderModelsResponse(providerName, providerType, configuredModels, upstream, fallbackCreated), configuredProviderModelsUpstreamEmpty + } + return upstream, configuredProviderModelsNotApplied +} + +func configuredProviderModelsResponse(providerName, providerType string, configuredModels []string, upstream *core.ModelsResponse, fallbackCreated int64) *core.ModelsResponse { + byID := make(map[string]core.Model) + if upstream != nil { + for _, model := range upstream.Data { + modelID := strings.TrimSpace(model.ID) + if modelID == "" { + continue + } + byID[modelID] = model + } + } + + owner := strings.TrimSpace(providerType) + if owner == "" { + owner = strings.TrimSpace(providerName) + } + if fallbackCreated <= 0 { + fallbackCreated = time.Now().Unix() + } + + data := make([]core.Model, 0, len(configuredModels)) + for _, modelID := range configuredModels { + model, ok := byID[modelID] + if !ok { + model = core.Model{ + ID: modelID, + Object: "model", + OwnedBy: owner, + Created: fallbackCreated, + } + } else { + model.ID = strings.TrimSpace(model.ID) + if model.ID == "" { + model.ID = modelID + } + if strings.TrimSpace(model.Object) == "" { + model.Object = "model" + } + if strings.TrimSpace(model.OwnedBy) == "" { + model.OwnedBy = owner + } + if model.Created == 0 { + model.Created = fallbackCreated + } + } + data = append(data, model) + } + + return &core.ModelsResponse{ + Object: "list", + Data: data, + } +} + +func modelsResponseFromProviderMap(providerModels map[string]*ModelInfo) *core.ModelsResponse { + if len(providerModels) == 0 { + return &core.ModelsResponse{Object: "list"} + } + modelIDs := make([]string, 0, len(providerModels)) + for modelID := range providerModels { + modelIDs = append(modelIDs, modelID) + } + sort.Strings(modelIDs) + + data := make([]core.Model, 0, len(modelIDs)) + for _, modelID := range modelIDs { + if info := providerModels[modelID]; info != nil { + data = append(data, info.Model) + } + } + return &core.ModelsResponse{ + Object: "list", + Data: data, + } +} + +func modelInfoMapFromResponse(resp *core.ModelsResponse, provider core.Provider, providerName, providerType string) map[string]*ModelInfo { + out := make(map[string]*ModelInfo) + if resp == nil { + return out + } + for _, model := range resp.Data { + modelID := strings.TrimSpace(model.ID) + if modelID == "" { + continue + } + model.ID = modelID + out[modelID] = &ModelInfo{ + Model: model, + Provider: provider, + ProviderName: providerName, + ProviderType: providerType, + } + } + return out +} + +func rebuildGlobalModelMap(modelsByProvider map[string]map[string]*ModelInfo, providerOrderNames []string) map[string]*ModelInfo { + global := make(map[string]*ModelInfo) + seenProvider := make(map[string]struct{}, len(providerOrderNames)) + for _, providerName := range providerOrderNames { + seenProvider[providerName] = struct{}{} + addProviderModels(global, modelsByProvider[providerName]) + } + + remaining := make([]string, 0, len(modelsByProvider)) + for providerName := range modelsByProvider { + if _, seen := seenProvider[providerName]; seen { + continue + } + remaining = append(remaining, providerName) + } + sort.Strings(remaining) + for _, providerName := range remaining { + addProviderModels(global, modelsByProvider[providerName]) + } + return global +} + +func addProviderModels(global map[string]*ModelInfo, providerModels map[string]*ModelInfo) { + if len(providerModels) == 0 { + return + } + modelIDs := make([]string, 0, len(providerModels)) + for modelID := range providerModels { + modelIDs = append(modelIDs, modelID) + } + sort.Strings(modelIDs) + for _, modelID := range modelIDs { + if _, exists := global[modelID]; exists { + continue + } + if info := providerModels[modelID]; info != nil { + global[modelID] = info + } + } +} diff --git a/internal/providers/configured_models_test.go b/internal/providers/configured_models_test.go new file mode 100644 index 000000000..4c2a2d0ba --- /dev/null +++ b/internal/providers/configured_models_test.go @@ -0,0 +1,38 @@ +package providers + +import ( + "testing" + + "gomodel/config" + "gomodel/internal/core" +) + +func TestApplyConfiguredProviderModels_BackfillsZeroCreatedForUpstreamMatch(t *testing.T) { + resp, reason := applyConfiguredProviderModels( + "test", + "test-type", + config.ConfiguredProviderModelsModeAllowlist, + []string{"configured-model"}, + &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "configured-model", Object: "model", OwnedBy: "upstream"}, + }, + }, + nil, + 123, + ) + + if reason != configuredProviderModelsAllowlist { + t.Fatalf("reason = %q, want %q", reason, configuredProviderModelsAllowlist) + } + if resp == nil || len(resp.Data) != 1 { + t.Fatalf("resp = %+v, want one configured model", resp) + } + if resp.Data[0].Created != 123 { + t.Fatalf("Created = %d, want fallback timestamp 123", resp.Data[0].Created) + } + if resp.Data[0].OwnedBy != "upstream" { + t.Fatalf("OwnedBy = %q, want upstream metadata preserved", resp.Data[0].OwnedBy) + } +} diff --git a/internal/providers/factory.go b/internal/providers/factory.go index 95ecc5c9a..efbf09e00 100644 --- a/internal/providers/factory.go +++ b/internal/providers/factory.go @@ -28,7 +28,6 @@ type DiscoveryConfig struct { RequireBaseURL bool AllowAPIKeyless bool SupportsAPIVersion bool - SupportsModelsEnv bool } // Registration contains metadata for registering a provider with the factory. diff --git a/internal/providers/init.go b/internal/providers/init.go index 0f754d254..2f9fe5738 100644 --- a/internal/providers/init.go +++ b/internal/providers/init.go @@ -89,6 +89,7 @@ func Init(ctx context.Context, result *config.LoadResult, factory *ProviderFacto registry := NewModelRegistry() registry.SetCache(modelCache) + registry.SetConfiguredProviderModelsMode(result.Config.Models.ConfiguredProviderModelsMode) count, err := initializeProviders(ctx, providerMap, factory, registry) if err != nil { @@ -239,6 +240,9 @@ func initializeProviders(ctx context.Context, providerMap map[string]ProviderCon } registry.RegisterProviderWithNameAndType(p, name, pCfg.Type) + if len(pCfg.Models) > 0 { + registry.SetProviderConfiguredModels(name, pCfg.Models) + } if len(pCfg.ModelMetadataOverrides) > 0 { registry.SetProviderMetadataOverrides(name, pCfg.ModelMetadataOverrides) } diff --git a/internal/providers/openai/compatible_provider.go b/internal/providers/openai/compatible_provider.go index be11c4deb..e8e3da9a6 100644 --- a/internal/providers/openai/compatible_provider.go +++ b/internal/providers/openai/compatible_provider.go @@ -3,11 +3,9 @@ package openai import ( "context" "io" - "log/slog" "net/http" "net/url" "strconv" - "strings" "gomodel/internal/core" "gomodel/internal/llmclient" @@ -17,27 +15,24 @@ import ( type RequestMutator func(*llmclient.Request) type CompatibleProviderConfig struct { - ProviderName string - BaseURL string - SetHeaders func(*http.Request, string) - RequestMutator RequestMutator - ConfiguredModels []string + ProviderName string + BaseURL string + SetHeaders func(*http.Request, string) + RequestMutator RequestMutator } type CompatibleProvider struct { - client *llmclient.Client - apiKey string - providerName string - requestMutator RequestMutator - configuredModels []string + client *llmclient.Client + apiKey string + providerName string + requestMutator RequestMutator } func NewCompatibleProvider(apiKey string, opts providers.ProviderOptions, cfg CompatibleProviderConfig) *CompatibleProvider { p := &CompatibleProvider{ - apiKey: apiKey, - providerName: cfg.ProviderName, - requestMutator: cfg.RequestMutator, - configuredModels: normalizeConfiguredModels(cfg.ConfiguredModels), + apiKey: apiKey, + providerName: cfg.ProviderName, + requestMutator: cfg.RequestMutator, } clientCfg := llmclient.Config{ ProviderName: cfg.ProviderName, @@ -59,10 +54,9 @@ func NewCompatibleProviderWithHTTPClient(apiKey string, httpClient *http.Client, httpClient = http.DefaultClient } p := &CompatibleProvider{ - apiKey: apiKey, - providerName: cfg.ProviderName, - requestMutator: cfg.RequestMutator, - configuredModels: normalizeConfiguredModels(cfg.ConfiguredModels), + apiKey: apiKey, + providerName: cfg.ProviderName, + requestMutator: cfg.RequestMutator, } clientCfg := llmclient.DefaultConfig(cfg.ProviderName, cfg.BaseURL) clientCfg.Hooks = hooks @@ -133,53 +127,6 @@ func (p *CompatibleProvider) StreamChatCompletion(ctx context.Context, req *core } func (p *CompatibleProvider) ListModels(ctx context.Context) (*core.ModelsResponse, error) { - if len(p.configuredModels) == 0 { - return p.doListModels(ctx) - } - - resp, err := p.doListModels(ctx) - if err != nil { - slog.Warn("openai-compatible upstream ListModels failed, using configured models fallback", - "provider", p.providerName, - "error", err, - "configured_models", len(p.configuredModels), - ) - } - - byID := make(map[string]core.Model, len(p.configuredModels)) - if resp != nil { - for _, model := range resp.Data { - byID[strings.TrimSpace(model.ID)] = model - } - } - - data := make([]core.Model, 0, len(p.configuredModels)) - for _, modelID := range p.configuredModels { - model, ok := byID[modelID] - if !ok { - model = core.Model{ - ID: modelID, - Object: "model", - OwnedBy: p.providerName, - } - } else { - if strings.TrimSpace(model.Object) == "" { - model.Object = "model" - } - if strings.TrimSpace(model.OwnedBy) == "" { - model.OwnedBy = p.providerName - } - } - data = append(data, model) - } - - return &core.ModelsResponse{ - Object: "list", - Data: data, - }, nil -} - -func (p *CompatibleProvider) doListModels(ctx context.Context) (*core.ModelsResponse, error) { var resp core.ModelsResponse err := p.Do(ctx, llmclient.Request{ Method: http.MethodGet, @@ -545,28 +492,3 @@ func responseInputItemsEndpoint(id string, params core.ResponseInputItemsParams) } return endpoint } - -// normalizeConfiguredModels deduplicates and trims model names. -func normalizeConfiguredModels(models []string) []string { - if len(models) == 0 { - return nil - } - - seen := make(map[string]struct{}, len(models)) - normalized := make([]string, 0, len(models)) - for _, model := range models { - model = strings.TrimSpace(model) - if model == "" { - continue - } - if _, exists := seen[model]; exists { - continue - } - seen[model] = struct{}{} - normalized = append(normalized, model) - } - if len(normalized) == 0 { - return nil - } - return normalized -} diff --git a/internal/providers/openai/compatible_provider_test.go b/internal/providers/openai/compatible_provider_test.go index 5343c2bd3..660c505c7 100644 --- a/internal/providers/openai/compatible_provider_test.go +++ b/internal/providers/openai/compatible_provider_test.go @@ -2,7 +2,6 @@ package openai import ( "context" - "encoding/json" "net/http" "net/http/httptest" "testing" @@ -11,96 +10,10 @@ import ( "gomodel/internal/llmclient" ) -func TestCompatibleProvider_ListModels_UsesConfiguredFallbackWhenUpstreamFailsWithHTML(t *testing.T) { - htmlBody := `ErrorNot Found` - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/models" { - w.Header().Set("Content-Type", "text/html") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(htmlBody)) - } - })) - defer server.Close() - - provider := NewCompatibleProviderWithHTTPClient( - "test-key", - server.Client(), - llmclient.Hooks{}, - CompatibleProviderConfig{ - ProviderName: "opencode-go", - BaseURL: server.URL, - ConfiguredModels: []string{"glm-5.1", "glm-5", "kimi-k2.5"}, - }, - ) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if resp == nil { - t.Fatal("expected response, got nil") - return - } - if len(resp.Data) != 3 { - t.Fatalf("len(resp.Data) = %d, want 3", len(resp.Data)) - } - expected := []string{"glm-5.1", "glm-5", "kimi-k2.5"} - for i, id := range expected { - if resp.Data[i].ID != id { - t.Errorf("resp.Data[%d].ID = %q, want %q", i, resp.Data[i].ID, id) - } - if resp.Data[i].Object != "model" { - t.Errorf("resp.Data[%d].Object = %q, want model", i, resp.Data[i].Object) - } - if resp.Data[i].OwnedBy != "opencode-go" { - t.Errorf("resp.Data[%d].OwnedBy = %q, want opencode-go", i, resp.Data[i].OwnedBy) - } - } -} - -func TestCompatibleProvider_ListModels_UsesConfiguredFallbackWhenUpstreamReturnsJSONError(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/models" { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte(`{"error":{"message":"Invalid API key"}}`)) - } - })) - defer server.Close() - - provider := NewCompatibleProviderWithHTTPClient( - "test-key", - server.Client(), - llmclient.Hooks{}, - CompatibleProviderConfig{ - ProviderName: "my-provider", - BaseURL: server.URL, - ConfiguredModels: []string{"custom-model-v1"}, - }, - ) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if resp == nil { - t.Fatal("expected response, got nil") - return - } - if len(resp.Data) != 1 || resp.Data[0].ID != "custom-model-v1" { - t.Fatalf("unexpected models: %+v", resp.Data) - } -} - -func TestCompatibleProvider_ListModels_MergesUpstreamMetadataWhenAvailable(t *testing.T) { +func TestCompatibleProvider_ListModels_ReturnsUpstreamOnSuccess(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(core.ModelsResponse{ - Object: "list", - Data: []core.Model{ - {ID: "shared-model", Object: "model", OwnedBy: "upstream", Created: 999}, - }, - }) + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-4o","object":"model","owned_by":"openai"}]}`)) })) defer server.Close() @@ -109,12 +22,8 @@ func TestCompatibleProvider_ListModels_MergesUpstreamMetadataWhenAvailable(t *te server.Client(), llmclient.Hooks{}, CompatibleProviderConfig{ - ProviderName: "merged-provider", + ProviderName: "upstream-only", BaseURL: server.URL, - ConfiguredModels: []string{ - "shared-model", - "only-configured", - }, }, ) @@ -122,40 +31,12 @@ func TestCompatibleProvider_ListModels_MergesUpstreamMetadataWhenAvailable(t *te if err != nil { t.Fatalf("ListModels() error = %v", err) } - if len(resp.Data) != 2 { - t.Fatalf("len(resp.Data) = %d, want 2", len(resp.Data)) - } - - // shared-model should carry upstream metadata - var shared, onlyConfigured core.Model - for i := range resp.Data { - if resp.Data[i].ID == "shared-model" { - shared = resp.Data[i] - } - if resp.Data[i].ID == "only-configured" { - onlyConfigured = resp.Data[i] - } - } - if shared.ID != "shared-model" || shared.Object != "model" { - t.Errorf("shared model: id=%q, object=%q", shared.ID, shared.Object) - } - // Shared model keeps upstream OwnedBy since it's non-empty - if shared.OwnedBy != "upstream" { - t.Errorf("shared-model.OwnedBy = %q, want upstream", shared.OwnedBy) - } - - if onlyConfigured.ID != "only-configured" { - t.Fatalf("only-configured model: id=%q", onlyConfigured.ID) - } - if onlyConfigured.Object != "model" { - t.Errorf("only-configured.Object = %q, want model", onlyConfigured.Object) - } - if onlyConfigured.OwnedBy != "merged-provider" { - t.Errorf("only-configured.OwnedBy = %q, want merged-provider", onlyConfigured.OwnedBy) + if len(resp.Data) != 1 || resp.Data[0].ID != "gpt-4o" { + t.Fatalf("unexpected models: %+v", resp.Data) } } -func TestCompatibleProvider_ListModels_NoConfiguredModels_OriginalBehaviorWhenUpstreamFails(t *testing.T) { +func TestCompatibleProvider_ListModels_ReturnsUpstreamError(t *testing.T) { server := httptest.NewServer(http.NotFoundHandler()) defer server.Close() @@ -164,109 +45,20 @@ func TestCompatibleProvider_ListModels_NoConfiguredModels_OriginalBehaviorWhenUp server.Client(), llmclient.Hooks{}, CompatibleProviderConfig{ - ProviderName: "test-provider", - BaseURL: server.URL, - ConfiguredModels: nil, + ProviderName: "test-provider", + BaseURL: server.URL, }, ) _, err := provider.ListModels(context.Background()) if err == nil { - t.Fatal("expected error when upstream fails and no configured models, got nil") + t.Fatal("expected error when upstream fails, got nil") } gatewayErr, ok := err.(*core.GatewayError) if !ok { t.Fatalf("error type = %T, want *core.GatewayError", err) } - // Upstream returns 404 so error type is not_found_error; the important - // invariant is that an error (not a fallback) is returned when no - // configured models are present. if gatewayErr.Type != core.ErrorTypeProvider && gatewayErr.Type != core.ErrorTypeNotFound { t.Errorf("gatewayErr.Type = %q, want provider_error or not_found_error", gatewayErr.Type) } } - -func TestCompatibleProvider_ListModels_NoConfiguredModels_ReturnsUpstreamOnSuccess(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-4o","object":"model","owned_by":"openai"}]}`)) - })) - defer server.Close() - - provider := NewCompatibleProviderWithHTTPClient( - "test-key", - server.Client(), - llmclient.Hooks{}, - CompatibleProviderConfig{ - ProviderName: "upstream-only", - BaseURL: server.URL, - ConfiguredModels: nil, - }, - ) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "gpt-4o" { - t.Fatalf("unexpected models: %+v", resp.Data) - } -} - -func TestCompatibleProvider_ListModels_EmptyConfiguredModels_OriginalBehavior(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-4o","object":"model"}]}`)) - })) - defer server.Close() - - provider := NewCompatibleProviderWithHTTPClient( - "test-key", - server.Client(), - llmclient.Hooks{}, - CompatibleProviderConfig{ - ProviderName: "test", - BaseURL: server.URL, - ConfiguredModels: []string{}, // explicitly empty - }, - ) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 || resp.Data[0].ID != "gpt-4o" { - t.Fatalf("unexpected models: %+v", resp.Data) - } -} - -func TestNormalizeConfiguredModels(t *testing.T) { - got := normalizeConfiguredModels([]string{ - " glm-5.1 ", - "", - "glm-5", - "glm-5.1", // duplicate - " ", // whitespace only - }) - - if len(got) != 2 { - t.Fatalf("len(got) = %d, want 2", len(got)) - } - if got[0] != "glm-5.1" || got[1] != "glm-5" { - t.Fatalf("got = %v, want [glm-5.1 glm-5]", got) - } -} - -func TestNormalizeConfiguredModels_AllEmpty(t *testing.T) { - got := normalizeConfiguredModels([]string{"", " ", ""}) - if got != nil { - t.Fatalf("got = %v, want nil", got) - } -} - -func TestNormalizeConfiguredModels_NilInput(t *testing.T) { - got := normalizeConfiguredModels(nil) - if got != nil { - t.Fatalf("got = %v, want nil", got) - } -} diff --git a/internal/providers/openai/openai.go b/internal/providers/openai/openai.go index 3429355a3..bae594eff 100644 --- a/internal/providers/openai/openai.go +++ b/internal/providers/openai/openai.go @@ -35,10 +35,9 @@ func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Prov baseURL := providers.ResolveBaseURL(cfg.BaseURL, defaultBaseURL) return &Provider{ CompatibleProvider: NewCompatibleProvider(cfg.APIKey, opts, CompatibleProviderConfig{ - ProviderName: "openai", - BaseURL: baseURL, - SetHeaders: setHeaders, - ConfiguredModels: opts.Models, + ProviderName: "openai", + BaseURL: baseURL, + SetHeaders: setHeaders, }), } } diff --git a/internal/providers/openrouter/openrouter.go b/internal/providers/openrouter/openrouter.go index 178c32b91..b7f4230a9 100644 --- a/internal/providers/openrouter/openrouter.go +++ b/internal/providers/openrouter/openrouter.go @@ -39,10 +39,9 @@ func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Prov appName: envOrDefault("OPENROUTER_APP_NAME", defaultAppName), } p.CompatibleProvider = openai.NewCompatibleProvider(cfg.APIKey, opts, openai.CompatibleProviderConfig{ - ProviderName: "openrouter", - BaseURL: baseURL, - SetHeaders: setHeaders, - ConfiguredModels: opts.Models, + ProviderName: "openrouter", + BaseURL: baseURL, + SetHeaders: setHeaders, }) p.SetRequestMutator(p.mutateRequest) return p diff --git a/internal/providers/oracle/oracle.go b/internal/providers/oracle/oracle.go index 75a228905..8e72fe714 100644 --- a/internal/providers/oracle/oracle.go +++ b/internal/providers/oracle/oracle.go @@ -3,9 +3,7 @@ package oracle import ( "context" "io" - "log/slog" "net/http" - "strings" "gomodel/internal/core" "gomodel/internal/llmclient" @@ -19,14 +17,12 @@ var Registration = providers.Registration{ Type: "oracle", New: New, Discovery: providers.DiscoveryConfig{ - RequireBaseURL: true, - SupportsModelsEnv: true, + RequireBaseURL: true, }, } type Provider struct { - compat *openai.CompatibleProvider - configuredModels []string + compat *openai.CompatibleProvider } func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Provider { @@ -37,18 +33,16 @@ func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Prov BaseURL: baseURL, SetHeaders: setHeaders, }), - configuredModels: normalizeConfiguredModels(opts.Models), } } -func NewWithHTTPClient(apiKey string, httpClient *http.Client, hooks llmclient.Hooks, models []string) *Provider { +func NewWithHTTPClient(apiKey string, httpClient *http.Client, hooks llmclient.Hooks) *Provider { return &Provider{ compat: openai.NewCompatibleProviderWithHTTPClient(apiKey, httpClient, hooks, openai.CompatibleProviderConfig{ ProviderName: "oracle", BaseURL: defaultBaseURL, SetHeaders: setHeaders, }), - configuredModels: normalizeConfiguredModels(models), } } @@ -65,56 +59,7 @@ func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatReque } func (p *Provider) ListModels(ctx context.Context) (*core.ModelsResponse, error) { - resp, err := p.compat.ListModels(ctx) - if len(p.configuredModels) == 0 { - if err != nil { - return nil, core.NewProviderError( - "oracle", - http.StatusBadGateway, - "oracle ListModels failed: "+err.Error()+"; set ORACLE_MODELS or add providers..models in config.yaml to use Oracle when upstream /models is unavailable", - err, - ) - } - return resp, nil - } - if err != nil { - slog.Warn("oracle upstream ListModels failed, using configured models fallback", - "error", err, - "configured_models", len(p.configuredModels), - ) - } - - byID := make(map[string]core.Model, len(p.configuredModels)) - if err == nil && resp != nil { - for _, model := range resp.Data { - byID[strings.TrimSpace(model.ID)] = model - } - } - - data := make([]core.Model, 0, len(p.configuredModels)) - for _, modelID := range p.configuredModels { - model, ok := byID[modelID] - if !ok { - model = core.Model{ - ID: modelID, - Object: "model", - OwnedBy: "oracle", - } - } else { - if strings.TrimSpace(model.Object) == "" { - model.Object = "model" - } - if strings.TrimSpace(model.OwnedBy) == "" { - model.OwnedBy = "oracle" - } - } - data = append(data, model) - } - - return &core.ModelsResponse{ - Object: "list", - Data: data, - }, nil + return p.compat.ListModels(ctx) } func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { @@ -132,27 +77,3 @@ func (p *Provider) Embeddings(_ context.Context, _ *core.EmbeddingRequest) (*cor func setHeaders(req *http.Request, apiKey string) { req.Header.Set("Authorization", "Bearer "+apiKey) } - -func normalizeConfiguredModels(models []string) []string { - if len(models) == 0 { - return nil - } - - seen := make(map[string]struct{}, len(models)) - normalized := make([]string, 0, len(models)) - for _, model := range models { - model = strings.TrimSpace(model) - if model == "" { - continue - } - if _, exists := seen[model]; exists { - continue - } - seen[model] = struct{}{} - normalized = append(normalized, model) - } - if len(normalized) == 0 { - return nil - } - return normalized -} diff --git a/internal/providers/oracle/oracle_test.go b/internal/providers/oracle/oracle_test.go index bce9c3efd..0018070f6 100644 --- a/internal/providers/oracle/oracle_test.go +++ b/internal/providers/oracle/oracle_test.go @@ -2,95 +2,15 @@ package oracle import ( "context" - "encoding/json" "net/http" "net/http/httptest" - "strings" "testing" "gomodel/internal/core" "gomodel/internal/llmclient" ) -func TestListModels_FallsBackToConfiguredModelsWhenUpstreamFails(t *testing.T) { - server := httptest.NewServer(http.NotFoundHandler()) - defer server.Close() - - provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}, []string{ - "openai.gpt-oss-120b", - "xai.grok-3", - }) - provider.SetBaseURL(server.URL) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if resp == nil { - t.Fatal("expected response, got nil") - return - } - if len(resp.Data) != 2 { - t.Fatalf("len(resp.Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].ID != "openai.gpt-oss-120b" { - t.Fatalf("resp.Data[0].ID = %q, want openai.gpt-oss-120b", resp.Data[0].ID) - } - if resp.Data[1].ID != "xai.grok-3" { - t.Fatalf("resp.Data[1].ID = %q, want xai.grok-3", resp.Data[1].ID) - } - for i, model := range resp.Data { - if model.Object != "model" { - t.Fatalf("resp.Data[%d].Object = %q, want model", i, model.Object) - } - if model.OwnedBy != "oracle" { - t.Fatalf("resp.Data[%d].OwnedBy = %q, want oracle", i, model.OwnedBy) - } - } -} - -func TestListModels_FiltersUpstreamModelsAndAddsMissingConfiguredModels(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/models" { - http.NotFound(w, r) - return - } - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(core.ModelsResponse{ - Object: "list", - Data: []core.Model{ - {ID: "xai.grok-3", Object: "model", OwnedBy: "oracle", Created: 123}, - {ID: "ignore-me", Object: "model", OwnedBy: "oracle", Created: 456}, - }, - }) - })) - defer server.Close() - - provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}, []string{ - "openai.gpt-oss-120b", - "xai.grok-3", - }) - provider.SetBaseURL(server.URL) - - resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 2 { - t.Fatalf("len(resp.Data) = %d, want 2", len(resp.Data)) - } - if resp.Data[0].ID != "openai.gpt-oss-120b" { - t.Fatalf("resp.Data[0].ID = %q, want openai.gpt-oss-120b", resp.Data[0].ID) - } - if resp.Data[1].ID != "xai.grok-3" { - t.Fatalf("resp.Data[1].ID = %q, want xai.grok-3", resp.Data[1].ID) - } - if resp.Data[1].Created != 123 { - t.Fatalf("resp.Data[1].Created = %d, want 123", resp.Data[1].Created) - } -} - -func TestListModels_ReturnsUpstreamInventoryWhenNoConfiguredModels(t *testing.T) { +func TestListModels_ReturnsUpstreamInventory(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/models" { http.NotFound(w, r) @@ -101,7 +21,7 @@ func TestListModels_ReturnsUpstreamInventoryWhenNoConfiguredModels(t *testing.T) })) defer server.Close() - provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}, nil) + provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}) provider.SetBaseURL(server.URL) resp, err := provider.ListModels(context.Background()) @@ -113,34 +33,8 @@ func TestListModels_ReturnsUpstreamInventoryWhenNoConfiguredModels(t *testing.T) } } -func TestListModels_ReturnsActionableErrorWhenUpstreamFailsWithoutConfiguredModels(t *testing.T) { - server := httptest.NewServer(http.NotFoundHandler()) - defer server.Close() - - provider := NewWithHTTPClient("oracle-key", server.Client(), llmclient.Hooks{}, nil) - provider.SetBaseURL(server.URL) - - _, err := provider.ListModels(context.Background()) - if err == nil { - t.Fatal("expected error, got nil") - } - gatewayErr, ok := err.(*core.GatewayError) - if !ok { - t.Fatalf("error type = %T, want *core.GatewayError", err) - } - if gatewayErr.Type != core.ErrorTypeProvider { - t.Fatalf("gatewayErr.Type = %q, want %q", gatewayErr.Type, core.ErrorTypeProvider) - } - if gatewayErr.Provider != "oracle" { - t.Fatalf("gatewayErr.Provider = %q, want oracle", gatewayErr.Provider) - } - if !strings.Contains(err.Error(), "set ORACLE_MODELS or add providers..models in config.yaml") { - t.Fatalf("err = %q, want mention of ORACLE_MODELS or providers..models in config.yaml", err) - } -} - func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { - provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}, nil) + provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}) _, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{Model: "text-embedding-3-small"}) if err == nil { @@ -159,7 +53,7 @@ func TestEmbeddings_ReturnsUnsupportedError(t *testing.T) { } func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}, nil) + provider := NewWithHTTPClient("oracle-key", nil, llmclient.Hooks{}) if _, ok := any(provider).(core.NativeBatchProvider); ok { t.Fatal("oracle provider should not implement native batch provider") @@ -171,26 +65,3 @@ func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) t.Fatal("oracle provider should not implement passthrough provider") } } - -func TestNormalizeConfiguredModels(t *testing.T) { - got := normalizeConfiguredModels([]string{ - " openai.gpt-oss-120b ", - "", - "xai.grok-3", - "openai.gpt-oss-120b", - }) - - if len(got) != 2 { - t.Fatalf("len(got) = %d, want 2", len(got)) - } - if got[0] != "openai.gpt-oss-120b" || got[1] != "xai.grok-3" { - t.Fatalf("got = %v, want [openai.gpt-oss-120b xai.grok-3]", got) - } -} - -func TestNormalizeConfiguredModels_AllEmpty(t *testing.T) { - got := normalizeConfiguredModels([]string{"", " ", ""}) - if got != nil { - t.Fatalf("got = %v, want nil", got) - } -} diff --git a/internal/providers/registry.go b/internal/providers/registry.go index 80ef4a9d0..3fb60c020 100644 --- a/internal/providers/registry.go +++ b/internal/providers/registry.go @@ -16,6 +16,7 @@ import ( "sync" "time" + "gomodel/config" "gomodel/internal/cache/modelcache" "gomodel/internal/core" "gomodel/internal/modeldata" @@ -51,6 +52,11 @@ type ModelRegistry struct { // instance name -> raw model ID. Applied after remote-registry enrichment as // a higher-priority layer. nil if no overrides declared. configMetadataOverrides map[string]map[string]*core.ModelMetadata + // configuredProviderModels holds operator-supplied model inventories keyed by + // configured provider instance name. The mode decides whether these entries + // are fallback-only or an allowlist over the discovered upstream inventory. + configuredProviderModels map[string][]string + configuredProviderModelsMode config.ConfiguredProviderModelsMode // Cached sorted slices, rebuilt lazily after models change. // nil means cache needs rebuilding. Protected by mu. @@ -76,12 +82,13 @@ func (s metadataEnrichmentStats) slogAttrs() []any { // NewModelRegistry creates a new model registry func NewModelRegistry() *ModelRegistry { return &ModelRegistry{ - models: make(map[string]*ModelInfo), - modelsByProvider: make(map[string]map[string]*ModelInfo), - providerTypes: make(map[core.Provider]string), - providerNames: make(map[core.Provider]string), - providerRuntime: make(map[string]providerRuntimeState), - refreshCh: make(chan struct{}, 1), + models: make(map[string]*ModelInfo), + modelsByProvider: make(map[string]map[string]*ModelInfo), + providerTypes: make(map[core.Provider]string), + providerNames: make(map[core.Provider]string), + providerRuntime: make(map[string]providerRuntimeState), + refreshCh: make(chan struct{}, 1), + configuredProviderModelsMode: config.ConfiguredProviderModelsModeFallback, } } @@ -138,6 +145,34 @@ func (r *ModelRegistry) SetProviderMetadataOverrides(providerName string, overri r.configMetadataOverrides[providerName] = clone } +// SetConfiguredProviderModelsMode controls how configured provider model lists +// affect the final registry inventory. +func (r *ModelRegistry) SetConfiguredProviderModelsMode(mode config.ConfiguredProviderModelsMode) { + r.mu.Lock() + defer r.mu.Unlock() + r.configuredProviderModelsMode = config.ResolveConfiguredProviderModelsMode(mode) +} + +// SetProviderConfiguredModels records the explicit model inventory declared for +// a configured provider instance. Call with an empty/nil slice to clear it. +func (r *ModelRegistry) SetProviderConfiguredModels(providerName string, models []string) { + providerName = strings.TrimSpace(providerName) + if providerName == "" { + return + } + normalized := normalizeConfiguredProviderModels(models) + r.mu.Lock() + defer r.mu.Unlock() + if len(normalized) == 0 { + delete(r.configuredProviderModels, providerName) + return + } + if r.configuredProviderModels == nil { + r.configuredProviderModels = make(map[string][]string) + } + r.configuredProviderModels[providerName] = normalized +} + // RegisterProviderWithNameAndType adds a provider with a configured provider instance name and type. // Name is used for unambiguous provider/model selection (e.g. "provider/model") and cache persistence. func (r *ModelRegistry) RegisterProviderWithNameAndType(provider core.Provider, providerName, providerType string) { @@ -195,6 +230,7 @@ func (r *ModelRegistry) initialize(ctx context.Context) error { maps.Copy(providerTypes, r.providerTypes) maps.Copy(providerNames, r.providerNames) r.mu.RUnlock() + configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() for _, provider := range providers { providerName := providerNames[provider] @@ -205,8 +241,33 @@ func (r *ModelRegistry) initialize(ctx context.Context) error { providerName = fmt.Sprintf("%p", provider) } - resp, err := provider.ListModels(ctx) - fetchAt := time.Now().UTC() + configuredModels := configuredProviderModels[providerName] + resp, configuredReason, fetchAt, err := fetchProviderInventory( + ctx, + provider, + providerName, + providerTypes[provider], + configuredProviderModelsMode, + configuredModels, + ) + var configuredUpstreamError string + if configuredReason != configuredProviderModelsNotApplied { + attrs := []any{ + "provider", providerName, + "reason", string(configuredReason), + "configured_models", len(configuredModels), + } + if err != nil { + configuredUpstreamError = err.Error() + attrs = append(attrs, "error", err) + slog.Warn("upstream ListModels failed, using configured provider models", attrs...) + } else if configuredReason == configuredProviderModelsAllowlist { + slog.Debug("using configured provider models", attrs...) + } else { + slog.Warn("using configured provider models", attrs...) + } + err = nil + } if err != nil { slog.Warn("failed to fetch models from provider", "provider", providerName, @@ -252,11 +313,15 @@ func (r *ModelRegistry) initialize(ctx context.Context) error { continue } - runtimeUpdates[providerName] = providerRuntimeState{ - registered: true, - lastModelFetchAt: fetchAt, - lastModelFetchSuccessAt: fetchAt, + runtimeUpdate := providerRuntimeState{ + registered: true, + lastModelFetchAt: fetchAt, + lastModelFetchError: configuredUpstreamError, + } + if configuredReason == configuredProviderModelsNotApplied { + runtimeUpdate.lastModelFetchSuccessAt = fetchAt } + runtimeUpdates[providerName] = runtimeUpdate if _, ok := newModelsByProvider[providerName]; !ok { newModelsByProvider[providerName] = make(map[string]*ModelInfo, len(resp.Data)) @@ -330,6 +395,42 @@ func (r *ModelRegistry) initialize(ctx context.Context) error { return nil } +func fetchProviderInventory( + ctx context.Context, + provider core.Provider, + providerName string, + providerType string, + mode config.ConfiguredProviderModelsMode, + configuredModels []string, +) (*core.ModelsResponse, configuredProviderModelsApplyReason, time.Time, error) { + fetchAt := time.Now().UTC() + if mode == config.ConfiguredProviderModelsModeAllowlist && len(configuredModels) > 0 { + resp, reason := applyConfiguredProviderModels( + providerName, + providerType, + mode, + configuredModels, + nil, + nil, + fetchAt.Unix(), + ) + return resp, reason, fetchAt, nil + } + + resp, err := provider.ListModels(ctx) + fetchAt = time.Now().UTC() + resp, reason := applyConfiguredProviderModels( + providerName, + providerType, + mode, + configuredModels, + resp, + err, + fetchAt.Unix(), + ) + return resp, reason, fetchAt, err +} + func (r *ModelRegistry) applyProviderRuntimeUpdates(updates map[string]providerRuntimeState) { if len(updates) == 0 { return @@ -350,8 +451,11 @@ func (r *ModelRegistry) applyProviderRuntimeUpdatesLocked(updates map[string]pro } if !update.lastModelFetchSuccessAt.IsZero() { current.lastModelFetchSuccessAt = update.lastModelFetchSuccessAt - current.lastModelFetchError = "" - } else if strings.TrimSpace(update.lastModelFetchError) != "" { + if strings.TrimSpace(update.lastModelFetchError) == "" { + current.lastModelFetchError = "" + } + } + if strings.TrimSpace(update.lastModelFetchError) != "" { current.lastModelFetchError = update.lastModelFetchError } r.providerRuntime[providerName] = current @@ -420,6 +524,14 @@ func (r *ModelRegistry) LoadFromCache(ctx context.Context) (int, error) { r.mu.RLock() nameToProvider := make(map[string]core.Provider, len(r.providerNames)) nameToProviderType := make(map[string]string, len(r.providerNames)) + providerOrderNames := make([]string, 0, len(r.providers)) + for _, provider := range r.providers { + providerName := r.providerNames[provider] + if providerName == "" { + continue + } + providerOrderNames = append(providerOrderNames, providerName) + } for provider, pName := range r.providerNames { nameToProvider[pName] = provider nameToProviderType[pName] = r.providerTypes[provider] @@ -429,12 +541,14 @@ func (r *ModelRegistry) LoadFromCache(ctx context.Context) (int, error) { // Populate model maps from grouped cache structure. Unqualified lookups keep "first provider wins". newModels := make(map[string]*ModelInfo) newModelsByProvider := make(map[string]map[string]*ModelInfo) + cachedProviderTypes := make(map[string]string, len(modelCache.Providers)) for providerName, cachedProv := range modelCache.Providers { provider, ok := nameToProvider[providerName] if !ok { // Provider not configured, skip all its models continue } + cachedProviderTypes[providerName] = strings.TrimSpace(cachedProv.ProviderType) providerType := strings.TrimSpace(nameToProviderType[providerName]) if providerType == "" { providerType = strings.TrimSpace(cachedProv.ProviderType) @@ -460,6 +574,28 @@ func (r *ModelRegistry) LoadFromCache(ctx context.Context) (int, error) { newModelsByProvider[providerName] = providerModels } + configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() + if len(configuredProviderModels) > 0 { + for providerName, configuredModels := range configuredProviderModels { + provider, ok := nameToProvider[providerName] + if !ok { + continue + } + providerType := strings.TrimSpace(nameToProviderType[providerName]) + if providerType == "" { + providerType = strings.TrimSpace(cachedProviderTypes[providerName]) + } + providerModels := newModelsByProvider[providerName] + upstream := modelsResponseFromProviderMap(providerModels) + resp, reason := applyConfiguredProviderModels(providerName, providerType, configuredProviderModelsMode, configuredModels, upstream, nil, modelCache.UpdatedAt.Unix()) + if reason == configuredProviderModelsNotApplied { + continue + } + newModelsByProvider[providerName] = modelInfoMapFromResponse(resp, provider, providerName, providerType) + } + } + newModels = rebuildGlobalModelMap(newModelsByProvider, providerOrderNames) + // Load model list data from cache if available var list *modeldata.ModelList if len(modelCache.ModelListData) > 0 { @@ -1352,6 +1488,20 @@ func (r *ModelRegistry) snapshotConfigOverrides() map[string]map[string]*core.Mo return out } +func (r *ModelRegistry) snapshotConfiguredProviderModels() (map[string][]string, config.ConfiguredProviderModelsMode) { + r.mu.RLock() + defer r.mu.RUnlock() + mode := config.ResolveConfiguredProviderModelsMode(r.configuredProviderModelsMode) + if len(r.configuredProviderModels) == 0 { + return nil, mode + } + out := make(map[string][]string, len(r.configuredProviderModels)) + for provider, models := range r.configuredProviderModels { + out[provider] = slices.Clone(models) + } + return out, mode +} + // collectionEmpty reports whether a reflect.Value representing a slice, array, // or map has no elements (covering both nil and non-nil-but-zero-length), and // falls back to reflect.Value.IsZero for other kinds. This lets override- diff --git a/internal/providers/registry_cache_test.go b/internal/providers/registry_cache_test.go index 3cf53659c..c569ad5b2 100644 --- a/internal/providers/registry_cache_test.go +++ b/internal/providers/registry_cache_test.go @@ -8,6 +8,7 @@ import ( "testing" "time" + "gomodel/config" "gomodel/internal/cache/modelcache" "gomodel/internal/core" ) @@ -203,6 +204,94 @@ func TestCacheFile(t *testing.T) { } }) + t.Run("LoadFromCacheConfiguredModelsAllowlistFiltersAndAdds", func(t *testing.T) { + tmpDir := t.TempDir() + cacheFile := filepath.Join(tmpDir, "models.json") + + modelCache := modelcache.ModelCache{ + UpdatedAt: time.Now().UTC(), + Providers: map[string]modelcache.CachedProvider{ + "openrouter": { + ProviderType: "openrouter", + OwnedBy: "openrouter", + Models: []modelcache.CachedModel{ + {ID: "configured-model", Created: 123}, + {ID: "cached-extra", Created: 456}, + }, + }, + }, + } + data, _ := json.Marshal(modelCache) + if err := os.WriteFile(cacheFile, data, 0o644); err != nil { + t.Fatalf("failed to write cache file: %v", err) + } + + registry := NewModelRegistry() + registry.SetCache(modelcache.NewLocalCache(cacheFile)) + registry.SetConfiguredProviderModelsMode(config.ConfiguredProviderModelsModeAllowlist) + registry.SetProviderConfiguredModels("openrouter", []string{"missing-configured", "configured-model"}) + + mock := ®istryMockProvider{name: "openrouter"} + registry.RegisterProviderWithNameAndType(mock, "openrouter", "openrouter") + + loaded, err := registry.LoadFromCache(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if loaded != 2 { + t.Fatalf("expected 2 models loaded, got %d", loaded) + } + if registry.Supports("cached-extra") { + t.Fatal("expected allowlist mode to hide cached-extra") + } + configured := registry.GetModel("configured-model") + if configured == nil { + t.Fatal("expected configured-model to resolve") + } + if configured.Model.Created != 123 || configured.Model.OwnedBy != "openrouter" { + t.Fatalf("configured metadata = %+v, want cached metadata preserved", configured.Model) + } + missing := registry.GetModel("missing-configured") + if missing == nil { + t.Fatal("expected missing-configured to resolve") + } + if missing.Model.OwnedBy != "openrouter" { + t.Fatalf("OwnedBy = %q, want openrouter", missing.Model.OwnedBy) + } + }) + + t.Run("LoadFromCacheConfiguredModelsFallbackUsesConfiguredWhenCachedProviderMissing", func(t *testing.T) { + tmpDir := t.TempDir() + cacheFile := filepath.Join(tmpDir, "models.json") + + modelCache := modelcache.ModelCache{ + UpdatedAt: time.Now().UTC(), + Providers: map[string]modelcache.CachedProvider{}, + } + data, _ := json.Marshal(modelCache) + if err := os.WriteFile(cacheFile, data, 0o644); err != nil { + t.Fatalf("failed to write cache file: %v", err) + } + + registry := NewModelRegistry() + registry.SetCache(modelcache.NewLocalCache(cacheFile)) + registry.SetProviderConfiguredModels("vllm", []string{"meta-llama/Llama-3.1-8B-Instruct"}) + + mock := ®istryMockProvider{name: "vllm"} + registry.RegisterProviderWithNameAndType(mock, "vllm", "vllm") + + loaded, err := registry.LoadFromCache(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if loaded != 1 { + t.Fatalf("expected 1 model loaded, got %d", loaded) + } + if !registry.Supports("meta-llama/Llama-3.1-8B-Instruct") { + t.Fatal("expected configured fallback model to be loaded") + } + }) + t.Run("LoadFromCacheBackfillsMissingProviderTypeFromConfiguredProvider", func(t *testing.T) { tmpDir := t.TempDir() cacheFile := filepath.Join(tmpDir, "models.json") diff --git a/internal/providers/registry_test.go b/internal/providers/registry_test.go index 24a1009df..abf2ae249 100644 --- a/internal/providers/registry_test.go +++ b/internal/providers/registry_test.go @@ -13,6 +13,7 @@ import ( "testing" "time" + "gomodel/config" "gomodel/internal/core" "gomodel/internal/modeldata" ) @@ -135,6 +136,177 @@ func TestModelRegistry(t *testing.T) { } }) + t.Run("ConfiguredModelsFallbackModeKeepsUpstreamWhenAvailable", func(t *testing.T) { + registry := NewModelRegistry() + mock := ®istryMockProvider{ + name: "test", + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "configured-model", Object: "model", OwnedBy: "upstream"}, + {ID: "upstream-extra", Object: "model", OwnedBy: "upstream"}, + }, + }, + } + registry.RegisterProviderWithNameAndType(mock, "test", "test") + registry.SetProviderConfiguredModels("test", []string{"configured-model"}) + + err := registry.Initialize(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if registry.ModelCount() != 2 { + t.Fatalf("ModelCount() = %d, want 2", registry.ModelCount()) + } + if !registry.Supports("upstream-extra") { + t.Fatal("expected fallback mode to keep upstream-extra when upstream models are available") + } + }) + + t.Run("ConfiguredModelsFallbackModeUsesConfiguredWhenUpstreamFails", func(t *testing.T) { + registry := NewModelRegistry() + mock := ®istryMockProvider{ + name: "test", + err: errors.New("models unavailable"), + } + registry.RegisterProviderWithNameAndType(mock, "test", "test") + registry.SetProviderConfiguredModels("test", []string{" configured-model ", "configured-model", "fallback-only"}) + + err := registry.Initialize(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if registry.ModelCount() != 2 { + t.Fatalf("ModelCount() = %d, want 2", registry.ModelCount()) + } + if !registry.Supports("configured-model") || !registry.Supports("fallback-only") { + t.Fatalf("expected configured fallback models to be registered, got %+v", registry.ListModels()) + } + model := registry.GetModel("configured-model") + if model == nil { + t.Fatal("expected configured-model to resolve") + } + if model.Model.Object != "model" { + t.Fatalf("Object = %q, want model", model.Model.Object) + } + if model.Model.OwnedBy != "test" { + t.Fatalf("OwnedBy = %q, want test", model.Model.OwnedBy) + } + if model.Model.Created <= 0 { + t.Fatalf("Created = %d, want non-zero configured fallback timestamp", model.Model.Created) + } + snapshots := registry.ProviderRuntimeSnapshots() + if len(snapshots) != 1 { + t.Fatalf("expected 1 provider runtime snapshot, got %d", len(snapshots)) + } + if !strings.Contains(snapshots[0].LastModelFetchError, "models unavailable") { + t.Fatalf("LastModelFetchError = %q, want upstream error preserved", snapshots[0].LastModelFetchError) + } + if snapshots[0].LastModelFetchSuccessAt != nil { + t.Fatalf("LastModelFetchSuccessAt = %v, want nil when configured fallback handles upstream failure", snapshots[0].LastModelFetchSuccessAt) + } + }) + + t.Run("ConfiguredModelsAllowlistModeSkipsUpstreamAndUsesConfiguredModels", func(t *testing.T) { + registry := NewModelRegistry() + registry.SetConfiguredProviderModelsMode(config.ConfiguredProviderModelsModeAllowlist) + var listCount atomic.Int32 + mock := &countingRegistryMockProvider{ + listCount: &listCount, + registryMockProvider: ®istryMockProvider{ + name: "test", + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "configured-model", Object: "model", OwnedBy: "upstream", Created: 123}, + {ID: "upstream-extra", Object: "model", OwnedBy: "upstream", Created: 456}, + }, + }, + }, + } + registry.RegisterProviderWithNameAndType(mock, "test", "test-type") + registry.SetProviderConfiguredModels("test", []string{"missing-configured", "configured-model"}) + + err := registry.Initialize(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if listCount.Load() != 0 { + t.Fatalf("ListModels calls = %d, want 0", listCount.Load()) + } + if registry.ModelCount() != 2 { + t.Fatalf("ModelCount() = %d, want 2", registry.ModelCount()) + } + if registry.Supports("upstream-extra") { + t.Fatal("expected allowlist mode to hide upstream-extra") + } + configured := registry.GetModel("configured-model") + if configured == nil { + t.Fatal("expected configured-model to resolve") + } + if configured.Model.Created <= 0 { + t.Fatalf("configured.Model.Created = %d in configured model %+v, want non-zero timestamp", configured.Model.Created, configured.Model) + } + if configured.Model.OwnedBy != "test-type" { + t.Fatalf("configured.Model.OwnedBy = %q in configured model %+v, want test-type", configured.Model.OwnedBy, configured.Model) + } + snapshots := registry.ProviderRuntimeSnapshots() + if len(snapshots) != 1 { + t.Fatalf("expected 1 provider runtime snapshot, got %d", len(snapshots)) + } + if snapshots[0].LastModelFetchSuccessAt != nil { + t.Fatalf("LastModelFetchSuccessAt = %v, want nil when allowlist skips upstream ListModels", snapshots[0].LastModelFetchSuccessAt) + } + missing := registry.GetModel("missing-configured") + if missing == nil { + t.Fatal("expected missing-configured to be added") + } + if missing.Model.OwnedBy != "test-type" { + t.Fatalf("OwnedBy = %q, want test-type", missing.Model.OwnedBy) + } + }) + + t.Run("ConfiguredModelsAllowlistModeUsesUpstreamWhenNoConfiguredModels", func(t *testing.T) { + registry := NewModelRegistry() + registry.SetConfiguredProviderModelsMode(config.ConfiguredProviderModelsModeAllowlist) + var listCount atomic.Int32 + mock := &countingRegistryMockProvider{ + listCount: &listCount, + registryMockProvider: ®istryMockProvider{ + name: "test", + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "upstream-model", Object: "model", OwnedBy: "upstream"}, + }, + }, + }, + } + registry.RegisterProviderWithNameAndType(mock, "test", "test-type") + + err := registry.Initialize(context.Background()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if listCount.Load() != 1 { + t.Fatalf("ListModels calls = %d, want 1", listCount.Load()) + } + if !registry.Supports("upstream-model") { + t.Fatal("expected upstream-model to resolve when provider has no configured models") + } + snapshots := registry.ProviderRuntimeSnapshots() + if len(snapshots) != 1 { + t.Fatalf("expected 1 provider runtime snapshot, got %d", len(snapshots)) + } + if snapshots[0].LastModelFetchSuccessAt == nil { + t.Fatal("expected LastModelFetchSuccessAt when upstream ListModels succeeds") + } + }) + t.Run("GetProvider", func(t *testing.T) { registry := NewModelRegistry() mock := ®istryMockProvider{ diff --git a/internal/providers/vllm/vllm.go b/internal/providers/vllm/vllm.go index 1ca7a29a0..765b5857d 100644 --- a/internal/providers/vllm/vllm.go +++ b/internal/providers/vllm/vllm.go @@ -38,10 +38,9 @@ func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Prov rootBaseURL := passthroughBaseURL(baseURL) return &Provider{ compatible: openai.NewCompatibleProvider(cfg.APIKey, opts, openai.CompatibleProviderConfig{ - ProviderName: "vllm", - BaseURL: baseURL, - SetHeaders: setHeaders, - ConfiguredModels: opts.Models, + ProviderName: "vllm", + BaseURL: baseURL, + SetHeaders: setHeaders, }), rootClient: llmclient.New(llmclient.Config{ ProviderName: "vllm",