diff --git a/litellm/__init__.py b/litellm/__init__.py
index e6c30e12286..9875e9054be 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -605,6 +605,7 @@ perplexity_models: Set = set()
watsonx_models: Set = set()
gemini_models: Set = set()
xai_models: Set = set()
+xai_oauth_models: Set = set()
zai_models: Set = set()
deepseek_models: Set = set()
runwayml_models: Set = set()
@@ -817,6 +818,7 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
text_completion_inception_models.add(key)
elif value.get("litellm_provider") == "xai":
xai_models.add(key)
+ xai_oauth_models.add(key.replace("xai/", "xai_oauth/", 1))
elif value.get("litellm_provider") == "zai":
zai_models.add(key)
elif value.get("litellm_provider") == "fal_ai":
@@ -1009,6 +1011,7 @@ model_list = list(
| text_completion_codestral_models
| text_completion_inception_models
| xai_models
+ | xai_oauth_models
| zai_models
| fal_ai_models
| deepseek_models
@@ -1106,6 +1109,7 @@ models_by_provider: dict = {
"text-completion-codestral": text_completion_codestral_models,
"text-completion-inception": text_completion_inception_models,
"xai": xai_models,
+ "xai_oauth": xai_oauth_models,
"zai": zai_models,
"fal_ai": fal_ai_models,
"deepseek": deepseek_models,
@@ -1898,6 +1902,10 @@ if TYPE_CHECKING:
JinaAIEmbeddingConfig as JinaAIEmbeddingConfig,
)
from .llms.xai.chat.transformation import XAIChatConfig as XAIChatConfig
+ from .llms.xai.oauth import XAIOAuthChatConfig as XAIOAuthChatConfig
+ from .llms.xai.oauth import (
+ XAIOAuthResponsesAPIConfig as XAIOAuthResponsesAPIConfig,
+ )
from .llms.zai.chat.transformation import ZAIChatConfig as ZAIChatConfig
from .llms.aiml.chat.transformation import AIMLChatConfig as AIMLChatConfig
from .llms.volcengine.chat.transformation import (
diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py
index de65ed93312..244a825b5b2 100644
--- a/litellm/litellm_core_utils/get_llm_provider_logic.py
+++ b/litellm/litellm_core_utils/get_llm_provider_logic.py
@@ -824,6 +824,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
) = litellm.XAIChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
)
+ elif custom_llm_provider == "xai_oauth":
+ from litellm.llms.xai.oauth import XAIOAuthChatConfig
+
+ (
+ api_base,
+ dynamic_api_key,
+ ) = XAIOAuthChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
elif custom_llm_provider == "zai":
(
api_base,
diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py
new file mode 100644
index 00000000000..2773733444b
--- /dev/null
+++ b/litellm/llms/xai/oauth.py
@@ -0,0 +1,440 @@
+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
+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.exceptions import AuthenticationError
+from litellm.llms.custom_httpx.http_handler import _get_httpx_client
+from litellm.llms.xai.chat.transformation import XAIChatConfig
+from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
+from litellm.secret_managers.main import get_secret_str
+from litellm.types.router import GenericLiteLLMParams
+from litellm.types.utils import LlmProviders
+
+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[httpx.Client] = 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) -> httpx.Client:
+ return self.http_client or _get_httpx_client()
+
+ def _ensure_token_dir(self) -> None:
+ os.makedirs(self.token_dir, exist_ok=True)
+
+ 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()
+ with open(self.auth_file, "w") as f:
+ json.dump(data, f)
+ try:
+ os.chmod(self.auth_file, 0o600)
+ except OSError:
+ verbose_logger.debug("Could not chmod xAI OAuth auth file")
+
+ 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]:
+ response = self._client().get(
+ XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"}
+ )
+ response.raise_for_status()
+ data = response.json()
+ 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]:
+ response = self._client().post(
+ token_endpoint,
+ headers={
+ "Accept": "application/json",
+ "Content-Type": "application/x-www-form-urlencoded",
+ },
+ data=data,
+ )
+ try:
+ 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
+ body = response.json()
+ 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
+
+
+class XAIOAuthChatConfig(XAIChatConfig):
+ def __init__(self) -> None:
+ super().__init__()
+ self.authenticator = XAIOAuthAuthenticator()
+
+ @property
+ def custom_llm_provider(self) -> Optional[str]:
+ return "xai_oauth"
+
+ def _get_openai_compatible_provider_info(
+ self, api_base: Optional[str], api_key: Optional[str]
+ ) -> Tuple[Optional[str], Optional[str]]:
+ try:
+ dynamic_api_key = api_key or self.authenticator.get_access_token()
+ except XAIOAuthError as exc:
+ raise AuthenticationError(
+ model="",
+ llm_provider=self.custom_llm_provider or "xai_oauth",
+ message=str(exc),
+ ) from exc
+ return (
+ api_base or self.authenticator.get_api_base(),
+ dynamic_api_key,
+ )
+
+
+class XAIOAuthResponsesAPIConfig(XAIResponsesAPIConfig):
+ def __init__(self) -> None:
+ super().__init__()
+ self.authenticator = XAIOAuthAuthenticator()
+
+ @property
+ def custom_llm_provider(self) -> LlmProviders:
+ return LlmProviders.XAI_OAUTH
+
+ def validate_environment(
+ self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
+ ) -> dict:
+ litellm_params = litellm_params or GenericLiteLLMParams()
+ try:
+ api_key = litellm_params.api_key or self.authenticator.get_access_token()
+ except XAIOAuthError as exc:
+ raise AuthenticationError(
+ model=model,
+ llm_provider=self.custom_llm_provider.value,
+ message=str(exc),
+ ) from exc
+ headers.update({"Authorization": f"Bearer {api_key}"})
+ return headers
+
+ def get_complete_url(
+ self,
+ api_base: Optional[str],
+ litellm_params: dict,
+ ) -> str:
+ api_base = api_base or self.authenticator.get_api_base()
+ return f"{api_base.rstrip('/')}/responses"
diff --git a/litellm/main.py b/litellm/main.py
index 64891e2def9..3539ce6d03d 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -2286,7 +2286,7 @@ def completion( # type: ignore # noqa: PLR0915
additional_args={"headers": headers},
)
raise e
- elif custom_llm_provider == "xai":
+ elif custom_llm_provider == "xai" or custom_llm_provider == "xai_oauth":
## COMPLETION CALL
try:
response = base_llm_http_handler.completion(
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/utils.py b/litellm/types/utils.py
index 9633cecf96c..18d930f8663 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -3261,6 +3261,7 @@ class LlmProviders(str, Enum):
OPENAI_LIKE = "openai_like" # embedding only
JINA_AI = "jina_ai"
XAI = "xai"
+ XAI_OAUTH = "xai_oauth"
ZAI = "zai"
CUSTOM_OPENAI = "custom_openai"
TEXT_COMPLETION_OPENAI = "text-completion-openai"
diff --git a/litellm/utils.py b/litellm/utils.py
index 7312e71bbd1..7eb6e3479f6 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -4594,7 +4594,7 @@ def get_optional_params( # noqa: PLR0915
else False
),
)
- elif custom_llm_provider == "xai":
+ elif custom_llm_provider == "xai" or custom_llm_provider == "xai_oauth":
optional_params = litellm.XAIChatConfig().map_openai_params(
model=model,
non_default_params=non_default_params,
@@ -5775,6 +5775,18 @@ def _get_model_info_helper( # noqa: PLR0915
]
split_model = potential_model_names["split_model"]
custom_llm_provider = potential_model_names["custom_llm_provider"]
+ xai_oauth_model_alias = custom_llm_provider == "xai_oauth" or model.startswith(
+ "xai_oauth/"
+ )
+ model_cost_custom_llm_provider = custom_llm_provider
+ if xai_oauth_model_alias:
+ model_cost_custom_llm_provider = "xai"
+ model = model.replace("xai_oauth/", "xai/", 1)
+ combined_model_name = combined_model_name.replace("xai_oauth/", "xai/", 1)
+ stripped_model_name = stripped_model_name.replace("xai_oauth/", "xai/", 1)
+ combined_stripped_model_name = combined_stripped_model_name.replace(
+ "xai_oauth/", "xai/", 1
+ )
#########################
provider_config: Optional[BaseLLMModelInfo] = None
if custom_llm_provider and custom_llm_provider in LlmProvidersSet:
@@ -5840,7 +5852,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 +5862,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 +5872,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 +5882,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 +5892,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,6 +5901,10 @@ 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"
)
+ if xai_oauth_model_alias:
+ _model_info = dict(_model_info)
+ _model_info["litellm_provider"] = "xai_oauth"
+ key = key.replace("xai/", "xai_oauth/", 1)
_input_cost_per_token: Optional[float] = _model_info.get(
"input_cost_per_token"
@@ -6634,6 +6655,13 @@ def validate_environment( # noqa: PLR0915
keys_in_environment = True
else:
missing_keys.append("XAI_API_KEY")
+ elif custom_llm_provider == "xai_oauth":
+ from litellm.llms.xai.oauth import XAIOAuthAuthenticator
+
+ if XAIOAuthAuthenticator()._read_auth_file() is not None:
+ keys_in_environment = True
+ else:
+ missing_keys.append("XAI OAuth credentials")
elif custom_llm_provider == "ai21_chat":
if "AI21_API_KEY" in os.environ:
keys_in_environment = True
@@ -8312,6 +8340,10 @@ class ProviderConfigManager:
LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False),
LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False),
LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False),
+ LlmProviders.XAI_OAUTH: (
+ lambda: ProviderConfigManager._get_xai_oauth_config(),
+ False,
+ ),
LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False),
LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False),
LlmProviders.INCEPTION: (lambda: litellm.InceptionChatConfig(), False),
@@ -8504,6 +8536,12 @@ class ProviderConfigManager:
return LangFlowConfig()
+ @staticmethod
+ def _get_xai_oauth_config() -> BaseConfig:
+ from litellm.llms.xai.oauth import XAIOAuthChatConfig
+
+ return XAIOAuthChatConfig()
+
@staticmethod
def get_provider_chat_config( # noqa: PLR0915
model: str,
@@ -8894,6 +8932,10 @@ class ProviderConfigManager:
return litellm.AzureOpenAIResponsesAPIConfig()
elif litellm.LlmProviders.XAI == provider:
return litellm.XAIResponsesAPIConfig()
+ elif litellm.LlmProviders.XAI_OAUTH == provider:
+ from litellm.llms.xai.oauth import XAIOAuthResponsesAPIConfig
+
+ return XAIOAuthResponsesAPIConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
return litellm.GithubCopilotResponsesAPIConfig()
elif litellm.LlmProviders.CHATGPT == provider:
diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json
index a1ad20fffd1..139e4f40bb9 100644
--- a/provider_endpoints_support.json
+++ b/provider_endpoints_support.json
@@ -2396,6 +2396,25 @@
"realtime": true
}
},
+ "xai_oauth": {
+ "display_name": "xAI OAuth (`xai_oauth`)",
+ "url": "https://docs.litellm.ai/docs/providers/xai_oauth",
+ "endpoints": {
+ "chat_completions": true,
+ "messages": true,
+ "responses": true,
+ "embeddings": false,
+ "image_generations": false,
+ "audio_transcriptions": false,
+ "audio_speech": false,
+ "moderations": false,
+ "batches": false,
+ "rerank": false,
+ "a2a": true,
+ "interactions": true,
+ "realtime": false
+ }
+ },
"xinference": {
"display_name": "Xinference (`xinference`)",
"url": "https://docs.litellm.ai/docs/providers/xinference",
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..3f394244e5d
--- /dev/null
+++ b/tests/test_litellm/llms/xai/test_xai_oauth.py
@@ -0,0 +1,676 @@
+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,
+ XAIOAuthChatConfig,
+ XAIOAuthLoginRequiredError,
+ XAIOAuthResponsesAPIConfig,
+)
+from litellm.types.router import GenericLiteLLMParams
+from litellm.types.utils import LlmProviders
+from litellm.utils import (
+ ProviderConfigManager,
+ 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()
+
+ authenticator._write_auth_file(
+ {
+ "access_token": "access-token",
+ "refresh_token": "refresh-token",
+ "expires_at": time.time() + 3600,
+ }
+ )
+
+ stored = json.loads((token_dir / "auth.json").read_text())
+ assert stored["access_token"] == "access-token"
+ 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_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_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_resolves_xai_oauth(tmp_path, monkeypatch):
+ token_dir, _ = _write_auth_file(
+ tmp_path,
+ {
+ "access_token": "provider-token",
+ "refresh_token": "refresh-token",
+ "expires_at": time.time() + 3600,
+ },
+ )
+ monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
+
+ model, provider, api_key, api_base = get_llm_provider("xai_oauth/grok-4")
+
+ assert model == "grok-4"
+ assert provider == "xai_oauth"
+ assert api_key == "provider-token"
+ assert api_base == "https://api.x.ai/v1"
+
+
+def test_chat_config_wraps_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:
+ XAIOAuthChatConfig()._get_openai_compatible_provider_info(None, None)
+
+ assert exc_info.value.llm_provider == "xai_oauth"
+ assert "litellm xai-oauth login" in str(exc_info.value)
+
+
+def test_chat_config_uses_attached_authenticator():
+ config = XAIOAuthChatConfig()
+ config.authenticator = MagicMock()
+ config.authenticator.get_access_token.return_value = "chat-token"
+ config.authenticator.get_api_base.return_value = "https://xai.example.com/v1"
+
+ api_base, api_key = config._get_openai_compatible_provider_info(None, None)
+
+ assert api_base == "https://xai.example.com/v1"
+ assert api_key == "chat-token"
+
+
+def test_get_model_info_reuses_xai_metadata_for_xai_oauth():
+ xai_info = litellm.get_model_info("xai/grok-4")
+ oauth_info = litellm.get_model_info("xai_oauth/grok-4")
+
+ assert oauth_info["litellm_provider"] == "xai_oauth"
+ assert oauth_info["max_input_tokens"] == xai_info["max_input_tokens"]
+ assert oauth_info["input_cost_per_token"] == xai_info["input_cost_per_token"]
+
+
+def test_validate_environment_reads_xai_oauth_auth_file(tmp_path, monkeypatch):
+ token_dir, _ = _write_auth_file(
+ tmp_path,
+ {
+ "access_token": "env-token",
+ "refresh_token": "refresh-token",
+ "expires_at": time.time() + 3600,
+ },
+ )
+ monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
+ monkeypatch.setattr(
+ "litellm.utils.get_llm_provider",
+ lambda model: ("grok-4", "xai_oauth", None, None),
+ )
+
+ assert validate_environment("xai_oauth/grok-4") == {
+ "keys_in_environment": True,
+ "missing_keys": [],
+ }
+
+
+def test_validate_environment_reports_missing_xai_oauth_credentials(
+ tmp_path, monkeypatch
+):
+ monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing"))
+ monkeypatch.setattr(
+ "litellm.utils.get_llm_provider",
+ lambda model: ("grok-4", "xai_oauth", None, None),
+ )
+
+ assert validate_environment("xai_oauth/grok-4") == {
+ "keys_in_environment": False,
+ "missing_keys": ["XAI OAuth credentials"],
+ }
+
+
+def test_xai_oauth_uses_xai_optional_param_mapping():
+ optional_params = get_optional_params(
+ model="grok-4",
+ custom_llm_provider="xai_oauth",
+ temperature=0.2,
+ max_tokens=8,
+ )
+
+ assert optional_params["temperature"] == 0.2
+ assert optional_params["max_tokens"] == 8
+
+
+def test_provider_config_manager_returns_xai_oauth_configs():
+ chat_config = ProviderConfigManager.get_provider_chat_config(
+ model="grok-4",
+ provider=LlmProviders.XAI_OAUTH,
+ )
+ responses_config = ProviderConfigManager.get_provider_responses_api_config(
+ provider=LlmProviders.XAI_OAUTH
+ )
+
+ assert isinstance(chat_config, XAIOAuthChatConfig)
+ assert isinstance(responses_config, XAIOAuthResponsesAPIConfig)
+
+
+def test_responses_config_injects_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 = XAIOAuthResponsesAPIConfig().validate_environment(
+ headers={},
+ model="grok-4",
+ litellm_params=GenericLiteLLMParams(),
+ )
+
+ assert headers["Authorization"] == "Bearer responses-token"
+
+
+def test_responses_config_endpoint_url_uses_oauth_authenticator():
+ config = XAIOAuthResponsesAPIConfig()
+ config.authenticator = MagicMock()
+ config.authenticator.get_api_base.return_value = "https://xai.example.com/v1/"
+
+ assert config.get_complete_url(api_base=None, litellm_params={}) == (
+ "https://xai.example.com/v1/responses"
+ )
+ assert (
+ config.get_complete_url(
+ api_base="https://custom.example.com/v1/", litellm_params={}
+ )
+ == "https://custom.example.com/v1/responses"
+ )
+
+
+def test_responses_config_wraps_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:
+ XAIOAuthResponsesAPIConfig().validate_environment(
+ headers={},
+ model="grok-4",
+ litellm_params=GenericLiteLLMParams(),
+ )
+
+ assert XAIOAuthResponsesAPIConfig().custom_llm_provider == LlmProviders.XAI_OAUTH
+ assert exc_info.value.llm_provider == "xai_oauth"
+
+
+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