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 */