diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index b32803b5dfc..f80cb41dc3f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -34,6 +34,7 @@ _OPTIONAL_KWARGS_KEYS = frozenset( "aws_bedrock_runtime_endpoint", "tpm", "rpm", + "use_xai_oauth", } ) diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index c06928516ef..8019bb67991 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -5,6 +5,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, strip_name_from_messages, @@ -39,6 +40,72 @@ class XAIChatConfig(OpenAIGPTConfig): dynamic_api_key = XAIModelInfo.get_api_key(api_key) return api_base, dynamic_api_key + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + dynamic_api_key = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(litellm_params) and not dynamic_api_key: + try: + headers["Authorization"] = ( + f"Bearer {XAIOAuthAuthenticator().get_access_token()}" + ) + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider=self.custom_llm_provider or "xai", + message=str(exc), + ) from exc + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + return headers + + return super().validate_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=dynamic_api_key, + api_base=api_base, + ) + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + dynamic_api_key = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(litellm_params) and not dynamic_api_key: + api_base = XAIOAuthAuthenticator().get_api_base() + + return super().get_complete_url( + api_base=api_base, + api_key=dynamic_api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) + def get_supported_openai_params(self, model: str) -> list: base_openai_params = [ "logit_bias", diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py new file mode 100644 index 00000000000..30c717b7ca0 --- /dev/null +++ b/litellm/llms/xai/oauth.py @@ -0,0 +1,421 @@ +import base64 +import hashlib +import json +import os +import secrets +import sys +import threading +import time +import uuid +import webbrowser +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any, Dict, Optional, Tuple, Union +from urllib.parse import parse_qs, urlencode, urlparse + +import httpx + +from litellm._logging import verbose_logger +from litellm.constants import XAI_API_BASE +from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client +from litellm.secret_managers.main import get_secret_str + +XAI_OAUTH_ISSUER = "https://auth.x.ai" +XAI_OAUTH_DISCOVERY_URL = f"{XAI_OAUTH_ISSUER}/.well-known/openid-configuration" +XAI_OAUTH_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828" +XAI_OAUTH_SCOPE = "openid profile email offline_access grok-cli:access api:access" +XAI_OAUTH_REDIRECT_HOST = "127.0.0.1" +XAI_OAUTH_REDIRECT_PORT = 56121 +XAI_OAUTH_REDIRECT_PATH = "/callback" +XAI_OAUTH_EXPIRY_SKEW_SECONDS = 120 +XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS = 180 +_XAI_OAUTH_REFRESH_LOCK = threading.Lock() + + +class XAIOAuthError(Exception): + pass + + +class XAIOAuthLoginRequiredError(XAIOAuthError): + pass + + +class _CallbackHandler(BaseHTTPRequestHandler): + server: "_CallbackServer" + + def do_GET(self) -> None: + parsed = urlparse(self.path) + if parsed.path != XAI_OAUTH_REDIRECT_PATH: + self.send_response(404) + self.end_headers() + return + + params = parse_qs(parsed.query) + result = { + "code": params.get("code", [None])[0], + "state": params.get("state", [None])[0], + "error": params.get("error", [None])[0], + "error_description": params.get("error_description", [None])[0], + } + self.server.callback_result = result + + if result["state"] != self.server.expected_state: + self.send_response(400) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.end_headers() + self.wfile.write( + b"

xAI authorization state mismatch.

" + ) + return + + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.end_headers() + body = ( + b"

xAI authorization failed.

You can close this tab." + if result["error"] + else b"

xAI authorization received.

You can close this tab." + ) + self.wfile.write(body) + + def log_message(self, format: str, *args: Any) -> None: + return + + +class _CallbackServer(HTTPServer): + expected_state: str + callback_result: Optional[Dict[str, Optional[str]]] + + +class XAIOAuthAuthenticator: + def __init__( + self, http_client: Optional[Union[httpx.Client, HTTPHandler]] = None + ) -> None: + self.token_dir = get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser( + "~/.config/litellm/xai_oauth" + ) + self.auth_file = os.path.join( + self.token_dir, get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json" + ) + self.http_client = http_client + + def get_api_base(self) -> str: + return ( + get_secret_str("XAI_OAUTH_API_BASE") + or get_secret_str("XAI_API_BASE") + or XAI_API_BASE + ) + + def get_access_token(self) -> str: + auth_data = self._read_auth_file() + if not auth_data: + raise XAIOAuthLoginRequiredError( + "xAI OAuth login required. Run `litellm xai-oauth login`." + ) + + access_token = auth_data.get("access_token") + if access_token and not self._is_expired(auth_data): + return access_token + + refresh_token = auth_data.get("refresh_token") + if not refresh_token: + raise XAIOAuthLoginRequiredError( + "xAI OAuth refresh token missing. Run `litellm xai-oauth login`." + ) + + with _XAI_OAUTH_REFRESH_LOCK: + locked_auth_data = self._read_auth_file() or auth_data + access_token = locked_auth_data.get("access_token") + if access_token and not self._is_expired(locked_auth_data): + return access_token + + refreshed = self._refresh_tokens(locked_auth_data) + return refreshed["access_token"] + + def login(self, force: bool = False, no_browser: bool = False) -> Dict[str, Any]: + existing = self._read_auth_file() + if existing and not force and existing.get("access_token"): + if not self._is_expired(existing): + return existing + if existing.get("refresh_token"): + try: + return self._refresh_tokens(existing) + except XAIOAuthError: + pass + + discovery = self._discover() + verifier, challenge = self._pkce_pair() + state = uuid.uuid4().hex + nonce = uuid.uuid4().hex + server, redirect_uri = self._start_callback_server(state) + authorize_url = self._build_authorize_url( + authorization_endpoint=discovery["authorization_endpoint"], + redirect_uri=redirect_uri, + challenge=challenge, + state=state, + nonce=nonce, + ) + + if no_browser or not webbrowser.open(authorize_url): + sys.stdout.write( + f"Open this URL to authenticate with xAI:\n{authorize_url}\n" + ) + sys.stdout.flush() + + result = self._wait_for_callback(server) + if result.get("state") != state: + raise XAIOAuthError("xAI OAuth state mismatch") + if result.get("error"): + description = result.get("error_description") or result["error"] + raise XAIOAuthError(f"xAI authorization failed: {description}") + code = result.get("code") + if not code: + raise XAIOAuthError("xAI authorization failed: no code returned") + + token_payload = self._exchange_token( + discovery["token_endpoint"], + { + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": XAI_OAUTH_CLIENT_ID, + "code_verifier": verifier, + }, + ) + auth_data = self._build_auth_record(token_payload, discovery["token_endpoint"]) + self._write_auth_file(auth_data) + return auth_data + + def _client(self) -> Union[httpx.Client, HTTPHandler]: + return self.http_client or _get_httpx_client() + + def _ensure_token_dir(self) -> None: + os.makedirs(self.token_dir, mode=0o700, exist_ok=True) + try: + os.chmod(self.token_dir, 0o700) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth token directory") + + def _read_auth_file(self) -> Optional[Dict[str, Any]]: + try: + with open(self.auth_file, "r") as f: + data = json.load(f) + return data if isinstance(data, dict) else None + except (IOError, json.JSONDecodeError): + return None + + def _write_auth_file(self, data: Dict[str, Any]) -> None: + self._ensure_token_dir() + tmp_file = os.path.join( + self.token_dir, + f".{os.path.basename(self.auth_file)}.{uuid.uuid4().hex}.tmp", + ) + flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL + if hasattr(os, "O_NOFOLLOW"): + flags |= os.O_NOFOLLOW + fd = os.open(tmp_file, flags, 0o600) + try: + with os.fdopen(fd, "w") as f: + json.dump(data, f) + f.flush() + os.fsync(f.fileno()) + os.replace(tmp_file, self.auth_file) + try: + os.chmod(self.auth_file, 0o600) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth auth file") + except Exception: + try: + os.close(fd) + except OSError: + pass + try: + os.unlink(tmp_file) + except OSError: + pass + raise + + def _is_expired(self, auth_data: Dict[str, Any]) -> bool: + expires_at = auth_data.get("expires_at") + if expires_at is None: + return True + try: + return time.time() >= float(expires_at) - XAI_OAUTH_EXPIRY_SKEW_SECONDS + except (TypeError, ValueError): + return True + + def _discover(self) -> Dict[str, str]: + try: + response = self._client().get( + XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"} + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + raise XAIOAuthError( + f"xAI OAuth discovery request failed: {exc.response.status_code} {exc.response.text}" + ) from exc + try: + data = response.json() + except ValueError as exc: + raise XAIOAuthError( + "xAI OAuth discovery response was not valid JSON" + ) from exc + authorization_endpoint = data.get("authorization_endpoint") + token_endpoint = data.get("token_endpoint") + if not authorization_endpoint or not token_endpoint: + raise XAIOAuthError("xAI OAuth discovery missing endpoints") + return { + "authorization_endpoint": self._validate_xai_endpoint( + authorization_endpoint + ), + "token_endpoint": self._validate_xai_endpoint(token_endpoint), + } + + def _validate_xai_endpoint(self, url: str) -> str: + parsed = urlparse(url) + host = (parsed.hostname or "").lower() + if parsed.scheme != "https" or (host != "x.ai" and not host.endswith(".x.ai")): + raise XAIOAuthError( + f"xAI OAuth discovery returned unexpected endpoint: {url}" + ) + return url + + def _pkce_pair(self) -> Tuple[str, str]: + verifier = ( + base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode() + ) + challenge = ( + base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()) + .rstrip(b"=") + .decode() + ) + return verifier, challenge + + def _start_callback_server(self, state: str) -> Tuple[_CallbackServer, str]: + last_error: Optional[OSError] = None + for port in (XAI_OAUTH_REDIRECT_PORT, 0): + try: + server = _CallbackServer( + (XAI_OAUTH_REDIRECT_HOST, port), _CallbackHandler + ) + server.expected_state = state + server.callback_result = None + actual_port = server.server_address[1] + redirect_uri = f"http://{XAI_OAUTH_REDIRECT_HOST}:{actual_port}{XAI_OAUTH_REDIRECT_PATH}" + return server, redirect_uri + except OSError as exc: + last_error = exc + raise XAIOAuthError(f"Could not start xAI OAuth callback server: {last_error}") + + def _build_authorize_url( + self, + authorization_endpoint: str, + redirect_uri: str, + challenge: str, + state: str, + nonce: str, + ) -> str: + params = { + "response_type": "code", + "client_id": XAI_OAUTH_CLIENT_ID, + "redirect_uri": redirect_uri, + "scope": XAI_OAUTH_SCOPE, + "code_challenge": challenge, + "code_challenge_method": "S256", + "state": state, + "nonce": nonce, + } + return f"{authorization_endpoint}?{urlencode(params)}" + + def _wait_for_callback(self, server: _CallbackServer) -> Dict[str, Optional[str]]: + server.timeout = 1 + deadline = time.time() + XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS + try: + while time.time() < deadline: + server.handle_request() + if server.callback_result is not None: + return server.callback_result + finally: + server.server_close() + raise XAIOAuthError("Timed out waiting for xAI OAuth callback") + + def _exchange_token( + self, token_endpoint: str, data: Dict[str, str] + ) -> Dict[str, Any]: + try: + response = self._client().post( + token_endpoint, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + }, + data=data, + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + raise XAIOAuthError( + f"xAI OAuth token request failed: {exc.response.status_code} {exc.response.text}" + ) from exc + try: + body = response.json() + except ValueError as exc: + raise XAIOAuthError("xAI OAuth token response was not valid JSON") from exc + if not isinstance(body, dict): + raise XAIOAuthError("xAI OAuth token response was not an object") + return body + + def _build_auth_record( + self, + token_payload: Dict[str, Any], + token_endpoint: str, + fallback_refresh_token: Optional[str] = None, + ) -> Dict[str, Any]: + access_token = token_payload.get("access_token") + refresh_token = token_payload.get("refresh_token") or fallback_refresh_token + if not access_token: + raise XAIOAuthError("xAI OAuth token response missing access_token") + if not refresh_token: + raise XAIOAuthError("xAI OAuth token response missing refresh_token") + expires_in = token_payload.get("expires_in") or 3600 + try: + expires_at = int(time.time() + int(expires_in)) + except (TypeError, ValueError): + expires_at = int(time.time() + 3600) + return { + "access_token": access_token, + "refresh_token": refresh_token, + "id_token": token_payload.get("id_token"), + "token_type": token_payload.get("token_type") or "Bearer", + "token_endpoint": token_endpoint, + "expires_at": expires_at, + } + + def _refresh_tokens(self, auth_data: Dict[str, Any]) -> Dict[str, Any]: + token_endpoint = auth_data.get("token_endpoint") + if not token_endpoint: + token_endpoint = self._discover()["token_endpoint"] + token_endpoint = self._validate_xai_endpoint(token_endpoint) + refresh_token = auth_data.get("refresh_token") + if not refresh_token: + raise XAIOAuthLoginRequiredError( + "xAI OAuth refresh token missing. Run `litellm xai-oauth login`." + ) + + token_payload = self._exchange_token( + token_endpoint, + { + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": XAI_OAUTH_CLIENT_ID, + }, + ) + refreshed = self._build_auth_record( + token_payload, + token_endpoint, + fallback_refresh_token=refresh_token, + ) + self._write_auth_file(refreshed) + return refreshed + + +def should_use_xai_oauth(litellm_params: Optional[Dict[str, Any]]) -> bool: + return bool((litellm_params or {}).get("use_xai_oauth")) diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index 55805ddaede..f81e860a8ce 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import litellm from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError 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 @@ -220,10 +221,27 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): litellm_params.api_key, legacy_generic_before_env=True ) + if not api_key: + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + if should_use_xai_oauth(litellm_params.model_dump()): + try: + api_key = XAIOAuthAuthenticator().get_access_token() + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider=self.custom_llm_provider.value, + message=str(exc), + ) from exc + if not api_key: raise ValueError( "XAI API key is required. Set api_key, litellm.xai_key, " - "litellm.api_key, or XAI_API_KEY." + "litellm.api_key, XAI_API_KEY, or use_xai_oauth=True." ) headers.update( @@ -244,12 +262,20 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): Returns: str: The full URL for the XAI /responses endpoint """ - api_base = ( - api_base - or litellm.api_base - or get_secret_str("XAI_API_BASE") - or XAI_API_BASE + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + api_key = XAIModelInfo.get_api_key( + litellm_params.get("api_key"), legacy_generic_before_env=True ) + if should_use_xai_oauth(litellm_params) and not api_key: + api_base = XAIOAuthAuthenticator().get_api_base() + else: + api_base = ( + api_base + or litellm.api_base + or get_secret_str("XAI_API_BASE") + or XAI_API_BASE + ) # Remove trailing slashes api_base = api_base.rstrip("/") diff --git a/litellm/main.py b/litellm/main.py index 1a0d0312d73..8e18d8b46cb 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1638,6 +1638,7 @@ def completion( # type: ignore # noqa: PLR0915 litellm_request_debug=kwargs.get("litellm_request_debug", False), tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), + use_xai_oauth=kwargs.get("use_xai_oauth", False), ) cast(LiteLLMLoggingObj, logging).update_environment_variables( model=model, diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index e4567b9f494..ae831ef1b53 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -555,6 +555,7 @@ class ProxyInitializationHelpers: @click.command() +@click.argument("cli_args", nargs=-1) @click.option( "--host", default="0.0.0.0", help="Host for the server to listen on.", envvar="HOST" ) @@ -808,6 +809,7 @@ class ProxyInitializationHelpers: help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) def run_server( # noqa: PLR0915 + cli_args, host, port, api_base, @@ -854,6 +856,20 @@ def run_server( # noqa: PLR0915 use_v2_migration_resolver: bool, reload: bool, ): + if cli_args: + if cli_args == ("xai-oauth", "login"): + from litellm.llms.xai.oauth import XAIOAuthAuthenticator + + authenticator = XAIOAuthAuthenticator() + auth_data = authenticator.login() + click.echo( + f"xAI OAuth login successful. Credentials saved to {authenticator.auth_file}." + ) + if auth_data.get("expires_at"): + click.echo(f"Access token expires at {auth_data['expires_at']}.") + return + raise click.UsageError(f"Unknown command: {' '.join(cli_args)}") + if setup: from litellm.setup_wizard import run_setup_wizard diff --git a/litellm/types/router.py b/litellm/types/router.py index ef7eb05d087..ed858557a61 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -220,6 +220,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): use_in_pass_through: Optional[bool] = False use_litellm_proxy: Optional[bool] = False use_chat_completions_api: Optional[bool] = None + use_xai_oauth: Optional[bool] = Field( + default=False, + description="Use stored xAI OAuth credentials when no xAI API key is configured.", + ) model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True) merge_reasoning_content_in_choices: Optional[bool] = False model_info: Optional[Dict] = None diff --git a/litellm/utils.py b/litellm/utils.py index 8d9d0a409c6..7957399ba89 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5775,6 +5775,7 @@ def _get_model_info_helper( # noqa: PLR0915 ] split_model = potential_model_names["split_model"] custom_llm_provider = potential_model_names["custom_llm_provider"] + model_cost_custom_llm_provider = custom_llm_provider ######################### provider_config: Optional[BaseLLMModelInfo] = None if custom_llm_provider and custom_llm_provider in LlmProvidersSet: @@ -5840,7 +5841,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5849,7 +5851,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5858,7 +5861,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5867,7 +5871,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5876,7 +5881,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None @@ -5884,7 +5890,6 @@ def _get_model_info_helper( # noqa: PLR0915 raise ValueError( "This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) - _input_cost_per_token: Optional[float] = _model_info.get( "input_cost_per_token" ) diff --git a/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py b/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py new file mode 100644 index 00000000000..03790b220eb --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py @@ -0,0 +1,81 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm import LlmProviders +from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.litellm_core_utils.get_llm_provider_logic import ( + _get_openai_compatible_provider_info, +) +from litellm.llms.xai.chat.transformation import XAIChatConfig +from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ( + ProviderConfigManager, + get_optional_params, + validate_environment, +) + + +def test_xai_provider_config_routing(): + chat_config = ProviderConfigManager.get_provider_chat_config( + model="grok-3-mini", + provider=LlmProviders.XAI, + ) + responses_config = ProviderConfigManager.get_provider_responses_api_config( + model="grok-3-mini", + provider=LlmProviders.XAI, + ) + + assert isinstance(chat_config, XAIChatConfig) + assert isinstance(responses_config, XAIResponsesAPIConfig) + + +def test_xai_openai_compatible_provider_info(): + model, custom_llm_provider, dynamic_api_key, api_base = ( + _get_openai_compatible_provider_info( + model="xai/grok-3-mini", + api_base="https://api.x.ai/v1", + api_key="api-key", + dynamic_api_key=None, + ) + ) + + assert model == "grok-3-mini" + assert custom_llm_provider == "xai" + assert api_base == "https://api.x.ai/v1" + assert dynamic_api_key == "api-key" + + +def test_xai_get_model_info_uses_xai_pricing_metadata(): + model_info = litellm.get_model_info("xai/grok-3-mini") + + assert model_info["litellm_provider"] == "xai" + assert model_info["key"] == "xai/grok-3-mini" + assert model_info["mode"] == "chat" + + +def test_xai_validate_environment_reads_api_key(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "api-key") + + result = validate_environment(model="xai/grok-3-mini") + + assert result == {"keys_in_environment": True, "missing_keys": []} + + +def test_xai_oauth_flag_is_generic_litellm_param(): + litellm_params = GenericLiteLLMParams(use_xai_oauth=True) + runtime_params = get_litellm_params(use_xai_oauth=True) + result = get_optional_params( + model="grok-3-mini", + custom_llm_provider="xai", + temperature=0.2, + drop_params=True, + ) + + assert result["temperature"] == 0.2 + assert litellm_params.use_xai_oauth is True + assert runtime_params["use_xai_oauth"] is True + assert "use_xai_oauth" not in result diff --git a/tests/test_litellm/llms/xai/test_xai_oauth.py b/tests/test_litellm/llms/xai/test_xai_oauth.py new file mode 100644 index 00000000000..45fa6a405f2 --- /dev/null +++ b/tests/test_litellm/llms/xai/test_xai_oauth.py @@ -0,0 +1,801 @@ +import base64 +import hashlib +import json +import os +import threading +import time +from urllib.parse import parse_qs, urlparse +from unittest.mock import MagicMock + +import httpx +import litellm +import pytest +from click.testing import CliRunner + +import litellm.llms.xai.oauth as xai_oauth_module +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.llms.xai.oauth import ( + XAI_OAUTH_CLIENT_ID, + XAI_OAUTH_SCOPE, + XAIOAuthError, + XAIOAuthAuthenticator, + XAIOAuthLoginRequiredError, +) +from litellm.llms.xai.chat.transformation import XAIChatConfig +from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import get_optional_params, validate_environment + + +def _write_auth_file(tmp_path, payload): + token_dir = tmp_path / "xai_oauth" + token_dir.mkdir() + auth_file = token_dir / "auth.json" + auth_file.write_text(json.dumps(payload)) + return token_dir, auth_file + + +def test_get_access_token_uses_fresh_local_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "fresh-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + assert XAIOAuthAuthenticator().get_access_token() == "fresh-token" + + +def test_get_access_token_refreshes_and_preserves_refresh_token(tmp_path, monkeypatch): + token_dir, auth_file = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + def handler(request: httpx.Request) -> httpx.Response: + body = dict(item.split("=") for item in request.content.decode().split("&")) + assert body["grant_type"] == "refresh_token" + assert body["refresh_token"] == "refresh-token" + assert body["client_id"] == XAI_OAUTH_CLIENT_ID + return httpx.Response( + 200, + json={ + "access_token": "new-token", + "expires_in": 3600, + "token_type": "Bearer", + }, + ) + + client = httpx.Client(transport=httpx.MockTransport(handler)) + + assert XAIOAuthAuthenticator(http_client=client).get_access_token() == "new-token" + stored = json.loads(auth_file.read_text()) + assert stored["access_token"] == "new-token" + assert stored["refresh_token"] == "refresh-token" + + +def test_get_access_token_reuses_token_refreshed_by_parallel_request(): + expired_auth_data = { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + } + refreshed_auth_data = { + "access_token": "already-refreshed-token", + "refresh_token": "rotated-refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() + 3600, + } + authenticator = XAIOAuthAuthenticator() + authenticator._read_auth_file = MagicMock( + side_effect=[expired_auth_data, refreshed_auth_data] + ) + authenticator._refresh_tokens = MagicMock() + + assert authenticator.get_access_token() == "already-refreshed-token" + authenticator._refresh_tokens.assert_not_called() + + +def test_get_access_token_requires_login_without_auth_file(tmp_path, monkeypatch): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(XAIOAuthLoginRequiredError): + XAIOAuthAuthenticator().get_access_token() + + +def test_get_access_token_ignores_invalid_auth_file(tmp_path, monkeypatch): + token_dir = tmp_path / "xai_oauth" + token_dir.mkdir() + (token_dir / "auth.json").write_text("{not-json") + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + with pytest.raises(XAIOAuthLoginRequiredError): + XAIOAuthAuthenticator().get_access_token() + + +def test_refresh_failure_surfaces_oauth_error(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + client = httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(401, text="invalid_grant", request=request) + ) + ) + + with pytest.raises(XAIOAuthError) as exc_info: + XAIOAuthAuthenticator(http_client=client).get_access_token() + + assert "401 invalid_grant" in str(exc_info.value) + + +def test_build_auth_record_requires_access_and_refresh_tokens(): + authenticator = XAIOAuthAuthenticator() + + with pytest.raises(XAIOAuthError, match="access_token"): + authenticator._build_auth_record( + {"refresh_token": "refresh-token"}, + "https://auth.x.ai/oauth/token", + ) + + with pytest.raises(XAIOAuthError, match="refresh_token"): + authenticator._build_auth_record( + {"access_token": "access-token"}, + "https://auth.x.ai/oauth/token", + ) + + +def test_build_auth_record_defaults_expiry_and_token_type(): + authenticator = XAIOAuthAuthenticator() + + auth_data = authenticator._build_auth_record( + { + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_in": "not-a-number", + }, + "https://auth.x.ai/oauth/token", + ) + + assert auth_data["token_type"] == "Bearer" + assert auth_data["expires_at"] > time.time() + + +def test_is_expired_treats_missing_or_invalid_expiry_as_expired(): + authenticator = XAIOAuthAuthenticator() + + assert authenticator._is_expired({}) is True + assert authenticator._is_expired({"expires_at": "not-a-number"}) is True + + +def test_write_auth_file_creates_private_file(tmp_path, monkeypatch): + token_dir = tmp_path / "xai_oauth" + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + authenticator = XAIOAuthAuthenticator() + old_umask = os.umask(0o022) + replace_calls = [] + real_replace = os.replace + + def assert_private_temp_file(src, dst): + replace_calls.append((src, dst)) + assert oct(os.stat(src).st_mode & 0o777) == "0o600" + with open(src) as f: + assert json.load(f)["refresh_token"] == "refresh-token" + real_replace(src, dst) + + monkeypatch.setattr(os, "replace", assert_private_temp_file) + + try: + authenticator._write_auth_file( + { + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + } + ) + finally: + os.umask(old_umask) + + stored = json.loads((token_dir / "auth.json").read_text()) + assert stored["access_token"] == "access-token" + assert replace_calls + assert oct(os.stat(token_dir).st_mode & 0o777) == "0o700" + assert oct(os.stat(token_dir / "auth.json").st_mode & 0o777) == "0o600" + + +def test_discovery_rejects_unexpected_endpoint(): + authenticator = XAIOAuthAuthenticator() + + with pytest.raises(XAIOAuthError, match="unexpected endpoint"): + authenticator._validate_xai_endpoint("https://evil.example.com/oauth/token") + + with pytest.raises(XAIOAuthError, match="unexpected endpoint"): + authenticator._validate_xai_endpoint("http://auth.x.ai/oauth/token") + + +def test_discover_returns_validated_xai_endpoints(): + def handler(request: httpx.Request) -> httpx.Response: + assert request.url == "https://auth.x.ai/.well-known/openid-configuration" + return httpx.Response( + 200, + json={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + }, + ) + + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client(transport=httpx.MockTransport(handler)) + ) + + assert authenticator._discover() == { + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + + +def test_discover_requires_authorization_and_token_endpoints(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport(lambda request: httpx.Response(200, json={})) + ) + ) + + with pytest.raises(XAIOAuthError, match="missing endpoints"): + authenticator._discover() + + +def test_discover_wraps_http_errors(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response( + 500, text="discovery failed", request=request + ) + ) + ) + ) + + with pytest.raises(XAIOAuthError) as exc_info: + authenticator._discover() + + assert "xAI OAuth discovery request failed: 500 discovery failed" in str( + exc_info.value + ) + + +def test_discover_wraps_invalid_json_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, text="not-json") + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="discovery response was not valid JSON"): + authenticator._discover() + + +def test_refresh_discovers_token_endpoint_when_auth_file_is_legacy( + tmp_path, monkeypatch +): + token_dir, auth_file = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response( + 200, + json={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + }, + ) + return httpx.Response( + 200, + json={ + "access_token": "discovered-token", + "refresh_token": "new-refresh-token", + "expires_in": 3600, + }, + ) + + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client(transport=httpx.MockTransport(handler)) + ) + + assert authenticator.get_access_token() == "discovered-token" + stored = json.loads(auth_file.read_text()) + assert stored["token_endpoint"] == "https://auth.x.ai/oauth/token" + + +def test_exchange_token_rejects_non_object_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, json=["not", "an", "object"]) + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="was not an object"): + authenticator._exchange_token("https://auth.x.ai/oauth/token", {}) + + +def test_exchange_token_wraps_invalid_json_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, text="not-json") + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="token response was not valid JSON"): + authenticator._exchange_token("https://auth.x.ai/oauth/token", {}) + + +def test_start_callback_server_falls_back_to_ephemeral_port(monkeypatch): + calls = [] + real_server = xai_oauth_module._CallbackServer + + class FirstPortFailsCallbackServer(real_server): + def __init__(self, server_address, handler_class): + calls.append(server_address[1]) + if server_address[1] == xai_oauth_module.XAI_OAUTH_REDIRECT_PORT: + raise OSError("port unavailable") + super().__init__(server_address, handler_class) + + monkeypatch.setattr( + xai_oauth_module, "_CallbackServer", FirstPortFailsCallbackServer + ) + + server, redirect_uri = XAIOAuthAuthenticator()._start_callback_server("state-value") + try: + assert calls == [xai_oauth_module.XAI_OAUTH_REDIRECT_PORT, 0] + assert redirect_uri.startswith("http://127.0.0.1:") + assert redirect_uri.endswith("/callback") + finally: + server.server_close() + + +def test_wait_for_callback_times_out_and_closes_server(monkeypatch): + server, _ = XAIOAuthAuthenticator()._start_callback_server("state-value") + monkeypatch.setattr(xai_oauth_module, "XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS", 0) + + with pytest.raises(XAIOAuthError, match="Timed out"): + XAIOAuthAuthenticator()._wait_for_callback(server) + + +def test_callback_handler_records_success_and_rejects_state_mismatch(): + authenticator = XAIOAuthAuthenticator() + server, redirect_uri = authenticator._start_callback_server("expected-state") + thread = threading.Thread(target=server.handle_request) + thread.start() + response = httpx.get(f"{redirect_uri}?code=auth-code&state=expected-state") + thread.join(timeout=5) + + assert response.status_code == 200 + assert server.callback_result == { + "code": "auth-code", + "state": "expected-state", + "error": None, + "error_description": None, + } + + server, redirect_uri = authenticator._start_callback_server("expected-state") + thread = threading.Thread(target=server.handle_request) + thread.start() + response = httpx.get(f"{redirect_uri}?code=auth-code&state=wrong-state") + thread.join(timeout=5) + + assert response.status_code == 400 + assert server.callback_result["state"] == "wrong-state" + + +def test_login_exchanges_authorization_code_and_persists_auth_record(monkeypatch): + authenticator = XAIOAuthAuthenticator() + fake_server = MagicMock() + written_records = [] + + class FakeUUID: + def __init__(self, value): + self.hex = value + + monkeypatch.setattr( + xai_oauth_module.uuid, + "uuid4", + MagicMock(side_effect=[FakeUUID("state-value"), FakeUUID("nonce-value")]), + ) + authenticator._read_auth_file = MagicMock(return_value=None) + authenticator._discover = MagicMock( + return_value={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + ) + authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge")) + authenticator._start_callback_server = MagicMock( + return_value=(fake_server, "http://127.0.0.1:56121/callback") + ) + authenticator._wait_for_callback = MagicMock( + return_value={"state": "state-value", "code": "auth-code"} + ) + authenticator._exchange_token = MagicMock( + return_value={ + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_in": 3600, + } + ) + authenticator._write_auth_file = MagicMock(side_effect=written_records.append) + + auth_data = authenticator.login(no_browser=True) + + authenticator._exchange_token.assert_called_once_with( + "https://auth.x.ai/oauth/token", + { + "grant_type": "authorization_code", + "code": "auth-code", + "redirect_uri": "http://127.0.0.1:56121/callback", + "client_id": XAI_OAUTH_CLIENT_ID, + "code_verifier": "verifier", + }, + ) + assert auth_data["access_token"] == "access-token" + assert written_records == [auth_data] + + +def test_login_raises_on_callback_error_or_missing_code(monkeypatch): + authenticator = XAIOAuthAuthenticator() + + class FakeUUID: + hex = "state-value" + + monkeypatch.setattr( + xai_oauth_module.uuid, "uuid4", MagicMock(return_value=FakeUUID()) + ) + authenticator._read_auth_file = MagicMock(return_value=None) + authenticator._discover = MagicMock( + return_value={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + ) + authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge")) + authenticator._start_callback_server = MagicMock( + return_value=(MagicMock(), "http://127.0.0.1:56121/callback") + ) + authenticator._wait_for_callback = MagicMock( + return_value={ + "state": "state-value", + "error": "access_denied", + "error_description": "denied", + } + ) + + with pytest.raises(XAIOAuthError, match="denied"): + authenticator.login(no_browser=True) + + authenticator._wait_for_callback = MagicMock(return_value={"state": "state-value"}) + + with pytest.raises(XAIOAuthError, match="no code returned"): + authenticator.login(no_browser=True) + + +def test_pkce_pair_generates_s256_challenge(): + verifier, challenge = XAIOAuthAuthenticator()._pkce_pair() + expected = ( + base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()) + .rstrip(b"=") + .decode() + ) + + assert challenge == expected + assert "=" not in verifier + assert "=" not in challenge + + +def test_build_authorize_url_contains_xai_oauth_parameters(): + authorize_url = XAIOAuthAuthenticator()._build_authorize_url( + authorization_endpoint="https://auth.x.ai/oauth/authorize", + redirect_uri="http://127.0.0.1:56121/callback", + challenge="pkce-challenge", + state="state-value", + nonce="nonce-value", + ) + parsed = urlparse(authorize_url) + params = parse_qs(parsed.query) + + assert parsed.scheme == "https" + assert parsed.netloc == "auth.x.ai" + assert params["response_type"] == ["code"] + assert params["client_id"] == [XAI_OAUTH_CLIENT_ID] + assert params["scope"] == [XAI_OAUTH_SCOPE] + assert params["code_challenge"] == ["pkce-challenge"] + assert params["code_challenge_method"] == ["S256"] + assert params["state"] == ["state-value"] + assert params["nonce"] == ["nonce-value"] + + +def test_get_llm_provider_uses_single_xai_provider(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "api-key") + + model, provider, api_key, api_base = get_llm_provider("xai/grok-4") + + assert model == "grok-4" + assert provider == "xai" + assert api_key == "api-key" + assert api_base == "https://api.x.ai/v1" + + +def test_xai_oauth_alias_is_not_a_provider(): + with pytest.raises(Exception): + get_llm_provider("xai_oauth/grok-4") + + +def test_chat_config_wraps_flagged_oauth_errors_as_authentication_error( + tmp_path, monkeypatch +): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + + assert exc_info.value.llm_provider == "xai" + assert "litellm xai-oauth login" in str(exc_info.value) + + +def test_chat_config_injects_flagged_oauth_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "chat-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + + assert headers["Authorization"] == "Bearer chat-token" + + +def test_chat_config_ignores_api_base_override_for_flagged_oauth(monkeypatch): + monkeypatch.setenv("XAI_OAUTH_API_BASE", "https://api.x.ai/v1") + + url = XAIChatConfig().get_complete_url( + api_base="https://attacker.example.com/v1", + api_key=None, + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert url == "https://api.x.ai/v1/chat/completions" + + +def test_chat_config_treats_blank_api_key_as_absent_for_flagged_oauth( + tmp_path, monkeypatch +): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "stored-oauth-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key="", + ) + + assert headers["Authorization"] == "Bearer stored-oauth-token" + + +def test_chat_config_allows_api_base_override_with_caller_api_key(): + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key="caller-api-key", + ) + url = XAIChatConfig().get_complete_url( + api_base="https://custom.example.com/v1", + api_key="caller-api-key", + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert headers["Authorization"] == "Bearer caller-api-key" + assert url == "https://custom.example.com/v1/chat/completions" + + +def test_chat_config_prioritizes_env_api_key_over_oauth_flag(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "env-api-key") + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + url = XAIChatConfig().get_complete_url( + api_base="https://custom.example.com/v1", + api_key=None, + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert headers["Authorization"] == "Bearer env-api-key" + assert url == "https://custom.example.com/v1/chat/completions" + + +def test_validate_environment_still_reports_xai_api_key(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "env-api-key") + + assert validate_environment("xai/grok-4") == { + "keys_in_environment": True, + "missing_keys": [], + } + + +def test_xai_oauth_flag_uses_xai_optional_param_mapping(): + litellm_params = GenericLiteLLMParams(use_xai_oauth=True) + optional_params = get_optional_params( + model="grok-4", + custom_llm_provider="xai", + temperature=0.2, + max_tokens=8, + ) + + assert optional_params["temperature"] == 0.2 + assert optional_params["max_tokens"] == 8 + assert litellm_params.use_xai_oauth is True + assert "use_xai_oauth" not in optional_params + + +def test_responses_config_injects_flagged_oauth_bearer_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "responses-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIResponsesAPIConfig().validate_environment( + headers={}, + model="grok-4", + litellm_params=GenericLiteLLMParams(use_xai_oauth=True), + ) + + assert headers["Authorization"] == "Bearer responses-token" + + +def test_responses_config_endpoint_url_uses_oauth_authenticator(monkeypatch): + monkeypatch.setenv("XAI_OAUTH_API_BASE", "https://xai.example.com/v1/") + config = XAIResponsesAPIConfig() + + assert config.get_complete_url( + api_base=None, litellm_params={"use_xai_oauth": True} + ) == ("https://xai.example.com/v1/responses") + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"use_xai_oauth": True}, + ) + == "https://xai.example.com/v1/responses" + ) + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"api_key": "", "use_xai_oauth": True}, + ) + == "https://xai.example.com/v1/responses" + ) + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"api_key": "caller-api-key"}, + ) + == "https://custom.example.com/v1/responses" + ) + + +def test_responses_config_wraps_flagged_oauth_errors_as_authentication_error( + tmp_path, monkeypatch +): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + XAIResponsesAPIConfig().validate_environment( + headers={}, + model="grok-4", + litellm_params=GenericLiteLLMParams(use_xai_oauth=True), + ) + + assert XAIResponsesAPIConfig().custom_llm_provider.value == "xai" + assert exc_info.value.llm_provider == "xai" + + +def test_proxy_cli_xai_oauth_login_uses_single_authenticator(monkeypatch): + from litellm.proxy.proxy_cli import run_server + + instances = [] + + class FakeAuthenticator: + auth_file = "/tmp/xai-oauth-auth.json" + + def __init__(self): + instances.append(self) + + def login(self): + return {"expires_at": 1234567890} + + monkeypatch.setattr( + "litellm.llms.xai.oauth.XAIOAuthAuthenticator", FakeAuthenticator + ) + + result = CliRunner().invoke(run_server, ["xai-oauth", "login"]) + + assert result.exit_code == 0 + assert len(instances) == 1 + assert "Credentials saved to /tmp/xai-oauth-auth.json" in result.output + assert "Access token expires at 1234567890" in result.output diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8379d0536a6..203a56f615b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -25081,6 +25081,12 @@ export interface components { * @default false */ use_litellm_proxy: boolean | null; + /** + * Use Xai Oauth + * @description Use stored xAI OAuth credentials when no xAI API key is configured. + * @default false + */ + use_xai_oauth: boolean | null; /** Vector Store Id */ vector_store_id?: string | null; /** Vertex Credentials */ @@ -32679,6 +32685,12 @@ export interface components { * @default false */ use_litellm_proxy: boolean | null; + /** + * Use Xai Oauth + * @description Use stored xAI OAuth credentials when no xAI API key is configured. + * @default false + */ + use_xai_oauth: boolean | null; /** Vector Store Id */ vector_store_id?: string | null; /** Vertex Credentials */