mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
55ee1e9825
commit
d30cfca382
11 changed files with 1447 additions and 12 deletions
|
|
@ -34,6 +34,7 @@ _OPTIONAL_KWARGS_KEYS = frozenset(
|
|||
"aws_bedrock_runtime_endpoint",
|
||||
"tpm",
|
||||
"rpm",
|
||||
"use_xai_oauth",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
421
litellm/llms/xai/oauth.py
Normal 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"))
|
||||
|
|
@ -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("/")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
801
tests/test_litellm/llms/xai/test_xai_oauth.py
Normal file
801
tests/test_litellm/llms/xai/test_xai_oauth.py
Normal 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
|
||||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue