diff --git a/litellm/router.py b/litellm/router.py index 6d8d4dfd032..3345be3f133 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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]]: diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index b8ebdacddb4..a3cbb7ef4fa 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -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", diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 4d4ca491742..0933b4d410e 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -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)