feat(sdk): add xAI OAuth provider (#29866)

* Add xAI OAuth provider

* Update oauth.py

Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>

* Fix xAI OAuth CI failures

* Add xAI OAuth coverage tests

* Move xAI OAuth coverage tests to core utils

* Address xAI OAuth review comments

* Prevent xAI OAuth api_base token exfiltration

* Treat blank xAI OAuth api keys as absent

* Wrap invalid xAI OAuth JSON responses

* Use xAI OAuth behind explicit flag

---------

Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
This commit is contained in:
Jeremy Chapeau 2026-06-10 02:54:25 -07:00 • committed by Sameer Kankute
parent 55ee1e9825
commit d30cfca382
No known key found for this signature in database
11 changed files with 1447 additions and 12 deletions

View file

@ -34,6 +34,7 @@ _OPTIONAL_KWARGS_KEYS = frozenset(
"aws_bedrock_runtime_endpoint",
"tpm",
"rpm",
"use_xai_oauth",
}
)

View file

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

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

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

View file

@ -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("/")

View file

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

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

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

View file

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

View file

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

View file

@ -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="<html>not-json</html>")
)
)
)
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="<html>not-json</html>")
)
)
)
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

View file

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