Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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. <PROVIDER> 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.
# --------------------------------------------------------------------
Expand Down
1 change: 1 addition & 0 deletions cmd/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -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_<PROVIDER>_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.
Expand Down
1 change: 1 addition & 0 deletions cmd/CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -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_<PROVIDER>_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.
Expand Down
35 changes: 11 additions & 24 deletions cmd/router/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 := ""
Expand Down Expand Up @@ -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)
Expand All @@ -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])
})
}

Expand All @@ -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])
})
}

Expand All @@ -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])
})
}

Expand All @@ -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])
})
}

Expand All @@ -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)
}
Expand Down Expand Up @@ -1474,9 +1479,6 @@ func envVarHint(provider string) string {
return "<unknown provider " + provider + ">"
}

// 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
Expand Down Expand Up @@ -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
}
99 changes: 99 additions & 0 deletions cmd/router/model_aliases.go
Original file line number Diff line number Diff line change
@@ -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_<PROVIDER>_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
}
122 changes: 122 additions & 0 deletions cmd/router/model_aliases_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
Loading