diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 6efee058a81..61901547f18 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -12,7 +12,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, strip_name_from_messages, ) -from litellm.llms.xai.common_utils import XAIModelInfo, xai_reported_cost_in_usd +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.xai.common_utils import XAIModelInfo, xai_error_status_code, xai_reported_cost_in_usd from litellm.llms.xai.cost_calculator import ( apply_server_side_tool_usage_details_to_usage, ) @@ -64,14 +65,14 @@ class XAIChatConfig(OpenAIGPTConfig): XAIOAuthAuthenticator, XAIOAuthError, should_use_xai_oauth, + xai_oauth_token_file, ) dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key) if should_use_xai_oauth(litellm_params) and not dynamic_api_key: - raw_token_file: Final = (litellm_params or {}).get("xai_oauth_token_file") - token_file: Final = raw_token_file if isinstance(raw_token_file, str) else None try: - headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator(auth_file=token_file).get_access_token()}" + authenticator: Final = XAIOAuthAuthenticator(auth_file=xai_oauth_token_file(litellm_params)) + headers["Authorization"] = f"Bearer {authenticator.get_access_token()}" except XAIOAuthError as exc: raise AuthenticationError( model=model, @@ -105,9 +106,7 @@ class XAIChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key) if should_use_xai_oauth(litellm_params) and not dynamic_api_key: - raw_token_file: Final = (litellm_params or {}).get("xai_oauth_token_file") - token_file: Final = raw_token_file if isinstance(raw_token_file, str) else None - api_base = XAIOAuthAuthenticator(auth_file=token_file).get_api_base() + api_base = XAIOAuthAuthenticator().get_api_base() return super().get_complete_url( api_base=api_base, @@ -118,6 +117,15 @@ class XAIChatConfig(OpenAIGPTConfig): stream=stream, ) + def get_error_class( + self, error_message: str, status_code: int, headers: dict[str, str] | httpx.Headers + ) -> BaseLLMException: + return super().get_error_class( + error_message=error_message, + status_code=xai_error_status_code(status_code, error_message), + headers=headers, + ) + def get_supported_openai_params(self, model: str) -> list: base_openai_params: Final = [ "logit_bias", diff --git a/litellm/llms/xai/common_utils.py b/litellm/llms/xai/common_utils.py index cf76a851a86..b7fb2b8da08 100644 --- a/litellm/llms/xai/common_utils.py +++ b/litellm/llms/xai/common_utils.py @@ -1,3 +1,4 @@ +from http import HTTPStatus from typing import Final import httpx @@ -9,6 +10,14 @@ from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ProviderSpecificModelInfo USD_TICKS_PER_DOLLAR: Final = 10_000_000_000 +XAI_SPENDING_LIMIT_ERROR_MARKER: Final = "spending-limit" + + +def xai_error_status_code(status_code: int, error_message: str) -> int: + """xAI rejects requests past an account's spending limit with a 403; report it as a 429 so routers fail over""" + if status_code == HTTPStatus.FORBIDDEN and XAI_SPENDING_LIMIT_ERROR_MARKER in error_message: + return HTTPStatus.TOO_MANY_REQUESTS.value + return status_code def xai_reported_cost_in_usd(cost_in_usd_ticks: object) -> float | None: diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py index 3476228fd85..4e8e26d4ad3 100644 --- a/litellm/llms/xai/oauth.py +++ b/litellm/llms/xai/oauth.py @@ -82,23 +82,23 @@ def _default_xai_oauth_token_dir() -> str: def resolve_xai_oauth_auth_file(auth_file: str | None, token_dir: str) -> str: - requested: Final = auth_file or get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json" - candidate: Final = requested if os.path.isabs(requested) else os.path.join(token_dir, requested) - resolved: Final = os.path.realpath(candidate) - allowed: Final = os.path.realpath(token_dir) - if resolved == allowed or resolved.startswith(allowed + os.sep): + if not auth_file: + return os.path.join(token_dir, get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json") + resolved: Final = os.path.realpath(os.path.join(token_dir, auth_file)) + if resolved.startswith(os.path.realpath(token_dir) + os.sep): return resolved raise XAIOAuthError("xAI OAuth auth file must stay inside the token directory") -def oauth_auth_file_for_account(account: str, token_dir: str) -> str: +def oauth_auth_file_for_account(account: str) -> str: if not _XAI_OAUTH_ACCOUNT_NAME_RE.fullmatch(account): raise ValueError("xAI OAuth account must match ^[A-Za-z0-9_-]+$") - return resolve_xai_oauth_auth_file(f"auth-{account}.json", token_dir) + return f"auth-{account}.json" -def is_xai_spending_limit_error(*, custom_llm_provider: str, status_code: int, error_str: str) -> bool: - return custom_llm_provider == "xai" and status_code == 403 and "spending-limit" in error_str +def xai_oauth_token_file(litellm_params: Mapping[str, object] | None) -> str | None: + token_file: Final = (litellm_params or {}).get("xai_oauth_token_file") + return token_file if isinstance(token_file, str) else None class _CallbackHandler(BaseHTTPRequestHandler): diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index 3b62f17fab6..26f725331c6 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -9,8 +9,9 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE from litellm.exceptions import AuthenticationError +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig -from litellm.llms.xai.common_utils import XAIModelInfo, xai_reported_cost_in_usd +from litellm.llms.xai.common_utils import XAIModelInfo, xai_error_status_code, xai_reported_cost_in_usd from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( ResponseAPIUsage, @@ -208,9 +209,8 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): ) if should_use_xai_oauth(litellm_params.model_dump()): - token_file: Final = litellm_params.xai_oauth_token_file try: - api_key = XAIOAuthAuthenticator(auth_file=token_file).get_access_token() + api_key = XAIOAuthAuthenticator(auth_file=litellm_params.xai_oauth_token_file).get_access_token() except XAIOAuthError as exc: raise AuthenticationError( model=model, @@ -246,9 +246,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): api_key: Final = 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: - raw_token_file: Final = litellm_params.get("xai_oauth_token_file") - token_file: Final = raw_token_file if isinstance(raw_token_file, str) else None - api_base = XAIOAuthAuthenticator(auth_file=token_file).get_api_base() + 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 @@ -257,6 +255,15 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return f"{api_base}/responses" + def get_error_class( + self, error_message: str, status_code: int, headers: dict[str, str] | httpx.Headers + ) -> BaseLLMException: + return super().get_error_class( + error_message=error_message, + status_code=xai_error_status_code(status_code, error_message), + headers=headers, + ) + def transform_response_api_response( self, model: str, diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 400b29c72f9..dc6a74ebef4 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1012,7 +1012,7 @@ class ProxyInitializationHelpers: envvar="PROMETHEUS_METRICS_PORT", ) def run_server( - cli_args, + cli_args: tuple[str, ...], host, port, api_base, @@ -1065,22 +1065,17 @@ def run_server( prometheus_metrics_port: int | None, ): if cli_args: - if len(cli_args) >= 2 and cli_args[0] == "xai-oauth" and cli_args[1] == "login": + if cli_args[:2] == ("xai-oauth", "login") and len(cli_args) <= 3: from litellm.llms.xai.oauth import ( XAIOAuthAuthenticator, oauth_auth_file_for_account, ) - from litellm.secret_managers.main import get_secret_str - account: Final = cli_args[2] if len(cli_args) >= 3 else None - token_dir: Final = get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser( - "~/.config/litellm/xai_oauth" - ) - authenticator: Final = ( - XAIOAuthAuthenticator(auth_file=oauth_auth_file_for_account(account, token_dir), token_dir=token_dir) - if account - else XAIOAuthAuthenticator() - ) + try: + auth_file: Final = oauth_auth_file_for_account(cli_args[2]) if len(cli_args) == 3 else None + except ValueError as exc: + raise click.UsageError(str(exc)) from exc + authenticator: Final = XAIOAuthAuthenticator(auth_file=auth_file) auth_data: Final = authenticator.login() click.echo(f"xAI OAuth login successful. Credentials saved to {authenticator.auth_file}.") if auth_data.get("expires_at"): diff --git a/tests/unit/llms/xai/test_xai_oauth.py b/tests/unit/llms/xai/test_xai_oauth.py index 3fc4c35052e..0456c993832 100644 --- a/tests/unit/llms/xai/test_xai_oauth.py +++ b/tests/unit/llms/xai/test_xai_oauth.py @@ -783,7 +783,8 @@ def test_proxy_cli_xai_oauth_login_uses_single_authenticator(monkeypatch): class FakeAuthenticator: auth_file = "/tmp/xai-oauth-auth.json" - def __init__(self): + def __init__(self, auth_file=None): + self.requested_auth_file = auth_file instances.append(self) def login(self): @@ -797,5 +798,6 @@ def test_proxy_cli_xai_oauth_login_uses_single_authenticator(monkeypatch): assert result.exit_code == 0 assert len(instances) == 1 + assert instances[0].requested_auth_file is None assert "Credentials saved to /tmp/xai-oauth-auth.json" in result.output assert "Access token expires at 1234567890" in result.output diff --git a/tests/unit/llms/xai/test_xai_oauth_multi_account.py b/tests/unit/llms/xai/test_xai_oauth_multi_account.py index 254b6cd151d..c611a98d949 100644 --- a/tests/unit/llms/xai/test_xai_oauth_multi_account.py +++ b/tests/unit/llms/xai/test_xai_oauth_multi_account.py @@ -1,287 +1,245 @@ import json import os -import time +from pathlib import Path +import httpx import pytest +import respx from click.testing import CliRunner import litellm -from litellm.litellm_core_utils.exception_mapping_utils import _map_openai_exception +from litellm import Router +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.xai.chat.transformation import XAIChatConfig -from litellm.llms.xai.oauth import XAIOAuthAuthenticator, XAIOAuthError +from litellm.llms.xai.oauth import XAIOAuthAuthenticator, XAIOAuthError, oauth_auth_file_for_account from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig from litellm.types.router import GenericLiteLLMParams +FAR_FUTURE_EXPIRY = 4_102_444_800 +SPENDING_LIMIT_BODY = '{"code":"personal-team-blocked:spending-limit","error":"spending limit reached"}' -def test_authenticator_accepts_explicit_absolute_auth_file(tmp_path): - custom = tmp_path / "custom-dir" / "alice.json" - custom.parent.mkdir(parents=True) - custom.write_text( - json.dumps( - { - "access_token": "alice-token", - "refresh_token": "refresh-token", - "expires_at": time.time() + 3600, - } - ) + +@pytest.fixture(autouse=True) +def _oauth_only(monkeypatch): + monkeypatch.delenv("XAI_API_KEY", raising=False) + monkeypatch.delenv("XAI_OAUTH_AUTH_FILE", raising=False) + monkeypatch.setattr(litellm, "xai_key", None) + monkeypatch.setattr(litellm, "api_key", None) + + +@pytest.fixture +def token_dir(tmp_path, monkeypatch): + directory = tmp_path / "xai_oauth" + directory.mkdir() + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(directory)) + return directory + + +def _write_token(path: Path, access_token: str) -> Path: + path.write_text( + json.dumps({"access_token": access_token, "refresh_token": "refresh", "expires_at": FAR_FUTURE_EXPIRY}) + ) + return path + + +def _chat_headers(token_file: str) -> dict: + return XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True, "xai_oauth_token_file": token_file}, + api_key=None, ) - auth = XAIOAuthAuthenticator(auth_file=str(custom), token_dir=str(custom.parent)) - assert auth.auth_file == os.path.realpath(str(custom)) +def test_absolute_auth_file_inside_token_dir_is_read(token_dir): + alice = _write_token(token_dir / "alice.json", "alice-token") + + auth = XAIOAuthAuthenticator(auth_file=str(alice)) + + assert auth.auth_file == os.path.realpath(alice) assert auth.get_access_token() == "alice-token" -def test_authenticator_relative_auth_file_joins_token_dir(tmp_path, monkeypatch): - token_dir = tmp_path / "xai_oauth" - token_dir.mkdir() - (token_dir / "auth-bob.json").write_text( - json.dumps( - { - "access_token": "bob-token", - "refresh_token": "refresh-token", - "expires_at": time.time() + 3600, - } - ) - ) - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) +def test_relative_auth_file_resolves_inside_token_dir(token_dir): + _write_token(token_dir / "auth-bob.json", "bob-token") - auth = XAIOAuthAuthenticator(auth_file="auth-bob.json") - - assert auth.auth_file == os.path.realpath(str(token_dir / "auth-bob.json")) - assert auth.get_access_token() == "bob-token" + assert XAIOAuthAuthenticator(auth_file="auth-bob.json").get_access_token() == "bob-token" -def test_authenticator_rejects_relative_path_escaping_token_dir(tmp_path, monkeypatch): - token_dir = tmp_path / "xai_oauth" - token_dir.mkdir() - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) +@pytest.mark.parametrize("auth_file", ["../outside.json", ".", "sub/../../outside.json"]) +def test_relative_auth_file_escaping_token_dir_is_rejected(token_dir, auth_file): + _write_token(token_dir.parent / "outside.json", "outside-token") with pytest.raises(XAIOAuthError, match="token directory"): - XAIOAuthAuthenticator(auth_file="../secret.json") + XAIOAuthAuthenticator(auth_file=auth_file) -def test_authenticator_rejects_absolute_path_outside_token_dir(tmp_path, monkeypatch): - token_dir = tmp_path / "xai_oauth" - token_dir.mkdir() - outside = tmp_path / "outside.json" - outside.write_text("{}") - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) +def test_absolute_auth_file_outside_token_dir_is_rejected(token_dir): + outside = _write_token(token_dir.parent / "outside.json", "outside-token") with pytest.raises(XAIOAuthError, match="token directory"): XAIOAuthAuthenticator(auth_file=str(outside)) -def test_authenticator_rejects_dotdot_absolute_path_outside_token_dir(tmp_path, monkeypatch): - token_dir = tmp_path / "xai_oauth" - token_dir.mkdir() - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) - escaped = str((token_dir / ".." / "outside.json").resolve()) +def test_symlink_inside_token_dir_pointing_outside_is_rejected(token_dir): + outside = _write_token(token_dir.parent / "outside.json", "outside-token") + (token_dir / "link.json").symlink_to(outside) with pytest.raises(XAIOAuthError, match="token directory"): - XAIOAuthAuthenticator(auth_file=escaped) + XAIOAuthAuthenticator(auth_file="link.json") -def test_oauth_auth_file_for_account_builds_path_inside_token_dir(tmp_path): - from litellm.llms.xai.oauth import oauth_auth_file_for_account - - assert oauth_auth_file_for_account("alice", str(tmp_path)) == os.path.realpath( - os.path.join(str(tmp_path), "auth-alice.json") - ) - assert oauth_auth_file_for_account("Bob_1-2", str(tmp_path)) == os.path.realpath( - os.path.join(str(tmp_path), "auth-Bob_1-2.json") - ) - - -@pytest.mark.parametrize( - "account", - ["", ".", "..", "foo/bar", "../alice", "alice.json", "foo\\bar", "alice/../bob"], -) -def test_oauth_auth_file_for_account_rejects_unsafe_names(account, tmp_path): - from litellm.llms.xai.oauth import oauth_auth_file_for_account - - with pytest.raises(ValueError, match="account"): - oauth_auth_file_for_account(account, str(tmp_path)) - - -def test_authenticator_explicit_auth_file_overrides_env(tmp_path, monkeypatch): - env_file = tmp_path / "env.json" - env_file.write_text( - json.dumps( - { - "access_token": "env-token", - "refresh_token": "r", - "expires_at": time.time() + 3600, - } - ) - ) - explicit_file = tmp_path / "explicit.json" - explicit_file.write_text( - json.dumps( - { - "access_token": "explicit-token", - "refresh_token": "r", - "expires_at": time.time() + 3600, - } - ) - ) +def test_operator_env_auth_file_outside_token_dir_keeps_working(token_dir, monkeypatch): + env_file = _write_token(token_dir.parent / "env.json", "env-token") monkeypatch.setenv("XAI_OAUTH_AUTH_FILE", str(env_file)) - auth = XAIOAuthAuthenticator(auth_file=str(explicit_file), token_dir=str(tmp_path)) - - assert auth.get_access_token() == "explicit-token" + assert XAIOAuthAuthenticator().get_access_token() == "env-token" -def test_chat_config_uses_per_deployment_token_file(tmp_path, monkeypatch): - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path)) - alice_file = tmp_path / "auth-alice.json" - alice_file.write_text( - json.dumps( - { - "access_token": "alice-chat-token", - "refresh_token": "refresh-token", - "expires_at": time.time() + 3600, - } - ) - ) +def test_explicit_auth_file_wins_over_env_auth_file(token_dir, monkeypatch): + monkeypatch.setenv("XAI_OAUTH_AUTH_FILE", str(_write_token(token_dir / "env.json", "env-token"))) + explicit = _write_token(token_dir / "explicit.json", "explicit-token") - headers = XAIChatConfig().validate_environment( - headers={}, - model="grok-4", - messages=[], - optional_params={}, - litellm_params={ - "use_xai_oauth": True, - "xai_oauth_token_file": str(alice_file), - }, - api_key=None, - ) - - assert headers["Authorization"] == "Bearer alice-chat-token" + assert XAIOAuthAuthenticator(auth_file=str(explicit)).get_access_token() == "explicit-token" -def test_chat_config_multi_deployment_token_isolation(tmp_path, monkeypatch): - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path)) - alice = tmp_path / "auth-alice.json" - alice.write_text( - json.dumps( - { - "access_token": "alice-token", - "refresh_token": "r", - "expires_at": time.time() + 3600, - } - ) - ) - bob = tmp_path / "auth-bob.json" - bob.write_text( - json.dumps( - { - "access_token": "bob-token", - "refresh_token": "r", - "expires_at": time.time() + 3600, - } - ) - ) +@pytest.mark.parametrize("account", ["alice", "Bob_1-2"]) +def test_account_login_file_lands_in_token_dir(token_dir, account): + auth = XAIOAuthAuthenticator(auth_file=oauth_auth_file_for_account(account)) - headers_alice = XAIChatConfig().validate_environment( - headers={}, - model="grok-4", - messages=[], - optional_params={}, - litellm_params={ - "use_xai_oauth": True, - "xai_oauth_token_file": str(alice), - }, - api_key=None, - ) - headers_bob = XAIChatConfig().validate_environment( - headers={}, - model="grok-4", - messages=[], - optional_params={}, - litellm_params={ - "use_xai_oauth": True, - "xai_oauth_token_file": str(bob), - }, - api_key=None, - ) - - assert headers_alice["Authorization"] == "Bearer alice-token" - assert headers_bob["Authorization"] == "Bearer bob-token" + assert auth.auth_file == os.path.realpath(token_dir / f"auth-{account}.json") -def test_responses_config_uses_per_deployment_token_file(tmp_path, monkeypatch): - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path)) - bob = tmp_path / "auth-bob.json" - bob.write_text( - json.dumps( - { - "access_token": "bob-responses-token", - "refresh_token": "r", - "expires_at": time.time() + 3600, - } - ) - ) +@pytest.mark.parametrize("account", ["", ".", "..", "foo/bar", "../alice", "alice.json", "foo\\bar", "alice/../bob"]) +def test_unsafe_account_names_are_rejected(account): + with pytest.raises(ValueError, match="account"): + oauth_auth_file_for_account(account) + + +def test_cli_login_rejects_unsafe_account_before_starting_oauth(token_dir): + from litellm.proxy.proxy_cli import run_server + + result = CliRunner().invoke(run_server, ["xai-oauth", "login", "../alice"]) + + assert result.exit_code == 2 + assert "xAI OAuth account must match" in result.output + assert list(token_dir.iterdir()) == [] + + +def test_cli_login_rejects_extra_arguments(token_dir): + from litellm.proxy.proxy_cli import run_server + + result = CliRunner().invoke(run_server, ["xai-oauth", "login", "alice", "bob"]) + + assert result.exit_code == 2 + assert "Unknown command" in result.output + + +def test_chat_deployments_authenticate_with_their_own_token_files(token_dir): + alice = _write_token(token_dir / "auth-alice.json", "alice-token") + bob = _write_token(token_dir / "auth-bob.json", "bob-token") + + assert _chat_headers(str(alice))["Authorization"] == "Bearer alice-token" + assert _chat_headers(str(bob))["Authorization"] == "Bearer bob-token" + + +def test_chat_token_file_outside_token_dir_raises_authentication_error(token_dir): + outside = _write_token(token_dir.parent / "outside.json", "outside-token") + + with pytest.raises(litellm.AuthenticationError, match="token directory"): + _chat_headers(str(outside)) + + +def test_responses_deployment_authenticates_with_its_token_file(token_dir): + bob = _write_token(token_dir / "auth-bob.json", "bob-responses-token") headers = XAIResponsesAPIConfig().validate_environment( headers={}, model="grok-4", - litellm_params=GenericLiteLLMParams( - use_xai_oauth=True, - xai_oauth_token_file=str(bob), - ), + litellm_params=GenericLiteLLMParams(use_xai_oauth=True, xai_oauth_token_file=str(bob)), ) assert headers["Authorization"] == "Bearer bob-responses-token" -def test_proxy_cli_xai_oauth_login_rejects_unsafe_account(monkeypatch, tmp_path): - from litellm.proxy.proxy_cli import run_server +@pytest.mark.parametrize("config", [XAIChatConfig(), XAIResponsesAPIConfig()]) +@pytest.mark.parametrize( + ("status_code", "message", "expected_status"), + [ + (403, SPENDING_LIMIT_BODY, 429), + (403, '{"error":"model access denied"}', 403), + (400, SPENDING_LIMIT_BODY, 400), + ], +) +def test_error_class_reports_spending_limit_403_as_rate_limit(config, status_code, message, expected_status): + try: + error = config.get_error_class(error_message=message, status_code=status_code, headers={}) + except BaseLLMException as raised: + error = raised - monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path)) - result = CliRunner().invoke(run_server, ["xai-oauth", "login", "../alice"]) - - assert result.exit_code != 0 - assert result.exception is not None + assert error.status_code == expected_status + assert error.message == message -def test_spending_limit_403_maps_to_rate_limit_error(): - class FakeHTTPError(Exception): - def __init__(self): - self.status_code = 403 - self.message = "personal-team-blocked:spending-limit" - self.response = None - super().__init__(self.message) - - with pytest.raises(litellm.RateLimitError): - _map_openai_exception( - model="grok-4", - original_exception=FakeHTTPError(), - custom_llm_provider="xai", - error_str="personal-team-blocked:spending-limit", - exception_type="APIStatusError", - exception_provider="XaiException", - extra_information="", - ) +def _chat_completion_body(content: str) -> dict: + return { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "grok-4", + "choices": [{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } -def test_spending_limit_403_on_non_xai_does_not_map_to_rate_limit_error(): - class FakeHTTPError(Exception): - def __init__(self): - self.status_code = 403 - self.message = "personal-team-blocked:spending-limit" - self.response = None - super().__init__(self.message) +def _xai_by_account(request: httpx.Request) -> httpx.Response: + if request.headers["Authorization"] == "Bearer alice-token": + return httpx.Response(403, text=SPENDING_LIMIT_BODY) + return httpx.Response(200, json=_chat_completion_body(f"served with {request.headers['Authorization']}")) - with pytest.raises(litellm.APIError) as exc_info: - _map_openai_exception( - model="gpt-4", - original_exception=FakeHTTPError(), - custom_llm_provider="openai", - error_str="personal-team-blocked:spending-limit", - exception_type="APIStatusError", - exception_provider="OpenAIException", - extra_information="", - ) - assert not isinstance(exc_info.value, litellm.RateLimitError) - assert exc_info.value.status_code == 403 + +def test_spending_limit_403_surfaces_as_rate_limit_error(token_dir): + alice = _write_token(token_dir / "auth-alice.json", "alice-token") + + with respx.mock(assert_all_called=True) as router: + router.post("https://api.x.ai/v1/chat/completions").mock(side_effect=_xai_by_account) + with pytest.raises(litellm.RateLimitError, match="spending-limit"): + litellm.completion( + model="xai/grok-4", + messages=[{"role": "user", "content": "hi"}], + use_xai_oauth=True, + xai_oauth_token_file=str(alice), + ) + + +def test_router_fails_over_to_next_account_when_one_hits_its_spending_limit(token_dir): + alice = _write_token(token_dir / "auth-alice.json", "alice-token") + bob = _write_token(token_dir / "auth-bob.json", "bob-token") + llm_router = Router( + model_list=[ + { + "model_name": "grok", + "litellm_params": { + "model": "xai/grok-4", + "use_xai_oauth": True, + "xai_oauth_token_file": str(account_file), + "order": order, + }, + } + for order, account_file in ((1, alice), (2, bob)) + ], + num_retries=1, + retry_after=0, + ) + + with respx.mock(assert_all_called=True) as mock: + route = mock.post("https://api.x.ai/v1/chat/completions").mock(side_effect=_xai_by_account) + response = llm_router.completion(model="grok", messages=[{"role": "user", "content": "hi"}]) + + assert response.choices[0].message.content == "served with Bearer bob-token" + assert [call.request.headers["Authorization"] for call in route.calls] == [ + "Bearer alice-token", + "Bearer bob-token", + ] diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index a2d944fcf39..99d214e415c 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -85,6 +85,7 @@ CONNECTION_NAMES: Final = ( "litellm_credential_name", "configurable_clientside_auth_params", "use_xai_oauth", + "xai_oauth_token_file", "aws_batch_role_arn", "s3_bucket_name", "s3_region_name",