mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Add xAI OAuth provider
This commit is contained in:
parent
aaf1e2444b
commit
634f448052
9 changed files with 1216 additions and 7 deletions
|
|
@ -605,6 +605,7 @@ perplexity_models: Set = set()
|
|||
watsonx_models: Set = set()
|
||||
gemini_models: Set = set()
|
||||
xai_models: Set = set()
|
||||
xai_oauth_models: Set = set()
|
||||
zai_models: Set = set()
|
||||
deepseek_models: Set = set()
|
||||
runwayml_models: Set = set()
|
||||
|
|
@ -817,6 +818,7 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
text_completion_inception_models.add(key)
|
||||
elif value.get("litellm_provider") == "xai":
|
||||
xai_models.add(key)
|
||||
xai_oauth_models.add(key.replace("xai/", "xai_oauth/", 1))
|
||||
elif value.get("litellm_provider") == "zai":
|
||||
zai_models.add(key)
|
||||
elif value.get("litellm_provider") == "fal_ai":
|
||||
|
|
@ -1009,6 +1011,7 @@ model_list = list(
|
|||
| text_completion_codestral_models
|
||||
| text_completion_inception_models
|
||||
| xai_models
|
||||
| xai_oauth_models
|
||||
| zai_models
|
||||
| fal_ai_models
|
||||
| deepseek_models
|
||||
|
|
@ -1106,6 +1109,7 @@ models_by_provider: dict = {
|
|||
"text-completion-codestral": text_completion_codestral_models,
|
||||
"text-completion-inception": text_completion_inception_models,
|
||||
"xai": xai_models,
|
||||
"xai_oauth": xai_oauth_models,
|
||||
"zai": zai_models,
|
||||
"fal_ai": fal_ai_models,
|
||||
"deepseek": deepseek_models,
|
||||
|
|
@ -1898,6 +1902,10 @@ if TYPE_CHECKING:
|
|||
JinaAIEmbeddingConfig as JinaAIEmbeddingConfig,
|
||||
)
|
||||
from .llms.xai.chat.transformation import XAIChatConfig as XAIChatConfig
|
||||
from .llms.xai.oauth import XAIOAuthChatConfig as XAIOAuthChatConfig
|
||||
from .llms.xai.oauth import (
|
||||
XAIOAuthResponsesAPIConfig as XAIOAuthResponsesAPIConfig,
|
||||
)
|
||||
from .llms.zai.chat.transformation import ZAIChatConfig as ZAIChatConfig
|
||||
from .llms.aiml.chat.transformation import AIMLChatConfig as AIMLChatConfig
|
||||
from .llms.volcengine.chat.transformation import (
|
||||
|
|
|
|||
|
|
@ -824,6 +824,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
) = litellm.XAIChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
elif custom_llm_provider == "xai_oauth":
|
||||
from litellm.llms.xai.oauth import XAIOAuthChatConfig
|
||||
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = XAIOAuthChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "zai":
|
||||
(
|
||||
api_base,
|
||||
|
|
|
|||
440
litellm/llms/xai/oauth.py
Normal file
440
litellm/llms/xai/oauth.py
Normal file
|
|
@ -0,0 +1,440 @@
|
|||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import webbrowser
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from urllib.parse import parse_qs, urlencode, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
XAI_OAUTH_ISSUER = "https://auth.x.ai"
|
||||
XAI_OAUTH_DISCOVERY_URL = f"{XAI_OAUTH_ISSUER}/.well-known/openid-configuration"
|
||||
XAI_OAUTH_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828"
|
||||
XAI_OAUTH_SCOPE = "openid profile email offline_access grok-cli:access api:access"
|
||||
XAI_OAUTH_REDIRECT_HOST = "127.0.0.1"
|
||||
XAI_OAUTH_REDIRECT_PORT = 56121
|
||||
XAI_OAUTH_REDIRECT_PATH = "/callback"
|
||||
XAI_OAUTH_EXPIRY_SKEW_SECONDS = 120
|
||||
XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS = 180
|
||||
_XAI_OAUTH_REFRESH_LOCK = threading.Lock()
|
||||
|
||||
|
||||
class XAIOAuthError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class XAIOAuthLoginRequiredError(XAIOAuthError):
|
||||
pass
|
||||
|
||||
|
||||
class _CallbackHandler(BaseHTTPRequestHandler):
|
||||
server: "_CallbackServer"
|
||||
|
||||
def do_GET(self) -> None:
|
||||
parsed = urlparse(self.path)
|
||||
if parsed.path != XAI_OAUTH_REDIRECT_PATH:
|
||||
self.send_response(404)
|
||||
self.end_headers()
|
||||
return
|
||||
|
||||
params = parse_qs(parsed.query)
|
||||
result = {
|
||||
"code": params.get("code", [None])[0],
|
||||
"state": params.get("state", [None])[0],
|
||||
"error": params.get("error", [None])[0],
|
||||
"error_description": params.get("error_description", [None])[0],
|
||||
}
|
||||
self.server.callback_result = result
|
||||
|
||||
if result["state"] != self.server.expected_state:
|
||||
self.send_response(400)
|
||||
self.send_header("Content-Type", "text/html; charset=utf-8")
|
||||
self.end_headers()
|
||||
self.wfile.write(
|
||||
b"<html><body><h1>xAI authorization state mismatch.</h1></body></html>"
|
||||
)
|
||||
return
|
||||
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/html; charset=utf-8")
|
||||
self.end_headers()
|
||||
body = (
|
||||
b"<html><body><h1>xAI authorization failed.</h1>You can close this tab.</body></html>"
|
||||
if result["error"]
|
||||
else b"<html><body><h1>xAI authorization received.</h1>You can close this tab.</body></html>"
|
||||
)
|
||||
self.wfile.write(body)
|
||||
|
||||
def log_message(self, format: str, *args: Any) -> None:
|
||||
return
|
||||
|
||||
|
||||
class _CallbackServer(HTTPServer):
|
||||
expected_state: str
|
||||
callback_result: Optional[Dict[str, Optional[str]]]
|
||||
|
||||
|
||||
class XAIOAuthAuthenticator:
|
||||
def __init__(self, http_client: Optional[httpx.Client] = None) -> None:
|
||||
self.token_dir = get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser(
|
||||
"~/.config/litellm/xai_oauth"
|
||||
)
|
||||
self.auth_file = os.path.join(
|
||||
self.token_dir, get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json"
|
||||
)
|
||||
self.http_client = http_client
|
||||
|
||||
def get_api_base(self) -> str:
|
||||
return (
|
||||
get_secret_str("XAI_OAUTH_API_BASE")
|
||||
or get_secret_str("XAI_API_BASE")
|
||||
or XAI_API_BASE
|
||||
)
|
||||
|
||||
def get_access_token(self) -> str:
|
||||
auth_data = self._read_auth_file()
|
||||
if not auth_data:
|
||||
raise XAIOAuthLoginRequiredError(
|
||||
"xAI OAuth login required. Run `litellm xai-oauth login`."
|
||||
)
|
||||
|
||||
access_token = auth_data.get("access_token")
|
||||
if access_token and not self._is_expired(auth_data):
|
||||
return access_token
|
||||
|
||||
refresh_token = auth_data.get("refresh_token")
|
||||
if not refresh_token:
|
||||
raise XAIOAuthLoginRequiredError(
|
||||
"xAI OAuth refresh token missing. Run `litellm xai-oauth login`."
|
||||
)
|
||||
|
||||
with _XAI_OAUTH_REFRESH_LOCK:
|
||||
locked_auth_data = self._read_auth_file() or auth_data
|
||||
access_token = locked_auth_data.get("access_token")
|
||||
if access_token and not self._is_expired(locked_auth_data):
|
||||
return access_token
|
||||
|
||||
refreshed = self._refresh_tokens(locked_auth_data)
|
||||
return refreshed["access_token"]
|
||||
|
||||
def login(self, force: bool = False, no_browser: bool = False) -> Dict[str, Any]:
|
||||
existing = self._read_auth_file()
|
||||
if existing and not force and existing.get("access_token"):
|
||||
if not self._is_expired(existing):
|
||||
return existing
|
||||
if existing.get("refresh_token"):
|
||||
try:
|
||||
return self._refresh_tokens(existing)
|
||||
except XAIOAuthError:
|
||||
pass
|
||||
|
||||
discovery = self._discover()
|
||||
verifier, challenge = self._pkce_pair()
|
||||
state = uuid.uuid4().hex
|
||||
nonce = uuid.uuid4().hex
|
||||
server, redirect_uri = self._start_callback_server(state)
|
||||
authorize_url = self._build_authorize_url(
|
||||
authorization_endpoint=discovery["authorization_endpoint"],
|
||||
redirect_uri=redirect_uri,
|
||||
challenge=challenge,
|
||||
state=state,
|
||||
nonce=nonce,
|
||||
)
|
||||
|
||||
if no_browser or not webbrowser.open(authorize_url):
|
||||
sys.stdout.write(
|
||||
f"Open this URL to authenticate with xAI:\n{authorize_url}\n"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
|
||||
result = self._wait_for_callback(server)
|
||||
if result.get("state") != state:
|
||||
raise XAIOAuthError("xAI OAuth state mismatch")
|
||||
if result.get("error"):
|
||||
description = result.get("error_description") or result["error"]
|
||||
raise XAIOAuthError(f"xAI authorization failed: {description}")
|
||||
code = result.get("code")
|
||||
if not code:
|
||||
raise XAIOAuthError("xAI authorization failed: no code returned")
|
||||
|
||||
token_payload = self._exchange_token(
|
||||
discovery["token_endpoint"],
|
||||
{
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"client_id": XAI_OAUTH_CLIENT_ID,
|
||||
"code_verifier": verifier,
|
||||
},
|
||||
)
|
||||
auth_data = self._build_auth_record(token_payload, discovery["token_endpoint"])
|
||||
self._write_auth_file(auth_data)
|
||||
return auth_data
|
||||
|
||||
def _client(self) -> httpx.Client:
|
||||
return self.http_client or _get_httpx_client()
|
||||
|
||||
def _ensure_token_dir(self) -> None:
|
||||
os.makedirs(self.token_dir, exist_ok=True)
|
||||
|
||||
def _read_auth_file(self) -> Optional[Dict[str, Any]]:
|
||||
try:
|
||||
with open(self.auth_file, "r") as f:
|
||||
data = json.load(f)
|
||||
return data if isinstance(data, dict) else None
|
||||
except (IOError, json.JSONDecodeError):
|
||||
return None
|
||||
|
||||
def _write_auth_file(self, data: Dict[str, Any]) -> None:
|
||||
self._ensure_token_dir()
|
||||
with open(self.auth_file, "w") as f:
|
||||
json.dump(data, f)
|
||||
try:
|
||||
os.chmod(self.auth_file, 0o600)
|
||||
except OSError:
|
||||
verbose_logger.debug("Could not chmod xAI OAuth auth file")
|
||||
|
||||
def _is_expired(self, auth_data: Dict[str, Any]) -> bool:
|
||||
expires_at = auth_data.get("expires_at")
|
||||
if expires_at is None:
|
||||
return True
|
||||
try:
|
||||
return time.time() >= float(expires_at) - XAI_OAUTH_EXPIRY_SKEW_SECONDS
|
||||
except (TypeError, ValueError):
|
||||
return True
|
||||
|
||||
def _discover(self) -> Dict[str, str]:
|
||||
response = self._client().get(
|
||||
XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
authorization_endpoint = data.get("authorization_endpoint")
|
||||
token_endpoint = data.get("token_endpoint")
|
||||
if not authorization_endpoint or not token_endpoint:
|
||||
raise XAIOAuthError("xAI OAuth discovery missing endpoints")
|
||||
return {
|
||||
"authorization_endpoint": self._validate_xai_endpoint(
|
||||
authorization_endpoint
|
||||
),
|
||||
"token_endpoint": self._validate_xai_endpoint(token_endpoint),
|
||||
}
|
||||
|
||||
def _validate_xai_endpoint(self, url: str) -> str:
|
||||
parsed = urlparse(url)
|
||||
host = (parsed.hostname or "").lower()
|
||||
if parsed.scheme != "https" or (host != "x.ai" and not host.endswith(".x.ai")):
|
||||
raise XAIOAuthError(
|
||||
f"xAI OAuth discovery returned unexpected endpoint: {url}"
|
||||
)
|
||||
return url
|
||||
|
||||
def _pkce_pair(self) -> Tuple[str, str]:
|
||||
verifier = (
|
||||
base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode()
|
||||
)
|
||||
challenge = (
|
||||
base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest())
|
||||
.rstrip(b"=")
|
||||
.decode()
|
||||
)
|
||||
return verifier, challenge
|
||||
|
||||
def _start_callback_server(self, state: str) -> Tuple[_CallbackServer, str]:
|
||||
last_error: Optional[OSError] = None
|
||||
for port in (XAI_OAUTH_REDIRECT_PORT, 0):
|
||||
try:
|
||||
server = _CallbackServer(
|
||||
(XAI_OAUTH_REDIRECT_HOST, port), _CallbackHandler
|
||||
)
|
||||
server.expected_state = state
|
||||
server.callback_result = None
|
||||
actual_port = server.server_address[1]
|
||||
redirect_uri = f"http://{XAI_OAUTH_REDIRECT_HOST}:{actual_port}{XAI_OAUTH_REDIRECT_PATH}"
|
||||
return server, redirect_uri
|
||||
except OSError as exc:
|
||||
last_error = exc
|
||||
raise XAIOAuthError(f"Could not start xAI OAuth callback server: {last_error}")
|
||||
|
||||
def _build_authorize_url(
|
||||
self,
|
||||
authorization_endpoint: str,
|
||||
redirect_uri: str,
|
||||
challenge: str,
|
||||
state: str,
|
||||
nonce: str,
|
||||
) -> str:
|
||||
params = {
|
||||
"response_type": "code",
|
||||
"client_id": XAI_OAUTH_CLIENT_ID,
|
||||
"redirect_uri": redirect_uri,
|
||||
"scope": XAI_OAUTH_SCOPE,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"state": state,
|
||||
"nonce": nonce,
|
||||
}
|
||||
return f"{authorization_endpoint}?{urlencode(params)}"
|
||||
|
||||
def _wait_for_callback(self, server: _CallbackServer) -> Dict[str, Optional[str]]:
|
||||
server.timeout = 1
|
||||
deadline = time.time() + XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS
|
||||
try:
|
||||
while time.time() < deadline:
|
||||
server.handle_request()
|
||||
if server.callback_result is not None:
|
||||
return server.callback_result
|
||||
finally:
|
||||
server.server_close()
|
||||
raise XAIOAuthError("Timed out waiting for xAI OAuth callback")
|
||||
|
||||
def _exchange_token(
|
||||
self, token_endpoint: str, data: Dict[str, str]
|
||||
) -> Dict[str, Any]:
|
||||
response = self._client().post(
|
||||
token_endpoint,
|
||||
headers={
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
},
|
||||
data=data,
|
||||
)
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise XAIOAuthError(
|
||||
f"xAI OAuth token request failed: {exc.response.status_code} {exc.response.text}"
|
||||
) from exc
|
||||
body = response.json()
|
||||
if not isinstance(body, dict):
|
||||
raise XAIOAuthError("xAI OAuth token response was not an object")
|
||||
return body
|
||||
|
||||
def _build_auth_record(
|
||||
self,
|
||||
token_payload: Dict[str, Any],
|
||||
token_endpoint: str,
|
||||
fallback_refresh_token: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
access_token = token_payload.get("access_token")
|
||||
refresh_token = token_payload.get("refresh_token") or fallback_refresh_token
|
||||
if not access_token:
|
||||
raise XAIOAuthError("xAI OAuth token response missing access_token")
|
||||
if not refresh_token:
|
||||
raise XAIOAuthError("xAI OAuth token response missing refresh_token")
|
||||
expires_in = token_payload.get("expires_in") or 3600
|
||||
try:
|
||||
expires_at = int(time.time() + int(expires_in))
|
||||
except (TypeError, ValueError):
|
||||
expires_at = int(time.time() + 3600)
|
||||
return {
|
||||
"access_token": access_token,
|
||||
"refresh_token": refresh_token,
|
||||
"id_token": token_payload.get("id_token"),
|
||||
"token_type": token_payload.get("token_type") or "Bearer",
|
||||
"token_endpoint": token_endpoint,
|
||||
"expires_at": expires_at,
|
||||
}
|
||||
|
||||
def _refresh_tokens(self, auth_data: Dict[str, Any]) -> Dict[str, Any]:
|
||||
token_endpoint = auth_data.get("token_endpoint")
|
||||
if not token_endpoint:
|
||||
token_endpoint = self._discover()["token_endpoint"]
|
||||
token_endpoint = self._validate_xai_endpoint(token_endpoint)
|
||||
refresh_token = auth_data.get("refresh_token")
|
||||
if not refresh_token:
|
||||
raise XAIOAuthLoginRequiredError(
|
||||
"xAI OAuth refresh token missing. Run `litellm xai-oauth login`."
|
||||
)
|
||||
|
||||
token_payload = self._exchange_token(
|
||||
token_endpoint,
|
||||
{
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
"client_id": XAI_OAUTH_CLIENT_ID,
|
||||
},
|
||||
)
|
||||
refreshed = self._build_auth_record(
|
||||
token_payload,
|
||||
token_endpoint,
|
||||
fallback_refresh_token=refresh_token,
|
||||
)
|
||||
self._write_auth_file(refreshed)
|
||||
return refreshed
|
||||
|
||||
|
||||
class XAIOAuthChatConfig(XAIChatConfig):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "xai_oauth"
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
try:
|
||||
dynamic_api_key = api_key or self.authenticator.get_access_token()
|
||||
except XAIOAuthError as exc:
|
||||
raise AuthenticationError(
|
||||
model="",
|
||||
llm_provider=self.custom_llm_provider or "xai_oauth",
|
||||
message=str(exc),
|
||||
) from exc
|
||||
return (
|
||||
api_base or self.authenticator.get_api_base(),
|
||||
dynamic_api_key,
|
||||
)
|
||||
|
||||
|
||||
class XAIOAuthResponsesAPIConfig(XAIResponsesAPIConfig):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.XAI_OAUTH
|
||||
|
||||
def validate_environment(
|
||||
self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
|
||||
) -> dict:
|
||||
litellm_params = litellm_params or GenericLiteLLMParams()
|
||||
try:
|
||||
api_key = litellm_params.api_key or self.authenticator.get_access_token()
|
||||
except XAIOAuthError as exc:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider.value,
|
||||
message=str(exc),
|
||||
) from exc
|
||||
headers.update({"Authorization": f"Bearer {api_key}"})
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
api_base = api_base or self.authenticator.get_api_base()
|
||||
return f"{api_base.rstrip('/')}/responses"
|
||||
|
|
@ -2286,7 +2286,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
additional_args={"headers": headers},
|
||||
)
|
||||
raise e
|
||||
elif custom_llm_provider == "xai":
|
||||
elif custom_llm_provider == "xai" or custom_llm_provider == "xai_oauth":
|
||||
## COMPLETION CALL
|
||||
try:
|
||||
response = base_llm_http_handler.completion(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -3261,6 +3261,7 @@ class LlmProviders(str, Enum):
|
|||
OPENAI_LIKE = "openai_like" # embedding only
|
||||
JINA_AI = "jina_ai"
|
||||
XAI = "xai"
|
||||
XAI_OAUTH = "xai_oauth"
|
||||
ZAI = "zai"
|
||||
CUSTOM_OPENAI = "custom_openai"
|
||||
TEXT_COMPLETION_OPENAI = "text-completion-openai"
|
||||
|
|
|
|||
|
|
@ -4594,7 +4594,7 @@ def get_optional_params( # noqa: PLR0915
|
|||
else False
|
||||
),
|
||||
)
|
||||
elif custom_llm_provider == "xai":
|
||||
elif custom_llm_provider == "xai" or custom_llm_provider == "xai_oauth":
|
||||
optional_params = litellm.XAIChatConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
|
|
@ -5775,6 +5775,18 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
]
|
||||
split_model = potential_model_names["split_model"]
|
||||
custom_llm_provider = potential_model_names["custom_llm_provider"]
|
||||
xai_oauth_model_alias = custom_llm_provider == "xai_oauth" or model.startswith(
|
||||
"xai_oauth/"
|
||||
)
|
||||
model_cost_custom_llm_provider = custom_llm_provider
|
||||
if xai_oauth_model_alias:
|
||||
model_cost_custom_llm_provider = "xai"
|
||||
model = model.replace("xai_oauth/", "xai/", 1)
|
||||
combined_model_name = combined_model_name.replace("xai_oauth/", "xai/", 1)
|
||||
stripped_model_name = stripped_model_name.replace("xai_oauth/", "xai/", 1)
|
||||
combined_stripped_model_name = combined_stripped_model_name.replace(
|
||||
"xai_oauth/", "xai/", 1
|
||||
)
|
||||
#########################
|
||||
provider_config: Optional[BaseLLMModelInfo] = None
|
||||
if custom_llm_provider and custom_llm_provider in LlmProvidersSet:
|
||||
|
|
@ -5840,7 +5852,8 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
key = _matched_key
|
||||
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
|
||||
if not _check_provider_match(
|
||||
model_info=_model_info, custom_llm_provider=custom_llm_provider
|
||||
model_info=_model_info,
|
||||
custom_llm_provider=model_cost_custom_llm_provider,
|
||||
):
|
||||
_model_info = None
|
||||
if _model_info is None:
|
||||
|
|
@ -5849,7 +5862,8 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
key = _matched_key
|
||||
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
|
||||
if not _check_provider_match(
|
||||
model_info=_model_info, custom_llm_provider=custom_llm_provider
|
||||
model_info=_model_info,
|
||||
custom_llm_provider=model_cost_custom_llm_provider,
|
||||
):
|
||||
_model_info = None
|
||||
if _model_info is None:
|
||||
|
|
@ -5858,7 +5872,8 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
key = _matched_key
|
||||
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
|
||||
if not _check_provider_match(
|
||||
model_info=_model_info, custom_llm_provider=custom_llm_provider
|
||||
model_info=_model_info,
|
||||
custom_llm_provider=model_cost_custom_llm_provider,
|
||||
):
|
||||
_model_info = None
|
||||
if _model_info is None:
|
||||
|
|
@ -5867,7 +5882,8 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
key = _matched_key
|
||||
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
|
||||
if not _check_provider_match(
|
||||
model_info=_model_info, custom_llm_provider=custom_llm_provider
|
||||
model_info=_model_info,
|
||||
custom_llm_provider=model_cost_custom_llm_provider,
|
||||
):
|
||||
_model_info = None
|
||||
if _model_info is None:
|
||||
|
|
@ -5876,7 +5892,8 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
key = _matched_key
|
||||
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
|
||||
if not _check_provider_match(
|
||||
model_info=_model_info, custom_llm_provider=custom_llm_provider
|
||||
model_info=_model_info,
|
||||
custom_llm_provider=model_cost_custom_llm_provider,
|
||||
):
|
||||
_model_info = None
|
||||
|
||||
|
|
@ -5884,6 +5901,10 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
raise ValueError(
|
||||
"This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
|
||||
)
|
||||
if xai_oauth_model_alias:
|
||||
_model_info = dict(_model_info)
|
||||
_model_info["litellm_provider"] = "xai_oauth"
|
||||
key = key.replace("xai/", "xai_oauth/", 1)
|
||||
|
||||
_input_cost_per_token: Optional[float] = _model_info.get(
|
||||
"input_cost_per_token"
|
||||
|
|
@ -6634,6 +6655,13 @@ def validate_environment( # noqa: PLR0915
|
|||
keys_in_environment = True
|
||||
else:
|
||||
missing_keys.append("XAI_API_KEY")
|
||||
elif custom_llm_provider == "xai_oauth":
|
||||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator
|
||||
|
||||
if XAIOAuthAuthenticator()._read_auth_file() is not None:
|
||||
keys_in_environment = True
|
||||
else:
|
||||
missing_keys.append("XAI OAuth credentials")
|
||||
elif custom_llm_provider == "ai21_chat":
|
||||
if "AI21_API_KEY" in os.environ:
|
||||
keys_in_environment = True
|
||||
|
|
@ -8312,6 +8340,10 @@ class ProviderConfigManager:
|
|||
LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False),
|
||||
LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False),
|
||||
LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False),
|
||||
LlmProviders.XAI_OAUTH: (
|
||||
lambda: ProviderConfigManager._get_xai_oauth_config(),
|
||||
False,
|
||||
),
|
||||
LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False),
|
||||
LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False),
|
||||
LlmProviders.INCEPTION: (lambda: litellm.InceptionChatConfig(), False),
|
||||
|
|
@ -8504,6 +8536,12 @@ class ProviderConfigManager:
|
|||
|
||||
return LangFlowConfig()
|
||||
|
||||
@staticmethod
|
||||
def _get_xai_oauth_config() -> BaseConfig:
|
||||
from litellm.llms.xai.oauth import XAIOAuthChatConfig
|
||||
|
||||
return XAIOAuthChatConfig()
|
||||
|
||||
@staticmethod
|
||||
def get_provider_chat_config( # noqa: PLR0915
|
||||
model: str,
|
||||
|
|
@ -8894,6 +8932,10 @@ class ProviderConfigManager:
|
|||
return litellm.AzureOpenAIResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.XAI == provider:
|
||||
return litellm.XAIResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.XAI_OAUTH == provider:
|
||||
from litellm.llms.xai.oauth import XAIOAuthResponsesAPIConfig
|
||||
|
||||
return XAIOAuthResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
return litellm.GithubCopilotResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.CHATGPT == provider:
|
||||
|
|
|
|||
|
|
@ -2396,6 +2396,25 @@
|
|||
"realtime": true
|
||||
}
|
||||
},
|
||||
"xai_oauth": {
|
||||
"display_name": "xAI OAuth (`xai_oauth`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/xai_oauth",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": true,
|
||||
"interactions": true,
|
||||
"realtime": false
|
||||
}
|
||||
},
|
||||
"xinference": {
|
||||
"display_name": "Xinference (`xinference`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/xinference",
|
||||
|
|
|
|||
676
tests/test_litellm/llms/xai/test_xai_oauth.py
Normal file
676
tests/test_litellm/llms/xai/test_xai_oauth.py
Normal file
|
|
@ -0,0 +1,676 @@
|
|||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import litellm
|
||||
import pytest
|
||||
from click.testing import CliRunner
|
||||
|
||||
import litellm.llms.xai.oauth as xai_oauth_module
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.llms.xai.oauth import (
|
||||
XAI_OAUTH_CLIENT_ID,
|
||||
XAI_OAUTH_SCOPE,
|
||||
XAIOAuthError,
|
||||
XAIOAuthAuthenticator,
|
||||
XAIOAuthChatConfig,
|
||||
XAIOAuthLoginRequiredError,
|
||||
XAIOAuthResponsesAPIConfig,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import (
|
||||
ProviderConfigManager,
|
||||
get_optional_params,
|
||||
validate_environment,
|
||||
)
|
||||
|
||||
|
||||
def _write_auth_file(tmp_path, payload):
|
||||
token_dir = tmp_path / "xai_oauth"
|
||||
token_dir.mkdir()
|
||||
auth_file = token_dir / "auth.json"
|
||||
auth_file.write_text(json.dumps(payload))
|
||||
return token_dir, auth_file
|
||||
|
||||
|
||||
def test_get_access_token_uses_fresh_local_token(tmp_path, monkeypatch):
|
||||
token_dir, _ = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "fresh-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() + 3600,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
assert XAIOAuthAuthenticator().get_access_token() == "fresh-token"
|
||||
|
||||
|
||||
def test_get_access_token_refreshes_and_preserves_refresh_token(tmp_path, monkeypatch):
|
||||
token_dir, auth_file = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "expired-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
"expires_at": time.time() - 1,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = dict(item.split("=") for item in request.content.decode().split("&"))
|
||||
assert body["grant_type"] == "refresh_token"
|
||||
assert body["refresh_token"] == "refresh-token"
|
||||
assert body["client_id"] == XAI_OAUTH_CLIENT_ID
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"access_token": "new-token",
|
||||
"expires_in": 3600,
|
||||
"token_type": "Bearer",
|
||||
},
|
||||
)
|
||||
|
||||
client = httpx.Client(transport=httpx.MockTransport(handler))
|
||||
|
||||
assert XAIOAuthAuthenticator(http_client=client).get_access_token() == "new-token"
|
||||
stored = json.loads(auth_file.read_text())
|
||||
assert stored["access_token"] == "new-token"
|
||||
assert stored["refresh_token"] == "refresh-token"
|
||||
|
||||
|
||||
def test_get_access_token_reuses_token_refreshed_by_parallel_request():
|
||||
expired_auth_data = {
|
||||
"access_token": "expired-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
"expires_at": time.time() - 1,
|
||||
}
|
||||
refreshed_auth_data = {
|
||||
"access_token": "already-refreshed-token",
|
||||
"refresh_token": "rotated-refresh-token",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
"expires_at": time.time() + 3600,
|
||||
}
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
authenticator._read_auth_file = MagicMock(
|
||||
side_effect=[expired_auth_data, refreshed_auth_data]
|
||||
)
|
||||
authenticator._refresh_tokens = MagicMock()
|
||||
|
||||
assert authenticator.get_access_token() == "already-refreshed-token"
|
||||
authenticator._refresh_tokens.assert_not_called()
|
||||
|
||||
|
||||
def test_get_access_token_requires_login_without_auth_file(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing"))
|
||||
|
||||
with pytest.raises(XAIOAuthLoginRequiredError):
|
||||
XAIOAuthAuthenticator().get_access_token()
|
||||
|
||||
|
||||
def test_get_access_token_ignores_invalid_auth_file(tmp_path, monkeypatch):
|
||||
token_dir = tmp_path / "xai_oauth"
|
||||
token_dir.mkdir()
|
||||
(token_dir / "auth.json").write_text("{not-json")
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
with pytest.raises(XAIOAuthLoginRequiredError):
|
||||
XAIOAuthAuthenticator().get_access_token()
|
||||
|
||||
|
||||
def test_refresh_failure_surfaces_oauth_error(tmp_path, monkeypatch):
|
||||
token_dir, _ = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "expired-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
"expires_at": time.time() - 1,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
client = httpx.Client(
|
||||
transport=httpx.MockTransport(
|
||||
lambda request: httpx.Response(401, text="invalid_grant", request=request)
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(XAIOAuthError) as exc_info:
|
||||
XAIOAuthAuthenticator(http_client=client).get_access_token()
|
||||
|
||||
assert "401 invalid_grant" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_build_auth_record_requires_access_and_refresh_tokens():
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="access_token"):
|
||||
authenticator._build_auth_record(
|
||||
{"refresh_token": "refresh-token"},
|
||||
"https://auth.x.ai/oauth/token",
|
||||
)
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="refresh_token"):
|
||||
authenticator._build_auth_record(
|
||||
{"access_token": "access-token"},
|
||||
"https://auth.x.ai/oauth/token",
|
||||
)
|
||||
|
||||
|
||||
def test_build_auth_record_defaults_expiry_and_token_type():
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
auth_data = authenticator._build_auth_record(
|
||||
{
|
||||
"access_token": "access-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_in": "not-a-number",
|
||||
},
|
||||
"https://auth.x.ai/oauth/token",
|
||||
)
|
||||
|
||||
assert auth_data["token_type"] == "Bearer"
|
||||
assert auth_data["expires_at"] > time.time()
|
||||
|
||||
|
||||
def test_is_expired_treats_missing_or_invalid_expiry_as_expired():
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
assert authenticator._is_expired({}) is True
|
||||
assert authenticator._is_expired({"expires_at": "not-a-number"}) is True
|
||||
|
||||
|
||||
def test_write_auth_file_creates_private_file(tmp_path, monkeypatch):
|
||||
token_dir = tmp_path / "xai_oauth"
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
authenticator._write_auth_file(
|
||||
{
|
||||
"access_token": "access-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() + 3600,
|
||||
}
|
||||
)
|
||||
|
||||
stored = json.loads((token_dir / "auth.json").read_text())
|
||||
assert stored["access_token"] == "access-token"
|
||||
assert oct(os.stat(token_dir / "auth.json").st_mode & 0o777) == "0o600"
|
||||
|
||||
|
||||
def test_discovery_rejects_unexpected_endpoint():
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="unexpected endpoint"):
|
||||
authenticator._validate_xai_endpoint("https://evil.example.com/oauth/token")
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="unexpected endpoint"):
|
||||
authenticator._validate_xai_endpoint("http://auth.x.ai/oauth/token")
|
||||
|
||||
|
||||
def test_discover_returns_validated_xai_endpoints():
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url == "https://auth.x.ai/.well-known/openid-configuration"
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"authorization_endpoint": "https://auth.x.ai/oauth/authorize",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
},
|
||||
)
|
||||
|
||||
authenticator = XAIOAuthAuthenticator(
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler))
|
||||
)
|
||||
|
||||
assert authenticator._discover() == {
|
||||
"authorization_endpoint": "https://auth.x.ai/oauth/authorize",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
}
|
||||
|
||||
|
||||
def test_discover_requires_authorization_and_token_endpoints():
|
||||
authenticator = XAIOAuthAuthenticator(
|
||||
http_client=httpx.Client(
|
||||
transport=httpx.MockTransport(lambda request: httpx.Response(200, json={}))
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="missing endpoints"):
|
||||
authenticator._discover()
|
||||
|
||||
|
||||
def test_refresh_discovers_token_endpoint_when_auth_file_is_legacy(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
token_dir, auth_file = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "expired-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() - 1,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
if request.method == "GET":
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"authorization_endpoint": "https://auth.x.ai/oauth/authorize",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
},
|
||||
)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"access_token": "discovered-token",
|
||||
"refresh_token": "new-refresh-token",
|
||||
"expires_in": 3600,
|
||||
},
|
||||
)
|
||||
|
||||
authenticator = XAIOAuthAuthenticator(
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler))
|
||||
)
|
||||
|
||||
assert authenticator.get_access_token() == "discovered-token"
|
||||
stored = json.loads(auth_file.read_text())
|
||||
assert stored["token_endpoint"] == "https://auth.x.ai/oauth/token"
|
||||
|
||||
|
||||
def test_exchange_token_rejects_non_object_response():
|
||||
authenticator = XAIOAuthAuthenticator(
|
||||
http_client=httpx.Client(
|
||||
transport=httpx.MockTransport(
|
||||
lambda request: httpx.Response(200, json=["not", "an", "object"])
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="was not an object"):
|
||||
authenticator._exchange_token("https://auth.x.ai/oauth/token", {})
|
||||
|
||||
|
||||
def test_start_callback_server_falls_back_to_ephemeral_port(monkeypatch):
|
||||
calls = []
|
||||
real_server = xai_oauth_module._CallbackServer
|
||||
|
||||
class FirstPortFailsCallbackServer(real_server):
|
||||
def __init__(self, server_address, handler_class):
|
||||
calls.append(server_address[1])
|
||||
if server_address[1] == xai_oauth_module.XAI_OAUTH_REDIRECT_PORT:
|
||||
raise OSError("port unavailable")
|
||||
super().__init__(server_address, handler_class)
|
||||
|
||||
monkeypatch.setattr(
|
||||
xai_oauth_module, "_CallbackServer", FirstPortFailsCallbackServer
|
||||
)
|
||||
|
||||
server, redirect_uri = XAIOAuthAuthenticator()._start_callback_server("state-value")
|
||||
try:
|
||||
assert calls == [xai_oauth_module.XAI_OAUTH_REDIRECT_PORT, 0]
|
||||
assert redirect_uri.startswith("http://127.0.0.1:")
|
||||
assert redirect_uri.endswith("/callback")
|
||||
finally:
|
||||
server.server_close()
|
||||
|
||||
|
||||
def test_wait_for_callback_times_out_and_closes_server(monkeypatch):
|
||||
server, _ = XAIOAuthAuthenticator()._start_callback_server("state-value")
|
||||
monkeypatch.setattr(xai_oauth_module, "XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS", 0)
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="Timed out"):
|
||||
XAIOAuthAuthenticator()._wait_for_callback(server)
|
||||
|
||||
|
||||
def test_callback_handler_records_success_and_rejects_state_mismatch():
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
server, redirect_uri = authenticator._start_callback_server("expected-state")
|
||||
thread = threading.Thread(target=server.handle_request)
|
||||
thread.start()
|
||||
response = httpx.get(f"{redirect_uri}?code=auth-code&state=expected-state")
|
||||
thread.join(timeout=5)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert server.callback_result == {
|
||||
"code": "auth-code",
|
||||
"state": "expected-state",
|
||||
"error": None,
|
||||
"error_description": None,
|
||||
}
|
||||
|
||||
server, redirect_uri = authenticator._start_callback_server("expected-state")
|
||||
thread = threading.Thread(target=server.handle_request)
|
||||
thread.start()
|
||||
response = httpx.get(f"{redirect_uri}?code=auth-code&state=wrong-state")
|
||||
thread.join(timeout=5)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert server.callback_result["state"] == "wrong-state"
|
||||
|
||||
|
||||
def test_login_exchanges_authorization_code_and_persists_auth_record(monkeypatch):
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
fake_server = MagicMock()
|
||||
written_records = []
|
||||
|
||||
class FakeUUID:
|
||||
def __init__(self, value):
|
||||
self.hex = value
|
||||
|
||||
monkeypatch.setattr(
|
||||
xai_oauth_module.uuid,
|
||||
"uuid4",
|
||||
MagicMock(side_effect=[FakeUUID("state-value"), FakeUUID("nonce-value")]),
|
||||
)
|
||||
authenticator._read_auth_file = MagicMock(return_value=None)
|
||||
authenticator._discover = MagicMock(
|
||||
return_value={
|
||||
"authorization_endpoint": "https://auth.x.ai/oauth/authorize",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
}
|
||||
)
|
||||
authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge"))
|
||||
authenticator._start_callback_server = MagicMock(
|
||||
return_value=(fake_server, "http://127.0.0.1:56121/callback")
|
||||
)
|
||||
authenticator._wait_for_callback = MagicMock(
|
||||
return_value={"state": "state-value", "code": "auth-code"}
|
||||
)
|
||||
authenticator._exchange_token = MagicMock(
|
||||
return_value={
|
||||
"access_token": "access-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
)
|
||||
authenticator._write_auth_file = MagicMock(side_effect=written_records.append)
|
||||
|
||||
auth_data = authenticator.login(no_browser=True)
|
||||
|
||||
authenticator._exchange_token.assert_called_once_with(
|
||||
"https://auth.x.ai/oauth/token",
|
||||
{
|
||||
"grant_type": "authorization_code",
|
||||
"code": "auth-code",
|
||||
"redirect_uri": "http://127.0.0.1:56121/callback",
|
||||
"client_id": XAI_OAUTH_CLIENT_ID,
|
||||
"code_verifier": "verifier",
|
||||
},
|
||||
)
|
||||
assert auth_data["access_token"] == "access-token"
|
||||
assert written_records == [auth_data]
|
||||
|
||||
|
||||
def test_login_raises_on_callback_error_or_missing_code(monkeypatch):
|
||||
authenticator = XAIOAuthAuthenticator()
|
||||
|
||||
class FakeUUID:
|
||||
hex = "state-value"
|
||||
|
||||
monkeypatch.setattr(
|
||||
xai_oauth_module.uuid, "uuid4", MagicMock(return_value=FakeUUID())
|
||||
)
|
||||
authenticator._read_auth_file = MagicMock(return_value=None)
|
||||
authenticator._discover = MagicMock(
|
||||
return_value={
|
||||
"authorization_endpoint": "https://auth.x.ai/oauth/authorize",
|
||||
"token_endpoint": "https://auth.x.ai/oauth/token",
|
||||
}
|
||||
)
|
||||
authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge"))
|
||||
authenticator._start_callback_server = MagicMock(
|
||||
return_value=(MagicMock(), "http://127.0.0.1:56121/callback")
|
||||
)
|
||||
authenticator._wait_for_callback = MagicMock(
|
||||
return_value={
|
||||
"state": "state-value",
|
||||
"error": "access_denied",
|
||||
"error_description": "denied",
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="denied"):
|
||||
authenticator.login(no_browser=True)
|
||||
|
||||
authenticator._wait_for_callback = MagicMock(return_value={"state": "state-value"})
|
||||
|
||||
with pytest.raises(XAIOAuthError, match="no code returned"):
|
||||
authenticator.login(no_browser=True)
|
||||
|
||||
|
||||
def test_pkce_pair_generates_s256_challenge():
|
||||
verifier, challenge = XAIOAuthAuthenticator()._pkce_pair()
|
||||
expected = (
|
||||
base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest())
|
||||
.rstrip(b"=")
|
||||
.decode()
|
||||
)
|
||||
|
||||
assert challenge == expected
|
||||
assert "=" not in verifier
|
||||
assert "=" not in challenge
|
||||
|
||||
|
||||
def test_build_authorize_url_contains_xai_oauth_parameters():
|
||||
authorize_url = XAIOAuthAuthenticator()._build_authorize_url(
|
||||
authorization_endpoint="https://auth.x.ai/oauth/authorize",
|
||||
redirect_uri="http://127.0.0.1:56121/callback",
|
||||
challenge="pkce-challenge",
|
||||
state="state-value",
|
||||
nonce="nonce-value",
|
||||
)
|
||||
parsed = urlparse(authorize_url)
|
||||
params = parse_qs(parsed.query)
|
||||
|
||||
assert parsed.scheme == "https"
|
||||
assert parsed.netloc == "auth.x.ai"
|
||||
assert params["response_type"] == ["code"]
|
||||
assert params["client_id"] == [XAI_OAUTH_CLIENT_ID]
|
||||
assert params["scope"] == [XAI_OAUTH_SCOPE]
|
||||
assert params["code_challenge"] == ["pkce-challenge"]
|
||||
assert params["code_challenge_method"] == ["S256"]
|
||||
assert params["state"] == ["state-value"]
|
||||
assert params["nonce"] == ["nonce-value"]
|
||||
|
||||
|
||||
def test_get_llm_provider_resolves_xai_oauth(tmp_path, monkeypatch):
|
||||
token_dir, _ = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "provider-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() + 3600,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider("xai_oauth/grok-4")
|
||||
|
||||
assert model == "grok-4"
|
||||
assert provider == "xai_oauth"
|
||||
assert api_key == "provider-token"
|
||||
assert api_base == "https://api.x.ai/v1"
|
||||
|
||||
|
||||
def test_chat_config_wraps_oauth_errors_as_authentication_error(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing"))
|
||||
|
||||
with pytest.raises(litellm.AuthenticationError) as exc_info:
|
||||
XAIOAuthChatConfig()._get_openai_compatible_provider_info(None, None)
|
||||
|
||||
assert exc_info.value.llm_provider == "xai_oauth"
|
||||
assert "litellm xai-oauth login" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_chat_config_uses_attached_authenticator():
|
||||
config = XAIOAuthChatConfig()
|
||||
config.authenticator = MagicMock()
|
||||
config.authenticator.get_access_token.return_value = "chat-token"
|
||||
config.authenticator.get_api_base.return_value = "https://xai.example.com/v1"
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
|
||||
assert api_base == "https://xai.example.com/v1"
|
||||
assert api_key == "chat-token"
|
||||
|
||||
|
||||
def test_get_model_info_reuses_xai_metadata_for_xai_oauth():
|
||||
xai_info = litellm.get_model_info("xai/grok-4")
|
||||
oauth_info = litellm.get_model_info("xai_oauth/grok-4")
|
||||
|
||||
assert oauth_info["litellm_provider"] == "xai_oauth"
|
||||
assert oauth_info["max_input_tokens"] == xai_info["max_input_tokens"]
|
||||
assert oauth_info["input_cost_per_token"] == xai_info["input_cost_per_token"]
|
||||
|
||||
|
||||
def test_validate_environment_reads_xai_oauth_auth_file(tmp_path, monkeypatch):
|
||||
token_dir, _ = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "env-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() + 3600,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
monkeypatch.setattr(
|
||||
"litellm.utils.get_llm_provider",
|
||||
lambda model: ("grok-4", "xai_oauth", None, None),
|
||||
)
|
||||
|
||||
assert validate_environment("xai_oauth/grok-4") == {
|
||||
"keys_in_environment": True,
|
||||
"missing_keys": [],
|
||||
}
|
||||
|
||||
|
||||
def test_validate_environment_reports_missing_xai_oauth_credentials(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing"))
|
||||
monkeypatch.setattr(
|
||||
"litellm.utils.get_llm_provider",
|
||||
lambda model: ("grok-4", "xai_oauth", None, None),
|
||||
)
|
||||
|
||||
assert validate_environment("xai_oauth/grok-4") == {
|
||||
"keys_in_environment": False,
|
||||
"missing_keys": ["XAI OAuth credentials"],
|
||||
}
|
||||
|
||||
|
||||
def test_xai_oauth_uses_xai_optional_param_mapping():
|
||||
optional_params = get_optional_params(
|
||||
model="grok-4",
|
||||
custom_llm_provider="xai_oauth",
|
||||
temperature=0.2,
|
||||
max_tokens=8,
|
||||
)
|
||||
|
||||
assert optional_params["temperature"] == 0.2
|
||||
assert optional_params["max_tokens"] == 8
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_xai_oauth_configs():
|
||||
chat_config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="grok-4",
|
||||
provider=LlmProviders.XAI_OAUTH,
|
||||
)
|
||||
responses_config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider=LlmProviders.XAI_OAUTH
|
||||
)
|
||||
|
||||
assert isinstance(chat_config, XAIOAuthChatConfig)
|
||||
assert isinstance(responses_config, XAIOAuthResponsesAPIConfig)
|
||||
|
||||
|
||||
def test_responses_config_injects_oauth_bearer_token(tmp_path, monkeypatch):
|
||||
token_dir, _ = _write_auth_file(
|
||||
tmp_path,
|
||||
{
|
||||
"access_token": "responses-token",
|
||||
"refresh_token": "refresh-token",
|
||||
"expires_at": time.time() + 3600,
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir))
|
||||
|
||||
headers = XAIOAuthResponsesAPIConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-4",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
)
|
||||
|
||||
assert headers["Authorization"] == "Bearer responses-token"
|
||||
|
||||
|
||||
def test_responses_config_endpoint_url_uses_oauth_authenticator():
|
||||
config = XAIOAuthResponsesAPIConfig()
|
||||
config.authenticator = MagicMock()
|
||||
config.authenticator.get_api_base.return_value = "https://xai.example.com/v1/"
|
||||
|
||||
assert config.get_complete_url(api_base=None, litellm_params={}) == (
|
||||
"https://xai.example.com/v1/responses"
|
||||
)
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base="https://custom.example.com/v1/", litellm_params={}
|
||||
)
|
||||
== "https://custom.example.com/v1/responses"
|
||||
)
|
||||
|
||||
|
||||
def test_responses_config_wraps_oauth_errors_as_authentication_error(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing"))
|
||||
|
||||
with pytest.raises(litellm.AuthenticationError) as exc_info:
|
||||
XAIOAuthResponsesAPIConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-4",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
)
|
||||
|
||||
assert XAIOAuthResponsesAPIConfig().custom_llm_provider == LlmProviders.XAI_OAUTH
|
||||
assert exc_info.value.llm_provider == "xai_oauth"
|
||||
|
||||
|
||||
def test_proxy_cli_xai_oauth_login_uses_single_authenticator(monkeypatch):
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
instances = []
|
||||
|
||||
class FakeAuthenticator:
|
||||
auth_file = "/tmp/xai-oauth-auth.json"
|
||||
|
||||
def __init__(self):
|
||||
instances.append(self)
|
||||
|
||||
def login(self):
|
||||
return {"expires_at": 1234567890}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.llms.xai.oauth.XAIOAuthAuthenticator", FakeAuthenticator
|
||||
)
|
||||
|
||||
result = CliRunner().invoke(run_server, ["xai-oauth", "login"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert len(instances) == 1
|
||||
assert "Credentials saved to /tmp/xai-oauth-auth.json" in result.output
|
||||
assert "Access token expires at 1234567890" in result.output
|
||||
Loading…
Add table
Reference in a new issue