diff --git a/.env.example b/.env.example index c71e65dcb..4fc95de12 100644 --- a/.env.example +++ b/.env.example @@ -44,6 +44,14 @@ POSTGRES_SSLMODE=require # set to `disable` for local Docker # GOOGLE_API_KEY=... # GOOGLE_BASE_URL=https://generativelanguage.googleapis.com/v1beta/openai +# Remap outbound model IDs per provider, for an OpenAI-compatible endpoint +# that publishes the catalog's models under its own names. JSON map of +# catalog model ID -> the name that endpoint expects. is one of +# OPENROUTER, FIREWORKS, MAKORA, TOGETHER, XAI, BEDROCK. Routing, pricing, +# and analytics stay keyed on the catalog ID. +# See docs/CONFIGURATION.md -> Deployment-level model aliases. +# ROUTER_OPENROUTER_MODEL_ALIASES={"deepseek/deepseek-v4-flash":"deepseek-v4-flash"} + # -------------------------------------------------------------------- # Server. # -------------------------------------------------------------------- diff --git a/cmd/AGENTS.md b/cmd/AGENTS.md index e19d9a8ec..41bd475b9 100644 --- a/cmd/AGENTS.md +++ b/cmd/AGENTS.md @@ -15,6 +15,7 @@ Composition root. Only place that constructs concrete adapters + wires them toge - `runSessionPinSweep` — TTL sweep loop - `resolveHardPinModel` / `resolveDefaultBaselineModel` / `resolveAvailableModels` — boot-time model resolution - `registerDeploymentKeyedProvider` — shared "resolve key → build client → log" registration for the providers whose gating collapses to that shape (Fireworks, Makora, Together, Bedrock, Google); OpenRouter and Anthropic/OpenAI stay bespoke + - `resolveModelAliases` ([`model_aliases.go`](router/model_aliases.go)) — each OpenAI-compatible provider's outbound model ID map: catalog `UpstreamID` bindings with `ROUTER__MODEL_ALIASES` overrides layered on top - small env parsers (e.g. `envVarHint`, `parseEnvInt`, `parseEnvFloat`, `parseEnvDurationMs`) - **No more heuristic-fallback router.** If cluster routing fails to boot, `main.go` panics. Misconfiguration must abort the process rather than silently degrade. - **Never introduce DI container, reflection-based wiring, or service locator.** Composition = plain Go function calls. diff --git a/cmd/CLAUDE.md b/cmd/CLAUDE.md index 6f7a86b3f..f6bd22dd2 100644 --- a/cmd/CLAUDE.md +++ b/cmd/CLAUDE.md @@ -15,6 +15,7 @@ Composition root. Only place that constructs concrete adapters + wires them toge - `runSessionPinSweep` — TTL sweep loop - `resolveHardPinModel` / `resolveDefaultBaselineModel` / `resolveAvailableModels` — boot-time model resolution - `registerDeploymentKeyedProvider` — shared "resolve key → build client → log" registration for the providers whose gating collapses to that shape (Fireworks, Makora, Together, Bedrock, Google); OpenRouter and Anthropic/OpenAI stay bespoke + - `resolveModelAliases` ([`model_aliases.go`](router/model_aliases.go)) — each OpenAI-compatible provider's outbound model ID map: catalog `UpstreamID` bindings with `ROUTER__MODEL_ALIASES` overrides layered on top - small env parsers (e.g. `envVarHint`, `parseEnvInt`, `parseEnvFloat`, `parseEnvDurationMs`) - **No more heuristic-fallback router.** If cluster routing fails to boot, `main.go` panics. Misconfiguration must abort the process rather than silently degrade. - **Never introduce DI container, reflection-based wiring, or service locator.** Composition = plain Go function calls. diff --git a/cmd/router/main.go b/cmd/router/main.go index 58f483c1e..429830340 100644 --- a/cmd/router/main.go +++ b/cmd/router/main.go @@ -178,6 +178,11 @@ func main() { // platform-key mode gated by balance checks. Self-hosted is never BYOK-only. byokOnly := deploymentMode == server.DeploymentModeManaged && billingSvc == nil + modelAliases, err := resolveModelAliases(logger) + if err != nil { + panic(err) + } + // Always registered. With ANTHROPIC_API_KEY (selfhosted only) the router // uses its own key; otherwise client auth headers pass through directly. anthropicKey := "" @@ -231,7 +236,7 @@ func main() { if !byokOnly && openRouterPlatformEnabled { openRouterKey = config.GetOr("OPENROUTER_API_KEY", "") } - providerMap[providers.ProviderOpenRouter] = openaiCompatProvider.NewClient(openRouterKey, openRouterBaseURL) + providerMap[providers.ProviderOpenRouter] = openaiCompatProvider.NewClientWithModelIDMap(openRouterKey, openRouterBaseURL, modelAliases[providers.ProviderOpenRouter]) switch { case byokOnly: logger.Info("OpenRouter provider enabled (BYOK only)", "base_url", openRouterBaseURL) @@ -250,7 +255,7 @@ func main() { registerDeploymentKeyedProvider(providerMap, envKeyedProviders, logger, providers.ProviderFireworks, "Fireworks", "FIREWORKS_API_KEY", fireworksBaseURL, byokOnly, func(key, baseURL string) providers.Client { - return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, upstreamIDsForProvider(providers.ProviderFireworks)) + return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, modelAliases[providers.ProviderFireworks]) }) } @@ -261,7 +266,7 @@ func main() { registerDeploymentKeyedProvider(providerMap, envKeyedProviders, logger, providers.ProviderMakora, "Makora", "MAKORA_API_KEY", makoraBaseURL, byokOnly, func(key, baseURL string) providers.Client { - return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, upstreamIDsForProvider(providers.ProviderMakora)) + return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, modelAliases[providers.ProviderMakora]) }) } @@ -274,7 +279,7 @@ func main() { registerDeploymentKeyedProvider(providerMap, envKeyedProviders, logger, providers.ProviderTogether, "Together", "TOGETHER_API_KEY", togetherBaseURL, byokOnly, func(key, baseURL string) providers.Client { - return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, upstreamIDsForProvider(providers.ProviderTogether)) + return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, modelAliases[providers.ProviderTogether]) }) } @@ -283,7 +288,7 @@ func main() { registerDeploymentKeyedProvider(providerMap, envKeyedProviders, logger, providers.ProviderXAI, "XAI", "XAI_API_KEY", xaiBaseURL, byokOnly, func(key, baseURL string) providers.Client { - return openaiCompatProvider.NewClient(key, baseURL) + return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, modelAliases[providers.ProviderXAI]) }) } @@ -297,7 +302,7 @@ func main() { registerDeploymentKeyedProvider(providerMap, envKeyedProviders, logger, providers.ProviderBedrock, "Bedrock", "AWS_BEARER_TOKEN_BEDROCK", bedrockBaseURL, byokOnly, func(key, baseURL string) providers.Client { - return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, upstreamIDsForProvider(providers.ProviderBedrock)) + return openaiCompatProvider.NewClientWithModelIDMap(key, baseURL, modelAliases[providers.ProviderBedrock]) }, "region", bedrockRegion) } @@ -1474,9 +1479,6 @@ func envVarHint(provider string) string { return "" } -// upstreamIDsForProvider maps public model ID -> upstream model ID for a -// provider's bindings with a non-empty UpstreamID; nil if no rewriting is -// needed (e.g. OpenRouter, where the slug IS the upstream ID). // registerDeploymentKeyedProvider resolves a provider's deployment-level API // key (respecting byokOnly), constructs its client via newClient, registers // it in providerMap, and logs its BYOK/keyed/passthrough state. Shared by the @@ -1509,18 +1511,3 @@ func registerDeploymentKeyedProvider( logger.Info(displayName+" provider registered (BYOK only — set "+keyEnvVar+" for deployment-level use)", "base_url", baseURL) } } - -func upstreamIDsForProvider(provider string) map[string]string { - out := make(map[string]string) - for _, m := range catalog.Models { - for _, b := range m.Providers { - if b.Provider == provider && b.UpstreamID != "" { - out[m.ID] = b.UpstreamID - } - } - } - if len(out) == 0 { - return nil - } - return out -} diff --git a/cmd/router/model_aliases.go b/cmd/router/model_aliases.go new file mode 100644 index 000000000..4aae55683 --- /dev/null +++ b/cmd/router/model_aliases.go @@ -0,0 +1,99 @@ +package main + +import ( + "encoding/json" + "errors" + "fmt" + "log/slog" + "strings" + + "workweave/router/internal/config" + "workweave/router/internal/providers" + "workweave/router/internal/router/catalog" +) + +// aliasableProviders are the OpenAI-compatible upstreams whose outbound model +// ID can be rewritten before dispatch. +var aliasableProviders = []string{ + providers.ProviderOpenRouter, + providers.ProviderFireworks, + providers.ProviderMakora, + providers.ProviderTogether, + providers.ProviderXAI, + providers.ProviderBedrock, +} + +// modelAliasEnvVar names the operator override for a provider's outbound model IDs. +func modelAliasEnvVar(provider string) string { + return "ROUTER_" + strings.ToUpper(provider) + "_MODEL_ALIASES" +} + +// resolveModelAliases returns each aliasable provider's public model ID -> +// upstream model ID map: the catalog's UpstreamID bindings with that provider's +// ROUTER__MODEL_ALIASES entries layered on top. +func resolveModelAliases(logger *slog.Logger) (map[string]map[string]string, error) { + out := make(map[string]map[string]string, len(aliasableProviders)) + for _, provider := range aliasableProviders { + envVar := modelAliasEnvVar(provider) + overrides, err := parseModelAliases(config.GetOr(envVar, "")) + if err != nil { + return nil, fmt.Errorf("%s: %w", envVar, err) + } + merged := upstreamIDsForProvider(provider) + for id, upstreamID := range overrides { + if _, known := catalog.ByID(id); !known { + // Retiring a catalog model must not turn a stale alias into a boot failure. + logger.Warn("Ignoring model alias for unknown catalog model", "env_var", envVar, "model", id) + continue + } + if merged == nil { + merged = make(map[string]string, len(overrides)) + } + merged[id] = upstreamID + } + if len(merged) > 0 { + out[provider] = merged + } + } + return out, nil +} + +// parseModelAliases decodes a JSON object of public model ID -> upstream model +// ID. An empty value yields no entries. +func parseModelAliases(raw string) (map[string]string, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil, nil + } + var aliases map[string]string + if err := json.Unmarshal([]byte(raw), &aliases); err != nil { + return nil, fmt.Errorf("expected a JSON object mapping model ID to upstream model ID: %w", err) + } + for id, upstreamID := range aliases { + if strings.TrimSpace(id) == "" { + return nil, errors.New("empty model ID") + } + if strings.TrimSpace(upstreamID) == "" { + return nil, fmt.Errorf("model %q maps to an empty upstream model ID", id) + } + } + return aliases, nil +} + +// upstreamIDsForProvider maps public model ID -> upstream model ID for a +// provider's bindings with a non-empty UpstreamID; nil if no rewriting is +// needed (e.g. OpenRouter, where the slug IS the upstream ID). +func upstreamIDsForProvider(provider string) map[string]string { + out := make(map[string]string) + for _, m := range catalog.Models { + for _, b := range m.Providers { + if b.Provider == provider && b.UpstreamID != "" { + out[m.ID] = b.UpstreamID + } + } + } + if len(out) == 0 { + return nil + } + return out +} diff --git a/cmd/router/model_aliases_test.go b/cmd/router/model_aliases_test.go new file mode 100644 index 000000000..fd2e26feb --- /dev/null +++ b/cmd/router/model_aliases_test.go @@ -0,0 +1,122 @@ +package main + +import ( + "io" + "log/slog" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "workweave/router/internal/providers" + "workweave/router/internal/router/catalog" +) + +func discardLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +// bindingWithUpstreamID returns the first catalog model bound to provider with a +// non-empty UpstreamID, so tests survive catalog churn instead of pinning an ID. +func bindingWithUpstreamID(t *testing.T, provider string) (modelID, upstreamID string) { + t.Helper() + for _, m := range catalog.Models { + for _, b := range m.Providers { + if b.Provider == provider && b.UpstreamID != "" { + return m.ID, b.UpstreamID + } + } + } + t.Skipf("catalog has no %s binding with an UpstreamID", provider) + return "", "" +} + +func TestModelAliasEnvVar(t *testing.T) { + assert.Equal(t, "ROUTER_OPENROUTER_MODEL_ALIASES", modelAliasEnvVar(providers.ProviderOpenRouter)) + assert.Equal(t, "ROUTER_BEDROCK_MODEL_ALIASES", modelAliasEnvVar(providers.ProviderBedrock)) +} + +func TestParseModelAliases(t *testing.T) { + t.Run("empty value yields no entries", func(t *testing.T) { + aliases, err := parseModelAliases(" ") + require.NoError(t, err) + assert.Empty(t, aliases) + }) + + t.Run("decodes a JSON object", func(t *testing.T) { + aliases, err := parseModelAliases(`{"deepseek/deepseek-v4-flash":"deepseek-v4-flash","z-ai/glm-5.2":"glm-5.2"}`) + require.NoError(t, err) + assert.Equal(t, map[string]string{ + "deepseek/deepseek-v4-flash": "deepseek-v4-flash", + "z-ai/glm-5.2": "glm-5.2", + }, aliases) + }) + + t.Run("rejects malformed JSON", func(t *testing.T) { + _, err := parseModelAliases(`{"deepseek/deepseek-v4-flash":`) + require.Error(t, err) + assert.Contains(t, err.Error(), "JSON object") + }) + + t.Run("rejects a JSON array", func(t *testing.T) { + _, err := parseModelAliases(`["deepseek-v4-flash"]`) + require.Error(t, err) + }) + + t.Run("rejects an empty upstream model ID", func(t *testing.T) { + _, err := parseModelAliases(`{"deepseek/deepseek-v4-flash":" "}`) + require.Error(t, err) + assert.Contains(t, err.Error(), "deepseek/deepseek-v4-flash") + }) + + t.Run("rejects an empty model ID", func(t *testing.T) { + _, err := parseModelAliases(`{"":"deepseek-v4-flash"}`) + require.Error(t, err) + }) +} + +func TestResolveModelAliasesAppliesEnvOverride(t *testing.T) { + require.NotEmpty(t, catalog.Models) + modelID := catalog.Models[0].ID + t.Setenv(modelAliasEnvVar(providers.ProviderOpenRouter), `{"`+modelID+`":"gateway-name"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + assert.Equal(t, "gateway-name", aliases[providers.ProviderOpenRouter][modelID]) +} + +func TestResolveModelAliasesOverridesCatalogUpstreamID(t *testing.T) { + modelID, catalogUpstreamID := bindingWithUpstreamID(t, providers.ProviderTogether) + t.Setenv(modelAliasEnvVar(providers.ProviderTogether), `{"`+modelID+`":"operator-name"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + assert.Equal(t, "operator-name", aliases[providers.ProviderTogether][modelID]) + assert.NotEqual(t, catalogUpstreamID, aliases[providers.ProviderTogether][modelID]) +} + +func TestResolveModelAliasesKeepsUnaliasedCatalogBindings(t *testing.T) { + modelID, catalogUpstreamID := bindingWithUpstreamID(t, providers.ProviderTogether) + require.NotEmpty(t, catalog.Models) + t.Setenv(modelAliasEnvVar(providers.ProviderTogether), `{"`+catalog.Models[0].ID+`":"operator-name"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + assert.Equal(t, catalogUpstreamID, aliases[providers.ProviderTogether][modelID]) +} + +func TestResolveModelAliasesSkipsUnknownModel(t *testing.T) { + t.Setenv(modelAliasEnvVar(providers.ProviderOpenRouter), `{"vendor/not-in-catalog":"whatever"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + assert.NotContains(t, aliases[providers.ProviderOpenRouter], "vendor/not-in-catalog") +} + +func TestResolveModelAliasesRejectsMalformedValue(t *testing.T) { + t.Setenv(modelAliasEnvVar(providers.ProviderFireworks), "not-json") + + _, err := resolveModelAliases(discardLogger()) + require.Error(t, err) + assert.Contains(t, err.Error(), "ROUTER_FIREWORKS_MODEL_ALIASES") +} diff --git a/cmd/router/model_aliases_wire_test.go b/cmd/router/model_aliases_wire_test.go new file mode 100644 index 000000000..1dd515763 --- /dev/null +++ b/cmd/router/model_aliases_wire_test.go @@ -0,0 +1,95 @@ +package main + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "workweave/router/internal/providers" + openaiCompatProvider "workweave/router/internal/providers/openaicompat" + "workweave/router/internal/router" +) + +// stubUpstream stands in for an OpenAI-compatible endpoint that publishes the +// catalog's models under its own names, and records the model it was sent. +func stubUpstream(t *testing.T, gotModel *string) *httptest.Server { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.NoError(t, err) + var payload struct { + Model string `json:"model"` + } + require.NoError(t, json.Unmarshal(body, &payload)) + *gotModel = payload.Model + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"1","object":"chat.completion","choices":[]}`)) + })) + t.Cleanup(srv.Close) + return srv +} + +func dispatch(t *testing.T, client providers.Client, modelID string) { + t.Helper() + prep := providers.PreparedRequest{ + Body: []byte(`{"model":"` + modelID + `","messages":[{"role":"user","content":"hi"}]}`), + Headers: http.Header{}, + } + decision := router.Decision{Provider: providers.ProviderOpenRouter, Model: modelID} + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", http.NoBody) + require.NoError(t, client.Proxy(context.Background(), decision, prep, httptest.NewRecorder(), req)) +} + +// Reproduces #541: the router put its catalog slash-form ID on the wire and a +// gateway publishing bare names rejected it. +func TestModelAliasRewritesModelOnTheWire(t *testing.T) { + const ( + catalogID = "deepseek/deepseek-v4-flash" + gatewayID = "deepseek-v4-flash" + otherModel = "xiaomi/mimo-v2.5-pro" + ) + + t.Run("without an alias the catalog ID reaches the endpoint", func(t *testing.T) { + var got string + srv := stubUpstream(t, &got) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + client := openaiCompatProvider.NewClientWithModelIDMap("k", srv.URL, aliases[providers.ProviderOpenRouter]) + + dispatch(t, client, catalogID) + assert.Equal(t, catalogID, got) + }) + + t.Run("the alias puts the endpoint's own name on the wire", func(t *testing.T) { + var got string + srv := stubUpstream(t, &got) + t.Setenv(modelAliasEnvVar(providers.ProviderOpenRouter), `{"`+catalogID+`":"`+gatewayID+`"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + client := openaiCompatProvider.NewClientWithModelIDMap("k", srv.URL, aliases[providers.ProviderOpenRouter]) + + dispatch(t, client, catalogID) + assert.Equal(t, gatewayID, got) + }) + + t.Run("an unaliased model is still sent unchanged", func(t *testing.T) { + var got string + srv := stubUpstream(t, &got) + t.Setenv(modelAliasEnvVar(providers.ProviderOpenRouter), `{"`+catalogID+`":"`+gatewayID+`"}`) + + aliases, err := resolveModelAliases(discardLogger()) + require.NoError(t, err) + client := openaiCompatProvider.NewClientWithModelIDMap("k", srv.URL, aliases[providers.ProviderOpenRouter]) + + dispatch(t, client, otherModel) + assert.Equal(t, otherModel, got) + }) +} diff --git a/docs/CONFIGURATION.md b/docs/CONFIGURATION.md index 67df0b07b..4b81fae4a 100644 --- a/docs/CONFIGURATION.md +++ b/docs/CONFIGURATION.md @@ -35,6 +35,7 @@ Claude Code keep using the user's logged-in plan. | `GOOGLE_BASE_URL` | `https://generativelanguage.googleapis.com/v1beta/openai` | Override for Gemini. | | `ANTHROPIC_GATEWAY_BASE_URL` | *(none)* | Base URL of an Anthropic-compatible gateway; `/v1/messages` is appended to it. | | `ANTHROPIC_GATEWAY_TOKEN` | *(none)* | Token for that gateway, sent as `Authorization: Bearer`. Only used when `ANTHROPIC_GATEWAY_BASE_URL` is also set. | +| `ROUTER__MODEL_ALIASES` | *(none)* | JSON map of catalog model ID to the name that provider's endpoint publishes. See [Deployment-level model aliases](#deployment-level-model-aliases). | **Anthropic-compatible gateway.** Some enterprises front Claude with their own gateway that speaks the Anthropic Messages spec but authenticates with a bearer @@ -45,6 +46,30 @@ is always registered so BYOK installations can point at their own gateway without deployment-level credentials; the env vars above are only for a deployment that has a gateway of its own. +### Deployment-level model aliases + +An OpenAI-compatible endpoint may publish the catalog's models under its own +names: a gateway that serves `deepseek-v4-flash` where the catalog says +`deepseek/deepseek-v4-flash`, or a Bedrock-backed proxy that expects AWS +dot-form IDs. Point a provider at the endpoint with its `*_BASE_URL` and remap +the outbound names with `ROUTER__MODEL_ALIASES`: + +```bash +OPENROUTER_BASE_URL=https://gateway.example.com/v1 +ROUTER_OPENROUTER_MODEL_ALIASES='{"deepseek/deepseek-v4-flash":"deepseek-v4-flash","xiaomi/mimo-v2.5-pro":"mimo-v2.5-pro"}' +``` + +`` is the upper-cased provider name — `OPENROUTER`, `FIREWORKS`, +`MAKORA`, `TOGETHER`, `XAI`, or `BEDROCK`. Keys are catalog model IDs and values +are what goes on the wire. Only the outbound model name changes: routing, +pricing, and analytics stay keyed on the catalog ID. + +Entries layer on top of the catalog's own per-binding upstream IDs, so aliasing +one model leaves that provider's other bindings alone, and a BYOK key's +`model_aliases` still wins over both. A malformed value aborts boot. An alias +naming a model outside the deployed catalog is logged and skipped, so retiring a +model can't turn a stale alias into a failed start. + **BYOK (per-installation keys).** Instead of (or in addition to) the env vars above, each installation can supply its own provider keys via the dashboard. Those are stored in Postgres and used only for that installation's traffic.