fix(router): resolve group aliases and mixed-format fallbacks for mid-stream prefill check

Greptile flagged two gaps in build_mid_stream_continuation_messages. Custom
router model-group names (e.g. an alias that resolves to claude-sonnet-4-6) are
absent from the registry, so the capability lookup silently fell back to the
legacy prefill path and the 400 persisted; the router now injects a resolver
that maps each candidate group to its deployment's registry model. Mixed
dict+string fallback lists also let a prefill-rejecting bare-string entry slip
through, because get_fallback_model_group only surfaces the dict target for the
group; candidate collection now gathers bare strings and dict targets uniformly
across every list format.
This commit is contained in:
mateo-berri 2026-06-18 13:34:58 +00:00
parent fd2fce1323
commit cb8f903bb9
No known key found for this signature in database
3 changed files with 134 additions and 31 deletions

View file

@ -2243,6 +2243,7 @@ class Router:
generated_content=e.generated_content,
model_group=model_group,
fallbacks=fallbacks,
resolve_underlying_model=self._underlying_model_for_group,
)
)
self._update_kwargs_before_fallbacks(
@ -2803,6 +2804,7 @@ class Router:
generated_content=e.generated_content,
model_group=model_group,
fallbacks=fallbacks,
resolve_underlying_model=router_self._underlying_model_for_group,
)
)
router_self._update_kwargs_before_fallbacks(
@ -9050,6 +9052,15 @@ class Router:
raise Exception("Model Name invalid - {}".format(type(model)))
return None
def _underlying_model_for_group(self, model_group_name: str) -> str | None:
"""Registry model name the group resolves to (e.g. the alias
``production-claude`` -> ``anthropic/claude-sonnet-4-6``), or None when
the group has no deployment."""
deployment = self.get_deployment_by_model_group_name(model_group_name)
if deployment is None:
return None
return deployment.litellm_params.model
def get_deployment_credentials_with_provider(
self, model_id: str
) -> Optional[Dict[str, Any]]:

View file

@ -1,3 +1,4 @@
from collections.abc import Callable
from enum import Enum
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
@ -19,8 +20,10 @@ else:
MID_STREAM_CONTINUATION_SYSTEM_PROMPT = "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: "
FallbackEntry = Union[str, Dict[str, List[str]]]
def _prefill_explicitly_unsupported(model_group: str | None) -> bool:
def _prefill_explicitly_unsupported(model: str | None) -> bool:
"""True only when the model registry explicitly marks the model as NOT
supporting assistant prefill (``supports_assistant_prefill: false``).
@ -28,22 +31,60 @@ def _prefill_explicitly_unsupported(model_group: str | None) -> bool:
the legacy prefill behavior — only models that would reject the prefill with
a 400 anyway are routed to the user-message continuation.
"""
if model_group is None:
if model is None:
return False
try:
from litellm.utils import get_model_info
model_info = get_model_info(model=model_group)
model_info = get_model_info(model=model)
return model_info.get("supports_assistant_prefill") is False
except Exception:
return False
def _candidate_model_groups(
model_group: str | None, fallbacks: list[FallbackEntry] | None
) -> tuple[str, ...]:
"""Every model group the reused continuation messages can reach for this hop:
the primary group plus its fallback targets. Bare-string fallback entries are
generic fallbacks tried for every group; dict-format entries are resolved for
this group. Both formats (and mixed lists) are collected directly so a
prefill-rejecting target anywhere in the list is never missed.
"""
if model_group is None:
return ()
if not fallbacks:
return (model_group,)
bare_strings = tuple(f for f in fallbacks if isinstance(f, str))
try:
# Copy: get_fallback_model_group pops string entries while iterating,
# which would mutate the caller's live fallbacks list.
dict_targets, _ = get_fallback_model_group(
fallbacks=list(fallbacks), model_group=model_group
)
except Exception:
dict_targets = None
return (model_group, *bare_strings, *(dict_targets or ()))
def _resolved_registry_model(
model_group: str,
resolve_underlying_model: Callable[[str], str | None] | None,
) -> str | None:
if resolve_underlying_model is None:
return None
try:
return resolve_underlying_model(model_group)
except Exception:
return None
def build_mid_stream_continuation_messages(
messages: list[Any],
generated_content: str,
model_group: str | None,
fallbacks: list[Any] | None = None,
resolve_underlying_model: Callable[[str], str | None] | None = None,
) -> list[Any]:
"""Build the message list a mid-stream fallback uses to resume an interrupted stream.
@ -60,35 +101,19 @@ def build_mid_stream_continuation_messages(
The same continuation messages are sent to every deployment the fallback
chain tries, so the safe pattern engages when the primary model OR any
configured fallback target for this model group is explicitly marked as
not supporting prefill.
configured fallback target is explicitly marked as not supporting prefill.
Router model-group names are often custom aliases absent from the registry,
so ``resolve_underlying_model`` maps each candidate group to the deployment
model the registry actually keys on.
"""
candidate_models: list[str | None] = [model_group]
if fallbacks is not None and model_group is not None:
if fallbacks and all(isinstance(f, str) for f in fallbacks):
# Flat string-format lists (["model-a", "model-b", ...]) are tried in
# order for every model_group. get_fallback_model_group surfaces only
# ONE entry from them — it pops a single string mid-iteration and
# leaves the rest hidden — so a prefill-rejecting model at a non-first
# position would slip through, and the once-built continuation that is
# reused across every hop would 400 when the chain reaches it. Check
# every entry directly instead.
candidate_models.extend(fallbacks)
else:
try:
# list() copy: get_fallback_model_group POPS string entries from a
# mixed fallbacks list while iterating — mutating the live list
# here would silently drop a fallback target before the actual
# fallback execution runs.
fallback_model_group, _ = get_fallback_model_group(
fallbacks=list(fallbacks), model_group=model_group
)
if fallback_model_group:
candidate_models.extend(fallback_model_group)
except Exception:
pass
if any(_prefill_explicitly_unsupported(m) for m in candidate_models):
groups = _candidate_model_groups(model_group, fallbacks)
models_to_check = frozenset(
model
for group in groups
for model in (group, _resolved_registry_model(group, resolve_underlying_model))
if model is not None
)
if any(_prefill_explicitly_unsupported(model) for model in models_to_check):
return messages + [
{
"role": "user",

View file

@ -174,3 +174,70 @@ def test_flat_string_fallback_all_prefill_supporting_keeps_legacy():
fallbacks=["gpt-4o", "gpt-3.5-turbo"],
)
_assert_legacy_prefill(result)
def test_mixed_format_fallback_prefill_rejecter_string_entry():
"""A mixed dict+string fallback list must still catch a prefill-rejecting
bare-string entry — get_fallback_model_group only surfaces the dict target
for the group, so the trailing string had been slipping through."""
result = build_mid_stream_continuation_messages(
messages=MESSAGES,
generated_content=PARTIAL,
model_group="gpt-4",
fallbacks=[{"gpt-4": ["gpt-4o"]}, "claude-sonnet-4-6"],
)
assert len(result) == 2
assert result[1]["role"] == "user"
assert PARTIAL in result[1]["content"]
assert all(m.get("prefix") is not True for m in result)
def test_custom_group_alias_resolved_to_prefill_rejecter():
"""A router group name is often a custom alias absent from the registry
(get_model_info(alias) raises). The injected resolver maps it to the
deployment's registry model so the prefill-rejecting capability is honored."""
aliases = {"production-claude": "anthropic/claude-sonnet-4-6"}
result = build_mid_stream_continuation_messages(
messages=MESSAGES,
generated_content=PARTIAL,
model_group="production-claude",
resolve_underlying_model=aliases.get,
)
assert len(result) == 2
assert result[1]["role"] == "user"
assert PARTIAL in result[1]["content"]
def test_custom_group_alias_without_resolver_falls_back_to_legacy():
"""Without a resolver an unregistered alias cannot be classified, so the
legacy prefill path is preserved (no regression vs. pre-fix behavior)."""
_assert_legacy_prefill(_build("production-claude"))
def test_resolver_applied_to_fallback_target_alias():
"""The resolver also resolves fallback-target aliases, not just the primary
group, so a prefill-rejecting fallback hidden behind an alias is caught."""
aliases = {"prod-anthropic": "claude-sonnet-4-6"}
result = build_mid_stream_continuation_messages(
messages=MESSAGES,
generated_content=PARTIAL,
model_group="gpt-4",
fallbacks=[{"gpt-4": ["prod-anthropic"]}],
resolve_underlying_model=aliases.get,
)
assert len(result) == 2
assert result[1]["role"] == "user"
assert PARTIAL in result[1]["content"]
def test_resolver_returning_prefill_supporting_model_keeps_legacy():
"""A resolver that maps to a prefill-supporting registry model must not flip
the behavior — only an explicit supports_assistant_prefill=false does."""
aliases = {"prod-gpt": "gpt-4o"}
result = build_mid_stream_continuation_messages(
messages=MESSAGES,
generated_content=PARTIAL,
model_group="prod-gpt",
resolve_underlying_model=aliases.get,
)
_assert_legacy_prefill(result)