diff --git a/.env.template b/.env.template index 9af3c3551..16fd43a52 100644 --- a/.env.template +++ b/.env.template @@ -9,6 +9,12 @@ # Accepts values like "10M", "1G", "500K" (default: 10M) # BODY_SIZE_LIMIT=10M +# Enable/disable Swagger UI at /swagger/index.html (default: true) +# SWAGGER_ENABLED=true + +# Enable/disable pprof profiling routes at /debug/pprof/* (default: false) +# PPROF_ENABLED=false + # Enable/disable provider-native passthrough routes under /p/{provider}/{endpoint} (default: true) # ENABLE_PASSTHROUGH_ROUTES=true @@ -37,20 +43,54 @@ # METRICS_ENDPOINT=/metrics # Cache Configuration -# Type: -# - "local" (default) for single instance, -# - "redis" for multiple instances -# CACHE_TYPE=local +# Model cache uses the local filesystem by default. +# Set REDIS_URL to use Redis-backed caching instead. -# Redis Configuration (only used when CACHE_TYPE=redis) +# Redis Configuration # REDIS_URL=redis://localhost:6379 # REDIS_KEY_MODELS=gomodel:models # REDIS_TTL_MODELS=86400 +# How often to refresh the model registry cache in seconds (default: 3600) +# CACHE_REFRESH_INTERVAL=3600 # REDIS_KEY_RESPONSES=gomodel:response: # REDIS_TTL_RESPONSES=3600 # Opt-in when config.yaml has no cache.response.simple block (e.g. env-only deploys). Omit otherwise. # RESPONSE_CACHE_SIMPLE_ENABLED=true +# Opt-in when config.yaml has no cache.response.semantic block (e.g. env-only deploys). Omit otherwise. +# SEMANTIC_CACHE_ENABLED=true +# Similarity threshold between 0 and 1 (default: 0.92) +# SEMANTIC_CACHE_THRESHOLD=0.92 +# Semantic cache entry TTL in seconds (default: 3600) +# SEMANTIC_CACHE_TTL=3600 +# Number of recent conversation messages to embed (default: 3) +# SEMANTIC_CACHE_MAX_CONV_MESSAGES=3 +# Exclude the system prompt from semantic cache keys (default: false) +# SEMANTIC_CACHE_EXCLUDE_SYSTEM_PROMPT=false +# Embedding provider name used for semantic cache +# SEMANTIC_CACHE_EMBEDDER_PROVIDER=openai +# Optional embedding model override +# SEMANTIC_CACHE_EMBEDDER_MODEL=text-embedding-3-small +# Vector store backend: qdrant, pgvector, pinecone, or weaviate +# SEMANTIC_CACHE_VECTOR_STORE_TYPE=qdrant +# Qdrant +# SEMANTIC_CACHE_QDRANT_URL=http://localhost:6333 +# SEMANTIC_CACHE_QDRANT_COLLECTION=gomodel_semantic +# SEMANTIC_CACHE_QDRANT_API_KEY= +# pgvector +# SEMANTIC_CACHE_PGVECTOR_URL=postgres://user:pass@localhost:5432/gomodel +# SEMANTIC_CACHE_PGVECTOR_TABLE=gomodel_semantic_cache +# SEMANTIC_CACHE_PGVECTOR_DIMENSION=1536 +# Pinecone +# SEMANTIC_CACHE_PINECONE_HOST=https://your-index.svc.region.pinecone.io +# SEMANTIC_CACHE_PINECONE_API_KEY= +# SEMANTIC_CACHE_PINECONE_NAMESPACE= +# SEMANTIC_CACHE_PINECONE_DIMENSION=1536 +# Weaviate +# SEMANTIC_CACHE_WEAVIATE_URL=http://localhost:8080 +# SEMANTIC_CACHE_WEAVIATE_CLASS=GomodelSemanticCache +# SEMANTIC_CACHE_WEAVIATE_API_KEY= + # Optional: Custom cache directory for local file cache # GOMODEL_CACHE_DIR=.cache @@ -58,6 +98,38 @@ # Set to empty string to disable (default: ENTERPILOT/ai-model-list on GitHub) # MODEL_LIST_URL=https://raw.githubusercontent.com/ENTERPILOT/ai-model-list/refs/heads/main/models.min.json +# Model Access Configuration +# Process-wide default for concrete provider models when no persisted override exists (default: true) +# Set to false to keep models unavailable until explicitly enabled by a model override +# MODELS_ENABLED_BY_DEFAULT=true + +# Fallback & Workflow Configuration +# Default translated-route fallback mode: auto, manual, or off (default: auto) +# FEATURE_FALLBACK_MODE=auto +# JSON file mapping model selectors to ordered fallback lists +# Required when FEATURE_FALLBACK_MODE=manual (default example: config/fallback.example.json) +# FALLBACK_MANUAL_RULES_PATH=config/fallback.example.json +# How often to refresh persisted execution plans from storage (default: 1m) +# EXECUTION_PLAN_REFRESH_INTERVAL=1m + +# LLM Client Resilience Configuration +# Retry attempts for upstream provider calls (default: 3) +# RETRY_MAX_RETRIES=3 +# Initial retry backoff duration (default: 1s) +# RETRY_INITIAL_BACKOFF=1s +# Maximum retry backoff duration (default: 30s) +# RETRY_MAX_BACKOFF=30s +# Exponential backoff factor (default: 2.0) +# RETRY_BACKOFF_FACTOR=2.0 +# Random jitter factor applied to retry delays (default: 0.1) +# RETRY_JITTER_FACTOR=0.1 +# Consecutive failures before opening the circuit breaker (default: 5) +# CIRCUIT_BREAKER_FAILURE_THRESHOLD=5 +# Consecutive successes required to close the circuit breaker (default: 2) +# CIRCUIT_BREAKER_SUCCESS_THRESHOLD=2 +# Circuit breaker open-state timeout duration (default: 30s) +# CIRCUIT_BREAKER_TIMEOUT=30s + # ============================================================================= # Admin API & Dashboard Configuration # ============================================================================= @@ -160,18 +232,23 @@ # ============================================================================= # OpenAI # OPENAI_API_KEY=sk-... +# OPENAI_BASE_URL=https://api.openai.com/v1 # Anthropic # ANTHROPIC_API_KEY=sk-ant-... +# ANTHROPIC_BASE_URL=https://api.anthropic.com/v1 # Google Gemini # GEMINI_API_KEY=... +# GEMINI_BASE_URL=https://generativelanguage.googleapis.com/v1beta/openai # xAI (Grok) # XAI_API_KEY=... +# XAI_BASE_URL=https://api.x.ai/v1 # Groq # GROQ_API_KEY=gsk_... +# GROQ_BASE_URL=https://api.groq.com/openai/v1 # OpenRouter (default base URL: https://openrouter.ai/api/v1) # OPENROUTER_API_KEY=sk-or-... @@ -189,6 +266,7 @@ # ORACLE_BASE_URL=https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/v1 # Ollama (local LLM server) -# Note: Ollama doesn't require an API key +# Note: Ollama doesn't require an API key, but one can be sent for secured deployments +# OLLAMA_API_KEY=... # Set base URL to enable (default: http://localhost:11434/v1) # OLLAMA_BASE_URL=http://localhost:11434/v1 diff --git a/CLAUDE.md b/CLAUDE.md index cc101ea6a..faff4d668 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -109,7 +109,7 @@ Full reference: `.env.template` and `config/config.yaml` - **Storage:** `STORAGE_TYPE` (sqlite), `SQLITE_PATH` (data/gomodel.db), `POSTGRES_URL`, `MONGODB_URL` - **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_TYPE` (local), `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`. +- **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`. - **HTTP client:** `HTTP_TIMEOUT` (600s), `HTTP_RESPONSE_HEADER_TIMEOUT` (600s) - **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) diff --git a/README.md b/README.md index 97d670fb1..96c4cf14c 100644 --- a/README.md +++ b/README.md @@ -180,7 +180,6 @@ Key settings: | `ENABLE_PASSTHROUGH_ROUTES` | `true` | Enable provider-native passthrough routes under `/p/{provider}/...` | | `ALLOW_PASSTHROUGH_V1_ALIAS` | `true` | Allow `/p/{provider}/v1/...` aliases while keeping `/p/{provider}/...` canonical | | `ENABLED_PASSTHROUGH_PROVIDERS` | `openai,anthropic` | Comma-separated list of enabled passthrough providers | -| `CACHE_TYPE` | `local` | Cache backend (`local` or `redis`) | | `STORAGE_TYPE` | `sqlite` | Storage backend (`sqlite`, `postgresql`, `mongodb`) | | `METRICS_ENABLED` | `false` | Enable Prometheus metrics | | `LOGGING_ENABLED` | `false` | Enable audit logging | diff --git a/config/config.example.yaml b/config/config.example.yaml index 8bd63fb28..4c238fed4 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -12,6 +12,9 @@ server: allow_passthrough_v1_alias: true # allow /p/{provider}/v1/... while keeping /p/{provider}/... canonical enabled_passthrough_providers: ["openai", "anthropic"] # providers enabled on /p/{provider}/... +models: + enabled_by_default: true # env: MODELS_ENABLED_BY_DEFAULT; when false, concrete models stay unavailable until explicitly enabled by a model override + cache: model: refresh_interval: 3600 # how often to refresh the model registry (seconds, default: 3600) diff --git a/config/config.go b/config/config.go index fb1678717..0584cdee8 100644 --- a/config/config.go +++ b/config/config.go @@ -31,6 +31,7 @@ var bodySizeLimitRegex = regexp.MustCompile(`(?i)^(\d+)([KMG])?B?$`) // Config holds the application configuration. type Config struct { Server ServerConfig `yaml:"server"` + Models ModelsConfig `yaml:"models"` Cache CacheConfig `yaml:"cache"` Storage StorageConfig `yaml:"storage"` Logging LogConfig `yaml:"logging"` @@ -128,6 +129,13 @@ type FallbackModelOverride struct { Mode FallbackMode `yaml:"mode" json:"mode"` } +// ModelsConfig holds global model access defaults. +type ModelsConfig struct { + // EnabledByDefault controls whether concrete provider models are available + // when no persisted override exists. Default: true. + EnabledByDefault bool `yaml:"enabled_by_default" env:"MODELS_ENABLED_BY_DEFAULT"` +} + // FallbackConfig holds translated-route model fallback policy. type FallbackConfig struct { // DefaultMode controls the fallback behavior when no per-model override exists. @@ -863,6 +871,9 @@ func buildDefaultConfig() *Config { "anthropic", }, }, + Models: ModelsConfig{ + EnabledByDefault: true, + }, Cache: CacheConfig{ Model: ModelCacheConfig{ RefreshInterval: 3600, diff --git a/docker-compose.yaml b/docker-compose.yaml index 8c3c161bd..4e5f24d1d 100644 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -14,7 +14,6 @@ services: - .env environment: # Cache configuration - - CACHE_TYPE=redis - REDIS_URL=redis://redis:6379 # Metrics - METRICS_ENABLED=true diff --git a/docs/2026-03-23_benchmark_scripts/gateway-comparison/run-benchmark.sh b/docs/2026-03-23_benchmark_scripts/gateway-comparison/run-benchmark.sh index b0f7601cd..9e4b97c95 100755 --- a/docs/2026-03-23_benchmark_scripts/gateway-comparison/run-benchmark.sh +++ b/docs/2026-03-23_benchmark_scripts/gateway-comparison/run-benchmark.sh @@ -231,7 +231,6 @@ LOGGING_ENABLED=false \ USAGE_ENABLED=false \ STORAGE_TYPE=sqlite \ SQLITE_PATH="/tmp/gomodel-bench.db" \ -CACHE_TYPE=local \ GOMODEL_CACHE_DIR="/tmp/gomodel-bench-cache" \ ADMIN_ENDPOINTS_ENABLED=false \ ADMIN_UI_ENABLED=false \ diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 383865a6c..5ce56c56e 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -47,7 +47,6 @@ The most common way to configure GOModel. Set any of the variables below to over | Variable | Description | Default | | ------------------- | --------------------------------- | ---------------- | -| `CACHE_TYPE` | Cache backend: `local` or `redis` | `local` | | `GOMODEL_CACHE_DIR` | Directory for local cache files | `.cache` | | `REDIS_URL` | Redis connection URL | _(empty)_ | | `REDIS_KEY_MODELS` | Redis key for model cache | `gomodel:models` | @@ -192,9 +191,9 @@ server: master_key: "my-secret-key" cache: - type: redis - redis: - url: "redis://my-redis:6379" + model: + redis: + url: "redis://my-redis:6379" providers: openai: diff --git a/helm/templates/_helpers.tpl b/helm/templates/_helpers.tpl index 5af376aa6..286edc2ba 100644 --- a/helm/templates/_helpers.tpl +++ b/helm/templates/_helpers.tpl @@ -84,17 +84,6 @@ Determine the Redis URL - either from values or auto-generated for subchart {{- end }} {{- end }} -{{/* -Determine the cache type - auto-set to redis if subchart is enabled -*/}} -{{- define "gomodel.cacheType" -}} -{{- if .Values.redis.enabled }} -{{- "redis" }} -{{- else }} -{{- .Values.cache.type }} -{{- end }} -{{- end }} - {{/* Create the image reference */}} diff --git a/helm/templates/configmap.yaml b/helm/templates/configmap.yaml index a003ff66b..e3fef84b0 100644 --- a/helm/templates/configmap.yaml +++ b/helm/templates/configmap.yaml @@ -7,8 +7,7 @@ metadata: data: PORT: {{ .Values.server.port | quote }} BODY_SIZE_LIMIT: {{ .Values.server.bodySizeLimit | quote }} - CACHE_TYPE: {{ include "gomodel.cacheType" . | quote }} - {{- if or .Values.redis.enabled (eq .Values.cache.type "redis") }} + {{- if or .Values.redis.enabled .Values.cache.redis.url }} REDIS_KEY_MODELS: {{ .Values.cache.redis.keyModels | default "gomodel:models" | quote }} REDIS_KEY_RESPONSES: {{ .Values.cache.redis.keyResponses | default "gomodel:response:" | quote }} REDIS_TTL_MODELS: {{ .Values.cache.redis.ttlModels | default 86400 | quote }} diff --git a/helm/templates/deployment.yaml b/helm/templates/deployment.yaml index ddcf4ce46..a4fdfa2b4 100644 --- a/helm/templates/deployment.yaml +++ b/helm/templates/deployment.yaml @@ -54,12 +54,7 @@ spec: name: {{ include "gomodel.fullname" . }} key: BODY_SIZE_LIMIT # Cache configuration - - name: CACHE_TYPE - valueFrom: - configMapKeyRef: - name: {{ include "gomodel.fullname" . }} - key: CACHE_TYPE - {{- if or .Values.redis.enabled (eq .Values.cache.type "redis") }} + {{- if or .Values.redis.enabled .Values.cache.redis.url }} - name: REDIS_URL valueFrom: secretKeyRef: diff --git a/helm/templates/secret.yaml b/helm/templates/secret.yaml index 1f4047c97..57f97babd 100644 --- a/helm/templates/secret.yaml +++ b/helm/templates/secret.yaml @@ -8,7 +8,7 @@ metadata: type: Opaque stringData: {{- include "gomodel.providerSecretData" . | nindent 2 }} - {{- if or .Values.redis.enabled (eq .Values.cache.type "redis") }} + {{- if or .Values.redis.enabled .Values.cache.redis.url }} REDIS_URL: {{ include "gomodel.redisUrl" . | quote }} {{- end }} {{- end }} diff --git a/helm/values.schema.json b/helm/values.schema.json index ca501f8d3..2e44f9d9c 100644 --- a/helm/values.schema.json +++ b/helm/values.schema.json @@ -27,33 +27,6 @@ } } }, - { - "if": { - "properties": { - "cache": { - "properties": { "type": { "const": "redis" } } - }, - "redis": { - "properties": { "enabled": { "const": false } } - } - } - }, - "then": { - "properties": { - "cache": { - "properties": { - "redis": { - "properties": { - "url": { "minLength": 1 } - }, - "required": ["url"] - } - }, - "required": ["redis"] - } - } - } - }, { "if": { "properties": { @@ -262,10 +235,6 @@ "cache": { "type": "object", "properties": { - "type": { - "type": "string", - "enum": ["local", "redis"] - }, "redis": { "type": "object", "properties": { diff --git a/helm/values.yaml b/helm/values.yaml index 2879d19dc..12e85043e 100644 --- a/helm/values.yaml +++ b/helm/values.yaml @@ -91,12 +91,10 @@ providers: # Cache configuration cache: - # -- Cache type: "local" or "redis" - type: "redis" - redis: # -- Redis connection URL (e.g., "redis://redis:6379") - # If redis.enabled is true, this is auto-configured to use the subchart + # If redis.enabled is true, this is auto-configured to use the subchart. + # Set this explicitly to use an external Redis instance instead. url: "" # -- Redis key prefix for storing the model cache keyModels: "gomodel:models" diff --git a/internal/admin/dashboard/static/css/dashboard.css b/internal/admin/dashboard/static/css/dashboard.css index e90102eb5..3c7a55339 100644 --- a/internal/admin/dashboard/static/css/dashboard.css +++ b/internal/admin/dashboard/static/css/dashboard.css @@ -1199,6 +1199,34 @@ td.col-price { color: var(--text-muted); } +.model-access-state-badge { + display: inline-flex; + align-items: center; + padding: 5px 10px; + border-radius: 999px; + border: 1px solid var(--border); + background: var(--bg); + color: var(--text); + font-size: 12px; + line-height: 1; + white-space: nowrap; +} + +.model-access-state-badge.is-enabled { + border-color: color-mix(in srgb, var(--success) 40%, var(--border)); + color: color-mix(in srgb, var(--success) 70%, var(--text)); +} + +.model-access-state-badge.is-restricted { + border-color: color-mix(in srgb, var(--accent) 40%, var(--border)); + color: color-mix(in srgb, var(--accent) 70%, var(--text)); +} + +.model-access-state-badge.is-disabled { + border-color: color-mix(in srgb, var(--danger) 40%, var(--border)); + color: color-mix(in srgb, var(--danger) 70%, var(--text)); +} + .data-table tr.alias-row.is-valid td { background: var(--alias-row-valid-bg); } @@ -1233,6 +1261,14 @@ td.col-price { opacity: 0.72; } +.data-table tr.model-access-disabled-row td { + background: color-mix(in srgb, var(--danger) 6%, var(--bg-surface)); +} + +.data-table tr.model-access-disabled-row:hover td { + background: color-mix(in srgb, var(--danger) 10%, var(--bg-surface-hover)); +} + /* Execution Plans */ .execution-plan-page-note { margin-top: 6px; diff --git a/internal/admin/dashboard/static/js/modules/aliases.js b/internal/admin/dashboard/static/js/modules/aliases.js index 20e670177..de939cb51 100644 --- a/internal/admin/dashboard/static/js/modules/aliases.js +++ b/internal/admin/dashboard/static/js/modules/aliases.js @@ -3,6 +3,7 @@ return { aliases: [], aliasesAvailable: true, + modelOverridesAvailable: true, displayModels: [], aliasLoading: false, aliasError: '', @@ -20,6 +21,20 @@ description: '', enabled: true }, + modelOverrideFormOpen: false, + modelOverrideSubmitting: false, + modelOverrideError: '', + modelOverrideNotice: '', + modelOverrideFormHasExistingOverride: false, + modelOverrideFormDefaultEnabled: true, + modelOverrideFormEffectiveEnabled: true, + modelOverrideFormDisplayName: '', + modelOverrideForm: { + selector: '', + enabled: false, + force_disabled: false, + allowed_only_for_user_paths: '' + }, buildDisplayModels() { const rows = this.models.map((model) => ({ @@ -31,6 +46,7 @@ model: model.model, is_alias: false, alias: null, + access: model && model.access ? model.access : null, kind_badge: '', masking_alias: null, alias_state_class: '', @@ -79,6 +95,7 @@ model: targetModel ? targetModel.model : { id: alias.name, object: 'model' }, is_alias: true, alias, + access: null, kind_badge: 'Alias', masking_alias: null, alias_state_class: this.aliasStateClass(alias), @@ -187,6 +204,9 @@ if (!row.is_alias && row.masking_alias) { classes.push('masked-model-row'); } + if (!row.is_alias && row.access && row.access.effective_enabled === false) { + classes.push('model-access-disabled-row'); + } return classes.join(' '); }, @@ -236,6 +256,59 @@ this.aliasForm = this.defaultAliasForm(); }, + defaultModelOverrideForm() { + return { + selector: '', + enabled: false, + force_disabled: false, + allowed_only_for_user_paths: '' + }; + }, + + normalizeModelOverridePaths(raw) { + return String(raw || '') + .split(/\r?\n|,/) + .map((value) => String(value || '').trim()) + .filter(Boolean); + }, + + openModelOverrideEdit(row) { + if (!row || row.is_alias) { + return; + } + + const access = row.access || {}; + const override = access.override || null; + const allowedPaths = override && Array.isArray(override.allowed_only_for_user_paths) + ? override.allowed_only_for_user_paths + : (Array.isArray(access.allowed_only_for_user_paths) ? access.allowed_only_for_user_paths : []); + + this.modelOverrideFormOpen = true; + this.modelOverrideError = ''; + this.modelOverrideNotice = ''; + this.modelOverrideFormHasExistingOverride = Boolean(override); + this.modelOverrideFormDefaultEnabled = access.default_enabled !== false; + this.modelOverrideFormEffectiveEnabled = access.effective_enabled !== false; + this.modelOverrideFormDisplayName = row.display_name || this.qualifiedModelName(row) || ''; + this.modelOverrideForm = { + selector: access.selector || this.qualifiedModelName(row), + enabled: Boolean(override && override.enabled === true), + force_disabled: Boolean(override && override.force_disabled), + allowed_only_for_user_paths: allowedPaths.join('\n') + }; + }, + + closeModelOverrideForm() { + this.modelOverrideFormOpen = false; + this.modelOverrideSubmitting = false; + this.modelOverrideError = ''; + this.modelOverrideFormHasExistingOverride = false; + this.modelOverrideFormDefaultEnabled = true; + this.modelOverrideFormEffectiveEnabled = true; + this.modelOverrideFormDisplayName = ''; + this.modelOverrideForm = this.defaultModelOverrideForm(); + }, + async aliasResponseMessage(res, fallback) { try { const payload = await res.json(); @@ -267,6 +340,52 @@ return 'Active'; }, + modelAccessStateText(access) { + if (!access) return 'Default'; + if (access.force_disabled) return 'Force Disabled'; + if (access.effective_enabled === false) { + return access.default_enabled === false ? 'Disabled by Default' : 'Disabled'; + } + if (access.override && access.override.enabled === true && access.default_enabled === false) { + return 'Explicitly Enabled'; + } + if (Array.isArray(access.allowed_only_for_user_paths) && access.allowed_only_for_user_paths.length > 0) { + return 'Restricted'; + } + return 'Enabled'; + }, + + modelAccessStateClass(access) { + if (!access) return ''; + if (access.force_disabled || access.effective_enabled === false) return 'is-disabled'; + if (Array.isArray(access.allowed_only_for_user_paths) && access.allowed_only_for_user_paths.length > 0) { + return 'is-restricted'; + } + return 'is-enabled'; + }, + + modelAccessSummary(access) { + if (!access) { + return ''; + } + + const parts = []; + if (access.force_disabled) { + parts.push('Force disabled globally'); + } else if (access.effective_enabled === false) { + parts.push(access.default_enabled === false ? 'Disabled until explicitly enabled' : 'Disabled'); + } else if (access.override && access.override.enabled === true && access.default_enabled === false) { + parts.push('Explicitly enabled'); + } + + const allowed = Array.isArray(access.allowed_only_for_user_paths) ? access.allowed_only_for_user_paths : []; + if (allowed.length > 0) { + parts.push('Allowed only for ' + allowed.join(', ')); + } + + return parts.join(' · '); + }, + async toggleAliasEnabled(alias) { if (!alias || !alias.name || this.aliasTogglingName === alias.name) { return; @@ -573,6 +692,115 @@ } finally { this.aliasDeletingName = ''; } + }, + + async submitModelOverrideForm() { + const selector = String(this.modelOverrideForm.selector || '').trim(); + const allowedOnlyForUserPaths = this.normalizeModelOverridePaths(this.modelOverrideForm.allowed_only_for_user_paths); + const forceDisabled = Boolean(this.modelOverrideForm.force_disabled); + const enabled = Boolean(this.modelOverrideForm.enabled) && !forceDisabled; + + if (!selector) { + this.modelOverrideError = 'Model selector is required.'; + return; + } + if (!enabled && !forceDisabled && allowedOnlyForUserPaths.length === 0) { + this.modelOverrideError = this.modelOverrideFormHasExistingOverride + ? 'Choose an access policy or remove the override.' + : 'Choose an access policy before saving.'; + return; + } + + this.modelOverrideSubmitting = true; + this.modelOverrideError = ''; + this.modelOverrideNotice = ''; + + const payload = { + force_disabled: forceDisabled, + allowed_only_for_user_paths: allowedOnlyForUserPaths + }; + if (enabled) { + payload.enabled = true; + } + + try { + const res = await fetch('/admin/api/v1/model-overrides/' + encodeURIComponent(selector), { + method: 'PUT', + headers: this.headers(), + body: JSON.stringify(payload) + }); + if (res.status === 503) { + this.modelOverridesAvailable = false; + this.modelOverrideError = 'Model overrides feature is unavailable.'; + return; + } + this.modelOverridesAvailable = true; + if (res.status === 401) { + this.authError = true; + this.needsAuth = true; + this.modelOverrideError = 'Authentication required.'; + return; + } + if (!res.ok) { + this.modelOverrideError = await this.aliasResponseMessage(res, 'Failed to save model access.'); + return; + } + + await this.fetchModels(); + this.closeModelOverrideForm(); + this.modelOverrideNotice = 'Model access saved.'; + } catch (e) { + console.error('Failed to save model override:', e); + this.modelOverrideError = 'Failed to save model access.'; + } finally { + this.modelOverrideSubmitting = false; + } + }, + + async deleteModelOverride() { + const selector = String(this.modelOverrideForm.selector || '').trim(); + if (!selector || !this.modelOverrideFormHasExistingOverride) { + return; + } + if (!window.confirm('Remove the model override for "' + selector + '"?')) { + return; + } + + this.modelOverrideSubmitting = true; + this.modelOverrideError = ''; + this.modelOverrideNotice = ''; + + try { + const res = await fetch('/admin/api/v1/model-overrides/' + encodeURIComponent(selector), { + method: 'DELETE', + headers: this.headers() + }); + if (res.status === 503) { + this.modelOverridesAvailable = false; + this.modelOverrideError = 'Model overrides feature is unavailable.'; + return; + } + this.modelOverridesAvailable = true; + if (res.status === 401) { + this.authError = true; + this.needsAuth = true; + this.modelOverrideError = 'Authentication required.'; + return; + } + if (res.status !== 404 && !res.ok) { + this.modelOverrideError = await this.aliasResponseMessage(res, 'Failed to remove model override.'); + return; + } + + await this.fetchModels(); + this.closeModelOverrideForm(); + this.modelOverrideNotice = 'Model override removed.'; + } catch (e) { + console.error('Failed to delete model override:', e); + this.modelOverrideError = 'Failed to remove model override.'; + } finally { + this.modelOverrideSubmitting = false; + } } }; } diff --git a/internal/admin/dashboard/templates/index.html b/internal/admin/dashboard/templates/index.html index 375273edd..a3f01c358 100644 --- a/internal/admin/dashboard/templates/index.html +++ b/internal/admin/dashboard/templates/index.html @@ -611,8 +611,13 @@

Registered Models

Aliases feature is unavailable.
+
+ Model overrides feature is unavailable. +
+
+
@@ -690,6 +695,57 @@

+
+
+
+

Model access override

+

+
+ +
+ + + +

+ The selector uses the concrete {provider_name}/{model} shape shown in the models list. allowed_only_for_user_paths is matched against the effective managed API key user_path. +

+

+ + + + + + + +

+ Leave blank to allow every user path. When set, descendant paths are also allowed. +

+ +
+ +
+ + + +
+
+

+
@@ -747,7 +803,16 @@

- +
+ + +
@@ -807,7 +872,16 @@

- +
+ + +
@@ -867,7 +941,16 @@

- +
+ + +
@@ -929,7 +1012,16 @@

- +
+ + +
@@ -991,7 +1083,16 @@

- +
+ + +
@@ -1053,7 +1154,16 @@

- +
+ + +
diff --git a/internal/admin/handler.go b/internal/admin/handler.go index 86bb8c443..733d39448 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -23,21 +23,23 @@ import ( "gomodel/internal/core" "gomodel/internal/executionplans" "gomodel/internal/guardrails" + "gomodel/internal/modeloverrides" "gomodel/internal/providers" "gomodel/internal/usage" ) // Handler serves admin API endpoints. type Handler struct { - usageReader usage.UsageReader - auditReader auditlog.Reader - registry *providers.ModelRegistry - authKeys *authkeys.Service - aliases *aliases.Service - plans *executionplans.Service - guardrails guardrails.Catalog - guardrailDefs *guardrails.Service - runtimeConfig DashboardConfigResponse + usageReader usage.UsageReader + auditReader auditlog.Reader + registry *providers.ModelRegistry + authKeys *authkeys.Service + aliases *aliases.Service + modelOverrides *modeloverrides.Service + plans *executionplans.Service + guardrails guardrails.Catalog + guardrailDefs *guardrails.Service + runtimeConfig DashboardConfigResponse mutationMu sync.Mutex } @@ -87,6 +89,13 @@ func WithAuthKeys(service *authkeys.Service) Option { } } +// WithModelOverrides enables model override administration endpoints. +func WithModelOverrides(service *modeloverrides.Service) Option { + return func(h *Handler) { + h.modelOverrides = service + } +} + // WithExecutionPlans enables execution-plan administration endpoints. func WithExecutionPlans(service *executionplans.Service) Option { return func(h *Handler) { @@ -667,9 +676,23 @@ func (h *Handler) AuditConversation(c *echo.Context) error { // @Success 200 {array} providers.ModelWithProvider // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/models [get] +type modelAccessResponse struct { + Selector string `json:"selector"` + DefaultEnabled bool `json:"default_enabled"` + EffectiveEnabled bool `json:"effective_enabled"` + ForceDisabled bool `json:"force_disabled"` + AllowedOnlyForUserPaths []string `json:"allowed_only_for_user_paths,omitempty"` + Override *modeloverrides.Override `json:"override,omitempty"` +} + +type modelInventoryResponse struct { + providers.ModelWithProvider + Access modelAccessResponse `json:"access"` +} + func (h *Handler) ListModels(c *echo.Context) error { if h.registry == nil { - return c.JSON(http.StatusOK, []providers.ModelWithProvider{}) + return c.JSON(http.StatusOK, []modelInventoryResponse{}) } cat := core.ModelCategory(c.QueryParam("category")) @@ -689,8 +712,50 @@ func (h *Handler) ListModels(c *echo.Context) error { if models == nil { models = []providers.ModelWithProvider{} } + if h.modelOverrides == nil { + response := make([]modelInventoryResponse, 0, len(models)) + for _, model := range models { + selector := core.ModelSelector{ + Provider: strings.TrimSpace(model.ProviderName), + Model: strings.TrimSpace(model.Model.ID), + } + response = append(response, modelInventoryResponse{ + ModelWithProvider: model, + Access: modelAccessResponse{ + Selector: selector.QualifiedModel(), + DefaultEnabled: true, + EffectiveEnabled: true, + }, + }) + } + return c.JSON(http.StatusOK, response) + } + + response := make([]modelInventoryResponse, 0, len(models)) + for _, model := range models { + selector := core.ModelSelector{ + Provider: strings.TrimSpace(model.ProviderName), + Model: strings.TrimSpace(model.Model.ID), + } + effective := h.modelOverrides.EffectiveState(selector) + access := modelAccessResponse{ + Selector: effective.Selector, + DefaultEnabled: effective.DefaultEnabled, + EffectiveEnabled: effective.Enabled, + ForceDisabled: effective.ForceDisabled, + AllowedOnlyForUserPaths: append([]string(nil), effective.AllowedOnlyForUserPaths...), + } + if override, ok := h.modelOverrides.Get(selector.QualifiedModel()); ok && override != nil { + overrideCopy := *override + access.Override = &overrideCopy + } + response = append(response, modelInventoryResponse{ + ModelWithProvider: model, + Access: access, + }) + } - return c.JSON(http.StatusOK, models) + return c.JSON(http.StatusOK, response) } // isValidCategory returns true if cat is a recognized model category. @@ -727,6 +792,12 @@ type upsertAliasRequest struct { Enabled *bool `json:"enabled,omitempty"` } +type upsertModelOverrideRequest struct { + Enabled *bool `json:"enabled,omitempty"` + ForceDisabled bool `json:"force_disabled,omitempty"` + AllowedOnlyForUserPaths []string `json:"allowed_only_for_user_paths,omitempty"` +} + type upsertGuardrailRequest struct { Type string `json:"type"` Description string `json:"description,omitempty"` @@ -760,6 +831,10 @@ func (h *Handler) aliasesUnavailableError() error { return featureUnavailableError("aliases feature is unavailable") } +func (h *Handler) modelOverridesUnavailableError() error { + return featureUnavailableError("model overrides feature is unavailable") +} + func (h *Handler) authKeysUnavailableError() error { return featureUnavailableError("auth keys feature is unavailable") } @@ -782,6 +857,16 @@ func aliasWriteError(err error) error { return err } +func modelOverrideWriteError(err error) error { + if err == nil { + return nil + } + if modeloverrides.IsValidationError(err) { + return core.NewInvalidRequestError(err.Error(), err) + } + return err +} + func executionPlanWriteError(err error) error { if err == nil { return nil @@ -839,6 +924,100 @@ func deactivateByID( return c.NoContent(http.StatusNoContent) } +func deleteByName( + c *echo.Context, + unavailableErr error, + paramName string, + decode func(string) (string, error), + deleteFunc func(context.Context, string) error, + notFoundErr error, + notFoundMessage string, + writeError func(error) error, +) error { + if unavailableErr != nil { + return handleError(c, unavailableErr) + } + + name, err := decode(c.Param(paramName)) + if err != nil { + return handleError(c, err) + } + + if err := deleteFunc(c.Request().Context(), name); err != nil { + if errors.Is(err, notFoundErr) { + return handleError(c, core.NewNotFoundError(notFoundMessage+name)) + } + return handleError(c, writeError(err)) + } + return c.NoContent(http.StatusNoContent) +} + +// ListModelOverrides handles GET /admin/api/v1/model-overrides. +func (h *Handler) ListModelOverrides(c *echo.Context) error { + if h.modelOverrides == nil { + return handleError(c, h.modelOverridesUnavailableError()) + } + views := h.modelOverrides.ListViews() + if views == nil { + views = []modeloverrides.View{} + } + return c.JSON(http.StatusOK, views) +} + +// UpsertModelOverride handles PUT /admin/api/v1/model-overrides/{selector}. +func (h *Handler) UpsertModelOverride(c *echo.Context) error { + if h.modelOverrides == nil { + return handleError(c, h.modelOverridesUnavailableError()) + } + + selector, err := decodeModelOverridePathSelector(c.Param("selector")) + if err != nil { + return handleError(c, err) + } + + var req upsertModelOverrideRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + if err := h.modelOverrides.Upsert(c.Request().Context(), modeloverrides.Override{ + Selector: selector, + Enabled: req.Enabled, + ForceDisabled: req.ForceDisabled, + AllowedOnlyForUserPaths: req.AllowedOnlyForUserPaths, + }); err != nil { + return handleError(c, modelOverrideWriteError(err)) + } + + override, ok := h.modelOverrides.Get(selector) + if !ok || override == nil { + slog.Error("model override service returned no override after upsert", "selector", selector) + return handleError(c, core.NewProviderError("model_overrides", http.StatusInternalServerError, "model override update failed unexpectedly", nil)) + } + return c.JSON(http.StatusOK, override) +} + +// DeleteModelOverride handles DELETE /admin/api/v1/model-overrides/{selector}. +func (h *Handler) DeleteModelOverride(c *echo.Context) error { + var unavailableErr error + var deleteFunc func(context.Context, string) error + if h.modelOverrides == nil { + unavailableErr = h.modelOverridesUnavailableError() + } else { + deleteFunc = h.modelOverrides.Delete + } + return deleteByName( + c, + unavailableErr, + "selector", + decodeModelOverridePathSelector, + deleteFunc, + modeloverrides.ErrNotFound, + "model override not found: ", + modelOverrideWriteError, + ) +} + // ListAuthKeys handles GET /admin/api/v1/auth-keys func (h *Handler) ListAuthKeys(c *echo.Context) error { if h.authKeys == nil { @@ -955,22 +1134,23 @@ func (h *Handler) UpsertAlias(c *echo.Context) error { // DeleteAlias handles DELETE /admin/api/v1/aliases/{name} func (h *Handler) DeleteAlias(c *echo.Context) error { + var unavailableErr error + var deleteFunc func(context.Context, string) error if h.aliases == nil { - return handleError(c, h.aliasesUnavailableError()) - } - - name, err := decodeAliasPathName(c.Param("name")) - if err != nil { - return handleError(c, err) - } - - if err := h.aliases.Delete(c.Request().Context(), name); err != nil { - if errors.Is(err, aliases.ErrNotFound) { - return handleError(c, core.NewNotFoundError("alias not found: "+name)) - } - return handleError(c, aliasWriteError(err)) + unavailableErr = h.aliasesUnavailableError() + } else { + deleteFunc = h.aliases.Delete } - return c.NoContent(http.StatusNoContent) + return deleteByName( + c, + unavailableErr, + "name", + decodeAliasPathName, + deleteFunc, + aliases.ErrNotFound, + "alias not found: ", + aliasWriteError, + ) } // ListGuardrailTypes handles GET /admin/api/v1/guardrails/types @@ -1305,3 +1485,15 @@ func decodeAliasPathName(raw string) (string, error) { } return name, nil } + +func decodeModelOverridePathSelector(raw string) (string, error) { + selector, err := url.PathUnescape(strings.TrimSpace(raw)) + if err != nil { + return "", core.NewInvalidRequestError("invalid model override selector", err) + } + selector = strings.TrimSpace(selector) + if selector == "" { + return "", core.NewInvalidRequestError("model override selector is required", nil) + } + return selector, nil +} diff --git a/internal/admin/handler_model_overrides_test.go b/internal/admin/handler_model_overrides_test.go new file mode 100644 index 000000000..0eb884860 --- /dev/null +++ b/internal/admin/handler_model_overrides_test.go @@ -0,0 +1,277 @@ +package admin + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/modeloverrides" + "gomodel/internal/providers" +) + +type modelOverrideTestStore struct { + items map[string]modeloverrides.Override +} + +func modelOverrideBoolPtr(value bool) *bool { + return &value +} + +func newModelOverrideTestStore(items ...modeloverrides.Override) *modelOverrideTestStore { + store := &modelOverrideTestStore{items: make(map[string]modeloverrides.Override, len(items))} + for _, item := range items { + store.items[item.Selector] = item + } + return store +} + +func (s *modelOverrideTestStore) List(_ context.Context) ([]modeloverrides.Override, error) { + result := make([]modeloverrides.Override, 0, len(s.items)) + for _, item := range s.items { + result = append(result, item) + } + return result, nil +} + +func (s *modelOverrideTestStore) Upsert(_ context.Context, override modeloverrides.Override) error { + s.items[override.Selector] = override + return nil +} + +func (s *modelOverrideTestStore) Delete(_ context.Context, selector string) error { + if _, ok := s.items[selector]; !ok { + return modeloverrides.ErrNotFound + } + delete(s.items, selector) + return nil +} + +func (s *modelOverrideTestStore) Close() error { return nil } + +type failingModelOverrideStore struct { + listErr error + upsertErr error + deleteErr error +} + +func (s *failingModelOverrideStore) List(_ context.Context) ([]modeloverrides.Override, error) { + return nil, s.listErr +} + +func (s *failingModelOverrideStore) Upsert(_ context.Context, _ modeloverrides.Override) error { + return s.upsertErr +} + +func (s *failingModelOverrideStore) Delete(_ context.Context, _ string) error { + return s.deleteErr +} + +func (s *failingModelOverrideStore) Close() error { return nil } + +func newModelOverrideRegistry(t *testing.T) *providers.ModelRegistry { + t.Helper() + registry := providers.NewModelRegistry() + mock := &handlerMockProvider{ + models: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "gpt-4o", Object: "model", OwnedBy: "openai"}, + }, + }, + } + registry.RegisterProviderWithNameAndType(mock, "openai", "openai") + if err := registry.Initialize(context.Background()); err != nil { + t.Fatalf("Initialize() error = %v", err) + } + return registry +} + +func newModelOverrideService(t *testing.T, store modeloverrides.Store, defaultEnabled bool) *modeloverrides.Service { + t.Helper() + service, err := modeloverrides.NewService(store, newModelOverrideRegistry(t), defaultEnabled) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + return service +} + +func TestListModels_IncludesModelAccessState(t *testing.T) { + registry := newModelOverrideRegistry(t) + service := newModelOverrideService(t, newModelOverrideTestStore(modeloverrides.Override{ + Selector: "openai/gpt-4o", + Enabled: modelOverrideBoolPtr(true), + AllowedOnlyForUserPaths: []string{"/team/alpha"}, + }), false) + + h := NewHandler(nil, registry, WithModelOverrides(service)) + c, rec := newHandlerContext("/admin/api/v1/models") + + if err := h.ListModels(c); err != nil { + t.Fatalf("ListModels() error = %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + + var body []modelInventoryResponse + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + if len(body) != 1 { + t.Fatalf("len(body) = %d, want 1", len(body)) + } + + row := body[0] + if row.Model.ID != "gpt-4o" { + t.Fatalf("row.Model.ID = %q, want gpt-4o", row.Model.ID) + } + if row.Access.Selector != "openai/gpt-4o" { + t.Fatalf("row.Access.Selector = %q, want openai/gpt-4o", row.Access.Selector) + } + if row.Access.DefaultEnabled { + t.Fatal("row.Access.DefaultEnabled = true, want false") + } + if !row.Access.EffectiveEnabled { + t.Fatal("row.Access.EffectiveEnabled = false, want true") + } + if len(row.Access.AllowedOnlyForUserPaths) != 1 || row.Access.AllowedOnlyForUserPaths[0] != "/team/alpha" { + t.Fatalf("row.Access.AllowedOnlyForUserPaths = %#v, want [/team/alpha]", row.Access.AllowedOnlyForUserPaths) + } + if row.Access.Override == nil || row.Access.Override.Selector != "openai/gpt-4o" { + t.Fatalf("row.Access.Override = %#v, want exact override", row.Access.Override) + } +} + +func TestModelOverrideEndpointsReturn503WhenServiceUnavailable(t *testing.T) { + h := NewHandler(nil, nil) + e := echo.New() + + assertUnavailable := func(name string, err error, rec *httptest.ResponseRecorder) { + t.Helper() + if err != nil { + t.Fatalf("%s error = %v", name, err) + } + if rec.Code != http.StatusServiceUnavailable { + t.Fatalf("%s status = %d, want 503", name, rec.Code) + } + + var body map[string]map[string]any + if decodeErr := json.Unmarshal(rec.Body.Bytes(), &body); decodeErr != nil { + t.Fatalf("%s decode error = %v", name, decodeErr) + } + if got := body["error"]["code"]; got != "feature_unavailable" { + t.Fatalf("%s error code = %v, want feature_unavailable", name, got) + } + } + + listCtx, listRec := newHandlerContext("/admin/api/v1/model-overrides") + assertUnavailable("ListModelOverrides", h.ListModelOverrides(listCtx), listRec) + + putReq := httptest.NewRequest(http.MethodPut, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", bytes.NewBufferString(`{"enabled":true}`)) + putReq.Header.Set("Content-Type", "application/json") + putRec := httptest.NewRecorder() + putCtx := e.NewContext(putReq, putRec) + putCtx.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + assertUnavailable("UpsertModelOverride", h.UpsertModelOverride(putCtx), putRec) + + deleteReq := httptest.NewRequest(http.MethodDelete, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", nil) + deleteRec := httptest.NewRecorder() + deleteCtx := e.NewContext(deleteReq, deleteRec) + deleteCtx.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + assertUnavailable("DeleteModelOverride", h.DeleteModelOverride(deleteCtx), deleteRec) +} + +func TestUpsertAndDeleteModelOverride(t *testing.T) { + service := newModelOverrideService(t, newModelOverrideTestStore(), true) + h := NewHandler(nil, nil, WithModelOverrides(service)) + e := echo.New() + + putReq := httptest.NewRequest(http.MethodPut, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", bytes.NewBufferString(`{"enabled":true,"allowed_only_for_user_paths":["team/alpha"]}`)) + putReq.Header.Set("Content-Type", "application/json") + putRec := httptest.NewRecorder() + putCtx := e.NewContext(putReq, putRec) + putCtx.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + + if err := h.UpsertModelOverride(putCtx); err != nil { + t.Fatalf("UpsertModelOverride() error = %v", err) + } + if putRec.Code != http.StatusOK { + t.Fatalf("put status = %d, want 200", putRec.Code) + } + + var body modeloverrides.Override + if err := json.Unmarshal(putRec.Body.Bytes(), &body); err != nil { + t.Fatalf("decode upsert response: %v", err) + } + if body.Selector != "openai/gpt-4o" { + t.Fatalf("body.Selector = %q, want openai/gpt-4o", body.Selector) + } + if body.Enabled == nil || !*body.Enabled { + t.Fatalf("body.Enabled = %#v, want true", body.Enabled) + } + if len(body.AllowedOnlyForUserPaths) != 1 || body.AllowedOnlyForUserPaths[0] != "/team/alpha" { + t.Fatalf("body.AllowedOnlyForUserPaths = %#v, want [/team/alpha]", body.AllowedOnlyForUserPaths) + } + + deleteReq := httptest.NewRequest(http.MethodDelete, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", nil) + deleteRec := httptest.NewRecorder() + deleteCtx := e.NewContext(deleteReq, deleteRec) + deleteCtx.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + + if err := h.DeleteModelOverride(deleteCtx); err != nil { + t.Fatalf("DeleteModelOverride() error = %v", err) + } + if deleteRec.Code != http.StatusNoContent { + t.Fatalf("delete status = %d, want 204", deleteRec.Code) + } +} + +func TestUpsertModelOverrideReturnsBadRequestForValidationErrors(t *testing.T) { + service := newModelOverrideService(t, newModelOverrideTestStore(), true) + h := NewHandler(nil, nil, WithModelOverrides(service)) + e := echo.New() + + req := httptest.NewRequest(http.MethodPut, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", bytes.NewBufferString(`{"enabled":false}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + c.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + + if err := h.UpsertModelOverride(c); err != nil { + t.Fatalf("UpsertModelOverride() error = %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want 400", rec.Code) + } +} + +func TestModelOverrideWriteErrorsBubbleProviderErrors(t *testing.T) { + service := newModelOverrideService(t, &failingModelOverrideStore{ + upsertErr: errors.New("boom"), + }, true) + h := NewHandler(nil, nil, WithModelOverrides(service)) + e := echo.New() + + req := httptest.NewRequest(http.MethodPut, "/admin/api/v1/model-overrides/openai%2Fgpt-4o", bytes.NewBufferString(`{"enabled":true}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + c.SetPathValues(echo.PathValues{{Name: "selector", Value: "openai/gpt-4o"}}) + + if err := h.UpsertModelOverride(c); err != nil { + t.Fatalf("UpsertModelOverride() error = %v", err) + } + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status = %d, want 500", rec.Code) + } +} diff --git a/internal/aliases/service.go b/internal/aliases/service.go index 44367db9b..5cbef358b 100644 --- a/internal/aliases/service.go +++ b/internal/aliases/service.go @@ -165,6 +165,16 @@ func (s *Service) GetProviderType(model string) string { // ExposedModels returns enabled aliases projected as model-list entries. func (s *Service) ExposedModels() []core.Model { + return s.exposedModelsFiltered(nil) +} + +// ExposedModelsFiltered returns enabled aliases projected as model-list entries +// while allowing callers to filter by the concrete target selector. +func (s *Service) ExposedModelsFiltered(allow func(core.ModelSelector) bool) []core.Model { + return s.exposedModelsFiltered(allow) +} + +func (s *Service) exposedModelsFiltered(allow func(core.ModelSelector) bool) []core.Model { aliases := s.List() result := make([]core.Model, 0, len(aliases)) for _, alias := range aliases { @@ -175,6 +185,9 @@ func (s *Service) ExposedModels() []core.Model { if err != nil { continue } + if allow != nil && !allow(selector) { + continue + } model, ok := s.catalog.LookupModel(selector.QualifiedModel()) if !ok || model == nil { continue diff --git a/internal/app/app.go b/internal/app/app.go index 3b8ba33b5..3588bb9e3 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -24,6 +24,7 @@ import ( "gomodel/internal/executionplans" "gomodel/internal/fallback" "gomodel/internal/guardrails" + "gomodel/internal/modeloverrides" "gomodel/internal/providers" "gomodel/internal/responsecache" "gomodel/internal/server" @@ -40,6 +41,7 @@ type App struct { usage *usage.Result batch *batch.Result aliases *aliases.Result + modelOverrides *modeloverrides.Result authKeys *authkeys.Result guardrails *guardrails.Result executionPlans *executionplans.Result @@ -168,6 +170,22 @@ func New(ctx context.Context, cfg Config) (*App, error) { } app.aliases = aliasResult + var modelOverrideResult *modeloverrides.Result + sharedModelOverrideStorage := firstSharedStorage(auditResult.Storage, usageResult.Storage, batchResult.Storage, aliasResult.Storage) + if sharedModelOverrideStorage != nil { + modelOverrideResult, err = modeloverrides.NewWithSharedStorage(ctx, appCfg, sharedModelOverrideStorage, providerResult.Registry) + } else { + modelOverrideResult, err = modeloverrides.New(ctx, appCfg, providerResult.Registry) + } + if err != nil { + closeErr := errors.Join(app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + if closeErr != nil { + return nil, fmt.Errorf("failed to initialize model overrides: %w (also: close error: %v)", err, closeErr) + } + return nil, fmt.Errorf("failed to initialize model overrides: %w", err) + } + app.modelOverrides = modelOverrideResult + refreshInterval := executionPlanRefreshInterval(appCfg) var guardrailExecutor guardrails.ChatCompletionExecutor = app.providers.Router if app.aliases != nil && app.aliases.Service != nil { @@ -176,14 +194,14 @@ func New(ctx context.Context, cfg Config) (*App, error) { // Initialize reusable guardrail definitions using shared storage when already available. var guardrailResult *guardrails.Result - sharedGuardrailStorage := firstSharedStorage(auditResult.Storage, usageResult.Storage, batchResult.Storage, aliasResult.Storage) + sharedGuardrailStorage := firstSharedStorage(auditResult.Storage, usageResult.Storage, batchResult.Storage, aliasResult.Storage, modelOverrideResult.Storage) if sharedGuardrailStorage != nil { guardrailResult, err = guardrails.NewWithSharedStorage(ctx, sharedGuardrailStorage, refreshInterval, guardrailExecutor) } else { guardrailResult, err = guardrails.New(ctx, appCfg, refreshInterval, guardrailExecutor) } if err != nil { - closeErr := errors.Join(app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to initialize guardrails: %w (also: close error: %v)", err, closeErr) } @@ -193,14 +211,14 @@ func New(ctx context.Context, cfg Config) (*App, error) { seedGuardrails, err := configGuardrailDefinitions(appCfg.Guardrails) if err != nil { - closeErr := errors.Join(app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to prepare guardrail definitions: %w (also: close error: %v)", err, closeErr) } return nil, fmt.Errorf("failed to prepare guardrail definitions: %w", err) } if err := guardrailResult.Service.UpsertDefinitions(ctx, seedGuardrails); err != nil { - closeErr := errors.Join(app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to upsert guardrails: %w (also: close error: %v)", err, closeErr) } @@ -215,7 +233,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { featureCaps := runtimeExecutionFeatureCaps(appCfg) var executionPlanResult *executionplans.Result - sharedExecutionPlanStorage := firstSharedStorage(auditResult.Storage, usageResult.Storage, batchResult.Storage, aliasResult.Storage, guardrailResult.Storage) + sharedExecutionPlanStorage := firstSharedStorage(auditResult.Storage, usageResult.Storage, batchResult.Storage, aliasResult.Storage, modelOverrideResult.Storage, guardrailResult.Storage) executionPlanCompiler := executionplans.NewCompilerWithFeatureCaps(guardrailResult.Service, featureCaps) if sharedExecutionPlanStorage != nil { executionPlanResult, err = executionplans.NewWithSharedStorage(ctx, sharedExecutionPlanStorage, executionPlanCompiler, refreshInterval) @@ -223,7 +241,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { executionPlanResult, err = executionplans.New(ctx, appCfg, executionPlanCompiler, refreshInterval) } if err != nil { - closeErr := errors.Join(app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to initialize execution plans: %w (also: close error: %v)", err, closeErr) } @@ -231,14 +249,14 @@ func New(ctx context.Context, cfg Config) (*App, error) { } defaultExecutionPlan := defaultExecutionPlanInput(appCfg, guardrailResult.Service.Names(), seedGuardrails) if err := executionPlanResult.Service.EnsureDefaultGlobal(ctx, defaultExecutionPlan); err != nil { - closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to seed execution plans: %w (also: close error: %v)", err, closeErr) } return nil, fmt.Errorf("failed to seed execution plans: %w", err) } if err := executionPlanResult.Service.Refresh(ctx); err != nil { - closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to load execution plans: %w (also: close error: %v)", err, closeErr) } @@ -252,6 +270,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { usageResult.Storage, batchResult.Storage, aliasResult.Storage, + modelOverrideResult.Storage, guardrailResult.Storage, executionPlanResult.Storage, ) @@ -261,7 +280,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { authKeyResult, err = authkeys.New(ctx, appCfg) } if err != nil { - closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(executionPlanResult.Close(), app.guardrails.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to initialize auth keys: %w (also: close error: %v)", err, closeErr) } @@ -291,6 +310,9 @@ func New(ctx context.Context, cfg Config) (*App, error) { aliases.NewBatchPreparer(provider, app.aliases.Service), }, batchRequestPreparers...) } + if app.modelOverrides != nil && app.modelOverrides.Service != nil { + batchRequestPreparers = append(batchRequestPreparers, modeloverrides.NewBatchPreparer(provider, app.modelOverrides.Service)) + } batchRequestPreparer := server.ComposeBatchRequestPreparers(providerAsNativeFileRouter(provider), batchRequestPreparers...) // Create server @@ -306,6 +328,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { UsageLogger: usageResult.Logger, PricingResolver: providerResult.Registry, ModelResolver: app.aliases.Service, + ModelAuthorizer: app.modelOverrides.Service, FallbackResolver: fallback.NewResolver(appCfg.Fallback, providerResult.Registry), ExecutionPolicyResolver: executionPlanResult.Service, TranslatedRequestPatcher: translatedRequestPatcher, @@ -334,6 +357,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { providerResult.Registry, authKeyResult.Service, app.aliases.Service, + app.modelOverrides.Service, executionPlanResult.Service, app.guardrails.Service, dashboardRuntimeConfig(appCfg, usageEnabledForDashboard), @@ -374,6 +398,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { guardrailsCloseErr error authKeysCloseErr error aliasCloseErr error + modelOverridesCloseErr error batchCloseErr error ) if app.executionPlans != nil { @@ -388,10 +413,13 @@ func New(ctx context.Context, cfg Config) (*App, error) { if app.aliases != nil { aliasCloseErr = app.aliases.Close() } + if app.modelOverrides != nil { + modelOverridesCloseErr = app.modelOverrides.Close() + } if app.batch != nil { batchCloseErr = app.batch.Close() } - closeErr := errors.Join(executionPlansCloseErr, guardrailsCloseErr, authKeysCloseErr, aliasCloseErr, batchCloseErr, app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(executionPlansCloseErr, guardrailsCloseErr, authKeysCloseErr, aliasCloseErr, modelOverridesCloseErr, batchCloseErr, app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to initialize response cache: %w (also: close error: %v)", err, closeErr) } @@ -401,6 +429,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { internalGuardrailExecutor := server.NewInternalChatCompletionExecutor(provider, server.InternalChatCompletionExecutorConfig{ ModelResolver: app.aliases.Service, + ModelAuthorizer: app.modelOverrides.Service, ExecutionPolicyResolver: executionPlanResult.Service, FallbackResolver: serverCfg.FallbackResolver, AuditLogger: auditResult.Logger, @@ -409,14 +438,14 @@ func New(ctx context.Context, cfg Config) (*App, error) { ResponseCache: rcm, }) if err := guardrailResult.Service.SetExecutor(ctx, internalGuardrailExecutor); err != nil { - closeErr := errors.Join(rcm.Close(), app.executionPlans.Close(), app.guardrails.Close(), app.authKeys.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(rcm.Close(), app.executionPlans.Close(), app.guardrails.Close(), app.authKeys.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to wire internal guardrail executor: %w (also: close error: %v)", err, closeErr) } return nil, fmt.Errorf("failed to wire internal guardrail executor: %w", err) } if err := executionPlanResult.Service.Refresh(ctx); err != nil { - closeErr := errors.Join(rcm.Close(), app.executionPlans.Close(), app.guardrails.Close(), app.authKeys.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) + closeErr := errors.Join(rcm.Close(), app.executionPlans.Close(), app.guardrails.Close(), app.authKeys.Close(), app.modelOverrides.Close(), app.aliases.Close(), app.batch.Close(), app.usage.Close(), app.audit.Close(), app.providers.Close()) if closeErr != nil { return nil, fmt.Errorf("failed to refresh execution plans after wiring internal guardrail executor: %w (also: close error: %v)", err, closeErr) } @@ -571,7 +600,15 @@ func (a *App) Shutdown(ctx context.Context) error { } } - // 5. Close reusable guardrails subsystem. + // 5. Close model overrides subsystem. + if a.modelOverrides != nil { + if err := a.modelOverrides.Close(); err != nil { + slog.Error("model overrides close error", "error", err) + errs = append(errs, fmt.Errorf("model overrides close: %w", err)) + } + } + + // 6. Close reusable guardrails subsystem. if a.guardrails != nil { if err := a.guardrails.Close(); err != nil { slog.Error("guardrails close error", "error", err) @@ -579,7 +616,7 @@ func (a *App) Shutdown(ctx context.Context) error { } } - // 6. Close managed auth keys subsystem. + // 7. Close managed auth keys subsystem. if a.authKeys != nil { if err := a.authKeys.Close(); err != nil { slog.Error("auth keys close error", "error", err) @@ -587,7 +624,7 @@ func (a *App) Shutdown(ctx context.Context) error { } } - // 7. Close batch store (flushes pending entries) + // 8. Close batch store (flushes pending entries) if a.batch != nil { if err := a.batch.Close(); err != nil { slog.Error("batch store close error", "error", err) @@ -595,7 +632,7 @@ func (a *App) Shutdown(ctx context.Context) error { } } - // 8. Close usage tracking (flushes pending entries) + // 9. Close usage tracking (flushes pending entries) if a.usage != nil { if err := a.usage.Close(); err != nil { slog.Error("usage logger close error", "error", err) @@ -603,7 +640,7 @@ func (a *App) Shutdown(ctx context.Context) error { } } - // 9. Close audit logging (flushes pending logs) + // 10. Close audit logging (flushes pending logs) if a.audit != nil { if err := a.audit.Close(); err != nil { slog.Error("audit logger close error", "error", err) @@ -679,6 +716,7 @@ func initAdmin( registry *providers.ModelRegistry, authKeyService *authkeys.Service, aliasService *aliases.Service, + modelOverrideService *modeloverrides.Service, executionPlanService *executionplans.Service, guardrailService *guardrails.Service, runtimeConfig admin.DashboardConfigResponse, @@ -719,6 +757,7 @@ func initAdmin( admin.WithAuditReader(auditReader), admin.WithAuthKeys(authKeyService), admin.WithAliases(aliasService), + admin.WithModelOverrides(modelOverrideService), admin.WithExecutionPlans(executionPlanService), admin.WithGuardrailService(guardrailService), admin.WithDashboardRuntimeConfig(runtimeConfig), diff --git a/internal/core/interfaces.go b/internal/core/interfaces.go index dd79e01a1..ebf14a986 100644 --- a/internal/core/interfaces.go +++ b/internal/core/interfaces.go @@ -116,6 +116,13 @@ type ProviderTypeNameResolver interface { GetProviderNameForType(providerType string) string } +// ProviderNameTypeResolver is an optional interface for components that can map +// a concrete configured provider instance name such as "openai_primary" back to +// its provider type, such as "openai". +type ProviderNameTypeResolver interface { + GetProviderTypeForName(providerName string) string +} + // AvailabilityChecker is an optional interface for providers that need // to verify service availability before registration. type AvailabilityChecker interface { diff --git a/internal/modeloverrides/batch_preparer.go b/internal/modeloverrides/batch_preparer.go new file mode 100644 index 000000000..b21a7ebea --- /dev/null +++ b/internal/modeloverrides/batch_preparer.go @@ -0,0 +1,106 @@ +package modeloverrides + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "gomodel/internal/core" +) + +type selectorResolver interface { + ResolveModel(requested core.RequestedModelSelector) (core.ModelSelector, bool, error) +} + +// BatchPreparer validates model access for native batch subrequests before provider submission. +type BatchPreparer struct { + provider core.RoutableProvider + service *Service +} + +// NewBatchPreparer creates an explicit model-override batch preparer. +func NewBatchPreparer(provider core.RoutableProvider, service *Service) *BatchPreparer { + return &BatchPreparer{ + provider: provider, + service: service, + } +} + +// PrepareBatchRequest validates inline and file-backed batch items without rewriting them. +func (p *BatchPreparer) PrepareBatchRequest(ctx context.Context, providerType string, req *core.BatchRequest) (*core.BatchRewriteResult, error) { + return core.RewriteBatchSource( + ctx, + providerType, + req, + p.batchFileTransport(), + []core.Operation{core.OperationChatCompletions, core.OperationResponses, core.OperationEmbeddings}, + func(ctx context.Context, item core.BatchRequestItem, decoded *core.DecodedBatchItemRequest) (json.RawMessage, error) { + requested, err := requestedSelectorForDecodedRequest(decoded.Request) + if err != nil { + return nil, err + } + resolved, err := p.resolveSelector(requested) + if err != nil { + return nil, err + } + if p.provider != nil && !p.provider.Supports(resolved.QualifiedModel()) { + return nil, core.NewInvalidRequestError("unsupported model: "+resolved.QualifiedModel(), nil) + } + if providerType != "" && p.provider != nil { + actualProviderType := strings.TrimSpace(p.provider.GetProviderType(resolved.QualifiedModel())) + if actualProviderType != "" && actualProviderType != providerType { + return nil, core.NewInvalidRequestError( + fmt.Sprintf( + "native batch supports a single provider per batch; resolved model %q targets provider %q but batch provider is %q", + resolved.QualifiedModel(), + actualProviderType, + providerType, + ), + nil, + ) + } + } + if p.service != nil { + if err := p.service.ValidateModelAccess(ctx, resolved); err != nil { + return nil, err + } + } + return core.CloneRawJSON(item.Body), nil + }, + ) +} + +func (p *BatchPreparer) batchFileTransport() core.BatchFileTransport { + if p == nil || p.provider == nil { + return nil + } + if files, ok := p.provider.(core.NativeFileRoutableProvider); ok { + return files + } + return nil +} + +func (p *BatchPreparer) resolveSelector(requested core.RequestedModelSelector) (core.ModelSelector, error) { + if p == nil || p.provider == nil { + return requested.Normalize() + } + if resolver, ok := p.provider.(selectorResolver); ok { + selector, _, err := resolver.ResolveModel(requested) + return selector, err + } + return requested.Normalize() +} + +func requestedSelectorForDecodedRequest(request any) (core.RequestedModelSelector, error) { + switch typed := request.(type) { + case *core.ChatRequest: + return core.NewRequestedModelSelector(typed.Model, typed.Provider), nil + case *core.ResponsesRequest: + return core.NewRequestedModelSelector(typed.Model, typed.Provider), nil + case *core.EmbeddingRequest: + return core.NewRequestedModelSelector(typed.Model, typed.Provider), nil + default: + return core.RequestedModelSelector{}, core.NewInvalidRequestError("unsupported batch item request", nil) + } +} diff --git a/internal/modeloverrides/factory.go b/internal/modeloverrides/factory.go new file mode 100644 index 000000000..2df1ad81c --- /dev/null +++ b/internal/modeloverrides/factory.go @@ -0,0 +1,119 @@ +package modeloverrides + +import ( + "context" + "database/sql" + "errors" + "fmt" + "sync" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "go.mongodb.org/mongo-driver/v2/mongo" + + "gomodel/config" + "gomodel/internal/storage" +) + +// Result holds the initialized model override service and any owned resources. +type Result struct { + Service *Service + Store Store + Storage storage.Storage + + stopRefresh func() + closeOnce sync.Once + closeErr error +} + +// Close releases resources held by the model override subsystem. +func (r *Result) Close() error { + if r == nil { + return nil + } + r.closeOnce.Do(func() { + if r.stopRefresh != nil { + r.stopRefresh() + r.stopRefresh = nil + } + + var errs []error + if r.Store != nil { + if err := r.Store.Close(); err != nil { + errs = append(errs, fmt.Errorf("store close: %w", err)) + } + } + if r.Storage != nil { + if err := r.Storage.Close(); err != nil { + errs = append(errs, fmt.Errorf("storage close: %w", err)) + } + } + if len(errs) > 0 { + r.closeErr = fmt.Errorf("close errors: %w", errors.Join(errs...)) + } + }) + return r.closeErr +} + +// New creates a model override subsystem with its own storage connection. +func New(ctx context.Context, cfg *config.Config, catalog Catalog) (*Result, error) { + if cfg == nil { + return nil, fmt.Errorf("config is required") + } + storeConn, err := storage.New(ctx, cfg.Storage.BackendConfig()) + if err != nil { + return nil, fmt.Errorf("failed to create storage: %w", err) + } + result, err := newResult(ctx, cfg, storeConn, catalog) + if err != nil { + _ = storeConn.Close() + return nil, err + } + result.Storage = storeConn + return result, nil +} + +// NewWithSharedStorage creates a model override subsystem using an existing storage connection. +func NewWithSharedStorage(ctx context.Context, cfg *config.Config, shared storage.Storage, catalog Catalog) (*Result, error) { + if shared == nil { + return nil, fmt.Errorf("shared storage is required") + } + if cfg == nil { + return nil, fmt.Errorf("config is required") + } + return newResult(ctx, cfg, shared, catalog) +} + +func newResult(ctx context.Context, cfg *config.Config, storeConn storage.Storage, catalog Catalog) (*Result, error) { + store, err := createStore(ctx, storeConn) + if err != nil { + return nil, err + } + service, err := NewService(store, catalog, cfg.Models.EnabledByDefault) + if err != nil { + return nil, err + } + if err := service.Refresh(ctx); err != nil { + return nil, err + } + + refreshInterval := time.Minute + if cfg.ExecutionPlans.RefreshInterval > 0 { + refreshInterval = cfg.ExecutionPlans.RefreshInterval + } + + return &Result{ + Service: service, + Store: store, + stopRefresh: service.StartBackgroundRefresh(refreshInterval), + }, nil +} + +func createStore(ctx context.Context, store storage.Storage) (Store, error) { + return storage.ResolveBackend[Store]( + store, + func(db *sql.DB) (Store, error) { return NewSQLiteStore(db) }, + func(pool *pgxpool.Pool) (Store, error) { return NewPostgreSQLStore(ctx, pool) }, + func(db *mongo.Database) (Store, error) { return NewMongoDBStore(db) }, + ) +} diff --git a/internal/modeloverrides/service.go b/internal/modeloverrides/service.go new file mode 100644 index 000000000..e075950d7 --- /dev/null +++ b/internal/modeloverrides/service.go @@ -0,0 +1,468 @@ +package modeloverrides + +import ( + "context" + "fmt" + "log/slog" + "net/http" + "slices" + "sort" + "strings" + "sync" + "sync/atomic" + "time" + + "gomodel/internal/core" +) + +type compiledOverride struct { + override Override +} + +type snapshot struct { + order []string + bySelector map[string]Override + modelWide map[string]compiledOverride + providerWide map[string]compiledOverride + exact map[string]compiledOverride + defaultEnable bool +} + +// Service keeps model access overrides cached in memory. +type Service struct { + store Store + catalog Catalog + defaultEnabled bool + current atomic.Value + refreshMu sync.Mutex +} + +// NewService creates a model override service backed by storage. +func NewService(store Store, catalog Catalog, defaultEnabled bool) (*Service, error) { + if store == nil { + return nil, fmt.Errorf("store is required") + } + if catalog == nil { + return nil, fmt.Errorf("catalog is required") + } + + service := &Service{ + store: store, + catalog: catalog, + defaultEnabled: defaultEnabled, + } + service.current.Store(snapshot{ + order: []string{}, + bySelector: map[string]Override{}, + modelWide: map[string]compiledOverride{}, + providerWide: map[string]compiledOverride{}, + exact: map[string]compiledOverride{}, + defaultEnable: defaultEnabled, + }) + return service, nil +} + +// EnabledByDefault reports the process-wide model availability default. +func (s *Service) EnabledByDefault() bool { + if s == nil { + return true + } + return s.defaultEnabled +} + +// Refresh reloads overrides from storage and atomically swaps the snapshot. +func (s *Service) Refresh(ctx context.Context) error { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + return s.refreshLocked(ctx) +} + +func (s *Service) refreshLocked(ctx context.Context) error { + overrides, err := s.store.List(ctx) + if err != nil { + return fmt.Errorf("list model overrides: %w", err) + } + next, err := s.buildSnapshot(overrides) + if err != nil { + return err + } + s.current.Store(next) + return nil +} + +func (s *Service) snapshot() snapshot { + if s == nil { + return snapshot{ + order: []string{}, + bySelector: map[string]Override{}, + modelWide: map[string]compiledOverride{}, + providerWide: map[string]compiledOverride{}, + exact: map[string]compiledOverride{}, + defaultEnable: true, + } + } + return s.current.Load().(snapshot) +} + +func (s *Service) buildSnapshot(overrides []Override) (snapshot, error) { + next := snapshot{ + order: make([]string, 0, len(overrides)), + bySelector: make(map[string]Override, len(overrides)), + modelWide: make(map[string]compiledOverride), + providerWide: make(map[string]compiledOverride), + exact: make(map[string]compiledOverride), + defaultEnable: s.defaultEnabled, + } + + for _, override := range overrides { + normalized, err := normalizeStoredOverride(override) + if err != nil { + return snapshot{}, fmt.Errorf("load model override %q: %w", override.Selector, err) + } + next.order = append(next.order, normalized.Selector) + next.bySelector[normalized.Selector] = normalized + + compiled := compiledOverride{override: normalized} + switch normalized.ScopeKind() { + case ScopeProviderModel: + next.exact[exactMatchKey(normalized.ProviderName, normalized.Model)] = compiled + case ScopeProvider: + next.providerWide[normalized.ProviderName] = compiled + default: + next.modelWide[normalized.Model] = compiled + } + } + sort.Strings(next.order) + return next, nil +} + +func cloneOverride(override Override) Override { + override.Enabled = cloneEnabled(override.Enabled) + override.AllowedOnlyForUserPaths = append([]string(nil), override.AllowedOnlyForUserPaths...) + return override +} + +func snapshotOverrides(snap snapshot) []Override { + result := make([]Override, 0, len(snap.order)) + for _, selector := range snap.order { + result = append(result, cloneOverride(snap.bySelector[selector])) + } + return result +} + +func upsertOverride(overrides []Override, next Override) []Override { + for i := range overrides { + if overrides[i].Selector == next.Selector { + overrides[i] = cloneOverride(next) + return overrides + } + } + return append(overrides, cloneOverride(next)) +} + +func deleteOverride(overrides []Override, selector string) []Override { + result := make([]Override, 0, len(overrides)) + for _, override := range overrides { + if override.Selector == selector { + continue + } + result = append(result, cloneOverride(override)) + } + return result +} + +func rollbackContext() (context.Context, context.CancelFunc) { + return context.WithTimeout(context.Background(), 30*time.Second) +} + +// List returns all cached overrides sorted by selector. +func (s *Service) List() []Override { + snap := s.snapshot() + result := make([]Override, 0, len(snap.order)) + for _, selector := range snap.order { + override := snap.bySelector[selector] + override.Enabled = cloneEnabled(override.Enabled) + override.AllowedOnlyForUserPaths = append([]string(nil), override.AllowedOnlyForUserPaths...) + result = append(result, override) + } + return result +} + +// ListViews returns all cached overrides with scope metadata. +func (s *Service) ListViews() []View { + overrides := s.List() + result := make([]View, 0, len(overrides)) + for _, override := range overrides { + result = append(result, View{ + Override: override, + ScopeKind: override.ScopeKind(), + }) + } + return result +} + +// Get returns one cached override by normalized selector. +func (s *Service) Get(selector string) (*Override, bool) { + normalized, _, _, err := normalizeSelectorInput(selectorProviderNames(s.catalog), selector) + if err != nil { + return nil, false + } + override, ok := s.snapshot().bySelector[normalized] + if !ok { + return nil, false + } + override.Enabled = cloneEnabled(override.Enabled) + override.AllowedOnlyForUserPaths = append([]string(nil), override.AllowedOnlyForUserPaths...) + return &override, true +} + +// Upsert validates and stores one override, then refreshes the in-memory snapshot. +func (s *Service) Upsert(ctx context.Context, override Override) error { + if s == nil { + return fmt.Errorf("model override service is required") + } + + normalized, err := normalizeOverrideInput(s.catalog, override) + if err != nil { + return err + } + if normalized.Enabled == nil && !normalized.ForceDisabled && len(normalized.AllowedOnlyForUserPaths) == 0 { + return newValidationError("override must enable, force disable, or set allowed_only_for_user_paths", nil) + } + if normalized.Enabled != nil && !*normalized.Enabled { + return newValidationError("enabled=false is not supported; use force_disabled=true or omit enabled", nil) + } + + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + + current := s.snapshot() + if _, err := s.buildSnapshot(upsertOverride(snapshotOverrides(current), normalized)); err != nil { + return fmt.Errorf("validate model overrides: %w", err) + } + previous, existed := current.bySelector[normalized.Selector] + if err := s.store.Upsert(ctx, normalized); err != nil { + return fmt.Errorf("upsert model override: %w", err) + } + if err := s.refreshLocked(ctx); err != nil { + rollbackCtx, cancel := rollbackContext() + defer cancel() + + var rollbackErr error + if existed { + rollbackErr = s.store.Upsert(rollbackCtx, previous) + } else { + rollbackErr = s.store.Delete(rollbackCtx, normalized.Selector) + } + if rollbackErr != nil { + return fmt.Errorf("refresh model overrides: %w (rollback failed: %v)", err, rollbackErr) + } + return fmt.Errorf("refresh model overrides: %w", err) + } + return nil +} + +// Delete removes one override and refreshes the in-memory snapshot. +func (s *Service) Delete(ctx context.Context, selector string) error { + if s == nil { + return fmt.Errorf("model override service is required") + } + + normalized, _, _, err := normalizeSelectorInput(selectorProviderNames(s.catalog), selector) + if err != nil { + return err + } + + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + + current := s.snapshot() + if _, err := s.buildSnapshot(deleteOverride(snapshotOverrides(current), normalized)); err != nil { + return fmt.Errorf("validate model overrides: %w", err) + } + previous, existed := current.bySelector[normalized] + if err := s.store.Delete(ctx, normalized); err != nil { + return fmt.Errorf("delete model override: %w", err) + } + if err := s.refreshLocked(ctx); err != nil { + if !existed { + return fmt.Errorf("refresh model overrides: %w", err) + } + + rollbackCtx, cancel := rollbackContext() + defer cancel() + if rollbackErr := s.store.Upsert(rollbackCtx, previous); rollbackErr != nil { + return fmt.Errorf("refresh model overrides: %w (rollback failed: %v)", err, rollbackErr) + } + return fmt.Errorf("refresh model overrides: %w", err) + } + return nil +} + +// StartBackgroundRefresh periodically reloads model overrides from storage until stopped. +func (s *Service) StartBackgroundRefresh(interval time.Duration) func() { + if interval <= 0 { + interval = time.Minute + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + var once sync.Once + + go func() { + defer close(done) + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + refreshCtx, refreshCancel := context.WithTimeout(ctx, 30*time.Second) + if err := s.Refresh(refreshCtx); err != nil { + slog.Error("failed to refresh model overrides", "error", err) + } + refreshCancel() + } + } + }() + + return func() { + once.Do(func() { + cancel() + <-done + }) + } +} + +// EffectiveState resolves the compiled access state for one concrete selector. +func (s *Service) EffectiveState(selector core.ModelSelector) EffectiveState { + return s.snapshot().effectiveState(selector) +} + +// AllowsModel reports whether selector is available for the effective request user path. +func (s *Service) AllowsModel(ctx context.Context, selector core.ModelSelector) bool { + state := s.EffectiveState(selector) + if !state.Enabled { + return false + } + if len(state.AllowedOnlyForUserPaths) == 0 { + return true + } + return userPathAllowed(core.UserPathFromContext(ctx), state.AllowedOnlyForUserPaths) +} + +// ValidateModelAccess returns a typed request error when selector is not available. +func (s *Service) ValidateModelAccess(ctx context.Context, selector core.ModelSelector) error { + state := s.EffectiveState(selector) + if !state.Enabled { + return core.NewInvalidRequestErrorWithStatus( + http.StatusBadRequest, + "requested model is not available", + nil, + ).WithCode("model_access_denied") + } + if len(state.AllowedOnlyForUserPaths) == 0 { + return nil + } + if userPathAllowed(core.UserPathFromContext(ctx), state.AllowedOnlyForUserPaths) { + return nil + } + return core.NewInvalidRequestErrorWithStatus( + http.StatusBadRequest, + "requested model is not available for this API key", + nil, + ).WithCode("model_access_denied") +} + +// FilterPublicModels removes models that are unavailable for the effective request user path. +func (s *Service) FilterPublicModels(ctx context.Context, models []core.Model) []core.Model { + if s == nil || len(models) == 0 { + return models + } + + result := make([]core.Model, 0, len(models)) + for _, model := range models { + selector, err := core.ParseModelSelector(model.ID, "") + if err != nil { + continue + } + if !s.AllowsModel(ctx, selector) { + continue + } + result = append(result, model) + } + return result +} + +func (snap snapshot) effectiveState(selector core.ModelSelector) EffectiveState { + model := strings.TrimSpace(selector.Model) + providerName := strings.TrimSpace(selector.Provider) + state := EffectiveState{ + Selector: selectorString(providerName, model), + ProviderName: providerName, + Model: model, + DefaultEnabled: snap.defaultEnable, + Enabled: snap.defaultEnable, + } + if model == "" && providerName == "" { + return state + } + + allowed := make([]string, 0) + seenAllowed := make(map[string]struct{}) + addAllowed := func(paths []string) { + for _, path := range paths { + if _, exists := seenAllowed[path]; exists { + continue + } + seenAllowed[path] = struct{}{} + allowed = append(allowed, path) + } + } + apply := func(rule compiledOverride, ok bool) { + if !ok { + return + } + if rule.override.Enabled != nil && *rule.override.Enabled { + state.Enabled = true + state.ForceDisabled = false + } + if rule.override.ForceDisabled { + state.Enabled = false + state.ForceDisabled = true + } + addAllowed(rule.override.AllowedOnlyForUserPaths) + } + + if modelWide, ok := snap.modelWide[model]; ok { + apply(modelWide, true) + } + if providerWide, ok := snap.providerWide[providerName]; ok { + apply(providerWide, true) + } + if exact, ok := snap.exact[exactMatchKey(providerName, model)]; ok { + apply(exact, true) + } + + sort.Strings(allowed) + state.AllowedOnlyForUserPaths = allowed + return state +} + +func userPathAllowed(userPath string, allowed []string) bool { + if len(allowed) == 0 { + return true + } + userPath, err := core.NormalizeUserPath(userPath) + if err != nil || userPath == "" { + return false + } + ancestors := core.UserPathAncestors(userPath) + for _, candidate := range ancestors { + if _, ok := slices.BinarySearch(allowed, candidate); ok { + return true + } + } + return false +} diff --git a/internal/modeloverrides/service_test.go b/internal/modeloverrides/service_test.go new file mode 100644 index 000000000..b1fa13145 --- /dev/null +++ b/internal/modeloverrides/service_test.go @@ -0,0 +1,280 @@ +package modeloverrides + +import ( + "context" + "errors" + "net/http" + "testing" + + "gomodel/internal/core" +) + +type testStore struct { + items map[string]Override +} + +func newTestStore(items ...Override) *testStore { + store := &testStore{items: make(map[string]Override, len(items))} + for _, item := range items { + store.items[item.Selector] = item + } + return store +} + +func (s *testStore) List(_ context.Context) ([]Override, error) { + result := make([]Override, 0, len(s.items)) + for _, item := range s.items { + result = append(result, item) + } + return result, nil +} + +func (s *testStore) Upsert(_ context.Context, override Override) error { + s.items[override.Selector] = override + return nil +} + +func (s *testStore) Delete(_ context.Context, selector string) error { + if _, ok := s.items[selector]; !ok { + return ErrNotFound + } + delete(s.items, selector) + return nil +} + +func (s *testStore) Close() error { return nil } + +type flakyListStore struct { + *testStore + listErr error +} + +func newFlakyListStore(items ...Override) *flakyListStore { + return &flakyListStore{testStore: newTestStore(items...)} +} + +func (s *flakyListStore) List(ctx context.Context) ([]Override, error) { + if s.listErr != nil { + return nil, s.listErr + } + return s.testStore.List(ctx) +} + +type testCatalog struct { + providerNames []string +} + +func (c testCatalog) ProviderNames() []string { + return append([]string(nil), c.providerNames...) +} + +func boolPtr(value bool) *bool { + return &value +} + +func TestNormalizeSelectorInput_UsesFirstSlashOnlyForKnownProviders(t *testing.T) { + providerNames := []string{"openai", "anthropic"} + + t.Run("known provider prefix becomes provider selector", func(t *testing.T) { + selector, providerName, model, err := normalizeSelectorInput(providerNames, "openai/gpt-4o") + if err != nil { + t.Fatalf("normalizeSelectorInput() error = %v", err) + } + if selector != "openai/gpt-4o" || providerName != "openai" || model != "gpt-4o" { + t.Fatalf("normalizeSelectorInput() = (%q, %q, %q), want (%q, %q, %q)", selector, providerName, model, "openai/gpt-4o", "openai", "gpt-4o") + } + }) + + t.Run("unknown provider prefix stays in raw model id", func(t *testing.T) { + selector, providerName, model, err := normalizeSelectorInput(providerNames, "vendor/model-with-slash") + if err != nil { + t.Fatalf("normalizeSelectorInput() error = %v", err) + } + if selector != "vendor/model-with-slash" || providerName != "" || model != "vendor/model-with-slash" { + t.Fatalf("normalizeSelectorInput() = (%q, %q, %q), want (%q, %q, %q)", selector, providerName, model, "vendor/model-with-slash", "", "vendor/model-with-slash") + } + }) + + t.Run("provider-wide selector keeps empty model", func(t *testing.T) { + selector, providerName, model, err := normalizeSelectorInput(providerNames, "anthropic/") + if err != nil { + t.Fatalf("normalizeSelectorInput() error = %v", err) + } + if selector != "anthropic/" || providerName != "anthropic" || model != "" { + t.Fatalf("normalizeSelectorInput() = (%q, %q, %q), want (%q, %q, %q)", selector, providerName, model, "anthropic/", "anthropic", "") + } + }) +} + +func TestService_DefaultDisabledRequiresExplicitEnableAndHonorsUserPaths(t *testing.T) { + service, err := NewService( + newTestStore(Override{ + Selector: "openai/gpt-4o", + Enabled: boolPtr(true), + AllowedOnlyForUserPaths: []string{"/team/alpha"}, + }), + testCatalog{providerNames: []string{"openai"}}, + false, + ) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + + enabledSelector := core.ModelSelector{Provider: "openai", Model: "gpt-4o"} + disabledSelector := core.ModelSelector{Provider: "openai", Model: "gpt-5"} + + state := service.EffectiveState(enabledSelector) + if !state.Enabled { + t.Fatal("EffectiveState().Enabled = false, want true") + } + if !state.DefaultEnabled { + // false is expected; explicit assertion below for clarity. + } else { + t.Fatal("EffectiveState().DefaultEnabled = true, want false") + } + if len(state.AllowedOnlyForUserPaths) != 1 || state.AllowedOnlyForUserPaths[0] != "/team/alpha" { + t.Fatalf("EffectiveState().AllowedOnlyForUserPaths = %#v, want [/team/alpha]", state.AllowedOnlyForUserPaths) + } + + allowedCtx := core.WithEffectiveUserPath(context.Background(), "/team/alpha/project-x") + if !service.AllowsModel(allowedCtx, enabledSelector) { + t.Fatal("AllowsModel() = false, want true for descendant user path") + } + if err := service.ValidateModelAccess(allowedCtx, enabledSelector); err != nil { + t.Fatalf("ValidateModelAccess() error = %v, want nil", err) + } + + deniedCtx := core.WithEffectiveUserPath(context.Background(), "/team/beta") + if service.AllowsModel(deniedCtx, enabledSelector) { + t.Fatal("AllowsModel() = true, want false for mismatched user path") + } + err = service.ValidateModelAccess(deniedCtx, enabledSelector) + if err == nil { + t.Fatal("ValidateModelAccess() error = nil, want access denial") + } + gatewayErr, ok := err.(*core.GatewayError) + if !ok { + t.Fatalf("ValidateModelAccess() error type = %T, want *core.GatewayError", err) + } + if gatewayErr.StatusCode != http.StatusBadRequest || gatewayErr.Code == nil || *gatewayErr.Code != "model_access_denied" { + t.Fatalf("ValidateModelAccess() = status %d code %#v, want 400/model_access_denied", gatewayErr.StatusCode, gatewayErr.Code) + } + + if service.AllowsModel(allowedCtx, disabledSelector) { + t.Fatal("AllowsModel() = true, want false for model without explicit enable when defaults are disabled") + } +} + +func TestService_ForceDisabledOverridesBroaderEnable(t *testing.T) { + service, err := NewService( + newTestStore( + Override{Selector: "openai/", Enabled: boolPtr(true)}, + Override{Selector: "openai/gpt-4o", ForceDisabled: true}, + ), + testCatalog{providerNames: []string{"openai"}}, + false, + ) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + + blocked := service.EffectiveState(core.ModelSelector{Provider: "openai", Model: "gpt-4o"}) + if blocked.Enabled { + t.Fatal("EffectiveState().Enabled = true, want false when exact force_disabled applies") + } + if !blocked.ForceDisabled { + t.Fatal("EffectiveState().ForceDisabled = false, want true") + } + + allowed := service.EffectiveState(core.ModelSelector{Provider: "openai", Model: "gpt-4.1"}) + if !allowed.Enabled { + t.Fatal("EffectiveState().Enabled = false, want true for provider-wide enable") + } + if allowed.ForceDisabled { + t.Fatal("EffectiveState().ForceDisabled = true, want false") + } +} + +func TestService_ExactEnableClearsBroaderForceDisabled(t *testing.T) { + service, err := NewService( + newTestStore( + Override{Selector: "openai/", ForceDisabled: true}, + Override{Selector: "openai/gpt-4o", Enabled: boolPtr(true)}, + ), + testCatalog{providerNames: []string{"openai"}}, + true, + ) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + + state := service.EffectiveState(core.ModelSelector{Provider: "openai", Model: "gpt-4o"}) + if !state.Enabled { + t.Fatal("EffectiveState().Enabled = false, want true when exact enable overrides broader force_disabled") + } + if state.ForceDisabled { + t.Fatal("EffectiveState().ForceDisabled = true, want false after exact enable override") + } + if err := service.ValidateModelAccess(context.Background(), core.ModelSelector{Provider: "openai", Model: "gpt-4o"}); err != nil { + t.Fatalf("ValidateModelAccess() error = %v, want nil", err) + } +} + +func TestService_UpsertRollsBackStorageOnRefreshFailure(t *testing.T) { + store := newFlakyListStore( + Override{Selector: "openai/gpt-4o", Enabled: boolPtr(true)}, + ) + service, err := NewService(store, testCatalog{providerNames: []string{"openai"}}, true) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + + store.listErr = errors.New("list failed") + err = service.Upsert(context.Background(), Override{Selector: "openai/gpt-5", Enabled: boolPtr(true)}) + if err == nil { + t.Fatal("Upsert() error = nil, want refresh failure") + } + if _, ok := store.items["openai/gpt-5"]; ok { + t.Fatal("store mutated after failed refresh; expected rollback to remove openai/gpt-5") + } + if _, ok := service.Get("openai/gpt-5"); ok { + t.Fatal("service cache mutated after failed refresh; expected openai/gpt-5 to remain absent") + } +} + +func TestService_DeleteRollsBackStorageOnRefreshFailure(t *testing.T) { + store := newFlakyListStore( + Override{Selector: "openai/gpt-4o", Enabled: boolPtr(true)}, + ) + service, err := NewService(store, testCatalog{providerNames: []string{"openai"}}, true) + if err != nil { + t.Fatalf("NewService() error = %v", err) + } + if err := service.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh() error = %v", err) + } + + store.listErr = errors.New("list failed") + err = service.Delete(context.Background(), "openai/gpt-4o") + if err == nil { + t.Fatal("Delete() error = nil, want refresh failure") + } + if _, ok := store.items["openai/gpt-4o"]; !ok { + t.Fatal("store lost openai/gpt-4o after failed refresh; expected rollback to restore it") + } + if _, ok := service.Get("openai/gpt-4o"); !ok { + t.Fatal("service cache lost openai/gpt-4o after failed refresh") + } +} diff --git a/internal/modeloverrides/store.go b/internal/modeloverrides/store.go new file mode 100644 index 000000000..db4d168eb --- /dev/null +++ b/internal/modeloverrides/store.go @@ -0,0 +1,66 @@ +package modeloverrides + +import ( + "context" + "errors" + "fmt" +) + +// ErrNotFound indicates a requested override was not found. +var ErrNotFound = errors.New("model override not found") + +// ValidationError indicates invalid override input or invalid override state. +type ValidationError struct { + Message string + Err error +} + +func (e *ValidationError) Error() string { + if e == nil { + return "" + } + return e.Message +} + +func (e *ValidationError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + +func newValidationError(message string, err error) error { + return &ValidationError{Message: message, Err: err} +} + +// IsValidationError reports whether err is a validation error. +func IsValidationError(err error) bool { + var target *ValidationError + return errors.As(err, &target) +} + +// Store defines persistence operations for model overrides. +type Store interface { + List(ctx context.Context) ([]Override, error) + Upsert(ctx context.Context, override Override) error + Delete(ctx context.Context, selector string) error + Close() error +} + +func collectOverrides(next func() (Override, bool, error), rowsErr func() error) ([]Override, error) { + result := make([]Override, 0) + for { + override, ok, err := next() + if err != nil { + return nil, err + } + if !ok { + break + } + result = append(result, override) + } + if err := rowsErr(); err != nil { + return nil, fmt.Errorf("iterate model overrides: %w", err) + } + return result, nil +} diff --git a/internal/modeloverrides/store_mongodb.go b/internal/modeloverrides/store_mongodb.go new file mode 100644 index 000000000..2fb2fe478 --- /dev/null +++ b/internal/modeloverrides/store_mongodb.go @@ -0,0 +1,133 @@ +package modeloverrides + +import ( + "context" + "fmt" + "strings" + "time" + + "go.mongodb.org/mongo-driver/v2/bson" + "go.mongodb.org/mongo-driver/v2/mongo" + "go.mongodb.org/mongo-driver/v2/mongo/options" +) + +type mongoOverrideDocument struct { + ID string `bson:"_id"` + ProviderName string `bson:"provider_name,omitempty"` + Model string `bson:"model,omitempty"` + Enabled *bool `bson:"enabled,omitempty"` + ForceDisabled bool `bson:"force_disabled,omitempty"` + AllowedOnlyForUserPaths []string `bson:"allowed_only_for_user_paths,omitempty"` + CreatedAt time.Time `bson:"created_at"` + UpdatedAt time.Time `bson:"updated_at"` +} + +type mongoOverrideIDFilter struct { + ID string `bson:"_id"` +} + +// MongoDBStore stores model overrides in MongoDB. +type MongoDBStore struct { + collection *mongo.Collection +} + +// NewMongoDBStore creates collection indexes if needed. +func NewMongoDBStore(database *mongo.Database) (*MongoDBStore, error) { + if database == nil { + return nil, fmt.Errorf("database is required") + } + coll := database.Collection("model_overrides") + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + indexes := []mongo.IndexModel{ + {Keys: bson.D{{Key: "provider_name", Value: 1}}}, + {Keys: bson.D{{Key: "model", Value: 1}}}, + {Keys: bson.D{{Key: "updated_at", Value: -1}}}, + } + if _, err := coll.Indexes().CreateMany(ctx, indexes); err != nil { + return nil, fmt.Errorf("create model_overrides indexes: %w", err) + } + return &MongoDBStore{collection: coll}, nil +} + +func (s *MongoDBStore) List(ctx context.Context) ([]Override, error) { + cursor, err := s.collection.Find(ctx, bson.M{}, options.Find().SetSort(bson.D{{Key: "_id", Value: 1}})) + if err != nil { + return nil, fmt.Errorf("list model overrides: %w", err) + } + defer cursor.Close(ctx) + + result := make([]Override, 0) + for cursor.Next(ctx) { + var doc mongoOverrideDocument + if err := cursor.Decode(&doc); err != nil { + return nil, fmt.Errorf("decode model override: %w", err) + } + result = append(result, overrideFromMongo(doc)) + } + if err := cursor.Err(); err != nil { + return nil, fmt.Errorf("iterate model overrides: %w", err) + } + return result, nil +} + +func (s *MongoDBStore) Upsert(ctx context.Context, override Override) error { + override, err := normalizeStoredOverride(override) + if err != nil { + return err + } + + now := time.Now().UTC() + if override.CreatedAt.IsZero() { + override.CreatedAt = now + } + override.UpdatedAt = now + + update := bson.M{ + "$set": bson.M{ + "provider_name": override.ProviderName, + "model": override.Model, + "enabled": override.Enabled, + "force_disabled": override.ForceDisabled, + "allowed_only_for_user_paths": override.AllowedOnlyForUserPaths, + "updated_at": override.UpdatedAt, + }, + "$setOnInsert": bson.M{ + "created_at": override.CreatedAt, + }, + } + _, err = s.collection.UpdateOne(ctx, mongoOverrideIDFilter{ID: override.Selector}, update, options.UpdateOne().SetUpsert(true)) + if err != nil { + return fmt.Errorf("upsert model override: %w", err) + } + return nil +} + +func (s *MongoDBStore) Delete(ctx context.Context, selector string) error { + result, err := s.collection.DeleteOne(ctx, mongoOverrideIDFilter{ID: strings.TrimSpace(selector)}) + if err != nil { + return fmt.Errorf("delete model override: %w", err) + } + if result.DeletedCount == 0 { + return ErrNotFound + } + return nil +} + +func (s *MongoDBStore) Close() error { + return nil +} + +func overrideFromMongo(doc mongoOverrideDocument) Override { + return Override{ + Selector: doc.ID, + ProviderName: doc.ProviderName, + Model: doc.Model, + Enabled: cloneEnabled(doc.Enabled), + ForceDisabled: doc.ForceDisabled, + AllowedOnlyForUserPaths: append([]string(nil), doc.AllowedOnlyForUserPaths...), + CreatedAt: doc.CreatedAt.UTC(), + UpdatedAt: doc.UpdatedAt.UTC(), + } +} diff --git a/internal/modeloverrides/store_postgresql.go b/internal/modeloverrides/store_postgresql.go new file mode 100644 index 000000000..a454a4c4a --- /dev/null +++ b/internal/modeloverrides/store_postgresql.go @@ -0,0 +1,162 @@ +package modeloverrides + +import ( + "context" + "encoding/json" + "fmt" + "strings" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +// PostgreSQLStore stores model overrides in PostgreSQL. +type PostgreSQLStore struct { + pool *pgxpool.Pool +} + +// NewPostgreSQLStore creates the model_overrides table and indexes if needed. +func NewPostgreSQLStore(ctx context.Context, pool *pgxpool.Pool) (*PostgreSQLStore, error) { + if ctx == nil { + return nil, fmt.Errorf("context is required") + } + if pool == nil { + return nil, fmt.Errorf("connection pool is required") + } + + _, err := pool.Exec(ctx, ` + CREATE TABLE IF NOT EXISTS model_overrides ( + selector TEXT PRIMARY KEY, + provider_name TEXT NOT NULL DEFAULT '', + model TEXT NOT NULL DEFAULT '', + enabled BOOLEAN NULL, + force_disabled BOOLEAN NOT NULL DEFAULT FALSE, + allowed_only_for_user_paths JSONB NOT NULL DEFAULT '[]'::jsonb, + created_at BIGINT NOT NULL, + updated_at BIGINT NOT NULL + ) + `) + if err != nil { + return nil, fmt.Errorf("failed to create model_overrides table: %w", err) + } + if _, err := pool.Exec(ctx, `CREATE INDEX IF NOT EXISTS idx_model_overrides_provider_name ON model_overrides(provider_name)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides provider_name index: %w", err) + } + if _, err := pool.Exec(ctx, `CREATE INDEX IF NOT EXISTS idx_model_overrides_model ON model_overrides(model)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides model index: %w", err) + } + if _, err := pool.Exec(ctx, `CREATE INDEX IF NOT EXISTS idx_model_overrides_updated_at ON model_overrides(updated_at DESC)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides updated_at index: %w", err) + } + return &PostgreSQLStore{pool: pool}, nil +} + +func (s *PostgreSQLStore) List(ctx context.Context) ([]Override, error) { + rows, err := s.pool.Query(ctx, ` + SELECT selector, provider_name, model, enabled, force_disabled, allowed_only_for_user_paths, created_at, updated_at + FROM model_overrides + ORDER BY selector ASC + `) + if err != nil { + return nil, fmt.Errorf("list model overrides: %w", err) + } + defer rows.Close() + return collectOverrides(func() (Override, bool, error) { + if !rows.Next() { + return Override{}, false, nil + } + override, err := scanPostgreSQLOverride(rows) + return override, true, err + }, rows.Err) +} + +func (s *PostgreSQLStore) Upsert(ctx context.Context, override Override) error { + override, err := normalizeStoredOverride(override) + if err != nil { + return err + } + + pathsJSON, err := json.Marshal(override.AllowedOnlyForUserPaths) + if err != nil { + return fmt.Errorf("encode allowed_only_for_user_paths: %w", err) + } + + now := time.Now().UTC().Unix() + if override.CreatedAt.IsZero() { + override.CreatedAt = time.Unix(now, 0).UTC() + } + override.UpdatedAt = time.Unix(now, 0).UTC() + + _, err = s.pool.Exec(ctx, ` + INSERT INTO model_overrides ( + selector, provider_name, model, enabled, force_disabled, allowed_only_for_user_paths, created_at, updated_at + ) + VALUES ($1, $2, $3, $4, $5, $6::jsonb, $7, $8) + ON CONFLICT(selector) DO UPDATE SET + provider_name = excluded.provider_name, + model = excluded.model, + enabled = excluded.enabled, + force_disabled = excluded.force_disabled, + allowed_only_for_user_paths = excluded.allowed_only_for_user_paths, + updated_at = excluded.updated_at + `, + override.Selector, + override.ProviderName, + override.Model, + override.Enabled, + override.ForceDisabled, + string(pathsJSON), + override.CreatedAt.Unix(), + override.UpdatedAt.Unix(), + ) + if err != nil { + return fmt.Errorf("upsert model override: %w", err) + } + return nil +} + +func (s *PostgreSQLStore) Delete(ctx context.Context, selector string) error { + cmd, err := s.pool.Exec(ctx, `DELETE FROM model_overrides WHERE selector = $1`, strings.TrimSpace(selector)) + if err != nil { + return fmt.Errorf("delete model override: %w", err) + } + if cmd.RowsAffected() == 0 { + return ErrNotFound + } + return nil +} + +func (s *PostgreSQLStore) Close() error { + return nil +} + +func scanPostgreSQLOverride(scanner interface{ Scan(dest ...any) error }) (Override, error) { + var override Override + var enabled *bool + var allowedOnlyForUserPaths []byte + var createdAt int64 + var updatedAt int64 + if err := scanner.Scan( + &override.Selector, + &override.ProviderName, + &override.Model, + &enabled, + &override.ForceDisabled, + &allowedOnlyForUserPaths, + &createdAt, + &updatedAt, + ); err != nil { + if err == pgx.ErrNoRows { + return Override{}, ErrNotFound + } + return Override{}, fmt.Errorf("scan model override: %w", err) + } + override.Enabled = enabled + if err := json.Unmarshal(allowedOnlyForUserPaths, &override.AllowedOnlyForUserPaths); err != nil { + return Override{}, fmt.Errorf("decode allowed_only_for_user_paths: %w", err) + } + override.CreatedAt = time.Unix(createdAt, 0).UTC() + override.UpdatedAt = time.Unix(updatedAt, 0).UTC() + return override, nil +} diff --git a/internal/modeloverrides/store_sqlite.go b/internal/modeloverrides/store_sqlite.go new file mode 100644 index 000000000..e6133761f --- /dev/null +++ b/internal/modeloverrides/store_sqlite.go @@ -0,0 +1,179 @@ +package modeloverrides + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "strings" + "time" +) + +// SQLiteStore stores model overrides in SQLite. +type SQLiteStore struct { + db *sql.DB +} + +// NewSQLiteStore creates the model_overrides table and indexes if needed. +func NewSQLiteStore(db *sql.DB) (*SQLiteStore, error) { + if db == nil { + return nil, fmt.Errorf("database connection is required") + } + + _, err := db.Exec(` + CREATE TABLE IF NOT EXISTS model_overrides ( + selector TEXT PRIMARY KEY, + provider_name TEXT NOT NULL DEFAULT '', + model TEXT NOT NULL DEFAULT '', + enabled INTEGER NULL, + force_disabled INTEGER NOT NULL DEFAULT 0, + allowed_only_for_user_paths TEXT NOT NULL DEFAULT '[]', + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ) + `) + if err != nil { + return nil, fmt.Errorf("failed to create model_overrides table: %w", err) + } + if _, err := db.Exec(`CREATE INDEX IF NOT EXISTS idx_model_overrides_provider_name ON model_overrides(provider_name)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides provider_name index: %w", err) + } + if _, err := db.Exec(`CREATE INDEX IF NOT EXISTS idx_model_overrides_model ON model_overrides(model)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides model index: %w", err) + } + if _, err := db.Exec(`CREATE INDEX IF NOT EXISTS idx_model_overrides_updated_at ON model_overrides(updated_at DESC)`); err != nil { + return nil, fmt.Errorf("failed to create model_overrides updated_at index: %w", err) + } + return &SQLiteStore{db: db}, nil +} + +func (s *SQLiteStore) List(ctx context.Context) ([]Override, error) { + rows, err := s.db.QueryContext(ctx, ` + SELECT selector, provider_name, model, enabled, force_disabled, allowed_only_for_user_paths, created_at, updated_at + FROM model_overrides + ORDER BY selector ASC + `) + if err != nil { + return nil, fmt.Errorf("list model overrides: %w", err) + } + defer rows.Close() + return collectOverrides(func() (Override, bool, error) { + if !rows.Next() { + return Override{}, false, nil + } + override, err := scanSQLiteOverride(rows) + return override, true, err + }, rows.Err) +} + +func (s *SQLiteStore) Upsert(ctx context.Context, override Override) error { + override, err := normalizeStoredOverride(override) + if err != nil { + return err + } + + pathsJSON, err := json.Marshal(override.AllowedOnlyForUserPaths) + if err != nil { + return fmt.Errorf("encode allowed_only_for_user_paths: %w", err) + } + + now := time.Now().UTC().Unix() + if override.CreatedAt.IsZero() { + override.CreatedAt = time.Unix(now, 0).UTC() + } + override.UpdatedAt = time.Unix(now, 0).UTC() + + _, err = s.db.ExecContext(ctx, ` + INSERT INTO model_overrides ( + selector, provider_name, model, enabled, force_disabled, allowed_only_for_user_paths, created_at, updated_at + ) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(selector) DO UPDATE SET + provider_name = excluded.provider_name, + model = excluded.model, + enabled = excluded.enabled, + force_disabled = excluded.force_disabled, + allowed_only_for_user_paths = excluded.allowed_only_for_user_paths, + updated_at = excluded.updated_at + `, + override.Selector, + override.ProviderName, + override.Model, + sqliteNullableBool(override.Enabled), + boolToSQLite(override.ForceDisabled), + string(pathsJSON), + override.CreatedAt.Unix(), + override.UpdatedAt.Unix(), + ) + if err != nil { + return fmt.Errorf("upsert model override: %w", err) + } + return nil +} + +func (s *SQLiteStore) Delete(ctx context.Context, selector string) error { + result, err := s.db.ExecContext(ctx, `DELETE FROM model_overrides WHERE selector = ?`, strings.TrimSpace(selector)) + if err != nil { + return fmt.Errorf("delete model override: %w", err) + } + affected, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read delete rows affected: %w", err) + } + if affected == 0 { + return ErrNotFound + } + return nil +} + +func (s *SQLiteStore) Close() error { + return nil +} + +func scanSQLiteOverride(scanner interface{ Scan(dest ...any) error }) (Override, error) { + var override Override + var enabled sql.NullBool + var forceDisabled int + var allowedOnlyForUserPaths string + var createdAt int64 + var updatedAt int64 + if err := scanner.Scan( + &override.Selector, + &override.ProviderName, + &override.Model, + &enabled, + &forceDisabled, + &allowedOnlyForUserPaths, + &createdAt, + &updatedAt, + ); err != nil { + return Override{}, fmt.Errorf("scan model override: %w", err) + } + if enabled.Valid { + override.Enabled = &enabled.Bool + } + override.ForceDisabled = forceDisabled != 0 + if err := json.Unmarshal([]byte(allowedOnlyForUserPaths), &override.AllowedOnlyForUserPaths); err != nil { + return Override{}, fmt.Errorf("decode allowed_only_for_user_paths: %w", err) + } + override.CreatedAt = time.Unix(createdAt, 0).UTC() + override.UpdatedAt = time.Unix(updatedAt, 0).UTC() + return override, nil +} + +func sqliteNullableBool(value *bool) any { + if value == nil { + return nil + } + if *value { + return 1 + } + return 0 +} + +func boolToSQLite(value bool) int { + if value { + return 1 + } + return 0 +} diff --git a/internal/modeloverrides/types.go b/internal/modeloverrides/types.go new file mode 100644 index 000000000..b24f06d36 --- /dev/null +++ b/internal/modeloverrides/types.go @@ -0,0 +1,256 @@ +package modeloverrides + +import ( + "sort" + "strings" + "time" + + "gomodel/internal/core" +) + +// Override stores one persisted access-policy override for a model selector. +// +// Selector syntax: +// - model +// - provider/model +// - provider/ +// +// The first slash separates provider name from model. When the prefix is not a +// configured provider name, the full value is treated as a raw model ID. +type Override struct { + Selector string `json:"selector" bson:"_id"` + ProviderName string `json:"provider_name,omitempty" bson:"provider_name,omitempty"` + Model string `json:"model,omitempty" bson:"model,omitempty"` + Enabled *bool `json:"enabled,omitempty" bson:"enabled,omitempty"` + ForceDisabled bool `json:"force_disabled,omitempty" bson:"force_disabled,omitempty"` + AllowedOnlyForUserPaths []string `json:"allowed_only_for_user_paths,omitempty" bson:"allowed_only_for_user_paths,omitempty"` + CreatedAt time.Time `json:"created_at" bson:"created_at"` + UpdatedAt time.Time `json:"updated_at" bson:"updated_at"` +} + +// ScopeKind identifies how broadly an override applies. +type ScopeKind string + +const ( + ScopeModel ScopeKind = "model" + ScopeProvider ScopeKind = "provider" + ScopeProviderModel ScopeKind = "provider_model" +) + +// ScopeKind reports the normalized selector scope for one override. +func (o Override) ScopeKind() ScopeKind { + switch { + case strings.TrimSpace(o.ProviderName) != "" && strings.TrimSpace(o.Model) != "": + return ScopeProviderModel + case strings.TrimSpace(o.ProviderName) != "": + return ScopeProvider + default: + return ScopeModel + } +} + +// View is the admin-facing representation of one persisted override. +type View struct { + Override + ScopeKind ScopeKind `json:"scope_kind"` +} + +// EffectiveState is the compiled access decision for one concrete selector. +type EffectiveState struct { + Selector string `json:"selector"` + ProviderName string `json:"provider_name,omitempty"` + Model string `json:"model,omitempty"` + DefaultEnabled bool `json:"default_enabled"` + Enabled bool `json:"enabled"` + ForceDisabled bool `json:"force_disabled"` + AllowedOnlyForUserPaths []string `json:"allowed_only_for_user_paths,omitempty"` +} + +// Catalog is the minimal configured-provider surface needed for selector validation. +type Catalog interface { + ProviderNames() []string +} + +func normalizeOverrideInput(catalog Catalog, override Override) (Override, error) { + selector, providerName, model, err := normalizeSelectorInput(selectorProviderNames(catalog), override.Selector) + if err != nil { + return Override{}, err + } + + override.Selector = selector + override.ProviderName = providerName + override.Model = model + + if override.ForceDisabled && override.Enabled != nil && *override.Enabled { + return Override{}, newValidationError("force_disabled cannot be combined with enabled=true", nil) + } + + paths, err := normalizeUserPaths(override.AllowedOnlyForUserPaths) + if err != nil { + return Override{}, err + } + override.AllowedOnlyForUserPaths = paths + return override, nil +} + +func normalizeStoredOverride(override Override) (Override, error) { + override.Selector = strings.TrimSpace(override.Selector) + override.ProviderName = strings.TrimSpace(override.ProviderName) + override.Model = strings.TrimSpace(override.Model) + + if override.Selector == "" { + override.Selector = selectorString(override.ProviderName, override.Model) + } + if override.Selector == "" { + return Override{}, newValidationError("selector is required", nil) + } + if override.ProviderName == "" && override.Model == "" { + providerName, model := parseStoredSelectorParts(override.Selector) + override.ProviderName = providerName + override.Model = model + } + if override.ProviderName == "" && override.Model == "" { + return Override{}, newValidationError("selector is required", nil) + } + if normalized := selectorString(override.ProviderName, override.Model); normalized != "" { + override.Selector = normalized + } + + paths, err := normalizeUserPaths(override.AllowedOnlyForUserPaths) + if err != nil { + return Override{}, err + } + override.AllowedOnlyForUserPaths = paths + return override, nil +} + +func normalizeSelectorInput(providerNames []string, raw string) (selector, providerName, model string, err error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", "", "", newValidationError("selector is required", nil) + } + + providerNameSet := make(map[string]struct{}, len(providerNames)) + for _, name := range providerNames { + name = strings.TrimSpace(name) + if name == "" { + continue + } + providerNameSet[name] = struct{}{} + } + + if prefix, rest, ok := splitFirst(raw); ok { + if _, exists := providerNameSet[prefix]; exists { + providerName = prefix + model = rest + } else { + model = raw + } + } else { + model = raw + } + + if providerName == "" && model == "" { + return "", "", "", newValidationError("selector is required", nil) + } + if providerName != "" { + if _, exists := providerNameSet[providerName]; !exists { + return "", "", "", newValidationError("unknown provider_name: "+providerName, nil) + } + } + return selectorString(providerName, model), providerName, model, nil +} + +func selectorProviderNames(catalog Catalog) []string { + if catalog == nil { + return nil + } + return append([]string(nil), catalog.ProviderNames()...) +} + +func normalizeUserPaths(paths []string) ([]string, error) { + if len(paths) == 0 { + return nil, nil + } + + seen := make(map[string]struct{}, len(paths)) + normalized := make([]string, 0, len(paths)) + for _, raw := range paths { + path, err := core.NormalizeUserPath(raw) + if err != nil { + return nil, newValidationError("invalid allowed_only_for_user_paths value", err) + } + if path == "" { + continue + } + if _, exists := seen[path]; exists { + continue + } + seen[path] = struct{}{} + normalized = append(normalized, path) + } + sort.Strings(normalized) + if len(normalized) == 0 { + return nil, nil + } + return normalized, nil +} + +func selectorString(providerName, model string) string { + providerName = strings.TrimSpace(providerName) + model = strings.TrimSpace(model) + switch { + case providerName != "" && model != "": + return providerName + "/" + model + case providerName != "": + return providerName + "/" + case model != "": + return model + default: + return "" + } +} + +func exactMatchKey(providerName, model string) string { + providerName = strings.TrimSpace(providerName) + model = strings.TrimSpace(model) + if providerName == "" || model == "" { + return "" + } + return providerName + "/" + model +} + +func splitFirst(value string) (prefix, rest string, ok bool) { + parts := strings.SplitN(strings.TrimSpace(value), "/", 2) + if len(parts) != 2 { + return "", "", false + } + prefix = strings.TrimSpace(parts[0]) + rest = strings.TrimSpace(parts[1]) + if prefix == "" { + return "", "", false + } + return prefix, rest, true +} + +func parseStoredSelectorParts(selector string) (providerName, model string) { + selector = strings.TrimSpace(selector) + if selector == "" { + return "", "" + } + if strings.HasSuffix(selector, "/") { + return strings.TrimSpace(strings.TrimSuffix(selector, "/")), "" + } + if providerName, model, ok := splitFirst(selector); ok { + return providerName, model + } + return "", selector +} + +func cloneEnabled(value *bool) *bool { + if value == nil { + return nil + } + enabled := *value + return &enabled +} diff --git a/internal/providers/registry.go b/internal/providers/registry.go index bbcc3d6b5..6ea4de726 100644 --- a/internal/providers/registry.go +++ b/internal/providers/registry.go @@ -680,6 +680,25 @@ func (r *ModelRegistry) GetProviderNameForType(providerType string) string { return "" } +// GetProviderTypeForName returns the provider type for the given concrete +// configured provider instance name. +func (r *ModelRegistry) GetProviderTypeForName(providerName string) string { + r.mu.RLock() + defer r.mu.RUnlock() + + providerName = strings.TrimSpace(providerName) + if providerName == "" { + return "" + } + for _, provider := range r.providers { + if strings.TrimSpace(r.providerNames[provider]) != providerName { + continue + } + return strings.TrimSpace(r.providerTypes[provider]) + } + return "" +} + // ProviderByType returns the first registered provider for the given provider type. // This lookup is independent of discovered models so provider-typed routes keep // working even when a provider currently exposes zero models. diff --git a/internal/providers/router.go b/internal/providers/router.go index d14bef96f..a9a4d99f2 100644 --- a/internal/providers/router.go +++ b/internal/providers/router.go @@ -543,6 +543,29 @@ func (r *Router) GetProviderNameForType(providerType string) string { return "" } +// GetProviderTypeForName returns the provider type for a concrete configured +// provider instance name. +func (r *Router) GetProviderTypeForName(providerName string) string { + providerName = strings.TrimSpace(providerName) + if providerName == "" { + return "" + } + if typed, ok := r.lookup.(core.ProviderNameTypeResolver); ok { + return strings.TrimSpace(typed.GetProviderTypeForName(providerName)) + } + if models, ok := r.lookup.(modelWithProviderLister); ok { + for _, entry := range models.ListModelsWithProvider() { + if strings.TrimSpace(entry.ProviderName) != providerName { + continue + } + if providerType := strings.TrimSpace(entry.ProviderType); providerType != "" { + return providerType + } + } + } + return "" +} + func (r *Router) providerByType(providerType string) core.Provider { models := r.lookup.ListModels() for _, model := range models { diff --git a/internal/server/execution_plan_helpers.go b/internal/server/execution_plan_helpers.go index 503223c7d..f9473740c 100644 --- a/internal/server/execution_plan_helpers.go +++ b/internal/server/execution_plan_helpers.go @@ -16,6 +16,18 @@ func ensureTranslatedRequestPlan( policyResolver RequestExecutionPolicyResolver, model, providerHint *string, +) (*core.ExecutionPlan, error) { + return ensureTranslatedRequestPlanWithAuthorizer(c, provider, resolver, nil, policyResolver, model, providerHint) +} + +func ensureTranslatedRequestPlanWithAuthorizer( + c *echo.Context, + provider core.RoutableProvider, + resolver RequestModelResolver, + authorizer RequestModelAuthorizer, + policyResolver RequestExecutionPolicyResolver, + model, + providerHint *string, ) (*core.ExecutionPlan, error) { if model == nil || providerHint == nil { return nil, core.NewInvalidRequestError("model selector targets are required", nil) @@ -27,8 +39,13 @@ func ensureTranslatedRequestPlan( } resolution := translatedPlanResolution(plan) + if resolution != nil && authorizer != nil { + if err := authorizer.ValidateModelAccess(c.Request().Context(), resolution.ResolvedSelector); err != nil { + return nil, err + } + } if resolution == nil { - resolution, err = resolveAndStoreRequestModelResolution(c, provider, resolver, *model, *providerHint) + resolution, err = resolveAndStoreRequestModelResolution(c, provider, resolver, authorizer, *model, *providerHint) if err != nil { return nil, err } diff --git a/internal/server/exposed_model_lister.go b/internal/server/exposed_model_lister.go index 4e0444bbd..83ccf8d10 100644 --- a/internal/server/exposed_model_lister.go +++ b/internal/server/exposed_model_lister.go @@ -11,6 +11,11 @@ type ExposedModelLister interface { ExposedModels() []core.Model } +// FilteredExposedModelLister optionally filters exposed models using their concrete targets. +type FilteredExposedModelLister interface { + ExposedModelsFiltered(allow func(core.ModelSelector) bool) []core.Model +} + func mergeExposedModelsResponse(base *core.ModelsResponse, exposed []core.Model) *core.ModelsResponse { if base == nil { base = &core.ModelsResponse{Object: "list", Data: []core.Model{}} diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 58c0c43bc..45e960295 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -18,6 +18,7 @@ import ( type Handler struct { provider core.RoutableProvider modelResolver RequestModelResolver + modelAuthorizer RequestModelAuthorizer fallbackResolver RequestFallbackResolver executionPolicyResolver RequestExecutionPolicyResolver translatedRequestPatcher TranslatedRequestPatcher @@ -50,10 +51,35 @@ func newHandler( executionPolicyResolver RequestExecutionPolicyResolver, fallbackResolver RequestFallbackResolver, translatedRequestPatcher TranslatedRequestPatcher, +) *Handler { + return newHandlerWithAuthorizer( + provider, + logger, + usageLogger, + pricingResolver, + modelResolver, + nil, + executionPolicyResolver, + fallbackResolver, + translatedRequestPatcher, + ) +} + +func newHandlerWithAuthorizer( + provider core.RoutableProvider, + logger auditlog.LoggerInterface, + usageLogger usage.LoggerInterface, + pricingResolver usage.PricingResolver, + modelResolver RequestModelResolver, + modelAuthorizer RequestModelAuthorizer, + executionPolicyResolver RequestExecutionPolicyResolver, + fallbackResolver RequestFallbackResolver, + translatedRequestPatcher TranslatedRequestPatcher, ) *Handler { return &Handler{ provider: provider, modelResolver: modelResolver, + modelAuthorizer: modelAuthorizer, fallbackResolver: fallbackResolver, executionPolicyResolver: executionPolicyResolver, translatedRequestPatcher: translatedRequestPatcher, @@ -80,6 +106,7 @@ func (h *Handler) translatedInference() *translatedInferenceService { s := &translatedInferenceService{ provider: h.provider, modelResolver: h.modelResolver, + modelAuthorizer: h.modelAuthorizer, executionPolicyResolver: h.executionPolicyResolver, fallbackResolver: h.fallbackResolver, translatedRequestPatcher: h.translatedRequestPatcher, @@ -99,6 +126,7 @@ func (h *Handler) nativeBatch() *nativeBatchService { return &nativeBatchService{ provider: h.provider, modelResolver: h.modelResolver, + modelAuthorizer: h.modelAuthorizer, executionPolicyResolver: h.executionPolicyResolver, batchRequestPreparer: h.batchRequestPreparer, batchStore: h.batchStore, @@ -116,6 +144,7 @@ func (h *Handler) nativeFiles() *nativeFileService { func (h *Handler) passthrough() *passthroughService { return &passthroughService{ provider: h.provider, + modelAuthorizer: h.modelAuthorizer, logger: h.logger, usageLogger: h.usageLogger, pricingResolver: h.pricingResolver, @@ -207,8 +236,32 @@ func (h *Handler) ListModels(c *echo.Context) error { if err != nil { return handleError(c, err) } + if h.modelAuthorizer != nil && resp != nil { + resp = &core.ModelsResponse{ + Object: resp.Object, + Data: h.modelAuthorizer.FilterPublicModels(c.Request().Context(), resp.Data), + } + } if h.exposedModelLister != nil { - resp = mergeExposedModelsResponse(resp, h.exposedModelLister.ExposedModels()) + if filtered, ok := h.exposedModelLister.(FilteredExposedModelLister); ok && h.modelAuthorizer != nil { + resp = mergeExposedModelsResponse(resp, filtered.ExposedModelsFiltered(func(selector core.ModelSelector) bool { + return h.modelAuthorizer.AllowsModel(c.Request().Context(), selector) + })) + } else { + exposed := h.exposedModelLister.ExposedModels() + if h.modelAuthorizer != nil { + filtered := make([]core.Model, 0, len(exposed)) + for _, model := range exposed { + selector, err := core.ParseModelSelector(model.ID, "") + if err != nil || !h.modelAuthorizer.AllowsModel(c.Request().Context(), selector) { + continue + } + filtered = append(filtered, model) + } + exposed = filtered + } + resp = mergeExposedModelsResponse(resp, exposed) + } } return c.JSON(http.StatusOK, resp) diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index a7d4924d2..9464a683a 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -357,6 +357,36 @@ type fileListCall struct { after string } +type recordingModelAuthorizer struct { + lastSelector core.ModelSelector + err error + allow func(core.ModelSelector) bool +} + +func (a *recordingModelAuthorizer) ValidateModelAccess(_ context.Context, selector core.ModelSelector) error { + a.lastSelector = selector + return a.err +} + +func (a *recordingModelAuthorizer) AllowsModel(_ context.Context, selector core.ModelSelector) bool { + if a.allow != nil { + return a.allow(selector) + } + return true +} + +func (a *recordingModelAuthorizer) FilterPublicModels(_ context.Context, models []core.Model) []core.Model { + return models +} + +type staticExposedModelLister struct { + models []core.Model +} + +func (l staticExposedModelLister) ExposedModels() []core.Model { + return append([]core.Model(nil), l.models...) +} + func readPassthroughRequestBody(t *testing.T, body io.ReadCloser) string { t.Helper() if body == nil { @@ -443,6 +473,22 @@ func (m *mockProvider) GetProviderNameForType(providerType string) string { return "" } +func (m *mockProvider) GetProviderTypeForName(providerName string) string { + providerName = strings.TrimSpace(providerName) + if providerName == "" || len(m.providerNames) == 0 { + return "" + } + for qualifiedModel, candidate := range m.providerNames { + if strings.TrimSpace(candidate) != providerName { + continue + } + if providerType := strings.TrimSpace(m.providerTypes[qualifiedModel]); providerType != "" { + return providerType + } + } + return "" +} + func inferQualifiedProviderValue(values map[string]string, model string) (string, bool) { model = strings.TrimSpace(model) if model == "" { @@ -2585,6 +2631,45 @@ func TestListModels_MergesExposedModelsWithoutAliasProviderDecorator(t *testing. require.Contains(t, body, `"id":"smart"`) } +func TestListModels_FiltersExposedModelsWhenAuthorizerIsPresent(t *testing.T) { + mock := &mockProvider{ + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{ + {ID: "gpt-4o", Object: "model", OwnedBy: "openai"}, + }, + }, + } + authorizer := &recordingModelAuthorizer{ + allow: func(selector core.ModelSelector) bool { + return selector.QualifiedModel() != "openai/gpt-5" + }, + } + + e := echo.New() + handler := NewHandler(mock, nil, nil, nil) + handler.modelAuthorizer = authorizer + handler.exposedModelLister = staticExposedModelLister{ + models: []core.Model{ + {ID: "openai/gpt-5", Object: "model", OwnedBy: "openai"}, + {ID: "openai/gpt-4o-mini", Object: "model", OwnedBy: "openai"}, + }, + } + + req := httptest.NewRequest(http.MethodGet, "/v1/models", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + err := handler.ListModels(c) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code) + + body := rec.Body.String() + require.Contains(t, body, `"id":"gpt-4o"`) + require.Contains(t, body, `"id":"openai/gpt-4o-mini"`) + require.NotContains(t, body, `"id":"openai/gpt-5"`) +} + func TestListModelsError(t *testing.T) { mock := &mockProvider{ err: io.EOF, // Simulate an error @@ -5670,6 +5755,102 @@ func TestProviderPassthrough_UsesPassthroughModelForAuditEntry(t *testing.T) { } } +func TestProviderPassthrough_UsesConfiguredProviderNameForAccessValidation(t *testing.T) { + provider := &mockProvider{ + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"ok":true}`)), + }, + providerTypes: map[string]string{ + "openai_test/gpt-5-mini": "openai", + }, + providerNames: map[string]string{ + "openai_test/gpt-5-mini": "openai_test", + }, + } + authorizer := &recordingModelAuthorizer{} + + e := echo.New() + handler := newHandlerWithAuthorizer(provider, nil, nil, nil, nil, authorizer, nil, nil, nil) + req := httptest.NewRequest(http.MethodPost, "/p/openai_test/chat/completions", strings.NewReader(`{"model":"gpt-5-mini"}`)) + req.Header.Set("Content-Type", "application/json") + req = req.WithContext(core.WithExecutionPlan(req.Context(), &core.ExecutionPlan{ + Mode: core.ExecutionModePassthrough, + ProviderType: "openai", + Passthrough: &core.PassthroughRouteInfo{ + Provider: "openai_test", + RawEndpoint: "chat/completions", + NormalizedEndpoint: "chat/completions", + Model: "gpt-5-mini", + AuditPath: "/p/openai_test/chat/completions", + }, + })) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + if err := handler.ProviderPassthrough(c); err != nil { + t.Fatalf("handler returned error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + if provider.lastPassthroughProvider != "openai" { + t.Fatalf("providerType = %q, want openai", provider.lastPassthroughProvider) + } + if authorizer.lastSelector.Provider != "openai_test" || authorizer.lastSelector.Model != "gpt-5-mini" { + t.Fatalf("validated selector = %#v, want openai_test/gpt-5-mini", authorizer.lastSelector) + } +} + +func TestProviderPassthrough_FallsBackFromProviderTypeToCanonicalProviderNameForAccessValidation(t *testing.T) { + provider := &mockProvider{ + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"ok":true}`)), + }, + providerTypes: map[string]string{ + "openai_test/gpt-5-mini": "openai", + }, + providerNames: map[string]string{ + "openai_test/gpt-5-mini": "openai_test", + }, + } + authorizer := &recordingModelAuthorizer{} + + e := echo.New() + handler := newHandlerWithAuthorizer(provider, nil, nil, nil, nil, authorizer, nil, nil, nil) + req := httptest.NewRequest(http.MethodPost, "/p/openai/chat/completions", strings.NewReader(`{"model":"gpt-5-mini"}`)) + req.Header.Set("Content-Type", "application/json") + req = req.WithContext(core.WithExecutionPlan(req.Context(), &core.ExecutionPlan{ + Mode: core.ExecutionModePassthrough, + ProviderType: "openai", + Passthrough: &core.PassthroughRouteInfo{ + Provider: "openai", + RawEndpoint: "chat/completions", + NormalizedEndpoint: "chat/completions", + Model: "gpt-5-mini", + AuditPath: "/p/openai/chat/completions", + }, + })) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + if err := handler.ProviderPassthrough(c); err != nil { + t.Fatalf("handler returned error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + if provider.lastPassthroughProvider != "openai" { + t.Fatalf("providerType = %q, want openai", provider.lastPassthroughProvider) + } + if authorizer.lastSelector.Provider != "openai_test" || authorizer.lastSelector.Model != "gpt-5-mini" { + t.Fatalf("validated selector = %#v, want openai_test/gpt-5-mini", authorizer.lastSelector) + } +} + func TestProviderPassthrough_V1AliasDisabledReturnsBadRequest(t *testing.T) { provider := &mockProvider{ passthroughResponse: &core.PassthroughResponse{ diff --git a/internal/server/http.go b/internal/server/http.go index 54eb858dc..46ad8a5fa 100644 --- a/internal/server/http.go +++ b/internal/server/http.go @@ -52,6 +52,7 @@ type Config struct { UsageLogger usage.LoggerInterface // Optional: Usage logger for token tracking PricingResolver usage.PricingResolver // Optional: Resolves pricing for cost calculation ModelResolver RequestModelResolver // Optional: explicit model resolver used during request planning + ModelAuthorizer RequestModelAuthorizer // Optional: request-scoped concrete model access controller ExecutionPolicyResolver RequestExecutionPolicyResolver // Optional: persisted execution-plan resolver used during request planning FallbackResolver RequestFallbackResolver // Optional: translated-route fallback resolver TranslatedRequestPatcher TranslatedRequestPatcher // Optional: request patcher for translated routes after planning @@ -88,17 +89,19 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { } var modelResolver RequestModelResolver + var modelAuthorizer RequestModelAuthorizer var executionPolicyResolver RequestExecutionPolicyResolver var fallbackResolver RequestFallbackResolver var translatedRequestPatcher TranslatedRequestPatcher if cfg != nil { modelResolver = cfg.ModelResolver + modelAuthorizer = cfg.ModelAuthorizer executionPolicyResolver = cfg.ExecutionPolicyResolver fallbackResolver = cfg.FallbackResolver translatedRequestPatcher = cfg.TranslatedRequestPatcher } - handler := newHandler(provider, auditLogger, usageLogger, pricingResolver, modelResolver, executionPolicyResolver, fallbackResolver, translatedRequestPatcher) + handler := newHandlerWithAuthorizer(provider, auditLogger, usageLogger, pricingResolver, modelResolver, modelAuthorizer, executionPolicyResolver, fallbackResolver, translatedRequestPatcher) if cfg != nil { handler.batchRequestPreparer = cfg.BatchRequestPreparer handler.exposedModelLister = cfg.ExposedModelLister @@ -215,7 +218,7 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { e.Use(RequestSnapshotCapture()) if cfg != nil && len(cfg.PassthroughSemanticEnrichers) > 0 { - e.Use(PassthroughSemanticEnrichment(cfg.PassthroughSemanticEnrichers, passthroughV1PrefixNormalizationEnabled(cfg))) + e.Use(PassthroughSemanticEnrichment(provider, cfg.PassthroughSemanticEnrichers, passthroughV1PrefixNormalizationEnabled(cfg))) } // Audit logging runs before request planning so early planning/validation @@ -295,6 +298,9 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { adminAPI.GET("/audit/conversation", cfg.AdminHandler.AuditConversation) adminAPI.GET("/models", cfg.AdminHandler.ListModels) adminAPI.GET("/models/categories", cfg.AdminHandler.ListCategories) + adminAPI.GET("/model-overrides", cfg.AdminHandler.ListModelOverrides) + adminAPI.PUT("/model-overrides/:selector", cfg.AdminHandler.UpsertModelOverride) + adminAPI.DELETE("/model-overrides/:selector", cfg.AdminHandler.DeleteModelOverride) adminAPI.GET("/auth-keys", cfg.AdminHandler.ListAuthKeys) adminAPI.POST("/auth-keys", cfg.AdminHandler.CreateAuthKey) adminAPI.POST("/auth-keys/:id/deactivate", cfg.AdminHandler.DeactivateAuthKey) diff --git a/internal/server/internal_chat_completion_executor.go b/internal/server/internal_chat_completion_executor.go index 6f951992d..9027fde6d 100644 --- a/internal/server/internal_chat_completion_executor.go +++ b/internal/server/internal_chat_completion_executor.go @@ -21,6 +21,7 @@ import ( // chat execution path used by gateway-owned workflows such as guardrails. type InternalChatCompletionExecutorConfig struct { ModelResolver RequestModelResolver + ModelAuthorizer RequestModelAuthorizer ExecutionPolicyResolver RequestExecutionPolicyResolver FallbackResolver RequestFallbackResolver AuditLogger auditlog.LoggerInterface @@ -37,6 +38,7 @@ type InternalChatCompletionExecutor struct { executionPolicyResolver RequestExecutionPolicyResolver logger auditlog.LoggerInterface service *translatedInferenceService + modelAuthorizer RequestModelAuthorizer } // NewInternalChatCompletionExecutor creates a transport-free translated chat @@ -45,6 +47,7 @@ func NewInternalChatCompletionExecutor(provider core.RoutableProvider, cfg Inter service := &translatedInferenceService{ provider: provider, modelResolver: cfg.ModelResolver, + modelAuthorizer: cfg.ModelAuthorizer, executionPolicyResolver: cfg.ExecutionPolicyResolver, fallbackResolver: cfg.FallbackResolver, logger: cfg.AuditLogger, @@ -56,6 +59,7 @@ func NewInternalChatCompletionExecutor(provider core.RoutableProvider, cfg Inter return &InternalChatCompletionExecutor{ provider: provider, modelResolver: cfg.ModelResolver, + modelAuthorizer: cfg.ModelAuthorizer, executionPolicyResolver: cfg.ExecutionPolicyResolver, logger: cfg.AuditLogger, service: service, @@ -87,7 +91,7 @@ func (e *InternalChatCompletionExecutor) ChatCompletion(ctx context.Context, req e.finishAuditEntry(ctx, entry, start, plan, req, resp, err, cacheType, providerType, providerName) }() - resolution, err := resolveRequestModel(e.provider, e.modelResolver, requested) + resolution, err := resolveRequestModelWithAuthorizer(ctx, e.provider, e.modelResolver, e.modelAuthorizer, requested) if err != nil { return nil, err } diff --git a/internal/server/model_access.go b/internal/server/model_access.go new file mode 100644 index 000000000..e9293c306 --- /dev/null +++ b/internal/server/model_access.go @@ -0,0 +1,14 @@ +package server + +import ( + "context" + + "gomodel/internal/core" +) + +// RequestModelAuthorizer validates request-scoped access to concrete models. +type RequestModelAuthorizer interface { + ValidateModelAccess(ctx context.Context, selector core.ModelSelector) error + AllowsModel(ctx context.Context, selector core.ModelSelector) bool + FilterPublicModels(ctx context.Context, models []core.Model) []core.Model +} diff --git a/internal/server/model_validation.go b/internal/server/model_validation.go index 62c2ac8dc..7ccb4fb97 100644 --- a/internal/server/model_validation.go +++ b/internal/server/model_validation.go @@ -80,18 +80,16 @@ func deriveExecutionPlanWithPolicy( switch desc.Operation { case core.OperationProviderPassthrough: passthrough := passthroughRouteInfo(c) - providerType, ok := providerPassthroughType(c) + providerType, providerName, ok := providerPassthroughType(c, provider) if !ok { return nil, nil } if passthrough == nil { passthrough = &core.PassthroughRouteInfo{} } - if strings.TrimSpace(passthrough.Provider) == "" { - cloned := *passthrough - cloned.Provider = providerType - passthrough = &cloned - } + cloned := *passthrough + cloned.Provider = providerType + passthrough = &cloned plan.Mode = core.ExecutionModePassthrough plan.ProviderType = providerType plan.Passthrough = passthrough @@ -99,7 +97,7 @@ func deriveExecutionPlanWithPolicy( c.Request().Context(), plan, policyResolver, - core.NewExecutionPlanSelector(workflowProviderNameForType(provider, providerType), passthrough.Model, userPath), + core.NewExecutionPlanSelector(providerName, passthrough.Model, userPath), ); err != nil { return nil, err } @@ -215,23 +213,26 @@ func selectorHintValueAllowed(result gjson.Result) bool { return result.Type == gjson.String || result.Type == gjson.Null } -func providerPassthroughType(c *echo.Context) (string, bool) { +func providerPassthroughType(c *echo.Context, provider core.RoutableProvider) (string, string, bool) { if info := passthroughRouteInfo(c); info != nil { - providerType := strings.TrimSpace(info.Provider) - if providerType != "" { - return providerType, true + resolved := resolvePassthroughProvider(provider, info.Provider) + if providerType := strings.TrimSpace(resolved.ProviderType); providerType != "" { + return providerType, strings.TrimSpace(resolved.ProviderName), true } } if env := core.GetWhiteBoxPrompt(c.Request().Context()); env != nil && env.OperationType == string(core.OperationProviderPassthrough) { - providerType := strings.TrimSpace(env.RouteHints.Provider) - if providerType != "" { - return providerType, true + resolved := resolvePassthroughProvider(provider, env.RouteHints.Provider) + if providerType := strings.TrimSpace(resolved.ProviderType); providerType != "" { + return providerType, strings.TrimSpace(resolved.ProviderName), true } } - if providerType, _, ok := core.ParseProviderPassthroughPath(c.Request().URL.Path); ok { - return providerType, true + if routeProvider, _, ok := core.ParseProviderPassthroughPath(c.Request().URL.Path); ok { + resolved := resolvePassthroughProvider(provider, routeProvider) + if providerType := strings.TrimSpace(resolved.ProviderType); providerType != "" { + return providerType, strings.TrimSpace(resolved.ProviderName), true + } } - return "", false + return "", "", false } func passthroughRouteInfo(c *echo.Context) *core.PassthroughRouteInfo { diff --git a/internal/server/model_validation_test.go b/internal/server/model_validation_test.go index 060ac6dae..1c7fe99cb 100644 --- a/internal/server/model_validation_test.go +++ b/internal/server/model_validation_test.go @@ -425,6 +425,55 @@ func TestExecutionPlanning_StoresPassthroughRouteInfo(t *testing.T) { } } +func TestExecutionPlanning_PassthroughProviderNameRouteUsesCanonicalProviderNameForPolicy(t *testing.T) { + provider := &mockProvider{ + providerTypes: map[string]string{ + "openai_test/gpt-5-mini": "openai", + }, + providerNames: map[string]string{ + "openai_test/gpt-5-mini": "openai_test", + }, + } + + e := echo.New() + var capturedSelector core.ExecutionPlanSelector + var capturedPlan *core.ExecutionPlan + + policyResolver := &staticExecutionPolicyResolver{ + match: func(selector core.ExecutionPlanSelector) (*core.ResolvedExecutionPolicy, error) { + capturedSelector = selector + return &core.ResolvedExecutionPolicy{ + VersionID: "plan-passthrough-v1", + Version: 1, + Name: "passthrough", + Features: core.DefaultExecutionFeatures(), + }, nil + }, + } + + middleware := RequestSnapshotCapture() + handler := middleware(ExecutionPlanningWithResolverAndPolicy(provider, nil, policyResolver)(func(c *echo.Context) error { + capturedPlan = core.GetExecutionPlan(c.Request().Context()) + return c.String(http.StatusOK, "ok") + })) + + req := httptest.NewRequest(http.MethodPost, "/p/openai_test/responses", strings.NewReader(`{"model":"gpt-5-mini"}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + err := handler(c) + require.NoError(t, err) + assert.Equal(t, "openai_test", capturedSelector.Provider) + assert.Equal(t, "gpt-5-mini", capturedSelector.Model) + if assert.NotNil(t, capturedPlan) { + assert.Equal(t, "openai", capturedPlan.ProviderType) + if assert.NotNil(t, capturedPlan.Passthrough) { + assert.Equal(t, "openai", capturedPlan.Passthrough.Provider) + } + } +} + func TestExecutionPlanning_PopulatesPassthroughProviderFromPathFallback(t *testing.T) { provider := &mockProvider{} diff --git a/internal/server/native_batch_service.go b/internal/server/native_batch_service.go index 3c122f0a1..e6b697906 100644 --- a/internal/server/native_batch_service.go +++ b/internal/server/native_batch_service.go @@ -22,6 +22,7 @@ import ( type nativeBatchService struct { provider core.RoutableProvider modelResolver RequestModelResolver + modelAuthorizer RequestModelAuthorizer executionPolicyResolver RequestExecutionPolicyResolver batchRequestPreparer BatchRequestPreparer batchStore batchstore.Store @@ -46,7 +47,7 @@ func (s *nativeBatchService) Batches(c *echo.Context) error { return handleError(c, core.NewInvalidRequestError("batch routing is not supported by the current provider router", nil)) } - selection, err := determineBatchExecutionSelection(s.provider, s.modelResolver, req) + selection, err := determineBatchExecutionSelectionWithAuthorizer(ctx, s.provider, s.modelResolver, s.modelAuthorizer, req) if err != nil { return handleError(c, err) } diff --git a/internal/server/native_batch_support.go b/internal/server/native_batch_support.go index 66bef4e05..38548d1b1 100644 --- a/internal/server/native_batch_support.go +++ b/internal/server/native_batch_support.go @@ -31,7 +31,21 @@ type batchExecutionSelection struct { selector core.ExecutionPlanSelector } -func determineBatchExecutionSelection(provider core.RoutableProvider, resolver RequestModelResolver, req *core.BatchRequest) (batchExecutionSelection, error) { +func determineBatchExecutionSelection( + provider core.RoutableProvider, + resolver RequestModelResolver, + req *core.BatchRequest, +) (batchExecutionSelection, error) { + return determineBatchExecutionSelectionWithAuthorizer(context.Background(), provider, resolver, nil, req) +} + +func determineBatchExecutionSelectionWithAuthorizer( + ctx context.Context, + provider core.RoutableProvider, + resolver RequestModelResolver, + authorizer RequestModelAuthorizer, + req *core.BatchRequest, +) (batchExecutionSelection, error) { if provider == nil { return batchExecutionSelection{}, core.NewInvalidRequestError("provider is not configured", nil) } @@ -76,6 +90,11 @@ func determineBatchExecutionSelection(provider core.RoutableProvider, resolver R if !provider.Supports(model) { return batchExecutionSelection{}, core.NewInvalidRequestError("unsupported model: "+model, nil) } + if authorizer != nil { + if err := authorizer.ValidateModelAccess(ctx, resolvedSelector); err != nil { + return batchExecutionSelection{}, err + } + } itemProvider := provider.GetProviderType(model) if providerType == "" { providerType = itemProvider diff --git a/internal/server/passthrough_execution_helpers.go b/internal/server/passthrough_execution_helpers.go index 04259dda9..061728fb5 100644 --- a/internal/server/passthrough_execution_helpers.go +++ b/internal/server/passthrough_execution_helpers.go @@ -8,7 +8,7 @@ import ( "gomodel/internal/core" ) -func passthroughExecutionTarget(c *echo.Context, allowPassthroughV1Alias bool) (string, string, *core.PassthroughRouteInfo, error) { +func passthroughExecutionTarget(c *echo.Context, provider core.RoutableProvider, allowPassthroughV1Alias bool) (string, string, *core.PassthroughRouteInfo, error) { if c == nil { return "", "", nil, core.NewInvalidRequestError("invalid provider passthrough path", nil) } @@ -18,7 +18,7 @@ func passthroughExecutionTarget(c *echo.Context, allowPassthroughV1Alias bool) ( return "", "", nil, core.NewInvalidRequestError("invalid provider passthrough path", nil) } - providerType := strings.TrimSpace(info.Provider) + providerType := strings.TrimSpace(resolvePassthroughProvider(provider, info.Provider).ProviderType) if providerType == "" { if plan := core.GetExecutionPlan(c.Request().Context()); plan != nil { providerType = strings.TrimSpace(plan.ProviderType) diff --git a/internal/server/passthrough_execution_helpers_test.go b/internal/server/passthrough_execution_helpers_test.go index c75f002b9..eda2d4d66 100644 --- a/internal/server/passthrough_execution_helpers_test.go +++ b/internal/server/passthrough_execution_helpers_test.go @@ -26,7 +26,7 @@ func TestPassthroughExecutionTarget_PrefersExecutionPlan(t *testing.T) { rec := httptest.NewRecorder() c := e.NewContext(req, rec) - providerType, endpoint, info, err := passthroughExecutionTarget(c, false) + providerType, endpoint, info, err := passthroughExecutionTarget(c, nil, false) if err != nil { t.Fatalf("passthroughExecutionTarget() error = %v", err) } @@ -50,7 +50,7 @@ func TestPassthroughExecutionTarget_NormalizesFallbackFromPath(t *testing.T) { rec := httptest.NewRecorder() c := e.NewContext(req, rec) - providerType, endpoint, info, err := passthroughExecutionTarget(c, true) + providerType, endpoint, info, err := passthroughExecutionTarget(c, nil, true) if err != nil { t.Fatalf("passthroughExecutionTarget() error = %v", err) } @@ -67,3 +67,33 @@ func TestPassthroughExecutionTarget_NormalizesFallbackFromPath(t *testing.T) { t.Fatalf("NormalizedEndpoint = %q, want responses", info.NormalizedEndpoint) } } + +func TestPassthroughExecutionTarget_ResolvesConfiguredProviderNameToType(t *testing.T) { + e := echo.New() + req := httptest.NewRequest(http.MethodPost, "/p/openai_test/v1/responses?trace=1", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + provider := &mockProvider{ + providerTypes: map[string]string{ + "openai_test/gpt-5-mini": "openai", + }, + providerNames: map[string]string{ + "openai_test/gpt-5-mini": "openai_test", + }, + } + + providerType, endpoint, info, err := passthroughExecutionTarget(c, provider, true) + if err != nil { + t.Fatalf("passthroughExecutionTarget() error = %v", err) + } + if providerType != "openai" { + t.Fatalf("providerType = %q, want openai", providerType) + } + if endpoint != "responses?trace=1" { + t.Fatalf("endpoint = %q, want responses?trace=1", endpoint) + } + if info == nil || info.Provider != "openai" { + t.Fatalf("info.Provider = %#v, want openai", info) + } +} diff --git a/internal/server/passthrough_provider_resolution.go b/internal/server/passthrough_provider_resolution.go new file mode 100644 index 000000000..8e8927016 --- /dev/null +++ b/internal/server/passthrough_provider_resolution.go @@ -0,0 +1,84 @@ +package server + +import ( + "strings" + + "gomodel/internal/core" +) + +type passthroughProviderResolution struct { + RouteProvider string + ProviderType string + ProviderName string +} + +func resolvePassthroughProvider(provider core.RoutableProvider, routeProvider string) passthroughProviderResolution { + routeProvider = strings.TrimSpace(routeProvider) + if routeProvider == "" { + return passthroughProviderResolution{} + } + + if provider != nil { + if named, ok := provider.(core.ProviderNameTypeResolver); ok { + if providerType := strings.TrimSpace(named.GetProviderTypeForName(routeProvider)); providerType != "" { + return passthroughProviderResolution{ + RouteProvider: routeProvider, + ProviderType: providerType, + ProviderName: routeProvider, + } + } + } + } + + return passthroughProviderResolution{ + RouteProvider: routeProvider, + ProviderType: routeProvider, + ProviderName: workflowProviderNameForType(provider, routeProvider), + } +} + +// passthroughAccessSelector derives an authorization selector from provider, +// which supplies provider name/type canonicalization, and info, which carries +// the passthrough route provider/model; it returns the selector and whether one +// could be built. It may intentionally return a core.ModelSelector with an +// empty Provider when resolvePassthroughProvider leaves ProviderName empty and +// none of the ProviderNameResolver.GetProviderName candidates resolve to a +// non-empty canonical name; downstream authorization/validation is expected to +// handle empty Provider values. +func passthroughAccessSelector(provider core.RoutableProvider, info *core.PassthroughRouteInfo) (core.ModelSelector, bool) { + if info == nil { + return core.ModelSelector{}, false + } + + model := strings.TrimSpace(info.Model) + if model == "" { + return core.ModelSelector{}, false + } + + routeProvider := strings.TrimSpace(info.Provider) + resolvedProvider := resolvePassthroughProvider(provider, routeProvider) + providerName := strings.TrimSpace(resolvedProvider.ProviderName) + + if named, ok := provider.(core.ProviderNameResolver); ok { + candidates := make([]string, 0, 3) + if routeProvider != "" { + candidates = append(candidates, routeProvider+"/"+model) + } + if resolvedProvider.ProviderType != "" && resolvedProvider.ProviderType != routeProvider { + candidates = append(candidates, resolvedProvider.ProviderType+"/"+model) + } + candidates = append(candidates, model) + + for _, candidate := range candidates { + if canonical := strings.TrimSpace(named.GetProviderName(candidate)); canonical != "" { + providerName = canonical + break + } + } + } + + return core.ModelSelector{ + Provider: providerName, + Model: model, + }, true +} diff --git a/internal/server/passthrough_semantic_enrichment.go b/internal/server/passthrough_semantic_enrichment.go index 6a44c832c..09dc0d5d2 100644 --- a/internal/server/passthrough_semantic_enrichment.go +++ b/internal/server/passthrough_semantic_enrichment.go @@ -10,7 +10,7 @@ import ( // PassthroughSemanticEnrichment applies provider-owned passthrough metadata // enrichment before execution planning runs. -func PassthroughSemanticEnrichment(enrichers []core.PassthroughSemanticEnricher, allowPassthroughV1Alias bool) echo.MiddlewareFunc { +func PassthroughSemanticEnrichment(provider core.RoutableProvider, enrichers []core.PassthroughSemanticEnricher, allowPassthroughV1Alias bool) echo.MiddlewareFunc { byProvider := make(map[string]core.PassthroughSemanticEnricher, len(enrichers)) for _, enricher := range enrichers { if enricher == nil { @@ -40,6 +40,9 @@ func PassthroughSemanticEnrichment(enrichers []core.PassthroughSemanticEnricher, return next(c) } info.NormalizedEndpoint = normalized + if resolved := resolvePassthroughProvider(provider, info.Provider); resolved.ProviderType != "" { + info.Provider = resolved.ProviderType + } if enricher := byProvider[strings.TrimSpace(info.Provider)]; enricher != nil { if enriched := enricher.Enrich(core.GetRequestSnapshot(c.Request().Context()), env, info); enriched != nil { diff --git a/internal/server/passthrough_semantic_enrichment_test.go b/internal/server/passthrough_semantic_enrichment_test.go index 56dd2519c..9c6b9f6ef 100644 --- a/internal/server/passthrough_semantic_enrichment_test.go +++ b/internal/server/passthrough_semantic_enrichment_test.go @@ -40,7 +40,7 @@ func TestPassthroughSemanticEnrichment_EnrichesPromptBeforePlanning(t *testing.T c := e.NewContext(req, rec) var capturedPlan *core.ExecutionPlan - handler := PassthroughSemanticEnrichment([]core.PassthroughSemanticEnricher{ + handler := PassthroughSemanticEnrichment(provider, []core.PassthroughSemanticEnricher{ passthroughSemanticEnricherStub{providerType: "openai"}, }, true)(ExecutionPlanning(provider)(func(c *echo.Context) error { capturedPlan = core.GetExecutionPlan(c.Request().Context()) diff --git a/internal/server/passthrough_service.go b/internal/server/passthrough_service.go index d151bb691..f49518295 100644 --- a/internal/server/passthrough_service.go +++ b/internal/server/passthrough_service.go @@ -10,6 +10,7 @@ import ( type passthroughService struct { provider core.RoutableProvider + modelAuthorizer RequestModelAuthorizer logger auditlog.LoggerInterface usageLogger usage.LoggerInterface pricingResolver usage.PricingResolver @@ -23,13 +24,20 @@ func (s *passthroughService) ProviderPassthrough(c *echo.Context) error { return handleError(c, core.NewInvalidRequestError("provider passthrough is not supported by the current provider router", nil)) } - providerType, endpoint, info, err := passthroughExecutionTarget(c, s.normalizePassthroughV1Prefix) + providerType, endpoint, info, err := passthroughExecutionTarget(c, s.provider, s.normalizePassthroughV1Prefix) if err != nil { return handleError(c, err) } if !isEnabledPassthroughProvider(providerType, s.enabledPassthroughProviders) { return handleError(c, s.unsupportedPassthroughProviderError(providerType)) } + if s.modelAuthorizer != nil { + if selector, ok := passthroughAccessSelector(s.provider, info); ok { + if err := s.modelAuthorizer.ValidateModelAccess(c.Request().Context(), selector); err != nil { + return handleError(c, err) + } + } + } ctx, _ := requestContextWithRequestID(c.Request()) c.SetRequest(c.Request().WithContext(ctx)) diff --git a/internal/server/request_model_resolution.go b/internal/server/request_model_resolution.go index 3248a6cb5..a86674a51 100644 --- a/internal/server/request_model_resolution.go +++ b/internal/server/request_model_resolution.go @@ -1,6 +1,7 @@ package server import ( + "context" "strings" "github.com/labstack/echo/v5" @@ -56,6 +57,16 @@ func workflowProviderNameForType(provider core.RoutableProvider, providerType st } func resolveRequestModel(provider core.RoutableProvider, resolver RequestModelResolver, requested core.RequestedModelSelector) (*core.RequestModelResolution, error) { + return resolveRequestModelWithAuthorizer(context.Background(), provider, resolver, nil, requested) +} + +func resolveRequestModelWithAuthorizer( + ctx context.Context, + provider core.RoutableProvider, + resolver RequestModelResolver, + authorizer RequestModelAuthorizer, + requested core.RequestedModelSelector, +) (*core.RequestModelResolution, error) { requested = core.NewRequestedModelSelector(requested.Model, requested.ProviderHint) resolvedSelector, aliasApplied, err := resolveExecutionSelector(provider, resolver, requested) @@ -76,6 +87,11 @@ func resolveRequestModel(provider core.RoutableProvider, resolver RequestModelRe if !provider.Supports(resolvedModel) { return nil, core.NewInvalidRequestError("unsupported model: "+resolvedModel, nil) } + if authorizer != nil { + if err := authorizer.ValidateModelAccess(ctx, resolvedSelector); err != nil { + return nil, err + } + } return &core.RequestModelResolution{ Requested: requested, @@ -156,7 +172,7 @@ func ensureRequestModelResolution(c *echo.Context, provider core.RoutableProvide if err != nil || !parsed { return nil, parsed, err } - resolution, err := resolveAndStoreRequestModelResolution(c, provider, resolver, model, providerHint) + resolution, err := resolveAndStoreRequestModelResolution(c, provider, resolver, nil, model, providerHint) return resolution, true, err } @@ -174,12 +190,13 @@ func resolveAndStoreRequestModelResolution( c *echo.Context, provider core.RoutableProvider, resolver RequestModelResolver, + authorizer RequestModelAuthorizer, model, providerHint string, ) (*core.RequestModelResolution, error) { requested := core.NewRequestedModelSelector(model, providerHint) enrichAuditEntryWithRequestedModel(c, requested) - resolution, err := resolveRequestModel(provider, resolver, requested) + resolution, err := resolveRequestModelWithAuthorizer(c.Request().Context(), provider, resolver, authorizer, requested) if err != nil { return nil, err } diff --git a/internal/server/request_model_resolution_test.go b/internal/server/request_model_resolution_test.go index 60c2fbfdf..0803a5484 100644 --- a/internal/server/request_model_resolution_test.go +++ b/internal/server/request_model_resolution_test.go @@ -3,6 +3,7 @@ package server import ( "context" "io" + "strings" "testing" "gomodel/internal/core" @@ -36,6 +37,20 @@ func (p *canonicalizingProvider) GetProviderName(model string) string { return p.names[model] } +func (p *canonicalizingProvider) GetProviderTypeForName(providerName string) string { + providerName = strings.TrimSpace(providerName) + if providerName == "" { + return "" + } + for qualifiedModel, candidate := range p.names { + if strings.TrimSpace(candidate) != providerName { + continue + } + return strings.TrimSpace(p.types[qualifiedModel]) + } + return "" +} + func (p *canonicalizingProvider) ChatCompletion(_ context.Context, _ *core.ChatRequest) (*core.ChatResponse, error) { return nil, nil } diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index d296a3479..6a8e76b9e 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -23,6 +23,7 @@ import ( type translatedInferenceService struct { provider core.RoutableProvider modelResolver RequestModelResolver + modelAuthorizer RequestModelAuthorizer executionPolicyResolver RequestExecutionPolicyResolver fallbackResolver RequestFallbackResolver translatedRequestPatcher TranslatedRequestPatcher @@ -134,7 +135,7 @@ func handleTranslatedInference[R any]( return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) } modelPtr, providerPtr := modelProvider(req) - plan, err := ensureTranslatedRequestPlan(c, s.provider, s.modelResolver, s.executionPolicyResolver, modelPtr, providerPtr) + plan, err := ensureTranslatedRequestPlanWithAuthorizer(c, s.provider, s.modelResolver, s.modelAuthorizer, s.executionPolicyResolver, modelPtr, providerPtr) if err != nil { return handleError(c, err) } @@ -329,7 +330,7 @@ func (s *translatedInferenceService) Embeddings(c *echo.Context) error { if err != nil { return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) } - plan, err := ensureTranslatedRequestPlan(c, s.provider, s.modelResolver, s.executionPolicyResolver, &req.Model, &req.Provider) + plan, err := ensureTranslatedRequestPlanWithAuthorizer(c, s.provider, s.modelResolver, s.modelAuthorizer, s.executionPolicyResolver, &req.Model, &req.Provider) if err != nil { return handleError(c, err) } @@ -769,6 +770,9 @@ func tryFallbackResponse[T any]( primaryModel := currentSelectorForPlan(plan, model, provider) lastErr := primaryErr for _, selector := range fallbacks { + if s.modelAuthorizer != nil && !s.modelAuthorizer.AllowsModel(ctx, selector) { + continue + } qualified := selector.QualifiedModel() providerType := s.providerTypeForSelector(selector, providerTypeFromPlan(plan)) providerName := resolvedProviderName(s.provider, selector, providerNameFromPlan(plan)) @@ -857,6 +861,9 @@ func tryFallbackStream( primaryModel := currentSelectorForPlan(plan, model, provider) lastErr := primaryErr for _, selector := range fallbacks { + if s.modelAuthorizer != nil && !s.modelAuthorizer.AllowsModel(ctx, selector) { + continue + } qualified := selector.QualifiedModel() providerType := s.providerTypeForSelector(selector, providerTypeFromPlan(plan)) providerName := resolvedProviderName(s.provider, selector, providerNameFromPlan(plan)) diff --git a/tests/integration/setup_test.go b/tests/integration/setup_test.go index 8ad9d1b88..890f95820 100644 --- a/tests/integration/setup_test.go +++ b/tests/integration/setup_test.go @@ -249,6 +249,9 @@ func buildAppConfig(t *testing.T, cfg TestServerConfig, mockLLMURL string, port Port: fmt.Sprintf("%d", port), MasterKey: cfg.MasterKey, }, + Models: config.ModelsConfig{ + EnabledByDefault: true, + }, Admin: config.AdminConfig{ EndpointsEnabled: cfg.AdminEndpointsEnabled, UIEnabled: cfg.AdminUIEnabled,