Add xAI OAuth provider

This commit is contained in:
Jeremy Chapeau 2026-06-06 15:52:19 -07:00 • committed by Jeremy Chapeau
parent aaf1e2444b
commit 634f448052
No known key found for this signature in database
9 changed files with 1216 additions and 7 deletions

View file

@ -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 (

View file

@ -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,

440
litellm/llms/xai/oauth.py Normal file
View file

@ -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"<html><body><h1>xAI authorization state mismatch.</h1></body></html>"
)
return
self.send_response(200)
self.send_header("Content-Type", "text/html; charset=utf-8")
self.end_headers()
body = (
b"<html><body><h1>xAI authorization failed.</h1>You can close this tab.</body></html>"
if result["error"]
else b"<html><body><h1>xAI authorization received.</h1>You can close this tab.</body></html>"
)
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"

View file

@ -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(

View file

@ -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

View file

@ -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"

View file

@ -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:

View file

@ -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",

View file

@ -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