mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
fd2fce1323
commit
cb8f903bb9
3 changed files with 134 additions and 31 deletions
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue