Address xAI OAuth review comments

This commit is contained in:
Jeremy Chapeau 2026-06-07 12:00:10 -07:00
parent b893e64f4a
commit 1950e957cd
No known key found for this signature in database
2 changed files with 84 additions and 17 deletions

View file

@ -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")

View file

@ -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
):