fix(compression): let explicit interrupts cancel safely
Makes interrupt-protected context compression cancellable by an explicit user or lifecycle stop, without weakening protection against ordinary incoming messages, voice interjections, or active-turn redirects. Separates explicit hard cancellation from ordinary interrupt/redirect state with a dedicated threading.Event; introduces AuxiliaryExplicitCancellation as an attempt-local frozen-cause signal; isolates the synchronous provider callback in a bounded daemon worker during protected compression; atomically linearizes Codex timeout cleanup against explicit cancellation; propagates hard cancellation through child agents and explicit stop surfaces; serializes hard-cancel admission against compression commit admission with CompressionCommitFence; aborts before session rotation or late DB commit, restores in-place transcript mutations and compressor state, and releases the heartbeat and compression lease. Based on #74449 by @suparious. Resolved merge conflicts in agent/context_compressor.py (feasibility check + try/except) and tui_gateway/methods_session.py.
This commit is contained in:
@@ -78,6 +78,7 @@ from agent.context_compressor import (
|
||||
COMPRESSED_SUMMARY_METADATA_KEY,
|
||||
ContextCompressor,
|
||||
)
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
from tools.approval import (
|
||||
reset_hermes_interactive_context,
|
||||
set_hermes_interactive_context,
|
||||
@@ -1547,8 +1548,8 @@ class HermesACPAgent(acp.Agent):
|
||||
# redirectable work.
|
||||
state.cancel_event.set()
|
||||
try:
|
||||
if getattr(state, "agent", None) and hasattr(state.agent, "interrupt"):
|
||||
state.agent.interrupt()
|
||||
if getattr(state, "agent", None):
|
||||
request_hard_interrupt(state.agent)
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Failed to interrupt ACP session %s",
|
||||
|
||||
@@ -765,6 +765,9 @@ def init_agent(
|
||||
# Interrupt mechanism for breaking out of tool loops
|
||||
agent._interrupt_requested = False
|
||||
agent._interrupt_message = None # Optional message that triggered interrupt
|
||||
# Explicit hard cancellation is separate from redirect/message state. A
|
||||
# thread-safe Event makes the cause atomic for auxiliary stream pollers.
|
||||
agent._hard_interrupt_requested = threading.Event()
|
||||
agent._execution_thread_id: int | None = None # Set at run_conversation() start
|
||||
agent._interrupt_thread_signal_pending = False
|
||||
agent._client_lock = threading.RLock()
|
||||
|
||||
+249
-8
@@ -228,30 +228,134 @@ def _create_openai_client(*, api_key: str, base_url: str, **kwargs: Any) -> Any:
|
||||
# part-way, compression falls back to a static "summary unavailable" marker
|
||||
# and the real handoff is lost (#23975). A thread-local flag lets such a
|
||||
# task mark its in-flight LLM call as interrupt-protected; the Codex
|
||||
# Responses stream's cancellation check honors it. TIMEOUTS still fire
|
||||
# Responses stream's cancellation check honors it. An explicit host cancel
|
||||
# (CLI Ctrl+C or /stop) may install a cancel check that overrides protection;
|
||||
# ordinary incoming-message interrupts remain protected. TIMEOUTS still fire
|
||||
# (a hung call must die), and all OTHER aux tasks (vision, web_extract,
|
||||
# title_generation, …) remain freely interruptible.
|
||||
_aux_interrupt_protection = threading.local()
|
||||
|
||||
|
||||
class AuxiliaryExplicitCancellation(BaseException):
|
||||
"""Frozen signal that an auxiliary attempt was explicitly hard-cancelled.
|
||||
|
||||
This deliberately follows ``asyncio.CancelledError`` and inherits directly
|
||||
from ``BaseException``: provider retry/fallback code catches ``Exception``
|
||||
broadly and must never reinterpret an explicit host stop as a transport
|
||||
failure. ``cause`` is immutable class data so downstream compression code
|
||||
does not re-query a mutable host Event after the transport has unwound.
|
||||
"""
|
||||
|
||||
cause = "explicit_host_cancel"
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__("auxiliary request explicitly cancelled by host")
|
||||
|
||||
|
||||
def _aux_interrupt_protected() -> bool:
|
||||
return bool(getattr(_aux_interrupt_protection, "active", False))
|
||||
|
||||
|
||||
def _aux_interrupt_cancel_requested() -> bool:
|
||||
"""Return whether an explicit host cancel overrides aux protection."""
|
||||
event = getattr(_aux_interrupt_protection, "cancel_event", None)
|
||||
if event is not None:
|
||||
try:
|
||||
return bool(event.is_set())
|
||||
except Exception:
|
||||
logger.debug("aux interrupt cancel event check failed", exc_info=True)
|
||||
return False
|
||||
check = getattr(_aux_interrupt_protection, "cancel_check", None)
|
||||
if not callable(check):
|
||||
return False
|
||||
try:
|
||||
return bool(check())
|
||||
except Exception:
|
||||
logger.debug("aux interrupt cancel check failed", exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def aux_interrupt_protection(active: bool = True):
|
||||
def aux_interrupt_protection(
|
||||
active: bool = True,
|
||||
cancel_check=None,
|
||||
cancel_event=None,
|
||||
):
|
||||
"""Mark the current thread's auxiliary LLM call as interrupt-protected.
|
||||
|
||||
Used by atomic aux tasks (compression) so a mid-flight gateway interrupt
|
||||
doesn't abort the call and trigger a degraded fallback. Re-entrant-safe:
|
||||
restores the previous value on exit.
|
||||
restores the previous value on exit. ``cancel_check`` lets the host retain
|
||||
an explicit hard-cancel path; ``cancel_event`` is preferred when the host
|
||||
already owns an Event. Nested protection scopes inherit both values.
|
||||
"""
|
||||
prev = getattr(_aux_interrupt_protection, "active", False)
|
||||
prev_cancel_check = getattr(_aux_interrupt_protection, "cancel_check", None)
|
||||
prev_cancel_event = getattr(_aux_interrupt_protection, "cancel_event", None)
|
||||
_aux_interrupt_protection.active = active
|
||||
if callable(cancel_check):
|
||||
_aux_interrupt_protection.cancel_check = cancel_check
|
||||
if cancel_event is not None and callable(getattr(cancel_event, "is_set", None)):
|
||||
_aux_interrupt_protection.cancel_event = cancel_event
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_aux_interrupt_protection.active = prev
|
||||
_aux_interrupt_protection.cancel_check = prev_cancel_check
|
||||
_aux_interrupt_protection.cancel_event = prev_cancel_event
|
||||
|
||||
|
||||
def _capture_aux_cancel_check() -> Optional[Callable[[], Any]]:
|
||||
"""Capture the current explicit-cancel source on the owning request thread."""
|
||||
event = getattr(_aux_interrupt_protection, "cancel_event", None)
|
||||
is_set = getattr(event, "is_set", None)
|
||||
if callable(is_set):
|
||||
return is_set
|
||||
check = getattr(_aux_interrupt_protection, "cancel_check", None)
|
||||
if callable(check):
|
||||
# Preserve callable identity so attempt-local decision objects retain
|
||||
# methods such as begin_timeout_cleanup() when captured by adapters.
|
||||
return check
|
||||
return None
|
||||
|
||||
|
||||
def _captured_aux_cancel_requested(cancel_check: Callable[[], Any]) -> bool:
|
||||
"""Read a request-thread cancellation source without leaking its failures."""
|
||||
try:
|
||||
return bool(cancel_check())
|
||||
except Exception:
|
||||
logger.debug("captured aux cancel check failed", exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
class _AuxiliaryCancellationDecision:
|
||||
"""Atomically choose explicit cancellation or provider timeout per attempt."""
|
||||
|
||||
def __init__(self, source_cancel_check: Callable[[], Any]) -> None:
|
||||
self._source_cancel_check = source_cancel_check
|
||||
self._lock = threading.Lock()
|
||||
self._outcome = "active"
|
||||
|
||||
def __call__(self) -> bool:
|
||||
with self._lock:
|
||||
if self._outcome == "cancelled":
|
||||
return True
|
||||
if self._outcome == "timed_out":
|
||||
return False
|
||||
if _captured_aux_cancel_requested(self._source_cancel_check):
|
||||
self._outcome = "cancelled"
|
||||
return True
|
||||
return False
|
||||
|
||||
def begin_timeout_cleanup(self) -> bool:
|
||||
"""Return whether timeout won and destructive cleanup is permitted."""
|
||||
with self._lock:
|
||||
if self._outcome == "active":
|
||||
if _captured_aux_cancel_requested(self._source_cancel_check):
|
||||
self._outcome = "cancelled"
|
||||
else:
|
||||
self._outcome = "timed_out"
|
||||
return self._outcome == "timed_out"
|
||||
|
||||
|
||||
# ── Forward-progress hook for streamed auxiliary calls ───────────────────
|
||||
@@ -298,6 +402,75 @@ def aux_progress_hook(hook):
|
||||
_aux_progress.hook = prev
|
||||
|
||||
|
||||
def _run_protected_sync_provider_call(
|
||||
callback: Callable[[dict[str, Any]], Any],
|
||||
kwargs: dict[str, Any],
|
||||
) -> Any:
|
||||
"""Run one protected provider callback in an attempt-isolated daemon.
|
||||
|
||||
A hard cancel must release the compression-owning thread promptly, but
|
||||
auxiliary clients are process-shared and cannot safely be closed or evicted
|
||||
to wake one request. Only protected calls with a captured hard-cancel source
|
||||
use this seam. Their provider callback (including stream aggregation) runs
|
||||
in a daemon worker while the owner polls cancellation. On cancel the owner
|
||||
unwinds immediately; the worker is left to finish under the provider timeout
|
||||
already present in ``kwargs``. It owns no transcript or compressor commit
|
||||
state and never holds the session lock.
|
||||
|
||||
Ordinary auxiliary calls, and protected calls without a cancellation source,
|
||||
retain the historical direct synchronous path with no extra thread.
|
||||
"""
|
||||
source_cancel_check = _capture_aux_cancel_check()
|
||||
if not _aux_interrupt_protected() or not callable(source_cancel_check):
|
||||
return callback(kwargs)
|
||||
|
||||
# Freeze one linearized outcome for this isolated attempt. The host Event is
|
||||
# reused and cleared on a later turn, while the Codex timeout Timer may race
|
||||
# owner polling. Both paths must decide under the same attempt-local lock.
|
||||
cancel_check = _AuxiliaryCancellationDecision(source_cancel_check)
|
||||
|
||||
if cancel_check():
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
progress_hook = getattr(_aux_progress, "hook", None)
|
||||
provider_context = contextvars.copy_context()
|
||||
done = threading.Event()
|
||||
outcome: dict[str, Any] = {}
|
||||
|
||||
def _provider_worker() -> None:
|
||||
try:
|
||||
with aux_progress_hook(progress_hook), aux_interrupt_protection(
|
||||
cancel_check=cancel_check
|
||||
):
|
||||
outcome["result"] = callback(kwargs)
|
||||
except BaseException as exc:
|
||||
outcome["exception"] = exc
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
threading.Thread(
|
||||
target=provider_context.run,
|
||||
args=(_provider_worker,),
|
||||
name="hermes-protected-aux-provider",
|
||||
daemon=True,
|
||||
).start()
|
||||
|
||||
while True:
|
||||
# Cancellation is checked before and after every completion wait so it
|
||||
# wins whenever result publication and the host Event become visible in
|
||||
# the same polling interval.
|
||||
if _captured_aux_cancel_requested(cancel_check):
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
if not done.wait(0.02):
|
||||
continue
|
||||
if _captured_aux_cancel_requested(cancel_check):
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
exception = outcome.get("exception")
|
||||
if exception is not None:
|
||||
raise exception
|
||||
return outcome.get("result")
|
||||
|
||||
|
||||
def _safe_isinstance(obj: Any, maybe_type: Any) -> bool:
|
||||
"""Return False instead of raising when a patched symbol is not a type."""
|
||||
try:
|
||||
@@ -1155,12 +1328,53 @@ class _CodexCompletionsAdapter:
|
||||
deadline = time.monotonic() + float(total_timeout) if total_timeout else None
|
||||
timed_out = threading.Event()
|
||||
timeout_timer: Optional[threading.Timer] = None
|
||||
# A protected provider call may outlive its owning compression attempt:
|
||||
# the owner returns promptly on hard cancellation while this adapter is
|
||||
# still blocked in the SDK stream on its isolated worker. Timer threads
|
||||
# do not inherit this worker's thread-local protection state, so freeze
|
||||
# the hard-cancel source here, before creating the timer.
|
||||
protected_cancel_check = (
|
||||
_capture_aux_cancel_check() if _aux_interrupt_protected() else None
|
||||
)
|
||||
attempt_stream_lock = threading.Lock()
|
||||
attempt_stream: List[Any] = []
|
||||
|
||||
def _timeout_message() -> str:
|
||||
return f"Codex auxiliary Responses stream exceeded {float(total_timeout):.1f}s total timeout"
|
||||
|
||||
def _close_client_on_timeout() -> None:
|
||||
begin_timeout_cleanup = getattr(
|
||||
protected_cancel_check, "begin_timeout_cleanup", None
|
||||
)
|
||||
if callable(begin_timeout_cleanup):
|
||||
timeout_won = bool(begin_timeout_cleanup())
|
||||
else:
|
||||
timeout_won = not (
|
||||
callable(protected_cancel_check)
|
||||
and _captured_aux_cancel_requested(protected_cancel_check)
|
||||
)
|
||||
# Publish transport timeout only after the attempt-local decision is
|
||||
# fixed, so owner polling cannot observe completion in between.
|
||||
timed_out.set()
|
||||
if not timeout_won:
|
||||
# The request owner already hard-cancelled this attempt. The
|
||||
# OpenAI client is process-shared, so closing/evicting it here
|
||||
# would disrupt unrelated sessions. Wake only this attempt's
|
||||
# event stream when responses.create() returned one in time;
|
||||
# otherwise rely on the bounded SDK/provider timeout.
|
||||
with attempt_stream_lock:
|
||||
stream = attempt_stream[0] if attempt_stream else None
|
||||
close_stream = getattr(stream, "close", None)
|
||||
if callable(close_stream):
|
||||
try:
|
||||
close_stream()
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Codex auxiliary: cancelled attempt stream close "
|
||||
"during timeout failed",
|
||||
exc_info=True,
|
||||
)
|
||||
return
|
||||
close = getattr(self._client, "close", None)
|
||||
if callable(close):
|
||||
try:
|
||||
@@ -1187,11 +1401,14 @@ class _CodexCompletionsAdapter:
|
||||
from tools.interrupt import is_interrupted
|
||||
# Honor interrupt protection for atomic aux tasks (compression):
|
||||
# a mid-flight gateway interrupt must NOT abort the summary call
|
||||
# and trigger a degraded fallback marker (#23975). Timeouts above
|
||||
# still fire; other aux tasks remain interruptible.
|
||||
# and trigger a degraded fallback marker (#23975). Explicit host
|
||||
# cancellation has its own frozen exception; timeouts above still
|
||||
# fire and other aux tasks remain interruptible.
|
||||
if _aux_interrupt_cancel_requested():
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
if is_interrupted() and not _aux_interrupt_protected():
|
||||
raise InterruptedError("Codex auxiliary Responses stream interrupted")
|
||||
except InterruptedError:
|
||||
except (InterruptedError, AuxiliaryExplicitCancellation):
|
||||
raise
|
||||
except Exception:
|
||||
# Interrupt state is a best-effort UX hook; never make it a
|
||||
@@ -1230,6 +1447,25 @@ class _CodexCompletionsAdapter:
|
||||
_check_cancelled()
|
||||
|
||||
event_stream = self._client.responses.create(**stream_kwargs)
|
||||
with attempt_stream_lock:
|
||||
attempt_stream.append(event_stream)
|
||||
# The timer can fire while responses.create() is blocked. If the
|
||||
# cancelled attempt had no stream to close at that instant, close it
|
||||
# now that it is safely attempt-owned; never touch the shared client.
|
||||
if (
|
||||
timed_out.is_set()
|
||||
and callable(protected_cancel_check)
|
||||
and _captured_aux_cancel_requested(protected_cancel_check)
|
||||
):
|
||||
close_fn = getattr(event_stream, "close", None)
|
||||
if callable(close_fn):
|
||||
try:
|
||||
close_fn()
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Codex auxiliary: late cancelled attempt stream close failed",
|
||||
exc_info=True,
|
||||
)
|
||||
try:
|
||||
# Some Codex-compatible hosts accept ``stream=True`` but return
|
||||
# a completed Responses object instead of an SSE iterator. Do
|
||||
@@ -1251,6 +1487,8 @@ class _CodexCompletionsAdapter:
|
||||
close_fn()
|
||||
except Exception:
|
||||
pass
|
||||
with attempt_stream_lock:
|
||||
attempt_stream.clear()
|
||||
|
||||
if final is None:
|
||||
raise RuntimeError("Codex auxiliary Responses stream did not return a final response")
|
||||
@@ -2672,14 +2910,17 @@ def _relay_sync_completion(
|
||||
) -> Any:
|
||||
callback = create or (lambda request: client.chat.completions.create(**request))
|
||||
route = _relay_auxiliary_metadata(provider=provider, api_mode=api_mode)
|
||||
# Protected compression calls isolate only the provider callback and stream
|
||||
# aggregation. The owning thread remains free to unwind its lease/DB
|
||||
# transaction on hard cancel without touching the process-shared client.
|
||||
if route is None:
|
||||
return callback(kwargs)
|
||||
return _run_protected_sync_provider_call(callback, kwargs)
|
||||
provider_name, fallback_model, metadata = route
|
||||
from agent import relay_llm
|
||||
|
||||
return relay_llm.execute_current(
|
||||
kwargs,
|
||||
callback,
|
||||
lambda request: _run_protected_sync_provider_call(callback, request),
|
||||
name=provider_name,
|
||||
model_name=str(kwargs.get("model") or fallback_model),
|
||||
metadata=metadata,
|
||||
|
||||
@@ -25,7 +25,12 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from agent.auxiliary_client import call_llm, _is_connection_error, aux_interrupt_protection
|
||||
from agent.auxiliary_client import (
|
||||
AuxiliaryExplicitCancellation,
|
||||
_is_connection_error,
|
||||
aux_interrupt_protection,
|
||||
call_llm,
|
||||
)
|
||||
from agent.context_engine import ContextEngine, sanitize_memory_context
|
||||
from agent.error_classifier import FailoverReason, classify_api_error
|
||||
from agent.model_metadata import (
|
||||
@@ -1828,6 +1833,11 @@ class ContextCompressor(ContextEngine):
|
||||
refresh: bool = False,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Return the live compression-failure cooldown for the bound session."""
|
||||
if refresh:
|
||||
# Transaction rollback must distinguish an authoritative empty row
|
||||
# from a failed/unavailable durable read. The public return value
|
||||
# cannot do so because it deliberately falls back to local state.
|
||||
self._last_cooldown_refresh_was_authoritative = None
|
||||
now_mono = time.monotonic()
|
||||
local_state = None
|
||||
if self._summary_failure_cooldown_until > now_mono:
|
||||
@@ -1852,10 +1862,16 @@ class ContextCompressor(ContextEngine):
|
||||
try:
|
||||
state = getter(session_id)
|
||||
except sqlite3.Error as exc:
|
||||
if refresh:
|
||||
self._last_cooldown_refresh_was_authoritative = False
|
||||
logger.debug("compression failure cooldown lookup failed: %s", exc)
|
||||
return local_state
|
||||
except Exception:
|
||||
if refresh:
|
||||
self._last_cooldown_refresh_was_authoritative = False
|
||||
return local_state
|
||||
if refresh:
|
||||
self._last_cooldown_refresh_was_authoritative = True
|
||||
if not state:
|
||||
if refresh:
|
||||
if local_state is not None and self._cooldown_persist_failed:
|
||||
@@ -6056,6 +6072,7 @@ This compaction should PRIORITISE preserving all information related to the focu
|
||||
# — take the narrow rescan, miss a beyond-window fossil, and discard the
|
||||
# rehydrated state as cross-session leakage (#57835).
|
||||
_previous_summary_before_scan = self._previous_summary
|
||||
_summary_has_user_turn_before_scan = getattr(self, "_summary_has_user_turn", None)
|
||||
# A persisted handoff summary can sit in the protected head after a
|
||||
# resume (commonly immediately after the system prompt). Search from
|
||||
# the first non-system message through the compression window. On the
|
||||
@@ -6264,11 +6281,19 @@ This compaction should PRIORITISE preserving all information related to the focu
|
||||
# Deriving the auto focus topic scans recent user turns — only pay
|
||||
# for it when a summary will actually be generated.
|
||||
summary_focus_topic = focus_topic or self._derive_auto_focus_topic(messages)
|
||||
summary = self._generate_summary(
|
||||
turns_to_summarize,
|
||||
focus_topic=summary_focus_topic,
|
||||
memory_context=memory_context,
|
||||
)
|
||||
try:
|
||||
summary = self._generate_summary(
|
||||
turns_to_summarize,
|
||||
focus_topic=summary_focus_topic,
|
||||
memory_context=memory_context,
|
||||
)
|
||||
except AuxiliaryExplicitCancellation:
|
||||
# Explicit cancellation is a true no-op. Restore state mutated by
|
||||
# the resume/handoff self-heal scan before the exception escapes to
|
||||
# the outer transaction, which restores the transcript and lease.
|
||||
self._previous_summary = _previous_summary_before_scan
|
||||
self._summary_has_user_turn = _summary_has_user_turn_before_scan
|
||||
raise
|
||||
|
||||
# If summary generation failed, behavior splits on
|
||||
# ``abort_on_summary_failure`` (config: compression.abort_on_summary_failure):
|
||||
|
||||
@@ -42,6 +42,7 @@ from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
from agent.context_engine import (
|
||||
automatic_compaction_status_message,
|
||||
sanitize_memory_context,
|
||||
@@ -207,6 +208,202 @@ def _cached_prompt_reflects_builtin_memory(agent: Any, cached_prompt: str) -> bo
|
||||
return True
|
||||
|
||||
|
||||
_COMPRESSOR_ATTEMPT_STATE_FIELDS = (
|
||||
"_previous_summary",
|
||||
"_summary_has_user_turn",
|
||||
"compression_count",
|
||||
"_last_compression_savings_pct",
|
||||
"_ineffective_compression_count",
|
||||
"_anti_thrash_recovery_deadline",
|
||||
"_fallback_compression_streak",
|
||||
"_verify_compaction_cleared_threshold",
|
||||
"_last_compression_made_progress",
|
||||
"_summary_failure_cooldown_until",
|
||||
"_cooldown_persist_failed",
|
||||
"_last_summary_error",
|
||||
"_consecutive_timeout_failures",
|
||||
"_last_summary_dropped_count",
|
||||
"_last_summary_fallback_used",
|
||||
"_last_compress_aborted",
|
||||
"_last_summary_auth_failure",
|
||||
"_last_summary_network_failure",
|
||||
"_last_aux_model_failure_error",
|
||||
"_last_aux_model_failure_model",
|
||||
"_summary_model_fallen_back",
|
||||
"summary_model",
|
||||
"_last_compression_telemetry",
|
||||
"_active_compression_telemetry",
|
||||
"_compression_telemetry_seed",
|
||||
)
|
||||
|
||||
_COMPRESSOR_COOLDOWN_STATE_FIELDS = (
|
||||
"_summary_failure_cooldown_until",
|
||||
"_last_summary_error",
|
||||
"_cooldown_persist_failed",
|
||||
)
|
||||
|
||||
|
||||
def _snapshot_compressor_attempt_state(compressor: Any) -> dict[str, Any]:
|
||||
"""Copy only mutable bookkeeping owned by one compression attempt.
|
||||
|
||||
The explicit allow-list avoids copying provider clients, SessionDB handles,
|
||||
locks, and plugin resources. Missing fields are intentionally ignored so
|
||||
legacy and third-party compressors keep their existing contract.
|
||||
"""
|
||||
try:
|
||||
values = vars(compressor)
|
||||
except TypeError:
|
||||
return {}
|
||||
selected = {
|
||||
name: values[name]
|
||||
for name in _COMPRESSOR_ATTEMPT_STATE_FIELDS
|
||||
if name in values
|
||||
}
|
||||
# Copy the collection as one object so aliases between fields (notably
|
||||
# _active_compression_telemetry and _last_compression_telemetry) survive.
|
||||
return copy.deepcopy(selected)
|
||||
|
||||
|
||||
def _restore_compressor_attempt_state(
|
||||
compressor: Any,
|
||||
snapshot: dict[str, Any],
|
||||
*,
|
||||
durable_cooldown_authoritative: Optional[bool] = None,
|
||||
durable_cooldown_state: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""Restore the safe per-attempt snapshot after a pre-commit hard cancel."""
|
||||
# A successful summary clears the durable cooldown before the outer commit
|
||||
# boundary. Recreate (or clear) that row before restoring exact in-memory
|
||||
# values, otherwise the next refresh would overwrite this rollback. Unknown
|
||||
# durable state and intentionally unpersisted local cooldowns are never
|
||||
# converted into destructive DB writes during cancellation.
|
||||
if (
|
||||
"_summary_failure_cooldown_until" in snapshot
|
||||
and durable_cooldown_authoritative is not False
|
||||
and (
|
||||
durable_cooldown_authoritative is True
|
||||
or not bool(snapshot.get("_cooldown_persist_failed", False))
|
||||
)
|
||||
):
|
||||
session_db = vars(compressor).get("_session_db")
|
||||
session_id = vars(compressor).get("_session_id")
|
||||
if session_db is not None and session_id:
|
||||
if durable_cooldown_authoritative is True:
|
||||
restorer = getattr(
|
||||
type(session_db),
|
||||
"restore_compression_failure_cooldown_row",
|
||||
None,
|
||||
)
|
||||
if not callable(restorer) or durable_cooldown_state is None:
|
||||
raise RuntimeError(
|
||||
"exact compression cooldown rollback API is unavailable"
|
||||
)
|
||||
# This API restores raw columns (including expired and null
|
||||
# combinations), verifies the read-back, and propagates failure.
|
||||
restorer(
|
||||
session_db,
|
||||
session_id,
|
||||
copy.deepcopy(durable_cooldown_state),
|
||||
)
|
||||
else:
|
||||
try:
|
||||
deadline = float(
|
||||
snapshot["_summary_failure_cooldown_until"] or 0.0
|
||||
)
|
||||
remaining = max(0.0, deadline - time.monotonic())
|
||||
durable_deadline = time.time() + remaining
|
||||
durable_error = snapshot.get("_last_summary_error")
|
||||
if remaining > 0:
|
||||
recorder = getattr(
|
||||
type(session_db),
|
||||
"record_compression_failure_cooldown",
|
||||
None,
|
||||
)
|
||||
if callable(recorder):
|
||||
recorder(
|
||||
session_db,
|
||||
session_id,
|
||||
durable_deadline,
|
||||
durable_error,
|
||||
)
|
||||
else:
|
||||
clearer = getattr(
|
||||
type(session_db),
|
||||
"clear_compression_failure_cooldown",
|
||||
None,
|
||||
)
|
||||
if callable(clearer):
|
||||
clearer(session_db, session_id)
|
||||
except Exception:
|
||||
# Legacy/third-party compatibility path: its existing APIs
|
||||
# do not provide a verifiable transaction contract.
|
||||
logger.debug(
|
||||
"compression cooldown persistence rollback failed",
|
||||
exc_info=True,
|
||||
)
|
||||
restored = copy.deepcopy(snapshot)
|
||||
for name, value in restored.items():
|
||||
setattr(compressor, name, value)
|
||||
|
||||
|
||||
def _capture_authoritative_cooldown_under_lease(
|
||||
compressor: Any,
|
||||
attempt_snapshot: dict[str, Any],
|
||||
) -> tuple[Optional[bool], Optional[dict[str, Any]]]:
|
||||
"""Refresh and snapshot built-in durable cooldown state under the lease.
|
||||
|
||||
Third-party compressors are deliberately not invoked here: arbitrary plugin
|
||||
callbacks must not run while the session lease is held. A durable read
|
||||
failure returns ``False`` so rollback cannot mistake unknown durable state
|
||||
for an authoritative empty row and clear it; an unavailable legacy API
|
||||
returns ``None`` and preserves the compatibility path.
|
||||
"""
|
||||
try:
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
if not isinstance(compressor, ContextCompressor):
|
||||
return None, None
|
||||
values = vars(compressor)
|
||||
session_db = values.get("_session_db")
|
||||
session_id = values.get("_session_id")
|
||||
raw_reader = (
|
||||
getattr(
|
||||
type(session_db), "get_compression_failure_cooldown_row", None
|
||||
)
|
||||
if session_db is not None
|
||||
else None
|
||||
)
|
||||
if session_db is None or not session_id:
|
||||
# Unbound compressors have no durable row to mutate or restore.
|
||||
return None, None
|
||||
if not callable(raw_reader):
|
||||
return False, None
|
||||
# Capture the exact persisted representation first. The active getter
|
||||
# intentionally filters expired rows and therefore cannot serve as a
|
||||
# lossless rollback snapshot.
|
||||
durable_state = raw_reader(session_db, session_id)
|
||||
if not isinstance(durable_state, dict):
|
||||
raise TypeError("raw compression cooldown snapshot must be a mapping")
|
||||
ContextCompressor.get_active_compression_failure_cooldown(
|
||||
compressor,
|
||||
refresh=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("authoritative compression cooldown capture failed: %s", exc)
|
||||
return False, None
|
||||
authoritative = getattr(
|
||||
compressor, "_last_cooldown_refresh_was_authoritative", None
|
||||
)
|
||||
if authoritative is not True:
|
||||
return authoritative, None
|
||||
|
||||
values = vars(compressor)
|
||||
for name in _COMPRESSOR_COOLDOWN_STATE_FIELDS:
|
||||
if name in values:
|
||||
attempt_snapshot[name] = copy.deepcopy(values[name])
|
||||
return True, copy.deepcopy(durable_state)
|
||||
|
||||
|
||||
class CompressionCommitFence:
|
||||
"""Fence timeout cancellation against post-summary session mutation.
|
||||
|
||||
@@ -241,7 +438,7 @@ class CompressionCommitFence:
|
||||
"""Seconds since the worker last reported forward progress."""
|
||||
return max(0.0, time.monotonic() - self._last_progress)
|
||||
|
||||
def cancel_before_commit(self) -> bool:
|
||||
def cancel_before_commit(self, cancel_event: Any = None) -> bool:
|
||||
"""Cancel a pending commit, or wait for an active commit to finish.
|
||||
|
||||
Returns ``True`` when cancellation won before the commit boundary.
|
||||
@@ -250,8 +447,12 @@ class CompressionCommitFence:
|
||||
"""
|
||||
with self._lock:
|
||||
if self._commit_started:
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
return False
|
||||
self._cancelled = True
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
return True
|
||||
|
||||
def try_cancel_before_commit(self) -> Optional[bool]:
|
||||
@@ -270,10 +471,13 @@ class CompressionCommitFence:
|
||||
finally:
|
||||
self._lock.release()
|
||||
|
||||
def begin_commit(self) -> bool:
|
||||
"""Enter the commit boundary unless cancellation already won."""
|
||||
def begin_commit(self, cancel_event: Any = None) -> bool:
|
||||
"""Atomically admit commit unless a hard cancellation already won."""
|
||||
self._lock.acquire()
|
||||
if self._cancelled:
|
||||
if self._cancelled or (
|
||||
cancel_event is not None and bool(cancel_event.is_set())
|
||||
):
|
||||
self._cancelled = True
|
||||
self._lock.release()
|
||||
return False
|
||||
self._commit_started = True
|
||||
@@ -307,13 +511,21 @@ def _lock_api_is_absent_on_session_db(lock_db: Any) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _refresh_persisted_compression_guards(compressor: Any) -> None:
|
||||
def _refresh_persisted_compression_guards(
|
||||
compressor: Any,
|
||||
*,
|
||||
include_cooldown: bool = True,
|
||||
) -> None:
|
||||
"""Refresh durable automatic-compression guards on a built-in compressor."""
|
||||
method_calls = (
|
||||
("get_active_compression_failure_cooldown", {"refresh": True}),
|
||||
method_calls = [
|
||||
("_load_fallback_compression_streak", {}),
|
||||
("_load_ineffective_compression_count", {}),
|
||||
)
|
||||
]
|
||||
if include_cooldown:
|
||||
method_calls.insert(
|
||||
0,
|
||||
("get_active_compression_failure_cooldown", {"refresh": True}),
|
||||
)
|
||||
for method_name, kwargs in method_calls:
|
||||
method = getattr(type(compressor), method_name, None)
|
||||
if not callable(method):
|
||||
@@ -1298,6 +1510,11 @@ def compress_context(
|
||||
prompt — the session is NOT rotated. Callers should detect the
|
||||
no-op via ``len(returned) == len(input)`` and stop the retry loop.
|
||||
"""
|
||||
_compressor_attempt_snapshot = _snapshot_compressor_attempt_state(
|
||||
agent.context_compressor
|
||||
)
|
||||
_durable_cooldown_authoritative: Optional[bool] = None
|
||||
_durable_cooldown_state: Optional[dict[str, Any]] = None
|
||||
if (
|
||||
defer_context_engine_notification
|
||||
and callable(getattr(agent, _PENDING_CONTEXT_ENGINE_NOTIFICATION, None))
|
||||
@@ -1343,8 +1560,13 @@ def compress_context(
|
||||
if getattr(agent, "api_mode", None) == "codex_app_server":
|
||||
_codex_fence_entered = False
|
||||
if commit_fence is not None:
|
||||
_codex_fence_entered = commit_fence.begin_commit()
|
||||
_codex_fence_entered = commit_fence.begin_commit(
|
||||
getattr(agent, "_hard_interrupt_requested", None)
|
||||
)
|
||||
if not _codex_fence_entered:
|
||||
_restore_compressor_attempt_state(
|
||||
agent.context_compressor, _compressor_attempt_snapshot
|
||||
)
|
||||
existing_prompt = getattr(agent, "_cached_system_prompt", None)
|
||||
if not existing_prompt:
|
||||
existing_prompt = agent._build_system_prompt(system_message)
|
||||
@@ -1671,13 +1893,36 @@ def compress_context(
|
||||
)
|
||||
return messages, _existing_sp
|
||||
|
||||
# Snapshot the authoritative durable cooldown only after this attempt owns
|
||||
# the session lease. This runs for force=True too, but does not apply the
|
||||
# automatic breaker gate: manual compression still retries immediately.
|
||||
_durable_cooldown_authoritative, _durable_cooldown_state = (
|
||||
_capture_authoritative_cooldown_under_lease(
|
||||
agent.context_compressor,
|
||||
_compressor_attempt_snapshot,
|
||||
)
|
||||
)
|
||||
if _durable_cooldown_authoritative is False:
|
||||
# A bound built-in compressor reached its durable getter and the read
|
||||
# failed. Proceeding with force=True could clear an unknown newer row
|
||||
# before cancellation has enough information to restore it. This is a
|
||||
# persistence-safety abort, not automatic breaker gating.
|
||||
_release_lock()
|
||||
existing_prompt = getattr(agent, "_cached_system_prompt", None)
|
||||
if not existing_prompt:
|
||||
existing_prompt = agent._build_system_prompt(system_message)
|
||||
return messages, existing_prompt
|
||||
|
||||
# The agent may have been constructed before another path completed an
|
||||
# in-place compaction on the same session. Re-read durable breaker state
|
||||
# after acquiring the session lock so this final gate cannot act on the
|
||||
# stale snapshot loaded by bind_session_state().
|
||||
if not force:
|
||||
compressor = agent.context_compressor
|
||||
_refresh_persisted_compression_guards(compressor)
|
||||
_refresh_persisted_compression_guards(
|
||||
compressor,
|
||||
include_cooldown=False,
|
||||
)
|
||||
blocked = getattr(
|
||||
type(compressor),
|
||||
"_automatic_compression_blocked",
|
||||
@@ -1691,6 +1936,7 @@ def compress_context(
|
||||
return messages, existing_prompt
|
||||
|
||||
_activity_heartbeat: Optional[_CompressionActivityHeartbeat] = None
|
||||
messages_before_compression = None
|
||||
try:
|
||||
if _lock_holder is not None:
|
||||
_lock_refresher = _CompressionLockLeaseRefresher(
|
||||
@@ -1799,13 +2045,74 @@ def compress_context(
|
||||
# provider that keeps the connection alive forever is cut off at the
|
||||
# streamed total ceiling (see _aux_stream_total_ceiling) instead of
|
||||
# outliving the SDK's inactivity timeout indefinitely.
|
||||
from agent.auxiliary_client import aux_progress_hook
|
||||
from agent.auxiliary_client import (
|
||||
aux_interrupt_protection,
|
||||
aux_progress_hook,
|
||||
)
|
||||
_progress_hook = (
|
||||
commit_fence.touch_progress if commit_fence is not None
|
||||
else (lambda: None)
|
||||
)
|
||||
with aux_progress_hook(_progress_hook):
|
||||
# Incoming-message interrupts and active-turn redirects must not tear an
|
||||
# atomic summary in half (#23975). Explicit stop surfaces set a separate
|
||||
# Event atomically; never infer cause from the racy message fields.
|
||||
_hard_cancel_event = getattr(agent, "_hard_interrupt_requested", None)
|
||||
with aux_progress_hook(_progress_hook), aux_interrupt_protection(
|
||||
cancel_event=_hard_cancel_event
|
||||
):
|
||||
compressed = compress_fn(messages, **compress_kwargs)
|
||||
# Freeze a hard stop that arrived after the final provider attempt
|
||||
# unwound but before this transaction can rotate session state.
|
||||
if _hard_cancel_event is not None and _hard_cancel_event.is_set():
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
except AuxiliaryExplicitCancellation:
|
||||
try:
|
||||
_restore_compressor_attempt_state(
|
||||
agent.context_compressor,
|
||||
_compressor_attempt_snapshot,
|
||||
durable_cooldown_authoritative=_durable_cooldown_authoritative,
|
||||
durable_cooldown_state=_durable_cooldown_state,
|
||||
)
|
||||
except BaseException as _rollback_exc:
|
||||
# Compensation failure must surface, but it must not strand the
|
||||
# session lease or retain an in-memory transcript mutation.
|
||||
if (
|
||||
messages_before_compression is not None
|
||||
and messages != messages_before_compression
|
||||
):
|
||||
messages[:] = copy.deepcopy(messages_before_compression)
|
||||
if _activity_heartbeat is not None:
|
||||
_activity_heartbeat.stop("context compression rollback failed")
|
||||
_activity_heartbeat = None
|
||||
_release_lock()
|
||||
_emit_compression_attempt_telemetry(
|
||||
agent,
|
||||
started_at=_attempt_started_at,
|
||||
commit_status="aborted",
|
||||
split_status="aborted",
|
||||
failure_class=f"rollback:{type(_rollback_exc).__name__}",
|
||||
)
|
||||
raise
|
||||
if (
|
||||
messages_before_compression is not None
|
||||
and messages != messages_before_compression
|
||||
):
|
||||
messages[:] = copy.deepcopy(messages_before_compression)
|
||||
if _activity_heartbeat is not None:
|
||||
_activity_heartbeat.stop("context compression cancelled")
|
||||
_activity_heartbeat = None
|
||||
_release_lock()
|
||||
_emit_compression_attempt_telemetry(
|
||||
agent,
|
||||
started_at=_attempt_started_at,
|
||||
commit_status="aborted",
|
||||
split_status="aborted",
|
||||
failure_class="explicit_interrupt",
|
||||
)
|
||||
_existing_sp = getattr(agent, "_cached_system_prompt", None)
|
||||
if not _existing_sp:
|
||||
_existing_sp = agent._build_system_prompt(system_message)
|
||||
return messages, _existing_sp
|
||||
except BaseException as _compress_exc:
|
||||
# ANY exception after lock acquisition — memory hook, capability
|
||||
# inspection, engine lookup, or compress() — must release the lock so
|
||||
@@ -1918,8 +2225,19 @@ def compress_context(
|
||||
return messages, _existing_sp
|
||||
|
||||
if commit_fence is not None:
|
||||
_commit_fence_entered = commit_fence.begin_commit()
|
||||
_commit_fence_entered = commit_fence.begin_commit(_hard_cancel_event)
|
||||
if not _commit_fence_entered:
|
||||
_restore_compressor_attempt_state(
|
||||
agent.context_compressor,
|
||||
_compressor_attempt_snapshot,
|
||||
durable_cooldown_authoritative=_durable_cooldown_authoritative,
|
||||
durable_cooldown_state=_durable_cooldown_state,
|
||||
)
|
||||
if (
|
||||
messages_before_compression is not None
|
||||
and messages != messages_before_compression
|
||||
):
|
||||
messages[:] = copy.deepcopy(messages_before_compression)
|
||||
logger.info(
|
||||
"Compression commit cancelled before session mutation "
|
||||
"(session=%s).",
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Compatibility helper for explicit agent stop producers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from typing import Any
|
||||
|
||||
|
||||
def request_hard_interrupt(agent: Any, message: str | None = None) -> bool:
|
||||
"""Request an explicit stop, falling back to the legacy interrupt ABI.
|
||||
|
||||
New agents expose ``hard_interrupt(message=None)``. Third-party agents and
|
||||
old test doubles may only expose ``interrupt(message=None)``; keep those
|
||||
usable without sending the newer ``hard_cancel=`` keyword they do not know.
|
||||
Returns ``False`` only when neither callable is available.
|
||||
"""
|
||||
# Avoid treating a dynamic ``__getattr__`` proxy (notably an unspecced
|
||||
# ``MagicMock`` or a third-party RPC facade) as if it genuinely implements
|
||||
# the new ABI. Static lookup proves the attribute exists on the instance or
|
||||
# its type before normal descriptor binding retrieves the callable.
|
||||
try:
|
||||
inspect.getattr_static(agent, "hard_interrupt")
|
||||
except AttributeError:
|
||||
interrupt = None
|
||||
else:
|
||||
interrupt = getattr(agent, "hard_interrupt", None)
|
||||
if not callable(interrupt):
|
||||
interrupt = getattr(agent, "interrupt", None)
|
||||
if not callable(interrupt):
|
||||
return False
|
||||
if message is None:
|
||||
interrupt()
|
||||
else:
|
||||
interrupt(message)
|
||||
return True
|
||||
@@ -21,6 +21,7 @@ from contextlib import contextmanager
|
||||
from concurrent.futures import Future, TimeoutError
|
||||
from typing import Any, Callable, Mapping, Optional
|
||||
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
PUBLIC_CONTRACT_VERSION = 1
|
||||
_MAX_GOAL_CHARS = 16_000
|
||||
@@ -300,16 +301,22 @@ class SubagentLifecycleService:
|
||||
agent = record.agent
|
||||
record.state = SubagentState.CANCEL_REQUESTED
|
||||
record.updated_at = time.time()
|
||||
if agent is None or not hasattr(agent, "interrupt"):
|
||||
if agent is None:
|
||||
return SubagentCancelResult(
|
||||
False, unsupported=True, state=SubagentState.CANCEL_REQUESTED
|
||||
)
|
||||
try:
|
||||
agent.interrupt(f"Lifecycle cancellation requested: {reason[:500]}")
|
||||
accepted = request_hard_interrupt(
|
||||
agent, f"Lifecycle cancellation requested: {reason[:500]}"
|
||||
)
|
||||
except Exception:
|
||||
return SubagentCancelResult(
|
||||
False, unsupported=True, state=SubagentState.CANCEL_REQUESTED
|
||||
)
|
||||
if not accepted:
|
||||
return SubagentCancelResult(
|
||||
False, unsupported=True, state=SubagentState.CANCEL_REQUESTED
|
||||
)
|
||||
return SubagentCancelResult(True, state=SubagentState.CANCEL_REQUESTED)
|
||||
|
||||
def result(self, handle: SubagentHandle) -> SubagentResult:
|
||||
|
||||
@@ -54,6 +54,7 @@ from hermes_cli.fallback_config import get_fallback_chain
|
||||
from hermes_cli.cli_agent_setup_mixin import CLIAgentSetupMixin
|
||||
from hermes_cli.cli_commands_mixin import CLICommandsMixin
|
||||
from hermes_cli.cli_billing_mixin import CLIBillingMixin
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
# prompt_toolkit for fixed input area TUI
|
||||
from prompt_toolkit.history import FileHistory
|
||||
@@ -15834,7 +15835,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
|
||||
self._last_ctrl_c_time = now
|
||||
print("\n⚡ Interrupting agent... (press Ctrl+C again to force exit)")
|
||||
self.agent.interrupt()
|
||||
request_hard_interrupt(self.agent)
|
||||
# If there's text or images, clear them (like bash).
|
||||
# If everything is already empty, exit.
|
||||
elif event.app.current_buffer.text or self._attached_images:
|
||||
@@ -15912,7 +15913,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
|
||||
if self._agent_running and self.agent:
|
||||
print("\n⚡ Interrupting agent...")
|
||||
self.agent.interrupt()
|
||||
request_hard_interrupt(self.agent)
|
||||
elif event.app.current_buffer.text or self._attached_images:
|
||||
event.app.current_buffer.reset()
|
||||
self._attached_images.clear()
|
||||
@@ -17492,8 +17493,11 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
# minutes (#65998 class). Never raises.
|
||||
_arm_exit_watchdog_on_shutdown_signal()
|
||||
try:
|
||||
if getattr(self, "agent", None) and getattr(self, "_agent_running", False):
|
||||
self.agent.interrupt(f"received signal {signum}")
|
||||
_signal_agent = getattr(self, "agent", None)
|
||||
if _signal_agent is not None and getattr(self, "_agent_running", False):
|
||||
request_hard_interrupt(
|
||||
_signal_agent, f"received signal {signum}"
|
||||
)
|
||||
try:
|
||||
_grace = float(os.getenv("HERMES_SIGTERM_GRACE", "1.5"))
|
||||
except (TypeError, ValueError):
|
||||
@@ -17684,7 +17688,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
# avoids wasted API calls and lets run_conversation clean up).
|
||||
if self.agent and getattr(self, '_agent_running', False):
|
||||
try:
|
||||
self.agent.interrupt()
|
||||
request_hard_interrupt(self.agent)
|
||||
except Exception:
|
||||
pass
|
||||
# Shut down voice recorder (release persistent audio stream)
|
||||
@@ -18105,7 +18109,7 @@ def main(
|
||||
try:
|
||||
_agent = getattr(cli, "agent", None)
|
||||
if _agent is not None:
|
||||
_agent.interrupt(f"received signal {signum}")
|
||||
request_hard_interrupt(_agent, f"received signal {signum}")
|
||||
try:
|
||||
_grace = float(os.getenv("HERMES_SIGTERM_GRACE", "1.5"))
|
||||
except (TypeError, ValueError):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
suparious
|
||||
+2
-2
@@ -48,6 +48,7 @@ from hermes_cli.config import (
|
||||
)
|
||||
from hermes_cli.fallback_config import get_fallback_chain
|
||||
from hermes_time import now as _hermes_now
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -3632,8 +3633,7 @@ def run_job(
|
||||
_last_desc, _iter_n, _iter_max,
|
||||
_cur_tool or "none",
|
||||
)
|
||||
if hasattr(agent, "interrupt"):
|
||||
agent.interrupt("Cron job timed out (inactivity)")
|
||||
request_hard_interrupt(agent, "Cron job timed out (inactivity)")
|
||||
raise TimeoutError(
|
||||
f"Cron job '{job_name}' idle for "
|
||||
f"{int(_secs_ago)}s (limit {int(_cron_inactivity_limit)}s) "
|
||||
|
||||
@@ -92,6 +92,7 @@ from gateway.platforms.base import (
|
||||
validate_media_delivery_path,
|
||||
)
|
||||
from agent.redact import redact_sensitive_text
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
from gateway.readiness import collect_runtime_readiness
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -4335,7 +4336,7 @@ class APIServerAdapter(BasePlatformAdapter):
|
||||
agent = agent_ref[0] if agent_ref else None
|
||||
if agent is not None:
|
||||
try:
|
||||
agent.interrupt("SSE client disconnected")
|
||||
request_hard_interrupt(agent, "SSE client disconnected")
|
||||
except Exception:
|
||||
pass
|
||||
_reap_disconnected_agent_processes(agent)
|
||||
@@ -4915,7 +4916,7 @@ class APIServerAdapter(BasePlatformAdapter):
|
||||
agent = agent_ref[0] if agent_ref else None
|
||||
if agent is not None:
|
||||
try:
|
||||
agent.interrupt("SSE client disconnected")
|
||||
request_hard_interrupt(agent, "SSE client disconnected")
|
||||
except Exception:
|
||||
pass
|
||||
_reap_disconnected_agent_processes(agent)
|
||||
@@ -4935,7 +4936,7 @@ class APIServerAdapter(BasePlatformAdapter):
|
||||
agent = agent_ref[0] if agent_ref else None
|
||||
if agent is not None:
|
||||
try:
|
||||
agent.interrupt("SSE task cancelled")
|
||||
request_hard_interrupt(agent, "SSE task cancelled")
|
||||
except Exception:
|
||||
pass
|
||||
# Same abandonment as a client disconnect: the run will never
|
||||
@@ -6788,7 +6789,7 @@ class APIServerAdapter(BasePlatformAdapter):
|
||||
|
||||
if agent is not None:
|
||||
try:
|
||||
agent.interrupt("Stop requested via API")
|
||||
request_hard_interrupt(agent, "Stop requested via API")
|
||||
except Exception:
|
||||
pass
|
||||
# The stopped run is abandoned — reap only the background
|
||||
|
||||
+5
-4
@@ -59,6 +59,7 @@ from agent.conversation_compression import (
|
||||
)
|
||||
from agent.conversation_loop import INTERRUPT_WAITING_FOR_MODEL_PREFIX
|
||||
from agent.i18n import t
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
from hermes_cli.config import cfg_get
|
||||
from hermes_cli.fallback_config import get_fallback_chain
|
||||
|
||||
@@ -9044,7 +9045,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
if agent is _AGENT_PENDING_SENTINEL:
|
||||
continue
|
||||
try:
|
||||
agent.interrupt(reason)
|
||||
request_hard_interrupt(agent, reason)
|
||||
logger.debug("Interrupted running agent for session %s during shutdown", session_key)
|
||||
except Exception as e:
|
||||
logger.debug("Failed interrupting agent during shutdown: %s", e)
|
||||
@@ -22289,7 +22290,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
_process_task_id = ""
|
||||
_process_baseline = None
|
||||
if running_agent and running_agent is not _AGENT_PENDING_SENTINEL:
|
||||
running_agent.interrupt(interrupt_reason)
|
||||
request_hard_interrupt(running_agent, interrupt_reason)
|
||||
_process_task_id = getattr(
|
||||
running_agent, "_gateway_turn_process_task_id", ""
|
||||
)
|
||||
@@ -24529,8 +24530,8 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
|
||||
# Interrupt the agent if it's still running so the thread
|
||||
# pool worker is freed.
|
||||
if _timed_out_agent and hasattr(_timed_out_agent, "interrupt"):
|
||||
_timed_out_agent.interrupt(_INTERRUPT_REASON_TIMEOUT)
|
||||
if _timed_out_agent:
|
||||
request_hard_interrupt(_timed_out_agent, _INTERRUPT_REASON_TIMEOUT)
|
||||
|
||||
_timeout_mins = int(_agent_timeout // 60) or 1
|
||||
|
||||
|
||||
@@ -3599,6 +3599,92 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"error": error,
|
||||
}
|
||||
|
||||
def get_compression_failure_cooldown_row(
|
||||
self,
|
||||
session_id: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""Return the exact stored cooldown columns without expiry filtering.
|
||||
|
||||
Compression cancellation uses this under its session lease so rollback
|
||||
can preserve an expired row, a partially-null row, or an absent session
|
||||
exactly instead of converting those states through the active-cooldown
|
||||
API.
|
||||
"""
|
||||
if not session_id:
|
||||
return {"session_exists": False, "cooldown_until": None, "error": None}
|
||||
with self._lock:
|
||||
row = self._conn.execute(
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return {"session_exists": False, "cooldown_until": None, "error": None}
|
||||
cooldown_until = (
|
||||
row["compression_failure_cooldown_until"]
|
||||
if isinstance(row, sqlite3.Row)
|
||||
else row[0]
|
||||
)
|
||||
error = (
|
||||
row["compression_failure_error"]
|
||||
if isinstance(row, sqlite3.Row)
|
||||
else row[1]
|
||||
)
|
||||
return {
|
||||
"session_exists": True,
|
||||
"cooldown_until": (
|
||||
float(cooldown_until) if cooldown_until is not None else None
|
||||
),
|
||||
"error": error,
|
||||
}
|
||||
|
||||
def restore_compression_failure_cooldown_row(
|
||||
self,
|
||||
session_id: str,
|
||||
snapshot: Dict[str, Any],
|
||||
) -> None:
|
||||
"""Restore and verify an exact cooldown-row snapshot.
|
||||
|
||||
Unlike the ordinary record/clear helpers, this transactional rollback
|
||||
API deliberately propagates write and verification failures. A caller
|
||||
must not report cancellation as mutation-free when compensation failed.
|
||||
"""
|
||||
expected_exists = bool(snapshot.get("session_exists", False))
|
||||
if not expected_exists:
|
||||
actual = self.get_compression_failure_cooldown_row(session_id)
|
||||
if actual.get("session_exists", False):
|
||||
raise RuntimeError(
|
||||
"cannot restore absent compression cooldown row: session now exists"
|
||||
)
|
||||
return
|
||||
|
||||
deadline = snapshot.get("cooldown_until")
|
||||
error = snapshot.get("error")
|
||||
|
||||
def _do(conn):
|
||||
cursor = conn.execute(
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = ?, "
|
||||
"compression_failure_error = ? WHERE id = ?",
|
||||
(deadline, error, session_id),
|
||||
)
|
||||
if cursor.rowcount != 1:
|
||||
raise RuntimeError(
|
||||
f"compression cooldown rollback session missing: {session_id}"
|
||||
)
|
||||
|
||||
self._execute_write(_do)
|
||||
actual = self.get_compression_failure_cooldown_row(session_id)
|
||||
expected = {
|
||||
"session_exists": True,
|
||||
"cooldown_until": float(deadline) if deadline is not None else None,
|
||||
"error": error,
|
||||
}
|
||||
if actual != expected:
|
||||
raise RuntimeError(
|
||||
f"compression cooldown rollback verification failed: "
|
||||
f"expected={expected!r}, actual={actual!r}"
|
||||
)
|
||||
|
||||
def clear_compression_failure_cooldown(self, session_id: str) -> None:
|
||||
"""Clear any persisted compression-failure cooldown for a session."""
|
||||
if not session_id:
|
||||
|
||||
+84
-15
@@ -115,6 +115,7 @@ from agent.process_bootstrap import (
|
||||
_get_proxy_for_base_url,
|
||||
)
|
||||
from agent.iteration_budget import IterationBudget
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
|
||||
from hermes_cli.env_loader import load_hermes_dotenv
|
||||
@@ -2976,7 +2977,7 @@ class AIAgent:
|
||||
logging.warning(f"Failed to save session log: {e}")
|
||||
|
||||
|
||||
def interrupt(self, message: str = None) -> None:
|
||||
def interrupt(self, message: Optional[str] = None, *, hard_cancel: bool = False) -> None:
|
||||
"""
|
||||
Request the agent to interrupt its current tool-calling loop.
|
||||
|
||||
@@ -2989,6 +2990,9 @@ class AIAgent:
|
||||
Args:
|
||||
message: Optional new message that triggered the interrupt.
|
||||
If provided, the agent will include this in its response context.
|
||||
hard_cancel: Mark this as an explicit stop rather than a redirect or
|
||||
incoming-message interrupt. Compression may honor this
|
||||
atomic signal even while ordinary interrupts are masked.
|
||||
|
||||
Example (CLI):
|
||||
# In a separate input thread:
|
||||
@@ -3002,15 +3006,41 @@ class AIAgent:
|
||||
"""
|
||||
# A hard stop and redirect share one lock so /stop cannot race with an
|
||||
# accepted correction and accidentally turn itself into a retry.
|
||||
def _admit_hard_cancel() -> None:
|
||||
event = getattr(self, "_hard_interrupt_requested", None)
|
||||
if event is None:
|
||||
return
|
||||
fence = vars(self).get("_active_compression_commit_fence")
|
||||
cancel_before_commit = getattr(
|
||||
type(fence), "cancel_before_commit", None
|
||||
)
|
||||
if callable(cancel_before_commit):
|
||||
try:
|
||||
# This sets the Event while holding the same lock used by
|
||||
# begin_commit(). If commit already won, it waits for that
|
||||
# tracked mutation to finish before publishing the stop.
|
||||
cancel_before_commit(fence, event)
|
||||
return
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Compression hard-cancel fence admission failed",
|
||||
exc_info=True,
|
||||
)
|
||||
event.set()
|
||||
|
||||
_redirect_lock = getattr(self, "_pending_redirect_lock", None)
|
||||
if _redirect_lock is not None:
|
||||
with _redirect_lock:
|
||||
self._interrupt_requested = True
|
||||
self._interrupt_message = message
|
||||
if hard_cancel:
|
||||
_admit_hard_cancel()
|
||||
self._pending_redirect = None
|
||||
else:
|
||||
self._interrupt_requested = True
|
||||
self._interrupt_message = message
|
||||
if hard_cancel:
|
||||
_admit_hard_cancel()
|
||||
self._pending_redirect = None
|
||||
|
||||
# Codex app-server owns its model/tool loop and watches a private
|
||||
@@ -3073,12 +3103,26 @@ class AIAgent:
|
||||
children_copy = list(self._active_children)
|
||||
for child in children_copy:
|
||||
try:
|
||||
child.interrupt(message)
|
||||
if hard_cancel:
|
||||
request_hard_interrupt(child, message)
|
||||
else:
|
||||
child.interrupt(message)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to propagate interrupt to child agent: %s", e)
|
||||
if not self.quiet_mode:
|
||||
print("\n⚡ Interrupt requested" + (f": '{message[:40]}...'" if message and len(message) > 40 else f": '{message}'" if message else ""))
|
||||
|
||||
def hard_interrupt(self, message: Optional[str] = None) -> None:
|
||||
"""Request an explicit stop while preserving ``interrupt()`` ABI.
|
||||
|
||||
Frontends can feature-detect this method and fall back to the legacy
|
||||
``interrupt()`` signature for synthetic or third-party agents.
|
||||
"""
|
||||
# Deliberately bypass dynamic dispatch: subclasses written against the
|
||||
# legacy interrupt(message=None) ABI may override interrupt without the
|
||||
# newer keyword-only hard_cancel argument.
|
||||
AIAgent.interrupt(self, message, hard_cancel=True)
|
||||
|
||||
def clear_interrupt(self, *, preserve_redirect: bool = False) -> bool:
|
||||
"""Clear the interrupt request and per-thread tool signal.
|
||||
|
||||
@@ -3093,6 +3137,7 @@ class AIAgent:
|
||||
return False
|
||||
self._interrupt_requested = False
|
||||
self._interrupt_message = None
|
||||
getattr(self, "_hard_interrupt_requested", threading.Event()).clear()
|
||||
if not preserve_redirect:
|
||||
self._pending_redirect = None
|
||||
else:
|
||||
@@ -3100,6 +3145,7 @@ class AIAgent:
|
||||
return False
|
||||
self._interrupt_requested = False
|
||||
self._interrupt_message = None
|
||||
getattr(self, "_hard_interrupt_requested", threading.Event()).clear()
|
||||
if not preserve_redirect:
|
||||
self._pending_redirect = None
|
||||
self._interrupt_thread_signal_pending = False
|
||||
@@ -6907,7 +6953,10 @@ class AIAgent:
|
||||
auto-compress abort. Auto-compress callers use the default
|
||||
``force=False``.
|
||||
"""
|
||||
from agent.conversation_compression import compress_context
|
||||
from agent.conversation_compression import (
|
||||
CompressionCommitFence,
|
||||
compress_context,
|
||||
)
|
||||
from agent.portal_tags import (
|
||||
get_conversation_context,
|
||||
reset_conversation_context,
|
||||
@@ -6932,19 +6981,39 @@ class AIAgent:
|
||||
root = self._conversation_root_id()
|
||||
if root:
|
||||
token = set_conversation_context(root)
|
||||
try:
|
||||
return compress_context(
|
||||
self, messages, system_message,
|
||||
approx_tokens=approx_tokens, task_id=task_id, focus_topic=focus_topic,
|
||||
force=force,
|
||||
defer_context_engine_notification=defer_context_engine_notification,
|
||||
commit_fence=commit_fence,
|
||||
# Every AIAgent compression has a fence, including ordinary in-turn and
|
||||
# manual paths. hard_interrupt() uses this exact instance to serialize
|
||||
# cancel admission against begin_commit().
|
||||
active_fence = commit_fence or CompressionCommitFence()
|
||||
# A single agent can receive overlapping automatic/manual entrypoints.
|
||||
# Serialize fence publication so a waiter cannot replace the fence of
|
||||
# the attempt currently generating/committing a summary.
|
||||
fence_registration_lock = vars(self).setdefault(
|
||||
"_compression_commit_fence_lock", threading.RLock()
|
||||
)
|
||||
with fence_registration_lock:
|
||||
missing_fence = object()
|
||||
previous_fence = vars(self).get(
|
||||
"_active_compression_commit_fence", missing_fence
|
||||
)
|
||||
finally:
|
||||
# Restore whatever the caller had, so a compaction never leaks its
|
||||
# tag into the surrounding scope.
|
||||
if token is not None:
|
||||
reset_conversation_context(token)
|
||||
self._active_compression_commit_fence = active_fence
|
||||
try:
|
||||
return compress_context(
|
||||
self, messages, system_message,
|
||||
approx_tokens=approx_tokens, task_id=task_id, focus_topic=focus_topic,
|
||||
force=force,
|
||||
defer_context_engine_notification=defer_context_engine_notification,
|
||||
commit_fence=active_fence,
|
||||
)
|
||||
finally:
|
||||
if previous_fence is missing_fence:
|
||||
vars(self).pop("_active_compression_commit_fence", None)
|
||||
else:
|
||||
self._active_compression_commit_fence = previous_fence
|
||||
# Restore whatever the caller had, so a compaction never leaks its
|
||||
# tag into the surrounding scope.
|
||||
if token is not None:
|
||||
reset_conversation_context(token)
|
||||
|
||||
def _set_tool_guardrail_halt(self, decision: ToolGuardrailDecision) -> None:
|
||||
"""Record the first guardrail decision that should stop this turn."""
|
||||
|
||||
@@ -0,0 +1,622 @@
|
||||
"""Deterministic cross-thread cancellation tests for compression aux transports."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextvars
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Callable
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import auxiliary_client as aux
|
||||
|
||||
|
||||
class _BlockingStream:
|
||||
def __init__(self, started: threading.Event) -> None:
|
||||
self.started = started
|
||||
self.closed = threading.Event()
|
||||
|
||||
def __iter__(self):
|
||||
self.started.set()
|
||||
self.closed.wait(timeout=5)
|
||||
raise RuntimeError("transport closed")
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
|
||||
def get_final_message(self) -> Any:
|
||||
self.started.set()
|
||||
self.closed.wait(timeout=5)
|
||||
raise RuntimeError("transport closed")
|
||||
|
||||
|
||||
class _GenericCompletions:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.stream = stream
|
||||
|
||||
def create(self, **_kwargs: Any) -> _BlockingStream:
|
||||
return self.stream
|
||||
|
||||
|
||||
class _GenericClient:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.chat = SimpleNamespace(completions=_GenericCompletions(stream))
|
||||
self.stream = stream
|
||||
self.closed = threading.Event()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
self.stream.close()
|
||||
|
||||
|
||||
class _CodexResponses:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.stream = stream
|
||||
|
||||
def create(self, **_kwargs: Any) -> _BlockingStream:
|
||||
return self.stream
|
||||
|
||||
|
||||
class _CodexRealClient:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.responses = _CodexResponses(stream)
|
||||
self.api_key = "test"
|
||||
self.base_url = "https://example.test/codex"
|
||||
self.stream = stream
|
||||
self.closed = threading.Event()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
self.stream.close()
|
||||
|
||||
|
||||
class _AnthropicStreamContext:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.stream = stream
|
||||
|
||||
def __enter__(self) -> _BlockingStream:
|
||||
return self.stream
|
||||
|
||||
def __exit__(self, *_args: Any) -> None:
|
||||
self.stream.close()
|
||||
|
||||
|
||||
class _AnthropicMessages:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.stream_obj = stream
|
||||
|
||||
def stream(self, **_kwargs: Any) -> _AnthropicStreamContext:
|
||||
return _AnthropicStreamContext(self.stream_obj)
|
||||
|
||||
|
||||
class _AnthropicRealClient:
|
||||
def __init__(self, stream: _BlockingStream) -> None:
|
||||
self.messages = _AnthropicMessages(stream)
|
||||
self.stream = stream
|
||||
self.closed = threading.Event()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
self.stream.close()
|
||||
|
||||
|
||||
class _BedrockRuntimeClient:
|
||||
def __init__(self, started: threading.Event, release: threading.Event) -> None:
|
||||
self.started = started
|
||||
self.release = release
|
||||
self.closed = threading.Event()
|
||||
|
||||
def converse(self, **_kwargs: Any) -> dict[str, Any]:
|
||||
self.started.set()
|
||||
self.release.wait(timeout=5)
|
||||
return {
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "cancelled response"}],
|
||||
}
|
||||
},
|
||||
"usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2},
|
||||
"stopReason": "end_turn",
|
||||
}
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
|
||||
|
||||
def _cancel_silent_request(
|
||||
client: Any,
|
||||
started: threading.Event,
|
||||
invoke: Callable[[Any], Any],
|
||||
) -> tuple[BaseException, float]:
|
||||
cancel_event = threading.Event()
|
||||
result: dict[str, BaseException] = {}
|
||||
|
||||
def _worker() -> None:
|
||||
try:
|
||||
with aux.aux_interrupt_protection(cancel_event=cancel_event):
|
||||
invoke(client)
|
||||
except BaseException as exc:
|
||||
result["exc"] = exc
|
||||
|
||||
worker = threading.Thread(target=_worker, daemon=True)
|
||||
worker.start()
|
||||
assert started.wait(timeout=1), "request never entered its silent transport"
|
||||
cancelled_at = time.monotonic()
|
||||
cancel_event.set()
|
||||
worker.join(timeout=1)
|
||||
elapsed = time.monotonic() - cancelled_at
|
||||
assert not worker.is_alive(), "explicit cancellation did not wake the silent request"
|
||||
return result["exc"], elapsed
|
||||
|
||||
|
||||
def _invoke_generic(client: Any) -> Any:
|
||||
return aux._relay_sync_completion(
|
||||
client,
|
||||
{"model": "test", "messages": [], "timeout": 30},
|
||||
create=lambda request: aux._create_with_progress(
|
||||
client, request, "compression", force_stream=True
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_protected_silent_provider_is_isolated_and_raises_frozen_explicit_cancel() -> None:
|
||||
started = threading.Event()
|
||||
stream = _BlockingStream(started)
|
||||
client = _GenericClient(stream)
|
||||
|
||||
exc, elapsed = _cancel_silent_request(client, started, _invoke_generic)
|
||||
|
||||
assert isinstance(exc, aux.AuxiliaryExplicitCancellation)
|
||||
assert exc.cause == "explicit_host_cancel"
|
||||
assert not client.closed.is_set()
|
||||
assert elapsed < 0.75
|
||||
stream.close() # release the bounded daemon provider worker
|
||||
|
||||
|
||||
def test_codex_silent_stream_is_isolated_without_closing_shared_client() -> None:
|
||||
started = threading.Event()
|
||||
stream = _BlockingStream(started)
|
||||
real_client = _CodexRealClient(stream)
|
||||
client = aux.CodexAuxiliaryClient(real_client, "gpt-test")
|
||||
|
||||
exc, elapsed = _cancel_silent_request(client, started, _invoke_generic)
|
||||
|
||||
assert isinstance(exc, aux.AuxiliaryExplicitCancellation)
|
||||
assert not real_client.closed.is_set()
|
||||
assert elapsed < 0.75
|
||||
stream.close()
|
||||
|
||||
|
||||
def test_cancelled_codex_orphan_timeout_preserves_cached_shared_client() -> None:
|
||||
"""A cancelled Codex worker's delayed timer owns only its event stream."""
|
||||
owner_started = threading.Event()
|
||||
|
||||
class _SilentOwnerStream:
|
||||
def __init__(self) -> None:
|
||||
self.closed = threading.Event()
|
||||
|
||||
def __iter__(self):
|
||||
owner_started.set()
|
||||
self.closed.wait(timeout=5)
|
||||
raise RuntimeError("owner stream closed")
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
|
||||
class _SuccessStream:
|
||||
def __iter__(self):
|
||||
message = SimpleNamespace(
|
||||
type="message",
|
||||
content=[SimpleNamespace(type="output_text", text="ok")],
|
||||
)
|
||||
return iter(
|
||||
[
|
||||
SimpleNamespace(type="response.output_item.done", item=message),
|
||||
SimpleNamespace(
|
||||
type="response.completed",
|
||||
response=SimpleNamespace(
|
||||
status="completed", id="success", usage=None
|
||||
),
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
pass
|
||||
|
||||
owner_stream = _SilentOwnerStream()
|
||||
|
||||
class _SharedResponses:
|
||||
def __init__(self, real_client: Any) -> None:
|
||||
self.real_client = real_client
|
||||
|
||||
def create(self, **kwargs: Any) -> Any:
|
||||
if self.real_client.closed.is_set():
|
||||
raise RuntimeError("shared client was closed")
|
||||
if kwargs["model"] == "owner":
|
||||
return owner_stream
|
||||
return _SuccessStream()
|
||||
|
||||
class _SharedRealClient:
|
||||
def __init__(self) -> None:
|
||||
self.closed = threading.Event()
|
||||
self.api_key = "test"
|
||||
self.base_url = "https://example.test/codex"
|
||||
self.responses = _SharedResponses(self)
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
owner_stream.close()
|
||||
|
||||
real_client = _SharedRealClient()
|
||||
wrapper = aux.CodexAuxiliaryClient(real_client, "gpt-test")
|
||||
cache_key = ("openai-codex", False, None, None, None)
|
||||
cancel_event = threading.Event()
|
||||
owner_outcome: dict[str, BaseException] = {}
|
||||
|
||||
def _run_owner() -> None:
|
||||
try:
|
||||
with aux.aux_interrupt_protection(cancel_event=cancel_event):
|
||||
aux._relay_sync_completion(
|
||||
wrapper,
|
||||
{"model": "owner", "messages": [], "timeout": 0.12},
|
||||
)
|
||||
except BaseException as exc:
|
||||
owner_outcome["exc"] = exc
|
||||
|
||||
with aux._client_cache_lock:
|
||||
aux._client_cache.clear()
|
||||
aux._client_cache[cache_key] = (wrapper, "gpt-test", None)
|
||||
owner = threading.Thread(target=_run_owner, daemon=True)
|
||||
try:
|
||||
owner.start()
|
||||
assert owner_started.wait(timeout=1)
|
||||
cancel_event.set()
|
||||
owner.join(timeout=1)
|
||||
assert not owner.is_alive()
|
||||
assert isinstance(owner_outcome["exc"], aux.AuxiliaryExplicitCancellation)
|
||||
# A real frontend clears the reusable host Event when the next turn
|
||||
# starts. The orphan must retain a frozen per-attempt cancellation cause.
|
||||
cancel_event.clear()
|
||||
|
||||
# A second user can use the shared client while the cancelled provider
|
||||
# worker is still orphaned and its total-timeout timer is still armed.
|
||||
assert not owner_stream.closed.is_set()
|
||||
concurrent = aux._relay_sync_completion(
|
||||
wrapper,
|
||||
{"model": "concurrent", "messages": [], "timeout": 1},
|
||||
)
|
||||
assert concurrent.choices[0].message.content == "ok"
|
||||
|
||||
# Let the orphan's real adapter timer fire. It may close the attempt's
|
||||
# event stream to wake that worker, but never the process-shared client.
|
||||
assert owner_stream.closed.wait(timeout=1)
|
||||
time.sleep(0.03)
|
||||
assert not real_client.closed.is_set()
|
||||
with aux._client_cache_lock:
|
||||
assert aux._client_cache[cache_key][0] is wrapper
|
||||
|
||||
successive = aux._relay_sync_completion(
|
||||
wrapper,
|
||||
{"model": "successive", "messages": [], "timeout": 1},
|
||||
)
|
||||
assert successive.choices[0].message.content == "ok"
|
||||
finally:
|
||||
owner_stream.close()
|
||||
with aux._client_cache_lock:
|
||||
aux._client_cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("winner", ["timeout", "cancel"])
|
||||
def test_codex_timeout_and_explicit_cancel_have_one_linearized_outcome(
|
||||
winner: str,
|
||||
) -> None:
|
||||
"""Timeout and explicit cancel can never produce a mixed owner/cleanup result."""
|
||||
timer_read_started = threading.Event()
|
||||
allow_timer_read_return = threading.Event()
|
||||
request_cancelled = threading.Event()
|
||||
stream_started = threading.Event()
|
||||
|
||||
class _RacingCancelSource:
|
||||
def is_set(self) -> bool:
|
||||
if winner == "timeout" and threading.current_thread().name.startswith(
|
||||
"Thread-"
|
||||
):
|
||||
# Take the timer's false snapshot, then hold it at the exact seam
|
||||
# where the historical implementation could race owner polling.
|
||||
was_set = request_cancelled.is_set()
|
||||
timer_read_started.set()
|
||||
assert allow_timer_read_return.wait(timeout=1)
|
||||
return was_set
|
||||
return request_cancelled.is_set()
|
||||
|
||||
class _SilentStream:
|
||||
def __init__(self) -> None:
|
||||
self.closed = threading.Event()
|
||||
|
||||
def __iter__(self):
|
||||
stream_started.set()
|
||||
self.closed.wait(timeout=5)
|
||||
raise RuntimeError("stream closed")
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
|
||||
stream = _SilentStream()
|
||||
|
||||
class _RealClient:
|
||||
def __init__(self) -> None:
|
||||
self.api_key = "test"
|
||||
self.base_url = "https://example.test/codex"
|
||||
self.responses = SimpleNamespace(create=lambda **_kwargs: stream)
|
||||
self.closed = threading.Event()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
stream.close()
|
||||
|
||||
real_client: Any = _RealClient()
|
||||
wrapper = aux.CodexAuxiliaryClient(real_client, "gpt-test")
|
||||
owner_outcome: dict[str, BaseException] = {}
|
||||
|
||||
def _run_owner() -> None:
|
||||
try:
|
||||
with aux.aux_interrupt_protection(cancel_event=_RacingCancelSource()):
|
||||
aux._relay_sync_completion(
|
||||
wrapper,
|
||||
{"model": "owner", "messages": [], "timeout": 0.08},
|
||||
)
|
||||
except BaseException as exc:
|
||||
owner_outcome["exc"] = exc
|
||||
|
||||
owner = threading.Thread(target=_run_owner, name="race-owner", daemon=True)
|
||||
owner.start()
|
||||
assert stream_started.wait(timeout=1)
|
||||
if winner == "timeout":
|
||||
assert timer_read_started.wait(timeout=1)
|
||||
request_cancelled.set()
|
||||
allow_timer_read_return.set()
|
||||
else:
|
||||
request_cancelled.set()
|
||||
owner.join(timeout=1)
|
||||
|
||||
assert not owner.is_alive()
|
||||
if winner == "timeout":
|
||||
assert real_client.closed.is_set()
|
||||
assert isinstance(owner_outcome["exc"], TimeoutError)
|
||||
assert not isinstance(owner_outcome["exc"], aux.AuxiliaryExplicitCancellation)
|
||||
else:
|
||||
assert isinstance(owner_outcome["exc"], aux.AuxiliaryExplicitCancellation)
|
||||
assert stream.closed.wait(timeout=1), "cancelled timer did not wake its stream"
|
||||
assert not real_client.closed.is_set()
|
||||
|
||||
|
||||
def test_anthropic_silent_stream_is_isolated_without_closing_shared_client() -> None:
|
||||
started = threading.Event()
|
||||
stream = _BlockingStream(started)
|
||||
real_client = _AnthropicRealClient(stream)
|
||||
client = aux.AnthropicAuxiliaryClient(
|
||||
real_client,
|
||||
"claude-test",
|
||||
"test-key",
|
||||
"https://api.anthropic.test",
|
||||
)
|
||||
|
||||
exc, elapsed = _cancel_silent_request(client, started, _invoke_generic)
|
||||
|
||||
assert isinstance(exc, aux.AuxiliaryExplicitCancellation)
|
||||
assert not real_client.closed.is_set()
|
||||
assert elapsed < 0.75
|
||||
stream.close()
|
||||
|
||||
|
||||
def test_cancelled_attempt_does_not_close_or_fail_concurrent_shared_client_call(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
a_started = threading.Event()
|
||||
a_release = threading.Event()
|
||||
b_started = threading.Event()
|
||||
b_release = threading.Event()
|
||||
closed = threading.Event()
|
||||
|
||||
class _SharedCompletions:
|
||||
def create(self, **kwargs: Any) -> Any:
|
||||
if kwargs["model"] == "session-a":
|
||||
a_started.set()
|
||||
a_release.wait(timeout=5)
|
||||
else:
|
||||
b_started.set()
|
||||
b_release.wait(timeout=5)
|
||||
if closed.is_set():
|
||||
raise RuntimeError("shared client was closed")
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=SimpleNamespace(content="ok"))]
|
||||
)
|
||||
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(completions=_SharedCompletions()),
|
||||
close=lambda: closed.set(),
|
||||
)
|
||||
cancel_event = threading.Event()
|
||||
outcomes: dict[str, Any] = {}
|
||||
evictions: list[Any] = []
|
||||
monkeypatch.setattr(
|
||||
aux, "_evict_cached_client_instance", lambda value: evictions.append(value)
|
||||
)
|
||||
|
||||
def _session_a() -> None:
|
||||
try:
|
||||
with aux.aux_interrupt_protection(cancel_event=cancel_event):
|
||||
aux._relay_sync_completion(
|
||||
client, {"model": "session-a", "messages": [], "timeout": 30}
|
||||
)
|
||||
except BaseException as exc:
|
||||
outcomes["a"] = exc
|
||||
|
||||
def _session_b() -> None:
|
||||
try:
|
||||
outcomes["b"] = aux._relay_sync_completion(
|
||||
client, {"model": "session-b", "messages": [], "timeout": 30}
|
||||
)
|
||||
except BaseException as exc: # pragma: no cover - asserted below
|
||||
outcomes["b"] = exc
|
||||
|
||||
a_thread = threading.Thread(target=_session_a, daemon=True)
|
||||
b_thread = threading.Thread(target=_session_b, daemon=True)
|
||||
a_thread.start()
|
||||
b_thread.start()
|
||||
assert a_started.wait(timeout=1)
|
||||
assert b_started.wait(timeout=1)
|
||||
cancel_event.set()
|
||||
a_thread.join(timeout=1)
|
||||
try:
|
||||
assert not a_thread.is_alive()
|
||||
assert isinstance(outcomes["a"], aux.AuxiliaryExplicitCancellation)
|
||||
assert not closed.is_set()
|
||||
assert evictions == []
|
||||
b_release.set()
|
||||
b_thread.join(timeout=1)
|
||||
assert not b_thread.is_alive()
|
||||
assert not isinstance(outcomes["b"], BaseException)
|
||||
assert outcomes["b"].choices[0].message.content == "ok"
|
||||
finally:
|
||||
a_release.set()
|
||||
b_release.set()
|
||||
|
||||
|
||||
def test_bedrock_silent_nonstream_request_is_isolated_without_close_wakeup() -> None:
|
||||
from agent.bedrock_adapter import _bedrock_runtime_client_cache, reset_client_cache
|
||||
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
runtime_client = _BedrockRuntimeClient(started, release)
|
||||
reset_client_cache()
|
||||
_bedrock_runtime_client_cache["us-test-1"] = runtime_client
|
||||
client = aux.BedrockAuxiliaryClient("us-test-1", "bedrock-test")
|
||||
try:
|
||||
exc, elapsed = _cancel_silent_request(client, started, _invoke_generic)
|
||||
finally:
|
||||
release.set()
|
||||
reset_client_cache()
|
||||
|
||||
assert isinstance(exc, aux.AuxiliaryExplicitCancellation)
|
||||
assert not runtime_client.closed.is_set()
|
||||
assert elapsed < 0.75
|
||||
|
||||
|
||||
def test_unprotected_sync_completion_stays_on_calling_thread() -> None:
|
||||
caller = threading.get_ident()
|
||||
observed: list[int] = []
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(
|
||||
create=lambda **_kwargs: (
|
||||
observed.append(threading.get_ident()),
|
||||
SimpleNamespace(choices=[]),
|
||||
)[1]
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
aux._relay_sync_completion(client, {"model": "test", "messages": []})
|
||||
|
||||
assert observed == [caller]
|
||||
|
||||
|
||||
def test_isolated_provider_worker_inherits_protection_and_progress_hook() -> None:
|
||||
caller = threading.get_ident()
|
||||
cancel_event = threading.Event()
|
||||
progress: list[str] = []
|
||||
observed: dict[str, Any] = {}
|
||||
|
||||
def _create(**_kwargs: Any) -> Any:
|
||||
observed["thread"] = threading.get_ident()
|
||||
observed["protected"] = aux._aux_interrupt_protected()
|
||||
aux._notify_aux_progress()
|
||||
return SimpleNamespace(choices=[])
|
||||
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(completions=SimpleNamespace(create=_create))
|
||||
)
|
||||
with aux.aux_progress_hook(lambda: progress.append("tick")), aux.aux_interrupt_protection(
|
||||
cancel_event=cancel_event
|
||||
):
|
||||
aux._relay_sync_completion(client, {"model": "test", "messages": []})
|
||||
|
||||
assert observed["protected"] is True
|
||||
assert observed["thread"] != caller
|
||||
assert progress == ["tick"]
|
||||
|
||||
|
||||
def test_isolated_provider_worker_inherits_caller_contextvars() -> None:
|
||||
from tools.approval import (
|
||||
get_current_session_key,
|
||||
reset_current_session_key,
|
||||
set_current_session_key,
|
||||
)
|
||||
|
||||
arbitrary = contextvars.ContextVar("isolated-provider-test", default="missing")
|
||||
arbitrary_token = arbitrary.set("caller-value")
|
||||
session_token = set_current_session_key("session-from-caller")
|
||||
observed: dict[str, str] = {}
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(
|
||||
create=lambda **_kwargs: (
|
||||
observed.update(
|
||||
arbitrary=arbitrary.get(),
|
||||
session_key=get_current_session_key(),
|
||||
),
|
||||
SimpleNamespace(choices=[]),
|
||||
)[1]
|
||||
)
|
||||
)
|
||||
)
|
||||
try:
|
||||
with aux.aux_interrupt_protection(cancel_event=threading.Event()):
|
||||
aux._relay_sync_completion(client, {"model": "test", "messages": []})
|
||||
finally:
|
||||
reset_current_session_key(session_token)
|
||||
arbitrary.reset(arbitrary_token)
|
||||
|
||||
assert observed == {
|
||||
"arbitrary": "caller-value",
|
||||
"session_key": "session-from-caller",
|
||||
}
|
||||
|
||||
|
||||
def test_hard_cancel_wins_when_provider_result_is_published_in_same_race() -> None:
|
||||
cancel_event = threading.Event()
|
||||
|
||||
def _create(**_kwargs: Any) -> Any:
|
||||
cancel_event.set()
|
||||
return SimpleNamespace(choices=[])
|
||||
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(completions=SimpleNamespace(create=_create))
|
||||
)
|
||||
with aux.aux_interrupt_protection(cancel_event=cancel_event):
|
||||
with pytest.raises(aux.AuxiliaryExplicitCancellation):
|
||||
aux._relay_sync_completion(client, {"model": "test", "messages": []})
|
||||
|
||||
|
||||
def test_unrelated_interrupted_error_is_not_reclassified_as_explicit_cancel() -> None:
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(
|
||||
create=lambda **_kwargs: (_ for _ in ()).throw(
|
||||
InterruptedError("provider syscall interrupted")
|
||||
)
|
||||
)
|
||||
),
|
||||
close=lambda: None,
|
||||
)
|
||||
|
||||
with aux.aux_interrupt_protection(cancel_event=threading.Event()):
|
||||
with pytest.raises(InterruptedError, match="provider syscall interrupted") as caught:
|
||||
aux._relay_sync_completion(client, {"model": "test", "messages": []})
|
||||
|
||||
assert not isinstance(caught.value, aux.AuxiliaryExplicitCancellation)
|
||||
@@ -28,8 +28,10 @@ fixture deterministically produces 2 children; with the lock, exactly 1.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import inspect
|
||||
import os
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
@@ -40,7 +42,12 @@ import pytest
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _build_agent_with_db(db: SessionDB, session_id: str):
|
||||
def _build_agent_with_db(
|
||||
db: SessionDB,
|
||||
session_id: str,
|
||||
*,
|
||||
stub_compressor: bool = True,
|
||||
):
|
||||
"""Build an AIAgent that's wired to ``db`` and pinned to ``session_id``."""
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
@@ -60,6 +67,9 @@ def _build_agent_with_db(db: SessionDB, session_id: str):
|
||||
# an LLM call. Sleep inside compress() so the two threads' rotations
|
||||
# actually overlap — without that the OS could happen to serialize them
|
||||
# and hide the bug.
|
||||
if not stub_compressor:
|
||||
return agent
|
||||
|
||||
compressor = MagicMock()
|
||||
|
||||
def _compress_with_overlap(*_a, **_kw):
|
||||
@@ -871,7 +881,504 @@ def test_lease_refresher_failure_window_is_bounded_by_ttl() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_hard_interrupt_aborts_compression_and_unblocks_session_writes(tmp_path: Path) -> None:
|
||||
"""Ctrl+C must abort an interrupt-protected summary without leaving the
|
||||
session write-blocked behind its compression lease."""
|
||||
from agent import auxiliary_client as aux
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "HARD_INTERRUPT_COMPRESSION_TEST"
|
||||
db.create_session(session_id, source="cli")
|
||||
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
original_messages = copy.deepcopy(messages)
|
||||
|
||||
def _cancelled_compress(*_args, **_kwargs):
|
||||
agent._hard_interrupt_requested.set()
|
||||
assert aux._aux_interrupt_cancel_requested() is True
|
||||
messages[0]["content"] = "must be rolled back"
|
||||
raise aux.AuxiliaryExplicitCancellation()
|
||||
|
||||
agent.context_compressor.compress.side_effect = _cancelled_compress
|
||||
|
||||
compressed, _prompt = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000
|
||||
)
|
||||
|
||||
assert compressed == original_messages
|
||||
assert messages == original_messages
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
db.append_message(session_id, "assistant", "writes recovered")
|
||||
|
||||
|
||||
def test_late_hard_interrupt_restores_full_compressor_attempt_state_and_retry(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""A stop after provider success but before compress() returns is a true no-op."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "LATE_HARD_INTERRUPT_STATE_TEST"
|
||||
db.create_session(session_id, source="cli")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent.compression_in_place = True
|
||||
agent._cached_system_prompt = "sys"
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
provider_returned = threading.Event()
|
||||
allow_compress_return = threading.Event()
|
||||
shared_telemetry = {"shared": [1, 2, 3]}
|
||||
state_fields = {
|
||||
"_previous_summary": "old-summary",
|
||||
"_summary_has_user_turn": False,
|
||||
"compression_count": 4,
|
||||
"_last_compression_savings_pct": 37.5,
|
||||
"_ineffective_compression_count": 1,
|
||||
"_anti_thrash_recovery_deadline": 123.0,
|
||||
"_fallback_compression_streak": 1,
|
||||
"_verify_compaction_cleared_threshold": False,
|
||||
"_last_compression_made_progress": False,
|
||||
"_summary_failure_cooldown_until": 456.0,
|
||||
"_cooldown_persist_failed": True,
|
||||
"_last_summary_error": "old-error",
|
||||
"_consecutive_timeout_failures": 2,
|
||||
"_last_summary_dropped_count": 3,
|
||||
"_last_summary_fallback_used": True,
|
||||
"_last_compress_aborted": False,
|
||||
"_last_summary_auth_failure": True,
|
||||
"_last_summary_network_failure": True,
|
||||
"_last_aux_model_failure_error": "old-aux-error",
|
||||
"_last_aux_model_failure_model": "old-aux-model",
|
||||
"_summary_model_fallen_back": True,
|
||||
"summary_model": "old-summary-model",
|
||||
"_last_compression_telemetry": shared_telemetry,
|
||||
"_active_compression_telemetry": shared_telemetry,
|
||||
"_compression_telemetry_seed": {"seed": [3]},
|
||||
}
|
||||
for name, value in state_fields.items():
|
||||
setattr(agent.context_compressor, name, copy.deepcopy(value))
|
||||
restored_shared_telemetry = copy.deepcopy(shared_telemetry)
|
||||
agent.context_compressor._last_compression_telemetry = restored_shared_telemetry
|
||||
agent.context_compressor._active_compression_telemetry = restored_shared_telemetry
|
||||
|
||||
def _provider_succeeded_then_waits(*_args, **_kwargs):
|
||||
for name in state_fields:
|
||||
setattr(agent.context_compressor, name, f"mutated-{name}")
|
||||
provider_returned.set()
|
||||
assert allow_compress_return.wait(timeout=5)
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] cancelled summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress.side_effect = _provider_succeeded_then_waits
|
||||
result: dict[str, tuple] = {}
|
||||
worker = threading.Thread(
|
||||
target=lambda: result.setdefault(
|
||||
"value", agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
),
|
||||
daemon=True,
|
||||
)
|
||||
worker.start()
|
||||
assert provider_returned.wait(timeout=2)
|
||||
agent.hard_interrupt("cancel after provider return")
|
||||
allow_compress_return.set()
|
||||
worker.join(timeout=5)
|
||||
|
||||
assert not worker.is_alive()
|
||||
assert result["value"][0] is messages
|
||||
assert {
|
||||
name: copy.deepcopy(getattr(agent.context_compressor, name))
|
||||
for name in state_fields
|
||||
} == state_fields
|
||||
assert (
|
||||
agent.context_compressor._active_compression_telemetry
|
||||
is agent.context_compressor._last_compression_telemetry
|
||||
)
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
agent.clear_interrupt()
|
||||
agent.context_compressor.compress.side_effect = lambda *_a, **_kw: [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] retry summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
retried, _prompt = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000
|
||||
)
|
||||
assert retried is not messages
|
||||
assert retried[0]["content"] == "[CONTEXT COMPACTION] retry summary"
|
||||
|
||||
|
||||
def test_force_cancel_restores_newer_durable_cooldown_captured_under_lease(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""A stale forced attempt rolls back to the lease-protected durable row."""
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "FORCE_CANCEL_DURABLE_COOLDOWN"
|
||||
db.create_session(session_id, source="cli")
|
||||
|
||||
# B binds first and therefore has no local cooldown. A then persists a
|
||||
# newer cooldown for the same durable session before B acquires its lease.
|
||||
stale_agent = _build_agent_with_db(
|
||||
db, session_id, stub_compressor=False
|
||||
)
|
||||
writer_agent = _build_agent_with_db(
|
||||
db, session_id, stub_compressor=False
|
||||
)
|
||||
stale = stale_agent.context_compressor
|
||||
writer = writer_agent.context_compressor
|
||||
assert isinstance(stale, ContextCompressor)
|
||||
assert isinstance(writer, ContextCompressor)
|
||||
assert stale._summary_failure_cooldown_until == 0.0
|
||||
|
||||
writer._record_compression_failure_cooldown(120.0, "newer durable failure")
|
||||
durable_before = tuple(
|
||||
db._conn.execute(
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
)
|
||||
assert durable_before[0] is not None
|
||||
|
||||
stale_seed = {"seed": ["truly-pre-attempt"]}
|
||||
stale._compression_telemetry_seed = copy.deepcopy(stale_seed)
|
||||
stale._previous_summary = "pre-attempt-summary"
|
||||
stale_agent._compression_feasibility_checked = True
|
||||
stale_agent.compression_in_place = True
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
real_clear = ContextCompressor._clear_compression_failure_cooldown
|
||||
|
||||
def _clear_then_hard_cancel() -> None:
|
||||
real_clear(stale)
|
||||
stale_agent._hard_interrupt_requested.set()
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
# Exercise the built-in force=True mutation point deterministically: force
|
||||
# clears the durable cooldown, then the frozen host cancellation unwinds it.
|
||||
stale._clear_compression_failure_cooldown = _clear_then_hard_cancel
|
||||
|
||||
compressed, _prompt = stale_agent._compress_context(
|
||||
messages,
|
||||
"sys",
|
||||
approx_tokens=120_000,
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert compressed is messages
|
||||
durable_after = tuple(
|
||||
db._conn.execute(
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
)
|
||||
assert durable_after == durable_before
|
||||
assert stale._summary_failure_cooldown_until > time.monotonic()
|
||||
assert stale._last_summary_error == "newer durable failure"
|
||||
assert stale._cooldown_persist_failed is False
|
||||
assert stale._compression_telemetry_seed == stale_seed
|
||||
assert stale._previous_summary == "pre-attempt-summary"
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
# A future compressor refresh must still observe the exact row rather than
|
||||
# the cancelled force attempt having permanently cleared it.
|
||||
future_agent = _build_agent_with_db(
|
||||
db, session_id, stub_compressor=False
|
||||
)
|
||||
future = future_agent.context_compressor.get_active_compression_failure_cooldown(
|
||||
refresh=True
|
||||
)
|
||||
assert future is not None
|
||||
assert future["error"] == "newer durable failure"
|
||||
|
||||
|
||||
def test_unrelated_interrupted_error_propagates_and_releases_compression_lease(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""A plugin/OS InterruptedError is a failure, not an explicit transaction abort."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "UNRELATED_INTERRUPT_COMPRESSION_TEST"
|
||||
db.create_session(session_id, source="cli")
|
||||
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
def _provider_interrupted(*_args, **_kwargs):
|
||||
messages[0]["content"] = "must be rolled back"
|
||||
raise InterruptedError("provider syscall interrupted")
|
||||
|
||||
agent.context_compressor.compress.side_effect = _provider_interrupted
|
||||
|
||||
with pytest.raises(InterruptedError, match="provider syscall interrupted"):
|
||||
agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
db.append_message(session_id, "assistant", "writes recovered")
|
||||
|
||||
|
||||
def test_redirect_interrupt_remains_protected_during_compression(tmp_path: Path) -> None:
|
||||
"""Redirects use interrupt_requested=True/message=None; only the atomic
|
||||
hard-cancel event may override summary protection."""
|
||||
from agent import auxiliary_client as aux
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "REDIRECT_COMPRESSION_TEST"
|
||||
db.create_session(session_id, source="cli")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent._interrupt_requested = True
|
||||
agent._interrupt_message = None
|
||||
agent._pending_redirect = "new correction"
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
def _protected_noop(current, **_kwargs):
|
||||
assert aux._aux_interrupt_cancel_requested() is False
|
||||
return copy.deepcopy(current)
|
||||
|
||||
agent.context_compressor.compress.side_effect = _protected_noop
|
||||
|
||||
compressed, _prompt = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000
|
||||
)
|
||||
|
||||
assert compressed == messages
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
|
||||
def test_hard_cancel_between_compress_return_and_commit_begin_wins_atomically(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""The hard-stop admission and commit admission share one fence lock."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "HARD_CANCEL_COMMIT_RACE"
|
||||
db.create_session(session_id, source="tui")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent.compression_in_place = True
|
||||
agent._cached_system_prompt = "sys"
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
before_commit = threading.Event()
|
||||
allow_commit_check = threading.Event()
|
||||
|
||||
class _CommitBarrierList(list):
|
||||
def __eq__(self, other):
|
||||
before_commit.set()
|
||||
assert allow_commit_check.wait(timeout=5)
|
||||
return super().__eq__(other)
|
||||
|
||||
agent.context_compressor.compress.side_effect = lambda *_a, **_kw: _CommitBarrierList(
|
||||
[
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
)
|
||||
archive_spy = MagicMock(wraps=db.archive_and_compact)
|
||||
db.archive_and_compact = archive_spy
|
||||
result: dict[str, tuple] = {}
|
||||
worker = threading.Thread(
|
||||
target=lambda: result.setdefault(
|
||||
"value", agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
),
|
||||
daemon=True,
|
||||
)
|
||||
worker.start()
|
||||
assert before_commit.wait(timeout=2)
|
||||
|
||||
agent.hard_interrupt("cancel before commit admission")
|
||||
allow_commit_check.set()
|
||||
worker.join(timeout=5)
|
||||
|
||||
assert not worker.is_alive()
|
||||
assert result["value"][0] is messages
|
||||
archive_spy.assert_not_called()
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
|
||||
def test_hard_stop_waits_for_commit_already_admitted(tmp_path: Path) -> None:
|
||||
"""A surfaced stop never races an untracked post-return transcript commit."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "HARD_CANCEL_AFTER_COMMIT_ADMISSION"
|
||||
db.create_session(session_id, source="tui")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent.compression_in_place = True
|
||||
agent._cached_system_prompt = "sys"
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
commit_started = threading.Event()
|
||||
allow_commit = threading.Event()
|
||||
stop_returned = threading.Event()
|
||||
original_archive = db.archive_and_compact
|
||||
|
||||
def _blocked_archive(*args, **kwargs):
|
||||
commit_started.set()
|
||||
assert allow_commit.wait(timeout=5)
|
||||
return original_archive(*args, **kwargs)
|
||||
|
||||
db.archive_and_compact = _blocked_archive
|
||||
agent.context_compressor.compress.side_effect = lambda *_a, **_kw: [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
compression_result: dict[str, tuple] = {}
|
||||
compression = threading.Thread(
|
||||
target=lambda: compression_result.setdefault(
|
||||
"value", agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
),
|
||||
daemon=True,
|
||||
)
|
||||
compression.start()
|
||||
assert commit_started.wait(timeout=2)
|
||||
|
||||
stop = threading.Thread(
|
||||
target=lambda: (
|
||||
agent.hard_interrupt("stop after commit admission"),
|
||||
stop_returned.set(),
|
||||
),
|
||||
daemon=True,
|
||||
)
|
||||
stop.start()
|
||||
assert not stop_returned.wait(timeout=0.1)
|
||||
allow_commit.set()
|
||||
compression.join(timeout=5)
|
||||
stop.join(timeout=5)
|
||||
|
||||
assert not compression.is_alive()
|
||||
assert not stop.is_alive()
|
||||
assert stop_returned.is_set()
|
||||
assert compression_result["value"][0][0]["content"] == (
|
||||
"[CONTEXT COMPACTION] summary"
|
||||
)
|
||||
assert agent._hard_interrupt_requested.is_set()
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("deadline_offset", [-10.0, 0.05, None])
|
||||
def test_force_cancel_restores_exact_expired_or_expiring_cooldown_row(
|
||||
tmp_path: Path,
|
||||
deadline_offset: float | None,
|
||||
) -> None:
|
||||
"""Cancellation preserves raw cooldown columns even after their deadline."""
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = f"RAW_COOLDOWN_{deadline_offset}"
|
||||
db.create_session(session_id, source="cli")
|
||||
deadline = time.time() + deadline_offset if deadline_offset is not None else None
|
||||
db.restore_compression_failure_cooldown_row(
|
||||
session_id,
|
||||
{
|
||||
"session_exists": True,
|
||||
"cooldown_until": deadline,
|
||||
"error": "expired-but-exact",
|
||||
},
|
||||
)
|
||||
before = db.get_compression_failure_cooldown_row(session_id)
|
||||
|
||||
agent = _build_agent_with_db(db, session_id, stub_compressor=False)
|
||||
compressor = agent.context_compressor
|
||||
assert isinstance(compressor, ContextCompressor)
|
||||
# A stale local persistence-failure marker must not suppress restoration
|
||||
# once the raw durable row was captured authoritatively under the lease.
|
||||
compressor._cooldown_persist_failed = True
|
||||
agent._compression_feasibility_checked = True
|
||||
agent.compression_in_place = True
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
real_clear = ContextCompressor._clear_compression_failure_cooldown
|
||||
|
||||
def _mutate_then_cancel() -> None:
|
||||
real_clear(compressor)
|
||||
if deadline_offset is not None and deadline_offset > 0:
|
||||
assert deadline is not None
|
||||
while time.time() <= deadline:
|
||||
time.sleep(0.005)
|
||||
agent._hard_interrupt_requested.set()
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
compressor._clear_compression_failure_cooldown = _mutate_then_cancel
|
||||
|
||||
compressed, _prompt = agent._compress_context(
|
||||
messages,
|
||||
"sys",
|
||||
approx_tokens=120_000,
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert compressed is messages
|
||||
assert db.get_compression_failure_cooldown_row(session_id) == before
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
|
||||
def test_cooldown_rollback_failure_surfaces_and_releases_lease(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""A failed compensating write cannot masquerade as a mutation-free cancel."""
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "COOLDOWN_ROLLBACK_WRITE_FAILURE"
|
||||
db.create_session(session_id, source="cli")
|
||||
db.record_compression_failure_cooldown(
|
||||
session_id,
|
||||
time.time() + 120.0,
|
||||
"must-restore",
|
||||
)
|
||||
agent = _build_agent_with_db(db, session_id, stub_compressor=False)
|
||||
compressor = agent.context_compressor
|
||||
assert isinstance(compressor, ContextCompressor)
|
||||
agent._compression_feasibility_checked = True
|
||||
agent.compression_in_place = True
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
real_clear = ContextCompressor._clear_compression_failure_cooldown
|
||||
|
||||
def _mutate_then_cancel() -> None:
|
||||
real_clear(compressor)
|
||||
agent._hard_interrupt_requested.set()
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
compressor._clear_compression_failure_cooldown = _mutate_then_cancel
|
||||
|
||||
def _rollback_write_fails(_self, _session_id, _snapshot) -> None:
|
||||
raise sqlite3.OperationalError("forced rollback write failure")
|
||||
|
||||
monkeypatch.setattr(
|
||||
SessionDB,
|
||||
"restore_compression_failure_cooldown_row",
|
||||
_rollback_write_fails,
|
||||
)
|
||||
|
||||
with pytest.raises(sqlite3.OperationalError, match="forced rollback write failure"):
|
||||
agent._compress_context(
|
||||
messages,
|
||||
"sys",
|
||||
approx_tokens=120_000,
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
|
||||
def test_exact_cooldown_restore_api_propagates_sqlite_write_failure(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "RAW_COOLDOWN_WRITE_FAILURE"
|
||||
db.create_session(session_id, source="test")
|
||||
|
||||
def _write_fails(_callback) -> None:
|
||||
raise sqlite3.OperationalError("forced low-level write failure")
|
||||
|
||||
monkeypatch.setattr(db, "_execute_write", _write_fails)
|
||||
|
||||
with pytest.raises(sqlite3.OperationalError, match="forced low-level write failure"):
|
||||
db.restore_compression_failure_cooldown_row(
|
||||
session_id,
|
||||
{
|
||||
"session_exists": True,
|
||||
"cooldown_until": time.time() + 10.0,
|
||||
"error": "must propagate",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -14,6 +14,8 @@ from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
import agent.auxiliary_client as aux
|
||||
|
||||
|
||||
@@ -38,6 +40,15 @@ class TestAuxInterruptProtection:
|
||||
assert aux._aux_interrupt_protected() is False
|
||||
|
||||
|
||||
def test_nested_protection_preserves_explicit_cancel_check(self):
|
||||
"""A hard-cancel hook installed by the compression host survives the
|
||||
compressor's nested protection scope."""
|
||||
with aux.aux_interrupt_protection(cancel_check=lambda: True):
|
||||
with aux.aux_interrupt_protection():
|
||||
assert aux._aux_interrupt_protected() is True
|
||||
assert aux._aux_interrupt_cancel_requested() is True
|
||||
assert aux._aux_interrupt_cancel_requested() is False
|
||||
|
||||
|
||||
class TestCompressionProtectsSummaryCall:
|
||||
"""The compressor must wrap its summary call_llm in aux_interrupt_protection
|
||||
@@ -82,3 +93,83 @@ class TestCompressionProtectsSummaryCall:
|
||||
)
|
||||
# Protection must be cleared after the call returns.
|
||||
assert aux._aux_interrupt_protected() is False
|
||||
|
||||
def test_explicit_interrupt_is_not_degraded_into_summary_fallback(self):
|
||||
"""Ctrl+C cancellation must escape summary fallback so the outer
|
||||
compression transaction can abort without rotating the session."""
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
with patch("agent.context_compressor.get_model_context_length", return_value=100000):
|
||||
c = ContextCompressor(model="test", quiet_mode=True)
|
||||
|
||||
msgs = [
|
||||
{"role": "user", "content": "do a thing"},
|
||||
{"role": "assistant", "content": "working"},
|
||||
{"role": "user", "content": "more"},
|
||||
{"role": "assistant", "content": "done"},
|
||||
]
|
||||
with aux.aux_interrupt_protection(cancel_check=lambda: True), patch(
|
||||
"agent.context_compressor.call_llm",
|
||||
side_effect=aux.AuxiliaryExplicitCancellation(),
|
||||
):
|
||||
try:
|
||||
c._generate_summary(msgs)
|
||||
except aux.AuxiliaryExplicitCancellation as exc:
|
||||
assert exc.cause == "explicit_host_cancel"
|
||||
else:
|
||||
raise AssertionError("compression swallowed an explicit interrupt")
|
||||
|
||||
def test_non_explicit_interrupted_error_remains_provider_failure(self):
|
||||
"""An unrelated provider/OS InterruptedError must keep the established
|
||||
summary-failure fallback semantics when no host cancel was requested."""
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
with patch("agent.context_compressor.get_model_context_length", return_value=100000):
|
||||
c = ContextCompressor(model="test", quiet_mode=True)
|
||||
|
||||
msgs = [
|
||||
{"role": "user", "content": "do a thing"},
|
||||
{"role": "assistant", "content": "working"},
|
||||
{"role": "user", "content": "more"},
|
||||
{"role": "assistant", "content": "done"},
|
||||
]
|
||||
with patch(
|
||||
"agent.context_compressor.call_llm",
|
||||
side_effect=InterruptedError("provider syscall interrupted"),
|
||||
):
|
||||
assert c._generate_summary(msgs) is None
|
||||
|
||||
def test_explicit_interrupt_restores_rehydration_state(self):
|
||||
"""Cancellation after the handoff scan must be a compressor no-op."""
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
with patch("agent.context_compressor.get_model_context_length", return_value=100000):
|
||||
c = ContextCompressor(
|
||||
model="test",
|
||||
protect_first_n=1,
|
||||
protect_last_n=1,
|
||||
quiet_mode=True,
|
||||
)
|
||||
c._previous_summary = "foreign-session-summary"
|
||||
c._summary_has_user_turn = False
|
||||
msgs = [
|
||||
{"role": "system", "content": "system"},
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"},
|
||||
{"role": "assistant", "content": "four"},
|
||||
{"role": "user", "content": "five"},
|
||||
{"role": "assistant", "content": "six"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
c,
|
||||
"_generate_summary",
|
||||
side_effect=aux.AuxiliaryExplicitCancellation(),
|
||||
):
|
||||
with pytest.raises(aux.AuxiliaryExplicitCancellation):
|
||||
c.compress(msgs)
|
||||
|
||||
assert c._previous_summary == "foreign-session-summary"
|
||||
assert c._summary_has_user_turn is False
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Compatibility contract for explicit hard-stop producers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
|
||||
class _ModernAgent:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, str | None]] = []
|
||||
|
||||
def hard_interrupt(self, message: str | None = None) -> None:
|
||||
self.calls.append(("hard", message))
|
||||
|
||||
def interrupt(self, message: str | None = None) -> None:
|
||||
self.calls.append(("soft", message))
|
||||
|
||||
|
||||
class _LegacyAgent:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, str | None]] = []
|
||||
|
||||
def interrupt(self, message: str | None = None) -> None:
|
||||
self.calls.append(("legacy", message))
|
||||
|
||||
|
||||
def test_explicit_producer_prefers_feature_detected_hard_interrupt() -> None:
|
||||
agent = _ModernAgent()
|
||||
|
||||
assert request_hard_interrupt(agent, "stop now") is True
|
||||
|
||||
assert agent.calls == [("hard", "stop now")]
|
||||
|
||||
|
||||
def test_explicit_producer_falls_back_to_old_interrupt_signature() -> None:
|
||||
agent = _LegacyAgent()
|
||||
|
||||
assert request_hard_interrupt(agent, "stop now") is True
|
||||
|
||||
assert agent.calls == [("legacy", "stop now")]
|
||||
|
||||
|
||||
def test_explicit_producer_reports_unsupported_agent() -> None:
|
||||
assert request_hard_interrupt(object(), "stop now") is False
|
||||
|
||||
|
||||
def test_dynamic_proxy_does_not_fabricate_hard_interrupt_support() -> None:
|
||||
agent = MagicMock()
|
||||
|
||||
assert request_hard_interrupt(agent, "stop now") is True
|
||||
|
||||
agent.interrupt.assert_called_once_with("stop now")
|
||||
agent.hard_interrupt.assert_not_called()
|
||||
|
||||
|
||||
def test_inherited_hard_interrupt_bypasses_legacy_subclass_override() -> None:
|
||||
from run_agent import AIAgent
|
||||
|
||||
class LegacySubclass(AIAgent):
|
||||
def __init__(self) -> None:
|
||||
self.legacy_calls: list[str | None] = []
|
||||
self._hard_interrupt_requested = threading.Event()
|
||||
self._pending_redirect_lock = threading.RLock()
|
||||
self._pending_redirect = None
|
||||
self._execution_thread_id = None
|
||||
self._interrupt_thread_signal_pending = False
|
||||
self._tool_worker_threads: set[int] = set()
|
||||
self._tool_worker_threads_lock = threading.Lock()
|
||||
self._active_children: list[object] = []
|
||||
self._active_children_lock = threading.Lock()
|
||||
self.quiet_mode = True
|
||||
self.api_mode = "test"
|
||||
|
||||
def interrupt(self, message: str | None = None) -> None: # type: ignore[override]
|
||||
self.legacy_calls.append(message)
|
||||
|
||||
agent = LegacySubclass()
|
||||
|
||||
assert request_hard_interrupt(agent, "stop now") is True
|
||||
|
||||
assert agent.legacy_calls == []
|
||||
assert agent._hard_interrupt_requested.is_set()
|
||||
assert agent._interrupt_requested is True
|
||||
assert agent._interrupt_message == "stop now"
|
||||
|
||||
|
||||
def test_tui_subagent_interrupt_is_an_explicit_hard_stop() -> None:
|
||||
import tools.delegate_tool as delegate_tool
|
||||
|
||||
agent = _ModernAgent()
|
||||
subagent_id = "sa-hard-stop-test"
|
||||
with delegate_tool._active_subagents_lock:
|
||||
delegate_tool._active_subagents[subagent_id] = {"agent": agent}
|
||||
try:
|
||||
assert delegate_tool.interrupt_subagent(subagent_id) is True
|
||||
finally:
|
||||
with delegate_tool._active_subagents_lock:
|
||||
delegate_tool._active_subagents.pop(subagent_id, None)
|
||||
|
||||
assert agent.calls == [("hard", f"Interrupted via TUI ({subagent_id})")]
|
||||
@@ -24,9 +24,15 @@ class FakeChild:
|
||||
self.provider = "test"
|
||||
self.model = "test-model"
|
||||
self.interrupted = False
|
||||
self.interrupt_kind = None
|
||||
|
||||
def interrupt(self, _reason):
|
||||
self.interrupted = True
|
||||
self.interrupt_kind = "soft"
|
||||
|
||||
def hard_interrupt(self, _reason):
|
||||
self.interrupted = True
|
||||
self.interrupt_kind = "hard"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -76,6 +82,17 @@ def test_cancel_is_cooperative_and_forged_handle_is_unknown(lifecycle):
|
||||
assert other_service.status(handle).state is SubagentState.UNKNOWN
|
||||
|
||||
|
||||
def test_cancel_uses_explicit_hard_interrupt(lifecycle):
|
||||
handle = lifecycle.launch(SubagentLaunchRequest(goal="x"))
|
||||
record = lifecycle._record(handle)
|
||||
assert record is not None and record.agent is not None
|
||||
|
||||
assert lifecycle.cancel(handle, reason="explicit user cancel").accepted
|
||||
|
||||
assert record.agent.interrupt_kind == "hard"
|
||||
lifecycle.wait(handle, timeout_seconds=1)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -27,6 +27,7 @@ class TestInterruptPropagationToChild(unittest.TestCase):
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
agent._interrupt_requested = False
|
||||
agent._interrupt_message = None
|
||||
agent._hard_interrupt_requested = threading.Event()
|
||||
agent._execution_thread_id = None
|
||||
agent._interrupt_thread_signal_pending = False
|
||||
agent._active_children = []
|
||||
@@ -54,6 +55,38 @@ class TestInterruptPropagationToChild(unittest.TestCase):
|
||||
assert is_interrupted() is False
|
||||
assert parent._interrupt_thread_signal_pending is True
|
||||
|
||||
def test_hard_cancel_is_explicit_atomic_and_propagated(self):
|
||||
parent = self._make_bare_agent()
|
||||
child = self._make_bare_agent()
|
||||
parent._active_children.append(child)
|
||||
|
||||
parent.interrupt("Stop requested", hard_cancel=True)
|
||||
|
||||
assert parent._hard_interrupt_requested.is_set()
|
||||
assert child._hard_interrupt_requested.is_set()
|
||||
parent.clear_interrupt()
|
||||
assert not parent._hard_interrupt_requested.is_set()
|
||||
|
||||
def test_message_interrupt_does_not_set_hard_cancel(self):
|
||||
agent = self._make_bare_agent()
|
||||
|
||||
agent.interrupt("new user message")
|
||||
|
||||
assert agent._interrupt_requested is True
|
||||
assert not agent._hard_interrupt_requested.is_set()
|
||||
|
||||
def test_active_turn_redirect_does_not_set_hard_cancel(self):
|
||||
agent = self._make_bare_agent()
|
||||
agent._model_request_active = threading.Event()
|
||||
agent._model_request_active.set()
|
||||
agent._pending_redirect = None
|
||||
|
||||
assert agent.redirect("new correction") is True
|
||||
|
||||
assert agent._interrupt_requested is True
|
||||
assert agent._interrupt_message is None
|
||||
assert not agent._hard_interrupt_requested.is_set()
|
||||
|
||||
def test_child_clear_interrupt_at_start_clears_thread(self):
|
||||
"""child.clear_interrupt() at start of run_conversation clears the
|
||||
bound execution thread's interrupt flag.
|
||||
|
||||
@@ -6,6 +6,10 @@ import sys
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from tui_gateway.compute_host import ComputeHost, HostSession
|
||||
|
||||
|
||||
def _stdout_queue(proc: subprocess.Popen) -> queue.Queue[dict]:
|
||||
out: queue.Queue[dict] = queue.Queue()
|
||||
@@ -83,3 +87,42 @@ def test_compute_host_line_json_seed_turn_interrupt():
|
||||
finally:
|
||||
if proc.poll() is None:
|
||||
proc.kill()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", ["legacy", "hard-only", "dynamic-getattr"])
|
||||
def test_compute_host_interrupt_uses_explicit_stop_compatibility(kind):
|
||||
calls = []
|
||||
|
||||
class _Legacy:
|
||||
def interrupt(self):
|
||||
calls.append("legacy")
|
||||
|
||||
class _HardOnly:
|
||||
def hard_interrupt(self):
|
||||
calls.append("hard")
|
||||
|
||||
class _Dynamic:
|
||||
def interrupt(self):
|
||||
calls.append("legacy")
|
||||
|
||||
def __getattr__(self, name):
|
||||
if name == "hard_interrupt":
|
||||
return lambda: calls.append("fabricated-hard")
|
||||
raise AttributeError(name)
|
||||
|
||||
agent = {
|
||||
"legacy": _Legacy(),
|
||||
"hard-only": _HardOnly(),
|
||||
"dynamic-getattr": _Dynamic(),
|
||||
}[kind]
|
||||
host = ComputeHost(heartbeat_secs=0)
|
||||
host._sessions["s1"] = HostSession(sid="s1", agent=agent)
|
||||
emitted = []
|
||||
host.emit = emitted.append
|
||||
try:
|
||||
host._handle_interrupt({"sid": "s1", "request_id": "stop"})
|
||||
finally:
|
||||
host.close()
|
||||
|
||||
assert calls == ["hard" if kind == "hard-only" else "legacy"]
|
||||
assert emitted[-1]["applied"] is True
|
||||
|
||||
@@ -118,7 +118,54 @@ def test_err_envelope(server):
|
||||
}
|
||||
|
||||
|
||||
# ── write_json ───────────────────────────────────────────────────────
|
||||
@pytest.mark.parametrize("kind", ["legacy", "hard-only", "dynamic-getattr"])
|
||||
def test_session_interrupt_uses_explicit_stop_compatibility(server, monkeypatch, kind):
|
||||
calls = []
|
||||
|
||||
class _Legacy:
|
||||
def interrupt(self):
|
||||
calls.append("legacy")
|
||||
|
||||
class _HardOnly:
|
||||
def hard_interrupt(self):
|
||||
calls.append("hard")
|
||||
|
||||
class _Dynamic:
|
||||
def interrupt(self):
|
||||
calls.append("legacy")
|
||||
|
||||
def __getattr__(self, name):
|
||||
if name == "hard_interrupt":
|
||||
return lambda: calls.append("fabricated-hard")
|
||||
raise AttributeError(name)
|
||||
|
||||
agent = {
|
||||
"legacy": _Legacy(),
|
||||
"hard-only": _HardOnly(),
|
||||
"dynamic-getattr": _Dynamic(),
|
||||
}[kind]
|
||||
session = {
|
||||
"agent": agent,
|
||||
"history_lock": threading.Lock(),
|
||||
"running": True,
|
||||
"queued_prompt": "later",
|
||||
"session_key": "session-key",
|
||||
"_run_thread": None,
|
||||
}
|
||||
monkeypatch.setattr(server, "_tts_stream_stop", lambda: None)
|
||||
monkeypatch.setattr(server, "_sess_nowait", lambda _params, _rid: (session, None))
|
||||
monkeypatch.setattr(server, "_sess", lambda _params, _rid: (session, None))
|
||||
monkeypatch.setattr(server, "_session_uses_compute_host", lambda _session: False)
|
||||
monkeypatch.setattr(server, "_clear_pending", lambda _sid: None)
|
||||
response = server._methods["session.interrupt"](
|
||||
"stop", {"session_id": "ui-session"}
|
||||
)
|
||||
|
||||
assert response["result"]["status"] == "interrupted"
|
||||
assert calls == ["hard" if kind == "hard-only" else "legacy"]
|
||||
|
||||
|
||||
# ── write_json ────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_write_json(capture):
|
||||
@@ -552,13 +599,13 @@ def test_skin_live_switch_end_to_end(server, tmp_path, monkeypatch):
|
||||
monkeypatch.setattr(server, "_emit", lambda ev, sid, payload=None: emitted.append((ev, payload)))
|
||||
|
||||
# Baseline (default) — seeds the signature.
|
||||
(tmp_path / "config.yaml").write_text("display:\n skin: default\n")
|
||||
(tmp_path / "config.yaml").write_text("display:\n skin: default\n", encoding="utf-8")
|
||||
server._broadcast_skin_if_changed()
|
||||
emitted.clear()
|
||||
|
||||
# Activate midnight, as `hermes config set display.skin midnight` would.
|
||||
time.sleep(0.01) # ensure the config mtime moves
|
||||
(tmp_path / "config.yaml").write_text("display:\n skin: midnight\n")
|
||||
(tmp_path / "config.yaml").write_text("display:\n skin: midnight\n", encoding="utf-8")
|
||||
server._broadcast_skin_if_changed()
|
||||
|
||||
assert [ev for ev, _ in emitted] == ["skin.changed"]
|
||||
|
||||
@@ -33,6 +33,7 @@ from typing import Any, Dict, List, Optional
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from toolsets import TOOLSETS
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
# Sentinel value used by the runtime provider system for providers that are
|
||||
# not natively known (named custom providers, third-party aggregators, etc.).
|
||||
@@ -196,7 +197,8 @@ def interrupt_subagent(subagent_id: str) -> bool:
|
||||
if agent is None:
|
||||
return False
|
||||
try:
|
||||
agent.interrupt(f"Interrupted via TUI ({subagent_id})")
|
||||
if not request_hard_interrupt(agent, f"Interrupted via TUI ({subagent_id})"):
|
||||
return False
|
||||
except Exception as exc:
|
||||
logger.debug("interrupt_subagent(%s) failed: %s", subagent_id, exc)
|
||||
return False
|
||||
@@ -2193,9 +2195,8 @@ def _run_single_child(
|
||||
except Exception as _timeout_exc:
|
||||
# Signal the child to stop so its thread can exit cleanly.
|
||||
try:
|
||||
if hasattr(child, "interrupt"):
|
||||
child.interrupt()
|
||||
elif hasattr(child, "_interrupt_requested"):
|
||||
interrupted = child is not None and request_hard_interrupt(child)
|
||||
if not interrupted and child is not None and hasattr(child, "_interrupt_requested"):
|
||||
child._interrupt_requested = True
|
||||
except Exception:
|
||||
pass
|
||||
@@ -3275,9 +3276,8 @@ def delegate_task(
|
||||
def _batch_interrupt():
|
||||
for _c in _child_agents:
|
||||
try:
|
||||
if hasattr(_c, "interrupt"):
|
||||
_c.interrupt("Async delegation cancelled")
|
||||
elif hasattr(_c, "_interrupt_requested"):
|
||||
interrupted = request_hard_interrupt(_c, "Async delegation cancelled")
|
||||
if not interrupted and hasattr(_c, "_interrupt_requested"):
|
||||
_c._interrupt_requested = True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -21,6 +21,8 @@ from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
|
||||
def now_ns() -> int:
|
||||
return time.perf_counter_ns()
|
||||
@@ -37,7 +39,7 @@ class SpikeAgent:
|
||||
def clear_interrupt(self) -> None:
|
||||
self._interrupt.clear()
|
||||
|
||||
def interrupt(self) -> None:
|
||||
def interrupt(self, *, hard_cancel: bool = False) -> None:
|
||||
self._interrupt.set()
|
||||
|
||||
def run_conversation(
|
||||
@@ -257,7 +259,7 @@ class ComputeHost:
|
||||
sid = str(frame.get("sid") or "")
|
||||
spike = self._sessions.get(sid)
|
||||
if spike is not None:
|
||||
spike.agent.interrupt()
|
||||
request_hard_interrupt(spike.agent)
|
||||
self.emit(
|
||||
{
|
||||
"type": "interrupt.ack",
|
||||
@@ -276,8 +278,8 @@ class ComputeHost:
|
||||
self.emit({"type": "interrupt.ack", "sid": sid, "request_id": frame.get("request_id"), "applied": False})
|
||||
return
|
||||
agent = session.get("agent")
|
||||
if agent is not None and hasattr(agent, "interrupt"):
|
||||
agent.interrupt()
|
||||
if agent is not None:
|
||||
request_hard_interrupt(agent)
|
||||
with session.get("history_lock", threading.Lock()):
|
||||
session["_turn_cancel_requested"] = True
|
||||
session["queued_prompt"] = None
|
||||
|
||||
@@ -2749,8 +2749,10 @@ def _(rid, params: dict) -> dict:
|
||||
session["queued_prompt"] = None
|
||||
session.pop("queued_prompts", None)
|
||||
session["_queued_prompt_generation"] = int(session.get("_queued_prompt_generation", 0)) + 1
|
||||
if should_interrupt and hasattr(session["agent"], "interrupt"):
|
||||
session["agent"].interrupt()
|
||||
if should_interrupt:
|
||||
from agent.interrupt_compat import request_hard_interrupt
|
||||
|
||||
request_hard_interrupt(session["agent"])
|
||||
if not run_thread_alive:
|
||||
with session["history_lock"]:
|
||||
if session.get("running"):
|
||||
|
||||
Reference in New Issue
Block a user