Add dedicated xai_key and fallback logic for xAI API key

Add a provider-specific litellm.xai_key fallback for xAI chat,
responses, and realtime requests.

Keep the Responses API and realtime fallback order compatible by
preserving litellm.api_key before XAI_API_KEY when no explicit
provider-specific key is set.
This commit is contained in:
Peter Dave Hello 2026-01-06 02:21:33 +08:00
parent 8afd3b5dec
commit b5542dc75a
6 changed files with 332 additions and 9 deletions

View file

@ -231,6 +231,7 @@ api_key: Optional[str] = None
openai_key: Optional[str] = None
groq_key: Optional[str] = None
gigachat_key: Optional[str] = None
xai_key: Optional[str] = None
databricks_key: Optional[str] = None
openai_like_key: Optional[str] = None
azure_key: Optional[str] = None

View file

@ -9,6 +9,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
strip_name_from_messages,
)
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
@ -34,7 +35,7 @@ class XAIChatConfig(OpenAIGPTConfig):
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
api_base = api_base or get_secret_str("XAI_API_BASE") or XAI_API_BASE # type: ignore
dynamic_api_key = api_key or get_secret_str("XAI_API_KEY")
dynamic_api_key = XAIModelInfo.get_api_key(api_key)
return api_base, dynamic_api_key
def get_supported_openai_params(self, model: str) -> list:

View file

@ -45,8 +45,28 @@ class XAIModelInfo(BaseLLMModelInfo):
return api_base or get_secret_str("XAI_API_BASE") or "https://api.x.ai"
@staticmethod
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
return api_key or get_secret_str("XAI_API_KEY")
def get_api_key(
api_key: Optional[str] = None,
legacy_generic_before_env: bool = False,
) -> Optional[str]:
"""
Resolve xAI API keys while preserving endpoint-specific legacy order.
Chat uses xai_key before XAI_API_KEY without adding a generic
litellm.api_key fallback. Responses and realtime historically
preferred litellm.api_key over XAI_API_KEY, so those paths opt into
the legacy order with legacy_generic_before_env=True. In both modes,
the provider-specific litellm.xai_key takes precedence over fallbacks.
"""
if legacy_generic_before_env:
return (
api_key
or litellm.xai_key
or litellm.api_key
or get_secret_str("XAI_API_KEY")
)
return api_key or litellm.xai_key or get_secret_str("XAI_API_KEY")
@staticmethod
def get_base_model(model: str) -> Optional[str]:
@ -59,7 +79,7 @@ class XAIModelInfo(BaseLLMModelInfo):
api_key = self.get_api_key(api_key)
if api_base is None or api_key is None:
raise ValueError(
"XAI_API_BASE or XAI_API_KEY is not set. Please set the environment variable, to query XAI's `/models` endpoint."
"XAI API base or key is not set. Set XAI_API_BASE and provide an xAI API key via api_key, litellm.xai_key, or XAI_API_KEY."
)
response = litellm.module_level_client.get(
url=f"{api_base}/v1/models",

View file

@ -4,6 +4,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm.constants import XAI_API_BASE
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.llms.xai import XAIWebSearchTool, XAIXSearchTool
@ -212,16 +213,17 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
Validate environment and set up headers for XAI API.
Uses XAI_API_KEY from environment or litellm_params.
Uses the shared xAI key resolver with Responses API legacy precedence.
"""
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = (
litellm_params.api_key or litellm.api_key or get_secret_str("XAI_API_KEY")
api_key = XAIModelInfo.get_api_key(
litellm_params.api_key, legacy_generic_before_env=True
)
if not api_key:
raise ValueError(
"XAI API key is required. Set XAI_API_KEY environment variable or pass api_key parameter."
"XAI API key is required. Set api_key, litellm.xai_key, "
"litellm.api_key, or XAI_API_KEY."
)
headers.update(

View file

@ -8,6 +8,7 @@ from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, request
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.secret_managers.main import get_secret_str
from litellm.types.realtime import (
RealtimeClientSecretRequest,
@ -383,7 +384,9 @@ async def _arealtime( # noqa: PLR0915
or "https://api.x.ai/v1"
)
# set API KEY
api_key = dynamic_api_key or litellm.api_key or get_secret_str("XAI_API_KEY")
api_key = XAIModelInfo.get_api_key(
dynamic_api_key, legacy_generic_before_env=True
)
await xai_realtime.async_realtime(
model=model,

View file

@ -0,0 +1,296 @@
import asyncio
import os
import sys
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import pytest
import litellm
from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.realtime_api import main as realtime_main
from litellm.types.router import GenericLiteLLMParams
class FakeLogging:
def update_from_kwargs(self, **kwargs):
pass
def test_get_api_key_prefers_xai_key_over_environment_and_generic_key(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
assert XAIModelInfo.get_api_key(None) == "xai_key_value"
def test_get_api_key_prefers_explicit_key_for_both_orderings(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
assert XAIModelInfo.get_api_key("param_api_key") == "param_api_key"
assert (
XAIModelInfo.get_api_key("param_api_key", legacy_generic_before_env=True)
== "param_api_key"
)
def test_get_api_key_prefers_environment_over_generic_key_by_default(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
assert XAIModelInfo.get_api_key(None) == "env_api_key"
def test_get_api_key_does_not_use_generic_key_by_default(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.delenv("XAI_API_KEY", raising=False)
assert XAIModelInfo.get_api_key(None) is None
def test_get_api_key_legacy_order_prefers_generic_key_over_env(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
assert (
XAIModelInfo.get_api_key(None, legacy_generic_before_env=True)
== "common_api_key"
)
def test_get_api_key_legacy_order_prefers_xai_key_over_generic_key(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
assert (
XAIModelInfo.get_api_key(None, legacy_generic_before_env=True)
== "xai_key_value"
)
def test_get_api_key_returns_none_when_no_key_is_available(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.delenv("XAI_API_KEY", raising=False)
assert XAIModelInfo.get_api_key(None) is None
def test_chat_config_uses_xai_key_fallback(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.delenv("XAI_API_KEY", raising=False)
_, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None)
assert api_key == "xai_key_value"
def test_chat_config_uses_environment_key_fallback(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
_, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None)
assert api_key == "env_api_key"
def test_chat_config_does_not_use_generic_key_fallback(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.delenv("XAI_API_KEY", raising=False)
_, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None)
assert api_key is None
def test_chat_config_prefers_explicit_api_key(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
_, api_key = XAIChatConfig()._get_openai_compatible_provider_info(
None, "param_api_key"
)
assert api_key == "param_api_key"
def test_responses_config_preserves_generic_key_precedence(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
headers = XAIResponsesAPIConfig().validate_environment({}, "xai/grok-3-mini", None)
assert headers["Authorization"] == "Bearer common_api_key"
def test_responses_config_prefers_litellm_params_api_key(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
headers = XAIResponsesAPIConfig().validate_environment(
{},
"xai/grok-3-mini",
GenericLiteLLMParams(api_key="param_api_key"),
)
assert headers["Authorization"] == "Bearer param_api_key"
def test_responses_config_uses_environment_key_fallback(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
headers = XAIResponsesAPIConfig().validate_environment({}, "xai/grok-3-mini", None)
assert headers["Authorization"] == "Bearer env_api_key"
def test_responses_config_raises_when_no_key_is_available(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.delenv("XAI_API_KEY", raising=False)
with pytest.raises(ValueError) as exc_info:
XAIResponsesAPIConfig().validate_environment({}, "xai/grok-3-mini", None)
error_message = str(exc_info.value)
assert "api_key" in error_message
assert "litellm.xai_key" in error_message
assert "litellm.api_key" in error_message
assert "XAI_API_KEY" in error_message
def test_responses_config_prefers_xai_key_over_generic_key(monkeypatch):
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
headers = XAIResponsesAPIConfig().validate_environment({}, "xai/grok-3-mini", None)
assert headers["Authorization"] == "Bearer xai_key_value"
def test_realtime_config_uses_xai_key_through_provider_resolution(monkeypatch):
captured_kwargs = {}
async def mock_async_realtime(**kwargs):
captured_kwargs.update(kwargs)
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
monkeypatch.setattr(
realtime_main.xai_realtime, "async_realtime", mock_async_realtime
)
asyncio.run(
realtime_main._arealtime(
model="xai/grok-4-1-fast-non-reasoning",
websocket=object(),
litellm_logging_obj=FakeLogging(),
)
)
assert captured_kwargs["api_key"] == "xai_key_value"
def test_realtime_config_uses_xai_key_when_provider_does_not_resolve_key(monkeypatch):
captured_kwargs = {}
async def mock_async_realtime(**kwargs):
captured_kwargs.update(kwargs)
def mock_get_llm_provider(model, api_base, api_key):
return model, "xai", None, api_base
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.setenv("XAI_API_KEY", "env_api_key")
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
monkeypatch.setattr(
realtime_main.xai_realtime, "async_realtime", mock_async_realtime
)
asyncio.run(
realtime_main._arealtime(
model="xai/grok-4-1-fast-non-reasoning",
websocket=object(),
litellm_logging_obj=FakeLogging(),
)
)
assert captured_kwargs["api_key"] == "xai_key_value"
def test_realtime_config_uses_generic_key_when_provider_does_not_resolve_key(
monkeypatch,
):
captured_kwargs = {}
async def mock_async_realtime(**kwargs):
captured_kwargs.update(kwargs)
def mock_get_llm_provider(model, api_base, api_key):
return model, "xai", None, api_base
monkeypatch.setattr(litellm, "xai_key", None)
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.delenv("XAI_API_KEY", raising=False)
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
monkeypatch.setattr(
realtime_main.xai_realtime, "async_realtime", mock_async_realtime
)
asyncio.run(
realtime_main._arealtime(
model="xai/grok-4-1-fast-non-reasoning",
websocket=object(),
litellm_logging_obj=FakeLogging(),
)
)
assert captured_kwargs["api_key"] == "common_api_key"
def test_get_models_uses_xai_key_fallback(monkeypatch):
captured_kwargs = {}
class FakeResponse:
status_code = 200
text = "{}"
def raise_for_status(self):
pass
def json(self):
return {"data": [{"id": "grok-test"}]}
def mock_get(**kwargs):
captured_kwargs.update(kwargs)
return FakeResponse()
monkeypatch.setattr(litellm, "xai_key", "xai_key_value")
monkeypatch.setattr(litellm, "api_key", "common_api_key")
monkeypatch.delenv("XAI_API_KEY", raising=False)
monkeypatch.setattr(litellm.module_level_client, "get", mock_get)
assert XAIModelInfo().get_models() == ["xai/grok-test"]
assert captured_kwargs["headers"]["Authorization"] == "Bearer xai_key_value"