fix: only apply NIM rotation to known NIM models

The NIM key rotation was incorrectly applying to all models when
NVIDIA_NIM_API_KEY_N are configured. This broke non-NIM models like
z-ai/glm-5.2 which should route through the Alias gateway, not
NVIDIA NIM. Add _is_nim_model() check with known NIM provider
prefixes to gate rotation. Fix applies to both LiteLLM path and
direct httpx path.
This commit is contained in:
UnaiAlias 2026-07-31 19:07:37 +00:00
parent c2408889d2
commit 929c2ea6c3
1 changed files with 35 additions and 10 deletions

View File

@ -3668,16 +3668,27 @@ class OpenAIChatCompletionsModel(Model):
# round-robin through them to stay under 40 req/min per key.
# Only applies when no explicit api_base was routed above (alias,
# ollama_cloud, or a custom provider fallback): those must keep
# their own endpoint + key.
# their own endpoint + key. Also only applies to known NIM models.
def _is_nim_model(model: str) -> bool:
"""Check if model is a known NIM model (provider is on NVIDIA NIM)."""
provider = model.split("/")[0].lower() if "/" in model else model.lower()
nim_providers = {
"nvidia", "meta", "google", "microsoft", "mistralai",
"llama", "nemotron", "codellama", "mixtral", "phi",
"qwen", "yi", "deepseek", "gemma", "stable-diffusion",
"sdxl", "cosxl", "proteus", "realistic-vision",
}
return provider in nim_providers
if is_nim_rotation_configured() and not kwargs.get("api_base"):
kwargs["api_key"] = get_next_nim_key()
kwargs["custom_llm_provider"] = "openai"
kwargs["api_base"] = (
resolve_llm_openai_compatible_base(
str(kwargs.get("model") or os.getenv("CAI_MODEL") or "")
).rstrip("/")
or os.getenv("OPENAI_API_BASE", "").rstrip("/")
)
model_str = str(kwargs.get("model") or os.getenv("CAI_MODEL") or "")
if _is_nim_model(model_str):
kwargs["api_key"] = get_next_nim_key()
kwargs["custom_llm_provider"] = "openai"
kwargs["api_base"] = (
resolve_llm_openai_compatible_base(model_str).rstrip("/")
or os.getenv("OPENAI_API_BASE", "").rstrip("/")
)
# Filter out NotGiven values to avoid JSON serialization issues
filtered_kwargs = {}
@ -3785,7 +3796,21 @@ class OpenAIChatCompletionsModel(Model):
request_body = {k: v for k, v in request_body.items() if v is not None}
api_url = f"{openai_api_base.rstrip('/')}/chat/completions"
if "api.nvidia.com" in openai_api_base.lower() and is_nim_rotation_configured():
model_str = str(kwargs.get("model") or os.getenv("CAI_MODEL") or "")
def _is_nim_model(model: str) -> bool:
provider = model.split("/")[0].lower() if "/" in model else model.lower()
nim_providers = {
"nvidia", "meta", "google", "microsoft", "mistralai",
"llama", "nemotron", "codellama", "mixtral", "phi",
"qwen", "yi", "deepseek", "gemma", "stable-diffusion",
"sdxl", "cosxl", "proteus", "realistic-vision",
}
return provider in nim_providers
if (
"api.nvidia.com" in openai_api_base.lower()
and is_nim_rotation_configured()
and _is_nim_model(model_str)
):
direct_api_key = get_next_nim_key() or "sk-placeholder"
else:
direct_api_key = (