mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Address xAI OAuth review comments
This commit is contained in:
parent
b893e64f4a
commit
1950e957cd
2 changed files with 84 additions and 17 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue