fix(router): map xai-oauth provider alias onto xai + use_xai_oauth

This commit is contained in:
Devin AI 2026-07-10 23:14:05 +00:00
parent eb7e4a567a
commit 30d98b9d58
2 changed files with 84 additions and 0 deletions

View file

@ -7357,6 +7357,7 @@ class Router:
- None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params)
"""
try:
_litellm_params = self._apply_xai_oauth_alias(_litellm_params)
litellm_params: LiteLLM_Params = LiteLLM_Params(**_litellm_params)
deployment = Deployment(
**deployment_info,
@ -7831,9 +7832,49 @@ class Router:
# deployments have been registered.
self._finalize_adaptive_router_if_configured()
@staticmethod
def _normalize_xai_oauth_alias_model(model: str, custom_llm_provider: Optional[str]) -> Optional[str]:
"""Map the ``xai-oauth`` / ``xai_oauth`` provider alias onto ``xai``.
xAI OAuth is enabled through the ``use_xai_oauth`` litellm param on an
``xai/`` deployment, so ``xai-oauth`` is not a real provider. Users who
follow the ``litellm xai-oauth login`` CLI naming and configure the
provider as ``xai-oauth`` (as a model prefix or ``custom_llm_provider``)
would otherwise hit "Unsupported provider - xai-oauth". Returns the
normalized ``xai/<model>`` string, or None when no alias is present.
"""
aliases = ("xai-oauth", "xai_oauth")
prefix = model.split("/", 1)[0] if "/" in model else None
if custom_llm_provider not in aliases and prefix not in aliases:
return None
bare_model = model.split("/", 1)[1] if prefix in (*aliases, "xai") else model
return f"xai/{bare_model}"
def _apply_xai_oauth_alias(self, litellm_params: dict) -> dict:
normalized_model = self._normalize_xai_oauth_alias_model(
litellm_params.get("model") or "",
litellm_params.get("custom_llm_provider"),
)
if normalized_model is None:
return litellm_params
return {
**{k: v for k, v in litellm_params.items() if k != "custom_llm_provider"},
"model": normalized_model,
"use_xai_oauth": True,
}
def _add_deployment(self, deployment: Deployment) -> Deployment:
import os
normalized_model = self._normalize_xai_oauth_alias_model(
deployment.litellm_params.model,
deployment.litellm_params.custom_llm_provider,
)
if normalized_model is not None:
deployment.litellm_params.model = normalized_model
deployment.litellm_params.custom_llm_provider = None
deployment.litellm_params.use_xai_oauth = True
#### VALIDATE MODEL ########
# Check if this is a prompt management model before validating as LLM provider
litellm_model = deployment.litellm_params.model

View file

@ -1068,6 +1068,49 @@ def test_add_invalid_provider_to_router():
assert router.pattern_router.patterns == {}
@pytest.mark.parametrize(
"litellm_params",
[
{"model": "xai-oauth/grok-4.5"},
{"model": "xai_oauth/grok-4.5"},
{"model": "grok-4.5", "custom_llm_provider": "xai-oauth"},
{"model": "xai/grok-4.5", "custom_llm_provider": "xai_oauth"},
{"model": "xai-oauth/grok-4.5", "custom_llm_provider": "xai-oauth"},
],
)
def test_router_normalizes_xai_oauth_alias(litellm_params):
"""
Regression for https://github.com/BerriAI/litellm/issues/32660
xai-oauth is not a real provider; it is the xai provider with the
use_xai_oauth flag. Configuring the provider as xai-oauth (matching the
`litellm xai-oauth login` CLI) used to fail with
"Unsupported provider - xai-oauth". The router should transparently rewrite
it to xai/<model> with use_xai_oauth=True.
"""
router = litellm.Router(
model_list=[{"model_name": "grok-4.5", "litellm_params": dict(litellm_params)}]
)
deployment = router.get_model_list()[0]["litellm_params"]
assert deployment["model"] == "xai/grok-4.5"
assert deployment.get("custom_llm_provider") is None
assert deployment["use_xai_oauth"] is True
def test_router_leaves_plain_xai_deployment_untouched():
"""A plain xai/ deployment must not have use_xai_oauth force-enabled."""
router = litellm.Router(
model_list=[
{"model_name": "grok-4.5", "litellm_params": {"model": "xai/grok-4.5"}}
]
)
deployment = router.get_model_list()[0]["litellm_params"]
assert deployment["model"] == "xai/grok-4.5"
assert not deployment.get("use_xai_oauth")
@pytest.mark.asyncio
async def test_router_ageneric_api_call_with_fallbacks_helper():
"""