mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
8afd3b5dec
commit
b5542dc75a
6 changed files with 332 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
296
tests/test_litellm/llms/xai/test_xai_key_fallback.py
Normal file
296
tests/test_litellm/llms/xai/test_xai_key_fallback.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue