mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 5cfaedb981 into 2c9b0e00ac
This commit is contained in:
commit
09204757fb
16 changed files with 421 additions and 12 deletions
|
|
@ -67,6 +67,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
|
|||
"otpm",
|
||||
"use_xai_oauth",
|
||||
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
|
||||
"xai_oauth_token_file",
|
||||
}
|
||||
)
|
||||
| AWS_CREDENTIAL_KWARGS_KEYS
|
||||
|
|
|
|||
|
|
@ -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,12 +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:
|
||||
try:
|
||||
headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().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,
|
||||
|
|
@ -114,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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import base64
|
|||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import sys
|
||||
import threading
|
||||
|
|
@ -30,6 +31,7 @@ XAI_OAUTH_REDIRECT_PORT: Final = 56121
|
|||
XAI_OAUTH_REDIRECT_PATH: Final = "/callback"
|
||||
XAI_OAUTH_EXPIRY_SKEW_SECONDS: Final = 120
|
||||
XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS: Final = 180
|
||||
_XAI_OAUTH_ACCOUNT_NAME_RE: Final = re.compile(r"^[A-Za-z0-9_-]+$")
|
||||
_XAI_OAUTH_REFRESH_LOCK: Final = threading.Lock()
|
||||
|
||||
|
||||
|
|
@ -75,6 +77,30 @@ class XAIOAuthLoginRequiredError(XAIOAuthError):
|
|||
pass
|
||||
|
||||
|
||||
def _default_xai_oauth_token_dir() -> str:
|
||||
return get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser("~/.config/litellm/xai_oauth")
|
||||
|
||||
|
||||
def resolve_xai_oauth_auth_file(auth_file: str | None, token_dir: str) -> str:
|
||||
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) -> str:
|
||||
if not _XAI_OAUTH_ACCOUNT_NAME_RE.fullmatch(account):
|
||||
raise ValueError("xAI OAuth account must match ^[A-Za-z0-9_-]+$")
|
||||
return f"auth-{account}.json"
|
||||
|
||||
|
||||
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):
|
||||
server: "_CallbackServer" # pyright: ignore[reportIncompatibleVariableOverride] # stdlib stubs type server as BaseServer; _CallbackServer is the only server this handler is registered on
|
||||
|
||||
|
|
@ -121,9 +147,14 @@ class _CallbackServer(HTTPServer):
|
|||
|
||||
|
||||
class XAIOAuthAuthenticator:
|
||||
def __init__(self, http_client: httpx.Client | HTTPHandler | None = 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")
|
||||
def __init__(
|
||||
self,
|
||||
http_client: httpx.Client | HTTPHandler | None = None,
|
||||
auth_file: str | None = None,
|
||||
token_dir: str | None = None,
|
||||
) -> None:
|
||||
self.token_dir = token_dir or _default_xai_oauth_token_dir()
|
||||
self.auth_file = resolve_xai_oauth_auth_file(auth_file, self.token_dir)
|
||||
self.http_client = http_client
|
||||
|
||||
def get_api_base(self) -> str:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -209,7 +210,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
if should_use_xai_oauth(litellm_params.model_dump()):
|
||||
try:
|
||||
api_key = XAIOAuthAuthenticator().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,
|
||||
|
|
@ -254,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,
|
||||
|
|
|
|||
|
|
@ -5692,6 +5692,7 @@ def completion(
|
|||
tpm=kwargs.get("tpm"),
|
||||
rpm=kwargs.get("rpm"),
|
||||
use_xai_oauth=kwargs.get("use_xai_oauth", False),
|
||||
xai_oauth_token_file=kwargs.get("xai_oauth_token_file"),
|
||||
gigachat_scope=kwargs.get("gigachat_scope"),
|
||||
gigachat_auth_url=kwargs.get("gigachat_auth_url"),
|
||||
gigachat_access_token=kwargs.get("gigachat_access_token"),
|
||||
|
|
|
|||
|
|
@ -303,6 +303,10 @@ def _build_banned_observability_params() -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
# Credential selectors that name an account stored on the proxy host, so no
|
||||
# client-side opt-in can hand them to a caller. Only the deployment config sets them.
|
||||
_OPERATOR_ONLY_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = ("xai_oauth_token_file",)
|
||||
|
||||
_BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
|
||||
"api_base",
|
||||
"base_url",
|
||||
|
|
@ -327,6 +331,7 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
|
|||
# caller-supplied value is the same exfil shape as
|
||||
# ``aws_web_identity_token`` on the Bedrock path.
|
||||
"azure_ad_token",
|
||||
*_OPERATOR_ONLY_REQUEST_BODY_PARAMS,
|
||||
# Endpoint-targeting fields that retarget the outbound request or
|
||||
# an observability callback. An attacker-controlled value either
|
||||
# exfiltrates the request payload (incl. messages + admin-set
|
||||
|
|
@ -386,6 +391,12 @@ def _check_banned_params(
|
|||
Shared between the root-level check and the nested-config check so a
|
||||
new banned param only needs to be added in one place.
|
||||
"""
|
||||
operator_only: Final = next((param for param in _OPERATOR_ONLY_REQUEST_BODY_PARAMS if param in body), None)
|
||||
if operator_only is not None:
|
||||
raise ValueError(
|
||||
f"Rejected Request: {operator_only} is not allowed in request body. "
|
||||
"Set it on the deployment's litellm_params in your proxy config.yaml instead."
|
||||
)
|
||||
for param in _BANNED_REQUEST_BODY_PARAMS:
|
||||
if param not in body:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -1012,7 +1012,7 @@ class ProxyInitializationHelpers:
|
|||
envvar="PROMETHEUS_METRICS_PORT",
|
||||
)
|
||||
def run_server(
|
||||
cli_args,
|
||||
cli_args: tuple[str, ...],
|
||||
host,
|
||||
port,
|
||||
api_base,
|
||||
|
|
@ -1065,10 +1065,17 @@ def run_server(
|
|||
prometheus_metrics_port: int | None,
|
||||
):
|
||||
if cli_args:
|
||||
if cli_args == ("xai-oauth", "login"):
|
||||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator
|
||||
if cli_args[:2] == ("xai-oauth", "login") and len(cli_args) <= 3:
|
||||
from litellm.llms.xai.oauth import (
|
||||
XAIOAuthAuthenticator,
|
||||
oauth_auth_file_for_account,
|
||||
)
|
||||
|
||||
authenticator: Final = 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"):
|
||||
|
|
|
|||
|
|
@ -87,6 +87,7 @@ class ProviderConnection:
|
|||
litellm_credential_name: str | None = None
|
||||
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
|
||||
use_xai_oauth: bool | None = None
|
||||
xai_oauth_token_file: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
|
|
|
|||
|
|
@ -423,6 +423,14 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
default=False,
|
||||
description="Use stored xAI OAuth credentials when no xAI API key is configured.",
|
||||
)
|
||||
xai_oauth_token_file: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Per-deployment xAI OAuth token file, relative to XAI_OAUTH_TOKEN_DIR or an absolute path inside it. "
|
||||
"With use_xai_oauth=True the deployment authenticates as that account, enabling multi-account "
|
||||
"SuperGrok routing. Deployment config only: the proxy rejects it in request bodies."
|
||||
),
|
||||
)
|
||||
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
|
||||
merge_reasoning_content_in_choices: bool | None = False
|
||||
model_info: dict | None = None
|
||||
|
|
|
|||
|
|
@ -3942,3 +3942,62 @@ class TestIsRequestBodySafeBlocksAwsIdentitySelectors:
|
|||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
class TestIsRequestBodySafeBlocksXaiOauthTokenFile:
|
||||
"""``xai_oauth_token_file`` picks which OAuth account stored on the proxy host signs the
|
||||
request, and router kwargs override deployment params, so no client-side opt-in may unlock it
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"request_body",
|
||||
[
|
||||
{"model": "grok-4", "xai_oauth_token_file": "auth-bob.json"},
|
||||
{"model": "grok-4", "extra_body": {"xai_oauth_token_file": "auth-bob.json"}},
|
||||
{"model": "grok-4", "metadata": {"xai_oauth_token_file": "auth-bob.json"}},
|
||||
{"model": "grok-4", "fallbacks": [{"model": "grok-4", "xai_oauth_token_file": "auth-bob.json"}]},
|
||||
],
|
||||
)
|
||||
def test_rejected_even_under_proxy_wide_opt_in(self, request_body):
|
||||
with pytest.raises(ValueError, match="xai_oauth_token_file"):
|
||||
is_request_body_safe(
|
||||
request_body=request_body,
|
||||
general_settings={"allow_client_side_credentials": True},
|
||||
llm_router=None,
|
||||
model="grok-4",
|
||||
)
|
||||
|
||||
def test_rejected_even_when_deployment_lists_it_as_clientside_configurable(self):
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "grok-4",
|
||||
"litellm_params": {
|
||||
"model": "xai/grok-4",
|
||||
"use_xai_oauth": True,
|
||||
"xai_oauth_token_file": "auth-alice.json",
|
||||
"configurable_clientside_auth_params": ["xai_oauth_token_file"],
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with pytest.raises(ValueError, match="xai_oauth_token_file"):
|
||||
is_request_body_safe(
|
||||
request_body={"model": "grok-4", "xai_oauth_token_file": "auth-bob.json"},
|
||||
general_settings={},
|
||||
llm_router=router,
|
||||
model="grok-4",
|
||||
)
|
||||
|
||||
def test_other_banned_params_still_honor_proxy_wide_opt_in(self):
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body={"model": "grok-4", "api_base": "https://example.invalid/v1"},
|
||||
general_settings={"allow_client_side_credentials": True},
|
||||
llm_router=None,
|
||||
model="grok-4",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.proxy.auth.auth_utils import is_request_body_safe # noqa: E402
|
|||
"base_url",
|
||||
"vertex_credentials",
|
||||
"azure_ad_token",
|
||||
"xai_oauth_token_file",
|
||||
],
|
||||
)
|
||||
def test_banned_param_under_extra_body_is_rejected(banned_param):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
245
tests/unit/llms/xai/test_xai_oauth_multi_account.py
Normal file
245
tests/unit/llms/xai/test_xai_oauth_multi_account.py
Normal file
|
|
@ -0,0 +1,245 @@
|
|||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from click.testing import CliRunner
|
||||
|
||||
import litellm
|
||||
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, 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"}'
|
||||
|
||||
|
||||
@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,
|
||||
)
|
||||
|
||||
|
||||
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_relative_auth_file_resolves_inside_token_dir(token_dir):
|
||||
_write_token(token_dir / "auth-bob.json", "bob-token")
|
||||
|
||||
assert XAIOAuthAuthenticator(auth_file="auth-bob.json").get_access_token() == "bob-token"
|
||||
|
||||
|
||||
@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=auth_file)
|
||||
|
||||
|
||||
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_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="link.json")
|
||||
|
||||
|
||||
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))
|
||||
|
||||
assert XAIOAuthAuthenticator().get_access_token() == "env-token"
|
||||
|
||||
|
||||
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")
|
||||
|
||||
assert XAIOAuthAuthenticator(auth_file=str(explicit)).get_access_token() == "explicit-token"
|
||||
|
||||
|
||||
@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))
|
||||
|
||||
assert auth.auth_file == os.path.realpath(token_dir / f"auth-{account}.json")
|
||||
|
||||
|
||||
@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)),
|
||||
)
|
||||
|
||||
assert headers["Authorization"] == "Bearer bob-responses-token"
|
||||
|
||||
|
||||
@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
|
||||
|
||||
assert error.status_code == expected_status
|
||||
assert error.message == message
|
||||
|
||||
|
||||
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 _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']}"))
|
||||
|
||||
|
||||
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",
|
||||
]
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -33328,6 +33328,11 @@ export interface components {
|
|||
vertex_project?: string | null;
|
||||
/** Watsonx Region Name */
|
||||
watsonx_region_name?: string | null;
|
||||
/**
|
||||
* Xai Oauth Token File
|
||||
* @description Per-deployment xAI OAuth token file, relative to XAI_OAUTH_TOKEN_DIR or an absolute path inside it. With use_xai_oauth=True the deployment authenticates as that account, enabling multi-account SuperGrok routing. Deployment config only: the proxy rejects it in request bodies.
|
||||
*/
|
||||
xai_oauth_token_file?: string | null;
|
||||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
|
|
@ -47147,6 +47152,11 @@ export interface components {
|
|||
vertex_project?: string | null;
|
||||
/** Watsonx Region Name */
|
||||
watsonx_region_name?: string | null;
|
||||
/**
|
||||
* Xai Oauth Token File
|
||||
* @description Per-deployment xAI OAuth token file, relative to XAI_OAUTH_TOKEN_DIR or an absolute path inside it. With use_xai_oauth=True the deployment authenticates as that account, enabling multi-account SuperGrok routing. Deployment config only: the proxy rejects it in request bodies.
|
||||
*/
|
||||
xai_oauth_token_file?: string | null;
|
||||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue