fix(coderd): preserve gateway model names (#26039)

OpenAI-compatible gateway providers such as OpenRouter require
slash-namespaced model IDs to reach the intended upstream model, but
native OpenAI routing strips those prefixes.

Preserve full model IDs for gateway provider types, reject
OpenRouter-like providers configured as native `openai` when a slash
model would be stripped, and validate chat model config changes under
the provider reference lock while still allowing unrelated edits to
existing configs.

Split from #26005.

> Mux created this PR on behalf of Mike.
This commit is contained in:
Michael Suchacz
2026-06-04 15:33:00 +02:00
committed by GitHub
parent 2cbce86eee
commit 502c5acca8
8 changed files with 529 additions and 4 deletions
@@ -3,6 +3,7 @@ package chatprovider
import (
"context"
"net/http"
neturl "net/url"
"sort"
"strings"
@@ -186,6 +187,30 @@ func (k ProviderAPIKeys) BaseURL(provider string) string {
return strings.TrimSpace(k.BaseURLByProvider[normalized])
}
// ProviderBaseURLHostname returns the normalized hostname from a provider base URL.
func ProviderBaseURLHostname(baseURL string) string {
parsed, ok := parseProviderBaseURL(baseURL)
if !ok {
return ""
}
return strings.ToLower(parsed.Hostname())
}
func parseProviderBaseURL(baseURL string) (*neturl.URL, bool) {
baseURL = strings.TrimSpace(baseURL)
if baseURL == "" {
return nil, false
}
parsed, err := neturl.Parse(baseURL)
if err == nil && parsed.Hostname() == "" && !strings.Contains(baseURL, "://") {
parsed, err = neturl.Parse("https://" + baseURL)
}
if err != nil {
return nil, false
}
return parsed, true
}
// MergeProviderAPIKeys overlays configured provider keys over fallback keys.
func MergeProviderAPIKeys(fallback ProviderAPIKeys, providers []ConfiguredProvider) ProviderAPIKeys {
merged := ProviderAPIKeys{
@@ -29,6 +29,28 @@ import (
"github.com/coder/coder/v2/testutil"
)
func TestProviderBaseURLHostname(t *testing.T) {
t.Parallel()
tests := []struct {
name string
baseURL string
want string
}{
{name: "URL", baseURL: "https://openrouter.ai/api/v1", want: "openrouter.ai"},
{name: "BareHost", baseURL: "openrouter.ai", want: "openrouter.ai"},
{name: "HostWithPort", baseURL: "https://openrouter.ai:443/api/v1", want: "openrouter.ai"},
{name: "Empty", baseURL: "", want: ""},
{name: "Invalid", baseURL: "://", want: ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
require.Equal(t, tt.want, chatprovider.ProviderBaseURLHostname(tt.baseURL))
})
}
}
func TestResolveUserProviderKeys(t *testing.T) {
t.Parallel()
@@ -1547,6 +1569,34 @@ func TestResolveModelWithProviderHint(t *testing.T) {
wantProvider: fantasyopenaicompat.Name,
wantModel: "anthropic/claude-4-5-sonnet",
},
{
name: "OpenRouterHintPreservesOpenRouterModelID",
modelName: "anthropic/claude-opus-4.6",
providerHint: fantasyopenrouter.Name,
wantProvider: fantasyopenrouter.Name,
wantModel: "anthropic/claude-opus-4.6",
},
{
name: "OpenAICompatHintPreservesOpenRouterModelID",
modelName: "anthropic/claude-opus-4.6",
providerHint: fantasyopenaicompat.Name,
wantProvider: fantasyopenaicompat.Name,
wantModel: "anthropic/claude-opus-4.6",
},
{
name: "OpenAIHintStripsCanonicalPrefix",
modelName: "anthropic/claude-opus-4.6",
providerHint: fantasyopenai.Name,
wantProvider: fantasyanthropic.Name,
wantModel: "claude-opus-4.6",
},
{
name: "OpenAIHintPreservesUnknownSlashNamespace",
modelName: "meta-llama/llama-3-70b",
providerHint: fantasyopenai.Name,
wantProvider: fantasyopenai.Name,
wantModel: "meta-llama/llama-3-70b",
},
{
name: "AnthropicHintStripsCanonicalPrefix",
modelName: "anthropic/claude-4-5-sonnet",
@@ -5,7 +5,6 @@ import (
"encoding/json"
"io"
"net/http"
"net/url"
"strings"
)
@@ -150,8 +149,8 @@ func rewriteOpenAICompatSingleToolChoice(payload map[string]any) bool {
// endpoints and Coder AI Bridge Gemini routes. Other gateways, such as Vercel,
// keep their own provider-specific compatibility behavior.
func shouldAddGoogleOpenAICompatThoughtSignatures(baseURL string, modelID string) bool {
parsed, err := url.Parse(baseURL)
if err != nil {
parsed, ok := parseProviderBaseURL(baseURL)
if !ok {
return false
}
host := strings.ToLower(parsed.Hostname())
+37
View File
@@ -87,6 +87,32 @@ func (t *aiGatewayRoundTripper) RoundTrip(req *http.Request) (*http.Response, er
return t.base.RoundTrip(cloned)
}
// ValidateAIGatewayProviderModel rejects slash-namespaced models on
// OpenRouter-like providers typed as openai, where the provider type
// strips the vendor prefix.
func ValidateAIGatewayProviderModel(provider database.AIProvider, model string) error {
if provider.Type != database.AiProviderTypeOpenai {
return nil
}
if !isSlashNamespacedAIGatewayModel(model) || !isOpenRouterLikeAIGatewayProvider(provider) {
return nil
}
return xerrors.New("OpenRouter-like provider configured as type openai does not support slash-namespaced models")
}
func isSlashNamespacedAIGatewayModel(model string) bool {
prefix, suffix, ok := strings.Cut(strings.TrimSpace(model), "/")
return ok && strings.TrimSpace(prefix) != "" && strings.TrimSpace(suffix) != ""
}
func isOpenRouterLikeAIGatewayProvider(provider database.AIProvider) bool {
if strings.EqualFold(strings.TrimSpace(provider.Name), "openrouter") {
return true
}
host := chatprovider.ProviderBaseURLHostname(provider.BaseUrl)
return host == "openrouter.ai" || strings.HasSuffix(host, ".openrouter.ai")
}
func (p *Server) newAIGatewayModel(
_ context.Context,
req modelClientRequest,
@@ -110,6 +136,17 @@ func (p *Server) newAIGatewayModel(
)
}
if err := ValidateAIGatewayProviderModel(route.Provider, req.ModelName); err != nil {
return nil, chaterror.WithClassification(
err,
chaterror.ClassifiedError{
Kind: codersdk.ChatErrorKindConfig,
Retryable: false,
Detail: "Ask an administrator to change the AI provider type to openrouter or openai-compat.",
},
)
}
factoryPtr := p.aibridgeTransportFactory
if factoryPtr == nil {
return nil, xerrors.New("AI Gateway transport factory is not configured")
@@ -2,6 +2,7 @@ package chatd
import (
"database/sql"
"encoding/json"
"fmt"
"io"
"net/http"
@@ -641,6 +642,29 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
require.False(t, classified.Retryable)
})
t.Run("OpenRouterMisconfiguredAsOpenAI", func(t *testing.T) {
t.Parallel()
factory := &aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) {
t.Fatal("transport must not be used for invalid provider config")
return nil, xerrors.New("unreachable")
})}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
provider := aibridgeTestAIProvider(providerID, "openrouter", database.AiProviderTypeOpenai)
_, err := server.newModel(
t.Context(),
aibridgeTestRequest(chat, "anthropic/claude-opus-4.6"),
aibridgeTestRoute(provider),
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.ErrorContains(t, err, "does not support slash-namespaced models")
classified := chaterror.Classify(err)
require.Equal(t, codersdk.ChatErrorKindConfig, classified.Kind)
require.False(t, classified.Retryable)
})
t.Run("StaticModel", func(t *testing.T) {
t.Parallel()
server := &Server{aiGatewayRoutingEnabled: true}
@@ -649,6 +673,112 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
})
}
func TestAIBridgeGatewayProviderTypesPreserveSlashModelID(t *testing.T) {
t.Parallel()
const modelName = "anthropic/claude-opus-4.6"
tests := []struct {
name string
providerName string
providerType database.AIProviderType
}{
{
name: "OpenRouter",
providerName: "openrouter",
providerType: database.AiProviderTypeOpenrouter,
},
{
name: "OpenAICompat",
providerName: "openai-compatible-relay",
providerType: database.AiProviderTypeOpenaiCompat,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
type seenRequest struct {
model string
path string
}
seen := make(chan seenRequest, 1)
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
body, err := io.ReadAll(req.Body)
require.NoError(t, err)
var payload struct {
Model string `json:"model"`
}
require.NoError(t, json.Unmarshal(body, &payload))
seen <- seenRequest{model: payload.Model, path: req.URL.Path}
var responsePayload map[string]any
if strings.Contains(req.URL.Path, "/responses") {
responsePayload = map[string]any{
"id": "resp_test",
"object": "response",
"created_at": 0,
"status": "completed",
"model": modelName,
"output": []map[string]any{{
"id": "msg_test",
"type": "message",
"role": "assistant",
"content": []map[string]any{{"type": "output_text", "text": "hello"}},
}},
"usage": map[string]any{"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
}
} else {
responsePayload = map[string]any{
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 0,
"model": modelName,
"choices": []map[string]any{{
"index": 0,
"message": map[string]any{"role": "assistant", "content": "hello"},
"finish_reason": "stop",
}},
"usage": map[string]any{"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
}
responseBody, err := json.Marshal(responsePayload)
require.NoError(t, err)
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(string(responseBody))),
Request: req,
}, nil
})}
chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
model, err := server.newModel(
t.Context(),
aibridgeTestRequest(chat, modelName),
aibridgeTestRoute(aibridgeTestAIProvider(uuid.New(), tt.providerName, tt.providerType)),
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.NoError(t, err)
_, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}},
}}})
require.NoError(t, err)
got := <-seen
require.NotEmpty(t, got.path)
require.Equal(t, modelName, got.model)
require.Equal(t, tt.providerName, factory.providerName)
require.Equal(t, aibridge.SourceAgents, factory.source)
})
}
}
func TestDirectModelBuildDoesNotRequireActiveAPIKeyID(t *testing.T) {
t.Parallel()