diff --git a/litellm/__init__.py b/litellm/__init__.py index 9eb3f075d5ec..d372c737a668 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -199,6 +199,7 @@ 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 diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 245e10e45c12..3f9b2085822f 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -4,6 +4,7 @@ import litellm from litellm._logging import verbose_logger +from litellm.llms.xai.common_utils import XAIModelInfo from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, strip_name_from_messages, @@ -26,7 +27,7 @@ def _get_openai_compatible_provider_info( 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: diff --git a/litellm/llms/xai/common_utils.py b/litellm/llms/xai/common_utils.py index df324cf3ee22..f2c6b935e366 100644 --- a/litellm/llms/xai/common_utils.py +++ b/litellm/llms/xai/common_utils.py @@ -46,7 +46,12 @@ def get_api_base(api_base: Optional[str] = None) -> Optional[str]: @staticmethod def get_api_key(api_key: Optional[str] = None) -> Optional[str]: - return api_key or get_secret_str("XAI_API_KEY") + return ( + api_key + or litellm.xai_key + or get_secret_str("XAI_API_KEY") + or litellm.api_key + ) @staticmethod def get_base_model(model: str) -> Optional[str]: diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index bd422c8d81e0..740e2a13aeb9 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -2,6 +2,7 @@ import litellm from litellm._logging import verbose_logger +from litellm.llms.xai.common_utils import XAIModelInfo from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams @@ -103,11 +104,7 @@ def validate_environment( Uses XAI_API_KEY from environment or litellm_params. """ 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) if not api_key: raise ValueError( @@ -143,4 +140,3 @@ def get_complete_url( api_base = api_base.rstrip("/") return f"{api_base}/responses" - diff --git a/tests/litellm/llms/xai/test_xai_key_fallback.py b/tests/litellm/llms/xai/test_xai_key_fallback.py new file mode 100644 index 000000000000..a678d961ad54 --- /dev/null +++ b/tests/litellm/llms/xai/test_xai_key_fallback.py @@ -0,0 +1,100 @@ +import os +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 + + +def test_get_api_key_priority(): + """ + Test the fallback order of XAI API key resolution: + 1. api_key parameter + 2. litellm.xai_key + 3. XAI_API_KEY environment variable + 4. litellm.api_key + 5. None + """ + + original_env = os.environ.get("XAI_API_KEY") + had_env = "XAI_API_KEY" in os.environ + original_xai_key = litellm.xai_key + original_api_key = litellm.api_key + + try: + # Case 1: api_key parameter is passed + litellm.xai_key = "xai_key_value" + litellm.api_key = "common_api_key" + os.environ["XAI_API_KEY"] = "env_api_key" + result = XAIModelInfo.get_api_key("param_api_key") + assert result == "param_api_key" + + # Case 2: api_key not passed, use litellm.xai_key + result = XAIModelInfo.get_api_key(None) + assert result == "xai_key_value" + + # Case 3: api_key and xai_key not set, prefer XAI_API_KEY over litellm.api_key + litellm.xai_key = None + result = XAIModelInfo.get_api_key(None) + assert result == "env_api_key" + + # Case 4: Empty XAI_API_KEY falls through to litellm.api_key + os.environ["XAI_API_KEY"] = "" + result = XAIModelInfo.get_api_key(None) + assert result == "common_api_key" + + # Case 5: None of the above, return None + os.environ.pop("XAI_API_KEY", None) + litellm.api_key = None + result = XAIModelInfo.get_api_key(None) + assert result is None + finally: + if had_env: + os.environ["XAI_API_KEY"] = original_env + else: + os.environ.pop("XAI_API_KEY", None) + litellm.xai_key = original_xai_key + litellm.api_key = original_api_key + + +def test_chat_config_uses_xai_key_fallback(): + original_env = os.environ.get("XAI_API_KEY") + had_env = "XAI_API_KEY" in os.environ + original_xai_key = litellm.xai_key + original_api_key = litellm.api_key + + try: + litellm.xai_key = "xai_key_value" + litellm.api_key = None + os.environ.pop("XAI_API_KEY", None) + _, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None) + assert api_key == "xai_key_value" + finally: + if had_env: + os.environ["XAI_API_KEY"] = original_env + else: + os.environ.pop("XAI_API_KEY", None) + litellm.xai_key = original_xai_key + litellm.api_key = original_api_key + + +def test_responses_config_uses_xai_key_fallback(): + original_env = os.environ.get("XAI_API_KEY") + had_env = "XAI_API_KEY" in os.environ + original_xai_key = litellm.xai_key + original_api_key = litellm.api_key + + try: + litellm.xai_key = "xai_key_value" + litellm.api_key = None + os.environ.pop("XAI_API_KEY", None) + headers = XAIResponsesAPIConfig().validate_environment( + {}, "xai/grok-3-mini", None + ) + assert headers["Authorization"] == "Bearer xai_key_value" + finally: + if had_env: + os.environ["XAI_API_KEY"] = original_env + else: + os.environ.pop("XAI_API_KEY", None) + litellm.xai_key = original_xai_key + litellm.api_key = original_api_key