Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 62 additions & 9 deletions openhands-sdk/openhands/sdk/agent/acp_file_credentials.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
CredentialAuthorizationRejected,
CredentialBindingError,
CredentialConflict,
CredentialInvalidResponse,
CredentialNeedsReauthentication,
CredentialSyncError,
ResolvedCredential,
Expand All @@ -30,6 +31,7 @@

_CHATGPT_AUTH_PATH = Path(".codex") / "auth.json"
_MONITOR_INTERVAL_SECONDS = 0.1
_MONITOR_MAX_RETRY_INTERVAL_SECONDS = 5.0
_MONITOR_JOIN_TIMEOUT_SECONDS = 2.0
_STABLE_READ_DELAY_SECONDS = 0.01
_SYNC_RETRY_DELAYS: tuple[float, ...] = (0.1, 0.5)
Expand Down Expand Up @@ -238,22 +240,44 @@ def _cleanup_runtime(self) -> None:
self._closed = True

def _monitor_loop(self) -> None:
while not self._stop.wait(_MONITOR_INTERVAL_SECONDS):
failure_logged = False
retry_interval = _MONITOR_INTERVAL_SECONDS
while not self._stop.wait(retry_interval):
try:
with self._sync_lock:
self._raise_sticky_error()
value = self._read_stable(attempts=1)
if value is not None:
self._sync_value(value)
except (CredentialNeedsReauthentication, CredentialSyncError) as exc:
failure_logged = False
retry_interval = _MONITOR_INTERVAL_SECONDS
except (
CredentialNeedsReauthentication,
CredentialConflict,
CredentialInvalidResponse,
) as exc:
self._set_error(exc)
return
except CredentialSyncError as exc:
self._set_error(exc)
if not failure_logged:
logger.warning("credential_binding_monitor_failed", exc_info=exc)
failure_logged = True
retry_interval = min(
retry_interval * 2,
_MONITOR_MAX_RETRY_INTERVAL_SECONDS,
)
except Exception as exc:
self._set_error(
CredentialSyncError("Codex credential monitoring failed.")
)
logger.warning("credential_binding_monitor_failed", exc_info=exc)
return
if not failure_logged:
logger.warning("credential_binding_monitor_failed", exc_info=exc)
failure_logged = True
retry_interval = min(
retry_interval * 2,
_MONITOR_MAX_RETRY_INTERVAL_SECONDS,
)

def _read_current(self) -> str | None:
with self._lock:
Expand Down Expand Up @@ -426,13 +450,42 @@ def _raise_sticky_error(self) -> None:

def _refresh_authorization_state(self) -> None:
revision = self._authorization_revision()
if revision is None:
return
with self._lock:
if revision == self._binding_authorization_revision:
error = self._error
if (
revision is not None
and revision != self._binding_authorization_revision
):
self._binding_authorization_revision = revision
if isinstance(error, CredentialAuthorizationRejected):
self._error = None
return
if error is None or isinstance(
error,
(
CredentialAuthorizationRejected,
CredentialConflict,
CredentialInvalidResponse,
CredentialNeedsReauthentication,
),
):
return
self._binding_authorization_revision = revision
if isinstance(self._error, CredentialAuthorizationRejected):
try:
self._load()
except (
CredentialAuthorizationRejected,
CredentialConflict,
CredentialInvalidResponse,
CredentialNeedsReauthentication,
) as exc:
with self._lock:
if self._error is error:
self._error = exc
return
except CredentialBindingError:
Comment thread
simonrosenberg marked this conversation as resolved.
return
Comment thread
simonrosenberg marked this conversation as resolved.
with self._lock:
if self._error is error:
self._error = None

def _authorization_revision(self) -> int | None:
Expand Down
104 changes: 103 additions & 1 deletion tests/sdk/agent/test_acp_file_credentials.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,9 +85,11 @@ def __init__(self, value: str) -> None:
super().__init__(value)
self.authorization_revision = 0
self.rejected = True
self.rejection_observed = threading.Event()

async def replace(self, expected_version: str, value: str) -> str:
if self.rejected:
self.rejection_observed.set()
raise CredentialAuthorizationRejected("rejected")
return await super().replace(expected_version, value)

Expand All @@ -98,6 +100,38 @@ def reauthorize(self) -> None:

class FailingBinding(MemoryBinding):
async def replace(self, expected_version: str, value: str) -> str:
self.replace_calls += 1
raise CredentialSyncError("unavailable")


class FlakyBinding(MemoryBinding):
def __init__(self, value: str) -> None:
super().__init__(value)
self.first_replace_failed = threading.Event()

async def replace(self, expected_version: str, value: str) -> str:
if not self.first_replace_failed.is_set():
self.first_replace_failed.set()
raise CredentialSyncError("temporarily unavailable")
return await super().replace(expected_version, value)


class DisappearingBinding(MemoryBinding):
def __init__(self, value: str) -> None:
super().__init__(value)
self.failed_replace = False
self.failed_loads = 0

async def load(self) -> ResolvedCredential:
if not self.failed_replace:
return await super().load()
self.failed_loads += 1
if self.failed_loads == 1:
raise CredentialSyncError("unavailable")
raise CredentialNeedsReauthentication("missing")

async def replace(self, expected_version: str, value: str) -> str:
self.failed_replace = True
raise CredentialSyncError("unavailable")


Expand Down Expand Up @@ -239,6 +273,57 @@ def test_unstable_read_does_not_poison_lifecycle() -> None:
lifecycle.close()


def test_monitor_recovers_after_transient_writeback_failure() -> None:
rotated = _auth("refresh-r1")
binding = FlakyBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
assert lifecycle.path is not None
runtime = cast(Any, lifecycle)
try:
lifecycle.path.write_text(rotated, encoding="utf-8")
assert binding.first_replace_failed.wait(2)
assert runtime._monitor.is_alive()
_wait_for_value(binding, rotated)
lifecycle.flush()
finally:
lifecycle.close()


def test_monitor_logs_persistent_writeback_failure_once() -> None:
binding = FailingBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
assert lifecycle.path is not None
runtime = cast(Any, lifecycle)
try:
with patch(
"openhands.sdk.agent.acp_file_credentials.logger.warning"
) as warning:
lifecycle.path.write_text(_auth("refresh-r1"), encoding="utf-8")
deadline = time.monotonic() + 2
while binding.replace_calls < 2 and time.monotonic() < deadline:
time.sleep(0.02)
assert binding.replace_calls >= 2
assert runtime._monitor.is_alive()
assert warning.call_count == 1
finally:
lifecycle.discard()


def test_monitor_stops_when_recovery_probe_requires_reauthentication() -> None:
binding = DisappearingBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
assert lifecycle.path is not None
runtime = cast(Any, lifecycle)
lifecycle.path.write_text(_auth("refresh-r1"), encoding="utf-8")
assert runtime._monitor is not None
runtime._monitor.join(timeout=2)

assert not runtime._monitor.is_alive()
with pytest.raises(CredentialNeedsReauthentication, match="missing"):
lifecycle.flush()
lifecycle.discard()


def test_unchanged_file_does_not_write() -> None:
binding = MemoryBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
Expand All @@ -259,7 +344,7 @@ def test_ambiguous_committed_write_converges() -> None:
lifecycle.close()


def test_exhausted_writeback_failure_is_sticky() -> None:
def test_writeback_failure_is_retried_after_successful_load() -> None:
binding = FailingBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
assert lifecycle.path is not None
Expand All @@ -274,6 +359,7 @@ def test_exhausted_writeback_failure_is_sticky() -> None:
with pytest.raises(CredentialSyncError, match="unavailable"):
lifecycle.close()

assert binding.replace_calls == 3
assert runtime_dir.exists()
lifecycle.discard()
assert not runtime_dir.exists()
Expand Down Expand Up @@ -331,6 +417,22 @@ def test_reauthorization_clears_authorization_rejection() -> None:
lifecycle.close()


def test_monitor_recovers_after_reauthorization() -> None:
rotated = _auth("refresh-r1")
binding = RevokedBinding(_auth("refresh-r0"))
lifecycle, _ = _lifecycle(binding, SecretRegistry())
assert lifecycle.path is not None
runtime = cast(Any, lifecycle)
try:
lifecycle.path.write_text(rotated, encoding="utf-8")
assert binding.rejection_observed.wait(2)
assert runtime._monitor.is_alive()
binding.reauthorize()
_wait_for_value(binding, rotated)
finally:
lifecycle.close()


def test_runtime_state_does_not_serialize_binding_values() -> None:
secret = _auth("never-serialize")
binding = MemoryBinding(secret)
Expand Down
Loading