diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py index a773e78c35a..55d90a78858 100644 --- a/litellm/llms/xai/oauth.py +++ b/litellm/llms/xai/oauth.py @@ -194,7 +194,11 @@ class XAIOAuthAuthenticator: return self.http_client or _get_httpx_client() def _ensure_token_dir(self) -> None: - os.makedirs(self.token_dir, exist_ok=True) + os.makedirs(self.token_dir, mode=0o700, exist_ok=True) + try: + os.chmod(self.token_dir, 0o700) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth token directory") def _read_auth_file(self) -> Optional[Dict[str, Any]]: try: @@ -206,12 +210,34 @@ class XAIOAuthAuthenticator: 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) + tmp_file = os.path.join( + self.token_dir, + f".{os.path.basename(self.auth_file)}.{uuid.uuid4().hex}.tmp", + ) + flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL + if hasattr(os, "O_NOFOLLOW"): + flags |= os.O_NOFOLLOW + fd = os.open(tmp_file, flags, 0o600) try: - os.chmod(self.auth_file, 0o600) - except OSError: - verbose_logger.debug("Could not chmod xAI OAuth auth file") + with os.fdopen(fd, "w") as f: + json.dump(data, f) + f.flush() + os.fsync(f.fileno()) + os.replace(tmp_file, self.auth_file) + try: + os.chmod(self.auth_file, 0o600) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth auth file") + except Exception: + try: + os.close(fd) + except OSError: + pass + try: + os.unlink(tmp_file) + except OSError: + pass + raise def _is_expired(self, auth_data: Dict[str, Any]) -> bool: expires_at = auth_data.get("expires_at") @@ -223,10 +249,15 @@ class XAIOAuthAuthenticator: return True def _discover(self) -> Dict[str, str]: - response = self._client().get( - XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"} - ) - response.raise_for_status() + try: + response = self._client().get( + XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"} + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + raise XAIOAuthError( + f"xAI OAuth discovery request failed: {exc.response.status_code} {exc.response.text}" + ) from exc data = response.json() authorization_endpoint = data.get("authorization_endpoint") token_endpoint = data.get("token_endpoint") diff --git a/tests/test_litellm/llms/xai/test_xai_oauth.py b/tests/test_litellm/llms/xai/test_xai_oauth.py index 3f394244e5d..ea137067886 100644 --- a/tests/test_litellm/llms/xai/test_xai_oauth.py +++ b/tests/test_litellm/llms/xai/test_xai_oauth.py @@ -195,17 +195,34 @@ def test_write_auth_file_creates_private_file(tmp_path, monkeypatch): token_dir = tmp_path / "xai_oauth" monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) authenticator = XAIOAuthAuthenticator() + old_umask = os.umask(0o022) + replace_calls = [] + real_replace = os.replace - authenticator._write_auth_file( - { - "access_token": "access-token", - "refresh_token": "refresh-token", - "expires_at": time.time() + 3600, - } - ) + def assert_private_temp_file(src, dst): + replace_calls.append((src, dst)) + assert oct(os.stat(src).st_mode & 0o777) == "0o600" + with open(src) as f: + assert json.load(f)["refresh_token"] == "refresh-token" + real_replace(src, dst) + + monkeypatch.setattr(os, "replace", assert_private_temp_file) + + try: + authenticator._write_auth_file( + { + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + } + ) + finally: + os.umask(old_umask) stored = json.loads((token_dir / "auth.json").read_text()) assert stored["access_token"] == "access-token" + assert replace_calls + assert oct(os.stat(token_dir).st_mode & 0o777) == "0o700" assert oct(os.stat(token_dir / "auth.json").st_mode & 0o777) == "0o600" @@ -251,6 +268,25 @@ def test_discover_requires_authorization_and_token_endpoints(): authenticator._discover() +def test_discover_wraps_http_errors(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response( + 500, text="discovery failed", request=request + ) + ) + ) + ) + + with pytest.raises(XAIOAuthError) as exc_info: + authenticator._discover() + + assert "xAI OAuth discovery request failed: 500 discovery failed" in str( + exc_info.value + ) + + def test_refresh_discovers_token_endpoint_when_auth_file_is_legacy( tmp_path, monkeypatch ):