-
Notifications
You must be signed in to change notification settings - Fork 116
fix(proxy): resolve force-model provider against live availability (#874) #876
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 3 commits
cc0e684
e417122
a17ce6c
219ed65
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -110,17 +110,17 @@ var forceModelAliases = map[string]string{ | |||||||||||
| "mistral": "mistralai/mistral-small-2603", | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // resolveForceModel is the legacy two-return surface. New pin-and-effort | ||||||||||||
| // callers use resolveForceModelWithEffort. | ||||||||||||
| func resolveForceModel(model string) (canonicalID, provider string, known bool) { | ||||||||||||
| canon, prov, kn, _ := resolveForceModelWithEffort(model) | ||||||||||||
| // resolveForceModel is the legacy two-return surface; nil available means | ||||||||||||
| // unrestricted. New pin-and-effort callers use resolveForceModelWithEffort. | ||||||||||||
| func resolveForceModel(model string, available map[string]struct{}) (canonicalID, provider string, known bool) { | ||||||||||||
| canon, prov, kn, _ := resolveForceModelWithEffort(model, available) | ||||||||||||
| return canon, prov, kn | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // resolveForceModelWithEffort is like resolveForceModel but also strips a | ||||||||||||
| // `:level` suffix. `known` is true only for catalog matches; known=false + | ||||||||||||
| // effort!="" lets callers surface "model not found" without losing the effort. | ||||||||||||
| func resolveForceModelWithEffort(model string) (canonicalID, provider string, known bool, effort string) { | ||||||||||||
| // resolveForceModelWithEffort strips a `:level` suffix and resolves catalog | ||||||||||||
| // bindings in order against available. known remains true when no binding is | ||||||||||||
| // available; nil available means unrestricted and selects the primary binding. | ||||||||||||
| func resolveForceModelWithEffort(model string, available map[string]struct{}) (canonicalID, provider string, known bool, effort string) { | ||||||||||||
| effortLevel, stripped := stripEffortSuffix(model) | ||||||||||||
| model = stripped | ||||||||||||
| model = strings.ToLower(strings.TrimSpace(model)) | ||||||||||||
|
|
@@ -135,21 +135,25 @@ func resolveForceModelWithEffort(model string) (canonicalID, provider string, kn | |||||||||||
| if alias, ok := forceModelAliases[model]; ok { | ||||||||||||
| model = alias | ||||||||||||
| } | ||||||||||||
| if m, ok := catalog.ByID(model); ok && len(m.Providers) > 0 && (requiredProvider == "" || m.Providers[0].Provider == requiredProvider) { | ||||||||||||
| return m.ID, m.Providers[0].Provider, true, effort | ||||||||||||
| if m, ok := catalog.ByID(model); ok && len(m.Providers) > 0 { | ||||||||||||
| if bound, nativeMismatch := resolveForcedBinding(m, requiredProvider, available); !nativeMismatch { | ||||||||||||
| return m.ID, bound, true, effort | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| if !strings.Contains(model, "/") { | ||||||||||||
| suffix := "/" + model | ||||||||||||
| var matched catalog.Model | ||||||||||||
| var matches int | ||||||||||||
| for _, m := range catalog.Models { | ||||||||||||
| if strings.HasSuffix(m.ID, suffix) && len(m.Providers) > 0 && (requiredProvider == "" || m.Providers[0].Provider == requiredProvider) { | ||||||||||||
| if strings.HasSuffix(m.ID, suffix) && len(m.Providers) > 0 { | ||||||||||||
| matched = m | ||||||||||||
| matches++ | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| if matches == 1 && len(matched.Providers) > 0 { | ||||||||||||
| return matched.ID, matched.Providers[0].Provider, true, effort | ||||||||||||
| if matches == 1 { | ||||||||||||
| if bound, nativeMismatch := resolveForcedBinding(matched, requiredProvider, available); !nativeMismatch { | ||||||||||||
| return matched.ID, bound, true, effort | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| if requiredProvider != "" { | ||||||||||||
|
|
@@ -171,6 +175,34 @@ func resolveForceModelWithEffort(model string) (canonicalID, provider string, kn | |||||||||||
| } | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // resolveForcedBinding narrows explicit openai/<id> inputs before walking | ||||||||||||
| // bindings in catalog order. nativeMismatch distinguishes no required-provider | ||||||||||||
| // binding from one that exists but is unavailable; nil available is unrestricted. | ||||||||||||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
Suggested change
Was 3 lines; first clause ("narrows explicit openai/ inputs") restates what the requiredProvider filter does rather than explaining why.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Applied verbatim in 219ed65. |
||||||||||||
| func resolveForcedBinding(m catalog.Model, requiredProvider string, available map[string]struct{}) (provider string, nativeMismatch bool) { | ||||||||||||
| bindings := m.Providers | ||||||||||||
| if requiredProvider != "" { | ||||||||||||
| filtered := make([]catalog.ProviderBinding, 0, len(bindings)) | ||||||||||||
| for _, b := range bindings { | ||||||||||||
| if b.Provider == requiredProvider { | ||||||||||||
| filtered = append(filtered, b) | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| if len(filtered) == 0 { | ||||||||||||
| return "", true | ||||||||||||
| } | ||||||||||||
| bindings = filtered | ||||||||||||
| } | ||||||||||||
| if available == nil { | ||||||||||||
| return bindings[0].Provider, false | ||||||||||||
| } | ||||||||||||
| for _, b := range bindings { | ||||||||||||
| if _, ok := available[b.Provider]; ok { | ||||||||||||
| return b.Provider, false | ||||||||||||
| } | ||||||||||||
| } | ||||||||||||
| return "", false | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // stripEffortSuffix splits a `:level` suffix off model, canonicalizes it via | ||||||||||||
| // CanonicalizeEffort, and returns ("", model) when no recognized suffix found. | ||||||||||||
| func stripEffortSuffix(model string) (effort string, modelOut string) { | ||||||||||||
|
|
@@ -236,26 +268,22 @@ func (s *Service) setForceModelPin( | |||||||||||
| return s.pinStore.Upsert(context.Background(), forced) | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // applyForceModelHeader honors the x-weave-force-model request header, | ||||||||||||
| // writing the same session pin the /force-model command writes. It's | ||||||||||||
| // (re)written on every request carrying the header. Unrecognized models are | ||||||||||||
| // ignored (routing proceeds automatically) rather than failing the request. | ||||||||||||
| // | ||||||||||||
| // A `:level` suffix is stashed on context as router.Overrides.ForceEffort | ||||||||||||
| // so pin + effort land in one header. | ||||||||||||
| // applyForceModelHeader pins against live provider availability (#874). A | ||||||||||||
| // `:level` suffix is stored in routing context so pin and effort apply together. | ||||||||||||
| func (s *Service) applyForceModelHeader( | ||||||||||||
| ctx context.Context, | ||||||||||||
| r *http.Request, | ||||||||||||
| env *translate.RequestEnvelope, | ||||||||||||
| installationID uuid.UUID, | ||||||||||||
| sessionKey [sessionpin.SessionKeyLen]byte, | ||||||||||||
| enabledProviders map[string]struct{}, | ||||||||||||
| ) string { | ||||||||||||
| raw := strings.TrimSpace(r.Header.Get(ForceModelHeader)) | ||||||||||||
| if raw == "" { | ||||||||||||
| return "" | ||||||||||||
| } | ||||||||||||
| log := observability.FromContext(ctx) | ||||||||||||
| canonicalModel, provider, known, effortLevel := resolveForceModelWithEffort(raw) | ||||||||||||
| canonicalModel, provider, known, effortLevel := resolveForceModelWithEffort(raw, enabledProviders) | ||||||||||||
| if effortLevel != "" { | ||||||||||||
| // Merge with any existing knobs so ForceEffort doesn't drop Alpha/QualityBias. | ||||||||||||
| merged := router.Overrides{ForceEffort: effortLevel} | ||||||||||||
|
|
@@ -278,6 +306,16 @@ func (s *Service) applyForceModelHeader( | |||||||||||
| ) | ||||||||||||
| return "" | ||||||||||||
| } | ||||||||||||
| if provider == "" { | ||||||||||||
| // No available binding; do not pin an unusable provider (#874). | ||||||||||||
| log.Warn("x-weave-force-model: no available provider for model; routing automatically", | ||||||||||||
| "input_model", raw, | ||||||||||||
| "canonical_model", canonicalModel, | ||||||||||||
| "enabled_providers", sortedEnabledKeys(enabledProviders), | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| ) | ||||||||||||
| return "" | ||||||||||||
| } | ||||||||||||
| if s.pinStore == nil { | ||||||||||||
| return canonicalModel | ||||||||||||
| } | ||||||||||||
|
|
@@ -297,11 +335,8 @@ func (s *Service) applyForceModelHeader( | |||||||||||
| return canonicalModel | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| // handleForceModelCommand processes a /force-model or /unforce-model directive: | ||||||||||||
| // writes (or expires) the session pin and returns a synthetic acknowledgment | ||||||||||||
| // response without dispatching upstream. inputTokens should be the request's | ||||||||||||
| // RoutingFeatures.Tokens so the token counter reflects actual turn input, not | ||||||||||||
| // just the synthetic response text. | ||||||||||||
| // handleForceModelCommand writes or expires a session pin and returns a | ||||||||||||
| // synthetic acknowledgment. nil enabledProviders means unrestricted. | ||||||||||||
| func (s *Service) handleForceModelCommand( | ||||||||||||
| ctx context.Context, | ||||||||||||
| w http.ResponseWriter, | ||||||||||||
|
|
@@ -310,6 +345,7 @@ func (s *Service) handleForceModelCommand( | |||||||||||
| installationID uuid.UUID, | ||||||||||||
| sessionKey [sessionpin.SessionKeyLen]byte, | ||||||||||||
| inputTokens int, | ||||||||||||
| enabledProviders map[string]struct{}, | ||||||||||||
| ) error { | ||||||||||||
| log := observability.FromContext(ctx) | ||||||||||||
| role := roleForTier(catalog.TierFor(env.Model())) | ||||||||||||
|
|
@@ -318,7 +354,8 @@ func (s *Service) handleForceModelCommand( | |||||||||||
| // StripRoutingMarkerFromMessages strips it from later inbound requests; | ||||||||||||
| // otherwise it'd persist in history and leak router internals upstream. | ||||||||||||
| var msg string | ||||||||||||
| if cmd.Clear { | ||||||||||||
| switch { | ||||||||||||
| case cmd.Clear: | ||||||||||||
| if s.pinStore != nil && installationID != uuid.Nil { | ||||||||||||
| if err := s.expireSessionPin(ctx, installationID, sessionKey, role, "user_unforced"); err != nil { | ||||||||||||
| log.Error("/unforce-model: pin store upsert failed", "err", err) | ||||||||||||
|
|
@@ -334,34 +371,50 @@ func (s *Service) handleForceModelCommand( | |||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| } else if canonicalModel, provider, known := resolveForceModel(cmd.Model); !known { | ||||||||||||
| // Not in the catalog (e.g. truncated "/force-model gpt-") — reject | ||||||||||||
| // rather than pin something we can't honor; prior pin left untouched. | ||||||||||||
| log.Info("/force-model: rejected unknown model", | ||||||||||||
| "input_model", cmd.Model, | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| msg = fmt.Sprintf("✦ **Weave Router** → force-model: %q isn't a recognized model · keeping automatic routing. Use a full model ID, e.g. claude-opus-5, gpt-5.5, or gemini-3-pro-preview.\n\n", cmd.Model) | ||||||||||||
| if env.SourceFormat() == translate.FormatOpenAI { | ||||||||||||
| msg = fmt.Sprintf("Weave Router: force-model: %q isn't a recognized model; keeping automatic routing. Use a full model ID, e.g. claude-opus-5, gpt-5.5, or gemini-3-pro-preview.", cmd.Model) | ||||||||||||
| } | ||||||||||||
| } else { | ||||||||||||
| if err := s.setForceModelPin(ctx, sessionKey, role, installationID, canonicalModel, provider); err != nil { | ||||||||||||
| log.Error("/force-model: pin store upsert failed", "err", err) | ||||||||||||
| return err | ||||||||||||
| } | ||||||||||||
| msg = fmt.Sprintf("✦ **Weave Router** → force-model applied: %s (%s) · Use /unforce-model to clear\n\n", canonicalModel, provider) | ||||||||||||
| if env.SourceFormat() == translate.FormatOpenAI { | ||||||||||||
| msg = fmt.Sprintf("Weave Router: force-model applied: %s (%s). Use /unforce-model to clear.", canonicalModel, provider) | ||||||||||||
| default: | ||||||||||||
| canonicalModel, provider, known := resolveForceModel(cmd.Model, enabledProviders) | ||||||||||||
| switch { | ||||||||||||
| case !known: | ||||||||||||
| // Reject unknown input rather than pinning it. | ||||||||||||
| log.Info("/force-model: rejected unknown model", | ||||||||||||
| "input_model", cmd.Model, | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| msg = fmt.Sprintf("✦ **Weave Router** → force-model: %q isn't a recognized model · keeping automatic routing. Use a full model ID, e.g. claude-opus-5, gpt-5.5, or gemini-3-pro-preview.\n\n", cmd.Model) | ||||||||||||
| if env.SourceFormat() == translate.FormatOpenAI { | ||||||||||||
| msg = fmt.Sprintf("Weave Router: force-model: %q isn't a recognized model; keeping automatic routing. Use a full model ID, e.g. claude-opus-5, gpt-5.5, or gemini-3-pro-preview.", cmd.Model) | ||||||||||||
| } | ||||||||||||
| case provider == "": | ||||||||||||
| // Reject recognized models with no available binding (#874). | ||||||||||||
| log.Warn("/force-model: rejected model with no available provider", | ||||||||||||
| "input_model", cmd.Model, | ||||||||||||
| "canonical_model", canonicalModel, | ||||||||||||
| "enabled_providers", sortedEnabledKeys(enabledProviders), | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| msg = fmt.Sprintf("✦ **Weave Router** → force-model: %s has no available provider on this deployment · keeping automatic routing\n\n", canonicalModel) | ||||||||||||
| if env.SourceFormat() == translate.FormatOpenAI { | ||||||||||||
| msg = fmt.Sprintf("Weave Router: force-model: %s has no available provider on this deployment; keeping automatic routing.", canonicalModel) | ||||||||||||
| } | ||||||||||||
| default: | ||||||||||||
| if err := s.setForceModelPin(ctx, sessionKey, role, installationID, canonicalModel, provider); err != nil { | ||||||||||||
| log.Error("/force-model: pin store upsert failed", "err", err) | ||||||||||||
| return err | ||||||||||||
| } | ||||||||||||
| msg = fmt.Sprintf("✦ **Weave Router** → force-model applied: %s (%s) · Use /unforce-model to clear\n\n", canonicalModel, provider) | ||||||||||||
| if env.SourceFormat() == translate.FormatOpenAI { | ||||||||||||
| msg = fmt.Sprintf("Weave Router: force-model applied: %s (%s). Use /unforce-model to clear.", canonicalModel, provider) | ||||||||||||
| } | ||||||||||||
| log.Debug("/force-model: session pin set", | ||||||||||||
| "input_model", cmd.Model, | ||||||||||||
| "canonical_model", canonicalModel, | ||||||||||||
| "provider", provider, | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| } | ||||||||||||
| log.Debug("/force-model: session pin set", | ||||||||||||
| "input_model", cmd.Model, | ||||||||||||
| "canonical_model", canonicalModel, | ||||||||||||
| "provider", provider, | ||||||||||||
| "session_key_hex", fmt.Sprintf("%x", sessionKey), | ||||||||||||
| "role", role, | ||||||||||||
| ) | ||||||||||||
| } | ||||||||||||
|
|
||||||||||||
| switch env.SourceFormat() { | ||||||||||||
|
|
||||||||||||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,49 @@ | ||
| package proxy | ||
|
|
||
| import ( | ||
| "context" | ||
| "testing" | ||
| "time" | ||
|
|
||
| "workweave/router/internal/providers" | ||
| "workweave/router/internal/router" | ||
| "workweave/router/internal/router/sessionpin" | ||
| "workweave/router/internal/translate" | ||
|
|
||
| "github.com/google/uuid" | ||
| "github.com/stretchr/testify/assert" | ||
| "github.com/stretchr/testify/require" | ||
| ) | ||
|
|
||
| // A pin with an unavailable provider falls through to automatic routing (#874). | ||
| func TestRunTurnLoop_ForcedPin_FallsThroughWhenProviderUnavailable(t *testing.T) { | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
fans out to 6 callees (efferent coupling). Grounded coupling-delta finding (deterministic), not an LLM guess. There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
fans out to 6 callees (efferent coupling). Grounded coupling-delta finding (deterministic), not an LLM guess. |
||
| fr := &tierProbeRouter{available: map[string]struct{}{ | ||
| "deepseek/deepseek-v4-flash": {}, | ||
| "claude-sonnet-4-6": {}, | ||
| }} | ||
| store := &forcedPinStore{pin: sessionpin.Pin{ | ||
| Provider: providers.ProviderMakora, | ||
| Model: "deepseek/deepseek-v4-flash", | ||
| Reason: translate.ReasonUserForceModel, | ||
| PinnedUntil: time.Now().Add(time.Hour), | ||
| }} | ||
| svc := NewService(fr, nil, nil, false, nil, store, false, | ||
| providers.ProviderAnthropic, "claude-haiku-4-5", nil). | ||
| WithAvailableModels(fr.available). | ||
| WithPlannerEnabled(false) | ||
|
|
||
| env, err := translate.ParseAnthropic([]byte(`{"model":"claude-sonnet-4-6","messages":[{"role":"user","content":"hi"}]}`)) | ||
| require.NoError(t, err) | ||
| feats := env.RoutingFeatures(false) | ||
|
|
||
| res, err := svc.runTurnLoop(context.Background(), env, feats, "key-1", uuid.New(), "", nil, router.Request{ | ||
| RequestedModel: feats.Model, | ||
| // Makora is excluded from this turn's enabled set. | ||
| EnabledProviders: map[string]struct{}{providers.ProviderOpenRouter: {}, providers.ProviderAnthropic: {}}, | ||
| }) | ||
| require.NoError(t, err) | ||
|
|
||
| assert.False(t, res.StickyHit, "a pin on an unavailable provider must not be served as a sticky hit") | ||
| assert.NotEqual(t, providers.ProviderMakora, res.Decision.Provider, | ||
| "routing must not serve the excluded provider") | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Was 3 lines; last clause ("selects the primary binding") restates the nil-unrestricted invariant already stated on
resolveForceModel.There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Applied verbatim in 219ed65.